anonymous-CAP commited on
Commit
905a517
·
verified ·
1 Parent(s): 23941a0

Upload 3 files

Browse files
Files changed (3) hide show
  1. config.json +20 -0
  2. model.safetensors +3 -0
  3. modeling_cap.py +37 -0
config.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "CAPModel"
4
+ ],
5
+ "base_model_name": "roberta-base",
6
+ "dropout": 0.1,
7
+ "dtype": "float32",
8
+ "id2label": {
9
+ "0": "LABEL_0",
10
+ "1": "LABEL_1",
11
+ "2": "LABEL_2"
12
+ },
13
+ "label2id": {
14
+ "LABEL_0": 0,
15
+ "LABEL_1": 1,
16
+ "LABEL_2": 2
17
+ },
18
+ "model_type": "cap",
19
+ "transformers_version": "5.14.1"
20
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9b25686d881a0fc00917267ea627895c56411634451e5d469320ef7099317b83
3
+ size 498616084
modeling_cap.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # modeling_cap.py
2
+ import torch
3
+ import torch.nn as nn
4
+ from transformers import AutoModel, PreTrainedModel, PretrainedConfig
5
+
6
+
7
+ class CAPConfig(PretrainedConfig):
8
+ model_type = "cap"
9
+
10
+ def __init__(self, base_model_name="roberta-base", num_labels=3, dropout=0.1, **kwargs):
11
+ super().__init__(**kwargs)
12
+ self.base_model_name = base_model_name
13
+ self.num_labels = num_labels
14
+ self.dropout = dropout
15
+
16
+
17
+ class CAPModel(PreTrainedModel):
18
+ config_class = CAPConfig
19
+ base_model_prefix = "backbone"
20
+
21
+ def __init__(self, config):
22
+ super().__init__(config)
23
+ self.backbone = AutoModel.from_pretrained(config.base_model_name)
24
+ hidden_size = self.backbone.config.hidden_size
25
+ self.num_labels = config.num_labels
26
+ self.dropout = nn.Dropout(config.dropout)
27
+ self.head = nn.Linear(hidden_size, config.num_labels)
28
+
29
+ def forward(self, input_ids, attention_mask, token_valid_mask):
30
+ outputs = self.backbone(input_ids=input_ids, attention_mask=attention_mask)
31
+ subword_states = self.dropout(outputs.last_hidden_state)
32
+ token_logits = self.head(subword_states)
33
+ valid_mask = token_valid_mask.unsqueeze(-1)
34
+ masked_token_logits = token_logits * valid_mask
35
+ valid_counts = token_valid_mask.sum(dim=1, keepdim=True).clamp(min=1e-9)
36
+ sequence_logits = masked_token_logits.sum(dim=1) / valid_counts
37
+ return sequence_logits, token_logits