diff --git a/src/pet/pet_model.py b/src/pet/pet_model.py index dbe2b6e..f71aa04 100644 --- a/src/pet/pet_model.py +++ b/src/pet/pet_model.py @@ -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, + ) diff --git a/src/pet/utils/config.py b/src/pet/utils/config.py index c6d4588..76f3617 100644 --- a/src/pet/utils/config.py +++ b/src/pet/utils/config.py @@ -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 diff --git a/src/pet/utils/save_and_load.py b/src/pet/utils/save_and_load.py index ea2160d..9ac53db 100644 --- a/src/pet/utils/save_and_load.py +++ b/src/pet/utils/save_and_load.py @@ -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: