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

PyTorch Transformer逐Token生成时如何解决批量形状不匹配问题?

修复Transformer自编码器逐Token生成的形状不匹配问题

你的代码存在两个核心问题导致非训练阶段形状不匹配:

  • 未给Transformer层设置batch_first=True,输入维度顺序与PyTorch默认要求不符
  • _decode_token_by_token方法中存在索引越界,且输入序列构建逻辑错误

以下是修复后的完整代码:

# model.py
import torch
import torch.nn as nn
import math


class TransformerAutoencoder(nn.Module):
    def __init__(self, d_model, nhead, num_layers, dim_feedforward, bottleneck_size, dropout=0.5):
        super(TransformerAutoencoder, self).__init__()

        # 添加batch_first=True,适配(batch_size, seq_len, d_model)格式的输入
        self.encoder = nn.TransformerEncoder(
            encoder_layer=nn.TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True),
            num_layers=num_layers
        )

        self.decoder = nn.TransformerDecoder(
            decoder_layer=nn.TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True),
            num_layers=num_layers
        )

        self.bottleneck = nn.Linear(d_model, bottleneck_size)
        self.bottleneck_expansion = nn.Linear(bottleneck_size, d_model)
        self.dropout = nn.Dropout(dropout)
        self.d_model = d_model
        self.relu = nn.ReLU()
        self.EOS_token = -1.0  # Define the EOS token as a constant

    def forward(self, src):
        num_time_frames = src.size(1)

        # Generate sinusoidal position embeddings
        position_embeddings = self._get_sinusoidal_position_embeddings(num_time_frames, self.d_model).to(src.device)

        # Add position embeddings to input, shape: (batch_size, num_time_frames, d_model)
        src = src + position_embeddings

        # Pass the input through the encoder, shape: (batch_size, num_time_frames, d_model)
        encoded = self.encoder(src)

        # Pass the encoded output through the bottleneck layer, shape: (batch_size, num_time_frames, bottleneck_size)
        bottleneck_output = self.bottleneck(encoded)
        bottleneck_output = self.dropout(bottleneck_output)

        # Expand the bottleneck output back to the original dimension, shape: (batch_size, num_time_frames, d_model)
        expanded = self.bottleneck_expansion(bottleneck_output)
        expanded = self.dropout(expanded)

        # Pass the expanded output through the decoder, shape: (batch_size, num_time_frames, d_model)
        if self.training:
            decoded = self.decoder(expanded, encoded)
        else:
            decoded = self._decode_token_by_token(expanded, encoded)

        # Apply the ReLU activation to the decoded output
        decoded = self.relu(decoded)

        return decoded, bottleneck_output

    def _decode_token_by_token(self, expanded, encoded):
        batch_size, num_time_frames, d_model = expanded.size()
        decoded = torch.full_like(expanded, self.EOS_token)
        
        # 初始输入取expanded的第一个时间步,生成第一个token
        current_input = expanded[:, :1]
        decoded[:, 0] = self.decoder(current_input, encoded)[:, 0]

        for t in range(1, num_time_frames):
            # 将之前生成的所有token作为当前输入
            current_input = decoded[:, :t]
            # 解码后取输出的最后一个token作为当前步结果,避免索引越界
            next_token = self.decoder(current_input, encoded)[:, -1]
            decoded[:, t] = next_token

        return decoded

    def _get_sinusoidal_position_embeddings(self, num_positions, d_model):
        position_embeddings = torch.zeros(num_positions, d_model)
        positions = torch.arange(0, num_positions, dtype=torch.float).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model))

        position_embeddings[:, 0::2] = torch.sin(positions * div_term)
        position_embeddings[:, 1::2] = torch.cos(positions * div_term)
        position_embeddings = position_embeddings.unsqueeze(0)

        return position_embeddings

关键修复说明:

  • Transformer层维度适配:在TransformerEncoderLayer和TransformerDecoderLayer中添加batch_first=True,确保模型接受(batch_size, seq_len, d_model)格式的输入,与代码中的张量维度一致,避免维度顺序错误导致的形状不匹配。
  • 逐Token生成逻辑修正:
    • 简化初始输入与赋值逻辑,直接生成第一个token
    • 后续每一步用之前所有生成的token作为输入,解码后取输出的最后一个token对应当前生成位置,彻底避免索引越界问题
    • 输入序列长度与当前生成步数严格匹配,符合自回归生成的逻辑

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 16:29:55