BERT情感分类训练报错TypeError: dropout(): argument 'input' (position 1) must be Tensor, not str
BERT情感分类训练报错TypeError: dropout(): argument 'input' (position 1) must be Tensor, not str
嗨,我看了你的代码,问题出在SentimentClassifier的forward函数里的变量名笔误,和Transformers版本关系不大,降级到3.x自然解决不了~
错误根源分析
你在forward函数里,先把BERT的返回值解构为_, pooled_output,但后面却试图调用一个从未定义过的bertOutput['pooler_output']作为dropout的输入。如果bertOutput在你的代码上下文里是个字符串(比如不小心把变量名写成了字符串字面量,或者之前误定义了这个变量为字符串),那dropout自然会收到字符串类型的输入,从而抛出这个错误。
修正后的代码
根据你使用的Transformers版本,有两种修复方式:
方式一:兼容新版Transformers(推荐)
新版Transformers中,BertModel的forward返回的是一个模型输出对象,我们直接访问它的pooler_output属性即可:
# Build the Sentiment Classifier class class SentimentClassifier(nn.Module): # Constructor class def __init__(self, n_classes): super(SentimentClassifier, self).__init__() self.bert = BertModel.from_pretrained(MODEL_NAME) self.drop = nn.Dropout(p=0.3) self.out = nn.Linear(self.bert.config.hidden_size, n_classes) # Forward propagation class def forward(self, input_ids, attention_mask): # 获取BERT的完整输出对象 bert_output = self.bert( input_ids=input_ids, attention_mask=attention_mask ) # 提取池化输出(Tensor类型) pooled_output = bert_output.pooler_output # 过dropout层 output = self.drop(pooled_output) return self.out(output)
方式二:适配旧版Transformers(3.x及更早)
旧版中BertModel的forward返回元组(last_hidden_state, pooled_output),直接使用你已经解构好的pooled_output即可:
# Build the Sentiment Classifier class class SentimentClassifier(nn.Module): # Constructor class def __init__(self, n_classes): super(SentimentClassifier, self).__init__() self.bert = BertModel.from_pretrained(MODEL_NAME) self.drop = nn.Dropout(p=0.3) self.out = nn.Linear(self.bert.config.hidden_size, n_classes) # Forward propagation class def forward(self, input_ids, attention_mask): # 旧版BERT返回(last_hidden_state, pooled_output) _, pooled_output = self.bert( input_ids=input_ids, attention_mask=attention_mask ) # 直接使用已获取的pooled_output,不要用未定义的bertOutput output = self.drop(pooled_output) return self.out(output)
额外提示
- 以后遇到这类类型错误,优先检查输入到函数的变量类型和变量名是否正确,尤其是PyTorch的层(比如dropout)只能接收Tensor类型输入。
- 如果你不确定BERT的返回值结构,可以在forward里加个
print(type(bert_output))或者print(dir(bert_output))来查看属性,避免踩坑。
备注:内容来源于stack exchange,提问作者Laura Valentini
相关产品推荐
相关产品推荐

