mirror of
https://github.com/wassname/peft.git
synced 2026-09-09 11:28:32 +08:00
Adding PETModelForTokenClassification and 🐛 fixes
This commit is contained in:
+155
-1
@@ -4,7 +4,7 @@ import warnings
|
||||
import torch
|
||||
from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
|
||||
from transformers import PreTrainedModel
|
||||
from transformers.modeling_outputs import SequenceClassifierOutput
|
||||
from transformers.modeling_outputs import SequenceClassifierOutput, TokenClassifierOutput
|
||||
|
||||
from .tuners import LoRAModel, PrefixEncoder, PromptEmbedding, PromptEncoder
|
||||
from .utils import PETConfig, PETType, TaskType, _set_trainable, shift_tokens_right
|
||||
@@ -502,3 +502,157 @@ class PETModelForSeq2SeqLM(PETModel):
|
||||
(prompts[:, self.pet_config.num_virtual_tokens :], decoder_inputs_embeds), dim=1
|
||||
)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, decoder_inputs_embeds=decoder_inputs_embeds, **kwargs)
|
||||
|
||||
|
||||
class PETModelForTokenClassification(PETModel):
|
||||
"""
|
||||
PET model for sequence classification tasks.
|
||||
|
||||
Args:
|
||||
model (:obj:`PreTrainedModel`): Base transformer model
|
||||
pet_config (:obj:`PETConfig`): PET config.
|
||||
|
||||
Attributes:
|
||||
config (:obj:`PretrainedConfig`): The configuration object of the base model. cls_layer_name (:obj:`str`): The
|
||||
name of the classification layer.
|
||||
|
||||
Example::
|
||||
|
||||
>>> from transformers import AutoModelForSequenceClassification >>> from pet import
|
||||
PETModelForTokenClassification, get_pet_config >>> config = {
|
||||
'pet_type': 'PREFIX_TUNING', 'task_type': 'TOKEN_CLS', 'inference_mode': False, 'num_virtual_tokens':
|
||||
20, 'token_dim': 768, 'num_transformer_submodules': 1, 'num_attention_heads': 12, 'num_layers': 12,
|
||||
'encoder_hidden_size': 768, 'prefix_projection': False, 'postprocess_past_key_value_function': None
|
||||
}
|
||||
>>> pet_config = get_pet_config(config) >>> model =
|
||||
AutoModelForSequenceClassification.from_pretrained("bert-base-cased") >>> pet_model =
|
||||
PETModelForSequenceClassification(model, pet_config) >>> pet_model.print_trainable_parameters() trainable
|
||||
params: 370178 || all params: 108680450 || trainable%: 0.3406113979101117
|
||||
"""
|
||||
|
||||
def __init__(self, model, pet_config: PETConfig):
|
||||
super().__init__(model, pet_config)
|
||||
self.modules_to_save = ["classifier"]
|
||||
|
||||
for name, module in self.base_model.named_children():
|
||||
if isinstance(module, torch.nn.Linear):
|
||||
self.cls_layer_name = name
|
||||
break
|
||||
|
||||
# to make sure classifier layer is trainable
|
||||
_set_trainable(self.base_model)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
inputs_embeds=None,
|
||||
labels=None,
|
||||
output_attentions=None,
|
||||
output_hidden_states=None,
|
||||
return_dict=None,
|
||||
**kwargs,
|
||||
):
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if self.pet_config.pet_type == PETType.LORA:
|
||||
return self.base_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
labels=labels,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=return_dict,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
batch_size = input_ids.shape[0]
|
||||
if attention_mask is not None:
|
||||
# concat prompt attention mask
|
||||
prefix_attention_mask = torch.ones(batch_size, self.pet_config.num_virtual_tokens).to(self.device)
|
||||
attention_mask = torch.cat((prefix_attention_mask, attention_mask), dim=1)
|
||||
if kwargs.get("position_ids", None) is not None:
|
||||
warnings.warn("Position ids are not supported for parameter efficient tuning. Ignoring position ids.")
|
||||
kwargs["position_ids"] = None
|
||||
kwargs.update(
|
||||
{
|
||||
"attention_mask": attention_mask,
|
||||
"labels": labels,
|
||||
"output_attentions": output_attentions,
|
||||
"output_hidden_states": output_hidden_states,
|
||||
"return_dict": return_dict,
|
||||
}
|
||||
)
|
||||
|
||||
if self.pet_config.pet_type == PETType.PREFIX_TUNING:
|
||||
return self._prefix_tuning_forward(input_ids=input_ids, **kwargs)
|
||||
else:
|
||||
if kwargs.get("token_type_ids", None) is not None:
|
||||
kwargs["token_type_ids"] = torch.cat(
|
||||
(
|
||||
torch.zeros(batch_size, self.pet_config.num_virtual_tokens).to(self.device),
|
||||
kwargs["token_type_ids"],
|
||||
),
|
||||
dim=1,
|
||||
).long()
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
prompts = self.get_prompt(batch_size=batch_size)
|
||||
inputs_embeds = torch.cat((prompts, inputs_embeds), dim=1)
|
||||
return self.base_model(inputs_embeds=inputs_embeds, **kwargs)
|
||||
|
||||
def _prefix_tuning_forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
inputs_embeds=None,
|
||||
labels=None,
|
||||
output_attentions=None,
|
||||
output_hidden_states=None,
|
||||
return_dict=None,
|
||||
**kwargs,
|
||||
):
|
||||
batch_size = input_ids.shape[0]
|
||||
past_key_values = self.get_prompt(batch_size)
|
||||
fwd_params = list(inspect.signature(self.base_model.forward).parameters.keys())
|
||||
kwargs.update(
|
||||
{
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"inputs_embeds": inputs_embeds,
|
||||
"output_attentions": output_attentions,
|
||||
"output_hidden_states": output_hidden_states,
|
||||
"return_dict": return_dict,
|
||||
"past_key_values": past_key_values,
|
||||
}
|
||||
)
|
||||
if "past_key_values" in fwd_params:
|
||||
return self.base_model(labels=labels, **kwargs)
|
||||
else:
|
||||
transformer_backbone_name = self.base_model.get_submodule(self.transformer_backbone_name)
|
||||
fwd_params = list(inspect.signature(transformer_backbone_name.forward).parameters.keys())
|
||||
if "past_key_values" not in fwd_params:
|
||||
raise ValueError("Model does not support past key values which are required for prefix tuning.")
|
||||
outputs = transformer_backbone_name(**kwargs)
|
||||
sequence_output = outputs[0]
|
||||
if "dropout" in [name for name, _ in list(self.base_model.named_children())]:
|
||||
sequence_output = self.base_model.dropout(sequence_output)
|
||||
logits = self.base_model.get_submodule(self.cls_layer_name)(sequence_output)
|
||||
|
||||
loss = None
|
||||
loss = None
|
||||
if labels is not None:
|
||||
loss_fct = CrossEntropyLoss()
|
||||
loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
|
||||
|
||||
if not return_dict:
|
||||
output = (logits,) + outputs[2:]
|
||||
return ((loss,) + output) if loss is not None else output
|
||||
|
||||
return TokenClassifierOutput(
|
||||
loss=loss,
|
||||
logits=logits,
|
||||
hidden_states=outputs.hidden_states,
|
||||
attentions=outputs.attentions,
|
||||
)
|
||||
|
||||
@@ -14,6 +14,7 @@ class TaskType(str, enum.Enum):
|
||||
SEQ_CLS = "SEQ_CLS"
|
||||
SEQ_2_SEQ_LM = "SEQ_2_SEQ_LM"
|
||||
CAUSAL_LM = "CAUSAL_LM"
|
||||
TOKEN_CLS = "TOKEN_CLS"
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -10,12 +10,11 @@ def get_pet_model_state_dict(model):
|
||||
Args:
|
||||
model (:obj:`PETModel`): The PET model.
|
||||
"""
|
||||
|
||||
state_dict = model.state_dict()
|
||||
if model.pet_config.pet_type == PETType.LORA:
|
||||
to_return = lora_state_dict(model, bias=model.pet_config.bias)
|
||||
else:
|
||||
to_return = {}
|
||||
state_dict = model.state_dict()
|
||||
prompt_embeddings = model.get_prompt_embedding_to_save()
|
||||
to_return["prompt_embeddings"] = prompt_embeddings
|
||||
if model.modules_to_save is not None:
|
||||
|
||||
Reference in New Issue
Block a user