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

训练时如何拼接不同形状的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 10:24:00