如何从DebertaModel获取池化输出(不使用DebertaForSequenceClassification)
从DebertaModel获取池化输出并用于自定义分类模型
你可以通过两种方式获取池化输出,灵活接入自定义分类层:
1. 使用模型自带的池化输出
DebertaModel的前向输出是BaseModelOutputWithPoolingAndCrossAttentions对象,其中的pooler_output是[CLS] token经过官方预定义的全连接+Tanh激活池化后的结果,可直接作为分类任务输入。
示例代码:
from transformers import DebertaModel, DebertaTokenizer import torch import torch.nn as nn # 加载预训练组件 tokenizer = DebertaTokenizer.from_pretrained('microsoft/deberta-base') deberta_model = DebertaModel.from_pretrained('microsoft/deberta-base') # 自定义分类模型 class CustomTextClassifier(nn.Module): def __init__(self, hidden_dim, num_classes): super().__init__() self.deberta = deberta_model self.classifier = nn.Linear(hidden_dim, num_classes) def forward(self, input_ids, attention_mask=None): # 提取官方池化输出 outputs = self.deberta(input_ids=input_ids, attention_mask=attention_mask) pooled_output = outputs.pooler_output # 传入自定义分类层 logits = self.classifier(pooled_output) return logits # 测试调用 sample_text = "测试文本内容" inputs = tokenizer(sample_text, return_tensors="pt") classifier = CustomTextClassifier(hidden_dim=768, num_classes=2) pred_logits = classifier(**inputs)
2. 自定义池化逻辑(更灵活)
如果不满足官方池化方式,可直接提取[CLS] token的最后一层隐藏状态,或对所有token的隐藏状态做均值/最大池化:
示例代码(以[CLS] token隐藏状态为例):
class CustomTextClassifier(nn.Module): def __init__(self, hidden_dim, num_classes): super().__init__() self.deberta = deberta_model self.classifier = nn.Linear(hidden_dim, num_classes) def forward(self, input_ids, attention_mask=None): outputs = self.deberta(input_ids=input_ids, attention_mask=attention_mask) # 提取[CLS] token的最后一层隐藏状态 cls_hidden = outputs.last_hidden_state[:, 0, :] # 也可改用均值池化: # cls_hidden = torch.mean(outputs.last_hidden_state, dim=1) logits = self.classifier(cls_hidden) return logits
内容的提问来源于stack exchange,提问作者tony
相关产品推荐
相关产品推荐

