Hate_bert_lambda1 / modeling_cap.py
anonymous-CAP's picture
Upload 3 files
02b733e verified
Raw
History Blame Contribute Delete
1.53 kB
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