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
相关产品推荐
相关产品推荐

