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

PyTorch中在多语言BERT上层搭建RNN层出现类型报错如何解决

报错原因

触发TypeError: relu(): argument 'input' (position 1) must be Tensor, not tuple的直接原因是torch.nn.RNN的前向传播返回值为元组而非单个张量:

  • 元组第一个元素是RNN所有时间步的输出特征,未设置batch_first=True时形状为(seq_len, batch_size, hidden_size)
  • 元组第二个元素是RNN最后一个时间步的隐状态,形状为(num_layers * num_directions, batch_size, hidden_size)
    你直接将RNN返回的元组传入ReLU激活层,自然会触发类型错误。

除此之外代码还存在隐藏的维度适配问题:当前你取的是BERT的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token池化输出,形状为二维张量(batch_size, 768),但RNN层默认要求输入为三维张量(seq_len, batch_size, input_size),就算修正元组问题,后续也会触发维度不匹配报错。且单条<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>聚合表征本身已经是整句全局特征,额外接RNN没有实际序列建模的意义。

修正方案

根据实际使用场景二选一即可:

场景1:使用BERT输出的全token序列表征喂入RNN(符合RNN设计逻辑)

这种场景需要取BERT返回的全token序列输出,同时给RNN添加batch_first=True参数适配维度顺序,解包RNN返回值后取最后一个时间步的输出作为整句表征,再接后续全连接层。修正后代码如下:

class BERTClass(torch.nn.Module):
    def __init__(self):
        super(BERTClass, self).__init__()
        self.l1 = BertModel.from_pretrained('bert-base-multilingual-cased', return_dict=False)
        # 需要冻结BERT参数时取消下方注释
        # for param in self.l1.parameters():
        #   param.requires_grad = False
        self.l2 = torch.nn.Dropout(0.4)
        # 添加batch_first=True适配(batch_size, seq_len, 768)的输入维度顺序
        self.l3 = torch.nn.RNN(768, 1028, batch_first=True)
        self.activation = torch.nn.ReLU()
        self.l4 = torch.nn.Dropout(0.2)
        self.l5 = torch.nn.Linear(1028, 128)
        self.activation2 = torch.nn.ReLU()
        self.l6 = torch.nn.Linear(128, 10)
        
    
    def forward(self, ids, mask, token_type_ids):
        # 取BERT第一个返回值:所有token的序列输出,形状为(batch_size, seq_len, 768)
        sequence_output, _ = self.l1(ids, attention_mask=mask, token_type_ids=token_type_ids)
        output_2 = self.l2(sequence_output)
        # 解包RNN返回的元组
        rnn_out, h_n = self.l3(output_2)
        # 取最后一个时间步的输出作为整句表征,形状为(batch_size, 1028)
        output3 = rnn_out[:, -1, :]
        act = self.activation(output3)
        output4 = self.l4(act)
        output5 = self.l5(output4)
        act2 = self.activation2(output5)
        output6 = self.l6(act2)
        return output6

model = BERTClass()

场景2:仅使用BERT的CLS池化输出做分类

这种场景下RNN层完全冗余,直接删除即可,避免不必要的计算和维度问题:

class BERTClass(torch.nn.Module):
    def __init__(self):
        super(BERTClass, self).__init__()
        self.l1 = BertModel.from_pretrained('bert-base-multilingual-cased', return_dict=False)
        self.l2 = torch.nn.Dropout(0.4)
        self.activation = torch.nn.ReLU()
        self.l4 = torch.nn.Dropout(0.2)
        self.l5 = torch.nn.Linear(768, 128)
        self.activation2 = torch.nn.ReLU()
        self.l6 = torch.nn.Linear(128, 10)
        
    
    def forward(self, ids, mask, token_type_ids):
        _, pooled_output = self.l1(ids, attention_mask=mask, token_type_ids=token_type_ids)
        output_2 = self.l2(pooled_output)
        act = self.activation(output_2)
        output4 = self.l4(act)
        output5 = self.l5(output4)
        act2 = self.activation2(output5)
        output6 = self.l6(act2)
        return output6

model = BERTClass()
补充说明
  • 短文本分类任务中,直接使用BERT的CLS输出、或对BERT全token输出做均值/最大池化的效果已经足够,堆叠RNN不一定能带来效果提升,反而会增加训练开销和过拟合风险。
  • 如果需要使用双向/多层RNN,初始化时可传入bidirectional=True、num_layers=2等参数,注意此时RNN输出维度会变为hidden_size * num_directions,后续全连接层的输入维度需要对应调整。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 10:45:49