训练时如何拼接不同形状的PyTorch张量且不影响反向传播?
问题解决说明
1. 张量形状不匹配的处理方案
你遇到的形状不匹配问题原因很明确:
hs(last_hidden_state)的形状是[batch_size, 序列长度, 768],对应每个token的输出特征cls_hs(pooler_output)的形状是[batch_size, 768],是整个句子的全局聚合特征
要拼接二者,只需要把cls_hs在序列维度上做扩展,让它和hs的序列长度一致,每个token位置都拼接上全局的cls特征即可,用PyTorch的unsqueeze + expand操作就能实现。
2. 拼接后反向传播是否正常
可以正常运行。PyTorch的拼接、维度扩展操作都是原生支持自动微分的,只要你保证拼接后的张量维度和entity_out线性层的输入维度匹配,反向传播过程不会有任何问题,梯度会正常回传到BERT以及所有涉及的层。
3. 完整修改后的代码
你需要额外调整entity_out层的输入维度,因为拼接后每个token的特征维度变成了768+768=1536,不能再用原来的768作为输入维度。
import torch import torch.nn as nn import transformers class NLUModel(nn.Module): def __init__(self, num_entity, num_intent, num_scenarios): super(NLUModel, self).__init__() self.num_entity = num_entity self.num_intent = num_intent self.num_scenario = num_scenarios self.bert = transformers.BertModel.from_pretrained(config.BASE_MODEL) self.dropout1 = nn.Dropout(0.3) self.dropout2 = nn.Dropout(0.3) self.dropout3 = nn.Dropout(0.3) # 输入维度调整为768*2,适配拼接后的特征长度 self.entity_out = nn.Linear(768 * 2, self.num_entity) self.intent_out = nn.Linear(768, self.num_intent) self.scenario_out = nn.Linear(768, self.num_scenario) def forward(self, ids, mask, token_type_ids): out = self.bert(input_ids=ids, attention_mask=mask, token_type_ids=token_type_ids) hs, cls_hs = out['last_hidden_state'], out['pooler_output'] entity_hs = self.dropout1(hs) intent_hs = self.dropout2(cls_hs) scenario_hs = self.dropout3(cls_hs) # 维度扩展+拼接逻辑 # 给intent_hs增加序列维度:[batch,768] -> [batch,1,768] intent_hs_expand = intent_hs.unsqueeze(1) # 扩展到和entity_hs一致的序列长度:[batch,1,768] -> [batch, seq_len,768] intent_hs_expand = intent_hs_expand.expand(-1, entity_hs.size(1), -1) # 最后一个维度拼接得到[batch, seq_len, 1536] concat_entity_input = torch.cat([entity_hs, intent_hs_expand], dim=-1) entity_hs = self.entity_out(concat_entity_input) intent_hs = self.intent_out(intent_hs) scenario_hs = self.scenario_out(scenario_hs) return entity_hs, intent_hs, scenario_hs
内容的提问来源于stack exchange,提问作者Chandra Irugalbandara
相关产品推荐
相关产品推荐

