File size: 1,534 Bytes
905a517
 
4a26f94
905a517
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4a26f94
 
905a517
 
 
 
4a26f94
905a517
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
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