Team Ai
Modelpublic

fifadxj/tiny-bert-sequence-classification

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes16downloads
modeling_tiny_bert_sequence_classification.py33 linesDownload Raw Back to root
1import torch.nn as nn
2from transformers import AutoModel, PreTrainedModel, BertConfig
3
4
5class TinyBertForSequenceClassification(PreTrainedModel):
6    config_class = BertConfig
7
8    def __init__(self, config):
9        super().__init__(config)
10        self.num_labels = config.num_labels
11        self.bert = AutoModel.from_config(config)
12        self.classifier = nn.Linear(config.hidden_size, config.num_labels)
13        self.post_init()
14
15    def forward(self, input_ids, attention_mask=None, token_type_ids=None, labels=None):
16        # 获取 BERT 的输出
17        outputs = self.bert(
18            input_ids=input_ids,
19            attention_mask=attention_mask,
20            token_type_ids=token_type_ids
21        )
22
23        cls_output = outputs.last_hidden_state[:, 0, :]
24        logits = self.classifier(cls_output)
25
26        # 计算损失(如果提供了标签)
27        loss = None
28        if labels is not None:
29            loss_fct = nn.CrossEntropyLoss()
30            loss = loss_fct(logits, labels)
31
32        return {"loss": loss, "logits": logits} if loss is not None else {"logits": logits}
33