使用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
关键修改点解释
- 替换模型类:用
BertModel替代BertForSequenceClassification,避免使用内置分类头,保留自定义分类逻辑的灵活性; - 正确提取特征:当
return_dict=False时,BertModel的返回值是二元组,通过_, pooled_output取出<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>的池化特征(维度为768,匹配你定义的Linear层输入维度); - 保留自定义层:原有的
dropout和classifier层可正常处理Tensor类型的pooled_output,不会再触发类型错误。
内容的提问来源于stack exchange,提问作者Konder
相关产品推荐
相关产品推荐

