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

单头Transformer模型输出与目标张量尺寸不匹配问题求助

单头Transformer尺寸不匹配问题排查与结构优化

一、核心问题排查方向

尺寸不匹配的根源几乎都来自forward方法中各模块的维度衔接错误,重点检查以下几点:

  • Embedding输出维度与Transformer注意力层的输入维度是否一致
  • 注意力层输出后,序列维度是否与后续全连接层的输入/目标张量维度对齐
  • 训练时输入序列、目标序列的形状是否和模型输出匹配

二、典型错误代码与修正示例

1. 常见错误模型实现(尺寸不匹配根源)

import torch
import torch.nn as nn

class BadSingleHeadTransformer(nn.Module):
    def __init__(self, vocab_size, embed_dim, num_classes):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_dim)
        self.attn = nn.MultiheadAttention(embed_dim, num_heads=1, batch_first=True)
        self.fc = nn.Linear(embed_dim, num_classes)
    
    def forward(self, x):
        x = self.embedding(x)  # 输出形状: [batch_size, seq_len, embed_dim]
        attn_out, _ = self.attn(x, x, x)  # 输出形状: [batch_size, seq_len, embed_dim]
        # 错误:直接将序列维度张量传给全连接,若目标是[batch_size, num_classes],维度完全不匹配
        output = self.fc(attn_out)
        return output

2. 修正后的模型实现

import torch
import torch.nn as nn

class SingleHeadTransformer(nn.Module):
    def __init__(self, vocab_size, embed_dim, seq_len, num_classes, pad_idx=0):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx)
        # 必须加入位置编码,Transformer无法天然捕捉序列顺序
        self.pos_encoding = nn.Parameter(torch.randn(1, seq_len, embed_dim))
        self.attn = nn.MultiheadAttention(embed_dim, num_heads=1, batch_first=True)
        # 残差连接+LayerNorm,避免梯度消失,提升训练稳定性
        self.norm = nn.LayerNorm(embed_dim)
        # 全连接层输入为聚合后的序列特征,维度匹配目标张量
        self.fc = nn.Linear(embed_dim, num_classes)
    
    def forward(self, x, padding_mask=None):
        batch_size, seq_len = x.shape
        # 嵌入层+位置编码
        x = self.embedding(x) + self.pos_encoding[:, :seq_len, :]
        # 计算注意力,传入padding mask避免关注无效pad token
        attn_out, _ = self.attn(x, x, x, key_padding_mask=padding_mask)
        # 残差连接+归一化
        x = self.norm(x + attn_out)
        # 聚合序列维度:用均值池化(也可选用cls token、最大池化)
        x = torch.mean(x, dim=1)  # 输出形状: [batch_size, embed_dim]
        # 全连接层输出匹配目标张量维度
        output = self.fc(x)  # 输出形状: [batch_size, num_classes]
        return output

三、训练循环的维度对齐检查

  • 分类任务:输入张量x形状为[batch_size, seq_len],目标张量y形状为[batch_size]或[batch_size, num_classes],需确保模型输出维度与之匹配
  • 序列生成任务:目标张量形状为[batch_size, target_seq_len],此时模型无需聚合序列维度,全连接层需输出[batch_size, target_seq_len, vocab_size],对应每个时间步的预测

四、针对性报错处理

若报错为RuntimeError: Expected target size [N, C], got [N, S]:

  1. 检查目标张量是否完成正确转换(比如分类任务中是否将序列标签压缩为单标签)
  2. 确认模型forward中是否对Transformer输出的序列维度做了聚合操作(如均值池化、取cls token)

五、结构优化建议

  • 强制加入位置编码:这是Transformer的核心组件,缺失会导致模型完全无法学习序列顺序特征
  • 增加Feed Forward Network(FFN):标准Transformer Encoder包含Attention+FFN,FFN可进一步提取特征,结构为nn.Linear(embed_dim, 4*embed_dim) -> nn.ReLU() -> nn.Linear(4*embed_dim, embed_dim)
  • 完善padding mask处理:避免注意力机制关注padding token,提升计算有效性
  • 加入 dropout:在Embedding层、Attention输出后加入dropout,防止过拟合

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 22:55:14