Merge pull request #7 from huggingface/smangrul/add-examples-fixes-docs

Adding `PETModelForTokenClassification` and 🐛 fixes
This commit is contained in:
Sourab Mangrulkar
2022-12-05 17:08:09 +05:30
committed by GitHub
3 changed files with 157 additions and 3 deletions
+155 -1
View File
@@ -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,
)
+1
View File
@@ -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
+1 -2
View File
@@ -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: