如何对BERT最后4层隐藏层计算均值并保持[2,256,768]维度?
问题分析与解决方案
你的核心问题是错误地使用torch.cat拼接后直接对所有维度求均值,导致输出维度完全不符合预期。
原代码的问题
你将4个[2,256,768]的张量在最后一维拼接,得到[2,256, 768*4]的张量,随后对dim=[0,1,2]求均值,这会把所有维度的数值平均成一个标量,和你想要的[2,256,768]完全不符。
正确实现方式
应该先把4个隐藏层张量在新的维度上堆叠,然后在这个堆叠维度上求均值,这样就能保留原有的batch、sequence、hidden维度:
def __init__(self, bert_model, num_labels): super(BERT_CRF, self).__init__() self.bert = bert_model self.dropout = nn.Dropout(0.25) self.classifier = nn.Linear(768, num_labels) self.crf = CRF(num_labels, batch_first = True) def forward(self, input_ids, attention_mask, labels=None, token_type_ids=None): outputs = self.bert(input_ids, attention_mask=attention_mask) # 堆叠4个隐藏层,新增维度0,得到形状[4, 2, 256, 768] stacked_hidden = torch.stack([outputs[1][-1], outputs[1][-2], outputs[1][-3], outputs[1][-4]], dim=0) # 在堆叠的维度(dim=0)上求均值,得到目标形状[2,256,768] sequence_output = stacked_hidden.mean(dim=0) sequence_output = self.dropout(sequence_output) emission = self.classifier(sequence_output) # 补充CRF的前向逻辑(训练/推理分支) if labels is not None: loss = self.crf(emission, labels, mask=attention_mask.bool(), reduction='mean') return loss else: preds = self.crf.decode(emission, mask=attention_mask.bool()) return preds
补充说明
torch.stack会创建新维度容纳同形状张量,这里用dim=0把4个层放在第一个维度,堆叠后形状为[4, 2, 256, 768]- 在
dim=0上求均值,就是对4个隐藏层的对应位置元素取平均,最终得到和单个隐藏层一致的形状[2,256,768] - 原代码缺少CRF的前向返回逻辑,我补充了常见的训练(返回loss)和推理(返回预测结果)分支,可根据需求调整
内容的提问来源于stack exchange,提问作者MAC
相关产品推荐
相关产品推荐

