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

训练BERT-base-uncased模型时遇Dropout层输入类型错误

问题分析与修复方案

错误根源

你遇到的dropout(): argument 'input' (position 1) must be Tensor, not str错误,主要由两个核心问题导致:

  • Transformers版本不兼容:新版BERT模型不再返回元组格式的输出,原代码通过_, o2 = self.bert(...)的方式会错误获取非Tensor类型的值;
  • 目标变量类型错误:训练数据中的情感标签是字符串("positive"/"negative"),未转换为数值类型,后续处理中引发连锁错误。

具体修复步骤

1. 修正BERT模型输出提取逻辑

新版Transformers的BertModel返回的是包含各类输出属性的对象,而非旧版的元组。修改模型的forward方法:

class BertBaseUncased(nn.Module):
    def __init__(self):
        super(BertBaseUncased,self).__init__()
        self.bert= transformers.BertModel.from_pretrained('bert-base-uncased')
        self.bert_drop=nn.Dropout(0.3)
        self.out= nn.Linear(768,1)
        
    def forward(self,ids,mask,token_type_ids):
        # 获取完整输出对象
        outputs = self.bert(ids,attention_mask=mask,token_type_ids=token_type_ids)
        # 提取池化后的输出(对应原代码的o2)
        o2 = outputs.pooler_output
        bo= self.bert_drop(o2)
        output= self.out(bo)
        return output

2. 正确转换目标标签为数值

取消train函数中注释的标签转换代码,并修正negative对应的数值为0(而非字符串):

def train():
    df= pd.read_csv('/kaggle/input/aamlp-text-data/imdb_folds.csv')
    # 修复标签转换:将字符串转为0/1数值
    df.sentiment= df.sentiment.apply(lambda x: 1 if x=="positive" else 0)
    df_train,df_valid=model_selection.train_test_split(df,test_size=0.1,random_state=42,shuffle=True,stratify=df.sentiment.values)
    # ... 后续代码保持不变 ...

3. 验证数据类型(可选)

可以在BertDataset的__getitem__方法中添加临时打印,确认target的类型:

def __getitem__(self,idx):
    review= str(self.review[idx])
    review= " ".join(review.split())
    inputs= self.tokenizer.encode_plus(review,
                                       None,
                                       add_special_tokens=True,
                                       max_length=self.max_len,
                                       pad_to_max_length=True
    )
    ids= inputs["input_ids"]
    mask= inputs["attention_mask"]
    token_type_ids= inputs["token_type_ids"]
    
    # 临时验证target类型
    target_val = self.target[idx]
    print(f"Target type: {type(target_val)}, value: {target_val}")
    
    return {
        "ids": torch.tensor(ids,dtype=torch.long),
        "mask": torch.tensor(mask,dtype=torch.long),
        "token_type_ids": torch.tensor(token_type_ids,dtype= torch.long),
        "targets": torch.tensor(target_val,dtype=torch.float)
    }

确认所有target都是数值类型后,删除该打印语句即可。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 02:07:13