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