You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Transformers 4.21.2构建BERT文本分类模型遇Tensor/tuple类型错误

解决Bert自定义分类模型中的TypeError问题

问题原因

你遇到的TypeError: dropout(): argument 'input' (position 1) must be Tensor, not tuple错误,核心原因是:

  • 你使用了BertForSequenceClassification模型且设置return_dict=False,它的forward方法返回的是包含logits及其他可选输出的元组,而非单个Tensor;
  • 你错误地将这个元组直接传给nn.Dropout层,而Dropout仅接受Tensor类型输入;
  • 额外说明:BertForSequenceClassification本身已内置分类头,你重复定义classifier层属于冗余设计,并非最佳实践。

修正方案

正确做法是使用BertModel(仅负责文本编码,不带分类头)构建自定义分类模型,修改后的代码如下:

class BertClassificationModel(nn.Module):
    def __init__(self, bert_model_name, num_labels, dropout=0.1):
        super(BertClassificationModel, self).__init__()
        # 替换为BertModel,获取原始文本编码特征
        self.bert = BertModel.from_pretrained(bert_model_name, return_dict=False)
        self.dropout = nn.Dropout(dropout)
        self.classifier = nn.Linear(768, num_labels)
        self.num_labels = num_labels
        
    def forward(self, input_ids, attention_mask=None, token_type_ids=None):
        # return_dict=False时,BertModel返回元组:(last_hidden_state, pooled_output)
        # 取第二个元素作为<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>标记的池化输出
        _, pooled_output = self.bert(input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)
        pooled_output = self.dropout(pooled_output)
        logits = self.classifier(pooled_output)
        return logits

关键修改点解释

  1. 替换模型类:用BertModel替代BertForSequenceClassification,避免使用内置分类头,保留自定义分类逻辑的灵活性;
  2. 正确提取特征:当return_dict=False时,BertModel的返回值是二元组,通过_, pooled_output取出<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>的池化特征(维度为768,匹配你定义的Linear层输入维度);
  3. 保留自定义层:原有的dropout和classifier层可正常处理Tensor类型的pooled_output,不会再触发类型错误。

内容的提问来源于stack exchange,提问作者Konder

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.20 10:01:57