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
相关产品推荐
相关产品推荐

