PyTorch Transformer位置编码器RuntimeError问题求助
Transformer位置编码器RuntimeError修复及参数设置说明
问题场景
输入数据形状为(128,7,21)(batch size=128,序列长度=7,特征数=21),实现位置编码器时触发RuntimeError:
RuntimeError: The expanded size of the tensor (10) must match the existing size (11) at non-singleton dimension 1. Target sizes: [7, 10]. Tensor sizes: [7, 11]
错误出现在pe[:, 1::2] = torch.cos(position * div_term)语句。
错误原因
你的d_model=21是奇数:
torch.arange(0, d_model, 2)生成的是[0,2,...,20],共11个元素,因此position * div_term的形状是(7,11)pe[:,1::2]选取的是索引1、3、...、19的列,共10个列,形状是(7,10)
两者维度不匹配,导致赋值失败。
修复方案
修改positional_encoding方法,处理奇数d_model的情况,同时修正forward方法的维度匹配问题:
import torch import torch.nn as nn import math class PositionalEncoder(nn.Module): def __init__(self, d_model: int, max_seq_len: int=7): super(PositionalEncoder, self).__init__() self.d_model = d_model # 创建位置编码矩阵 pe = self.positional_encoding(max_seq_len, d_model) self.register_buffer('pe', pe) def positional_encoding(self, max_seq_len, d_model): position = torch.arange(0, max_seq_len).unsqueeze(1).float() div_term = torch.exp(torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)) pe = torch.zeros(max_seq_len, d_model) # 填充偶数索引列(0,2,...) pe[:, 0::2] = torch.sin(position * div_term) # 填充奇数索引列(1,3,...),奇数d_model时截断最后一个元素匹配维度 pe[:, 1::2] = torch.cos(position * div_term[:-1]) if d_model % 2 != 0 else torch.cos(position * div_term) return pe def forward(self, x): seq_len = x.size(1) # 适配batch维度,确保位置编码与输入维度匹配后相加 x = x + self.pe[:seq_len, :].unsqueeze(0) return x
max_seq_len参数设置说明
- 如果你的输入序列长度固定为7,当前设置为7完全合理,不会浪费显存。
- 如果后续会处理更长的序列,需要将
max_seq_len设为训练/推理数据中最长的序列长度,或预留少量余量(比如预期最长序列为10,就设为10)。 - 位置编码矩阵在初始化时生成,运行时无法处理超过
max_seq_len的序列,务必提前预估最大序列长度。
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

