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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 10:25:15