| import torch |
| import torch.nn as nn |
| from transformers import AutoModel, AutoConfig, PreTrainedModel, PretrainedConfig |
|
|
|
|
| class CAPConfig(PretrainedConfig): |
| model_type = "cap" |
|
|
| def __init__(self, base_model_name="roberta-base", num_labels=3, dropout=0.1, **kwargs): |
| super().__init__(**kwargs) |
| self.base_model_name = base_model_name |
| self.num_labels = num_labels |
| self.dropout = dropout |
|
|
|
|
| class CAPModel(PreTrainedModel): |
| config_class = CAPConfig |
| base_model_prefix = "backbone" |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| backbone_config = AutoConfig.from_pretrained(config.base_model_name) |
| self.backbone = AutoModel.from_config(backbone_config) |
| hidden_size = self.backbone.config.hidden_size |
| self.num_labels = config.num_labels |
| self.dropout = nn.Dropout(config.dropout) |
| self.head = nn.Linear(hidden_size, config.num_labels) |
| self.post_init() |
|
|
| def forward(self, input_ids, attention_mask, token_valid_mask): |
| outputs = self.backbone(input_ids=input_ids, attention_mask=attention_mask) |
| subword_states = self.dropout(outputs.last_hidden_state) |
| token_logits = self.head(subword_states) |
| valid_mask = token_valid_mask.unsqueeze(-1) |
| masked_token_logits = token_logits * valid_mask |
| valid_counts = token_valid_mask.sum(dim=1, keepdim=True).clamp(min=1e-9) |
| sequence_logits = masked_token_logits.sum(dim=1) / valid_counts |
| return sequence_logits, token_logits |