在Python中如何利用PyTorch实现Transformer位置编码的计算?

冬伟吖_5970

冬伟吖_5970

2026-06-03

465人浏览

原创

torch.nn.embedding不适合实现正弦位置编码,因其引入可学习参数且无法保证公式结构;正确做法是手动向量化计算sin/cos编码并注册为buffer。

在python中如何利用pytorch实现transformer位置编码的计算?

PyTorch中torch.nn.Embedding不适合直接实现正弦位置编码

正弦位置编码(Sinusoidal Positional Encoding)是原始Transformer论文中提出的固定编码方式,它不参与训练、无参数、依赖序列长度和维度。用torch.nn.Embedding加载预计算的位置索引会引入不必要的可学习权重,且无法保证公式结构(如指数衰减的波长)。实际项目中常见错误是误把位置ID喂给Embedding层后当作PE使用——这本质上是学习式位置编码(Learned PE),和Sinusoidal PE行为不同,会影响模型复现与推理一致性。

正确做法是手动构造:按公式 PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) 和 PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model)) 计算,用torch.zeros初始化后逐项填充或向量化赋值。

  • 必须用torch.float32(或torch.float64)精度计算,避免half下指数运算溢出
  • pos范围应覆盖最大可能序列长度(如512),但可在运行时截取,无需每次重算
  • 结果需.unsqueeze(0)扩展为[1, seq_len, d_model],方便与词嵌入相加

如何用向量化方式高效生成sin/cos位置编码矩阵

手动循环每个位置和维度会严重拖慢,尤其在GPU上。PyTorch支持全张量运算,只需两步:先生成所有pos和所有i的网格,再统一套公式。关键在于用torch.arange和torch.unsqueeze构造广播维度。

import torch
<p>def get_sinusoid_encoding_table(seq_len: int, d_model: int):
position = torch.arange(seq_len, dtype=torch.float32).unsqueeze(1)  # [seq_len, 1]
div_term = torch.exp(torch.arange(0, d_model, 2, dtype=torch.float32) <em> 
(-torch.log(torch.tensor(10000.0)) / d_model))  # [d_model//2]
pe = torch.zeros(seq_len, d_model)
pe[:, 0::2] = torch.sin(position </em> div_term)  # 偶数列
pe[:, 1::2] = torch.cos(position * div_term)  # 奇数列
return pe.unsqueeze(0)  # [1, seq_len, d_model]</p><p>pe_table = get_sinusoid_encoding_table(512, 512)</p>

注意div_term只计算一次,长度为d_model//2;position * div_term触发广播,产出[seq_len, d_model//2],再分别填入偶/奇列。该函数返回的pe_table可缓存复用,无需每次forward重建。

python-pro
python-pro

高级 Python 特性、异步编程、性能调优、静态类型、内存管理、Python 内部机制及生态库方面的专家。

下载

位置编码是否需要随输入动态padding?

不需要。标准做法是在模型初始化时预先生成最长所需长度的PE表(如max_len=512),forward时根据当前input_ids.shape[1]切片使用:pe[:, :seq_len, :]。强行对每个batch动态生成会浪费显存和时间,且破坏梯度图稳定性。

  • 若序列超长(如1024),需提前设好足够大的max_len,否则切片会越界
  • PE表本身不带设备信息,记得调用.to(x.device)再相加,否则CPU/GPU不匹配报错Expected all tensors to be on the same device
  • 不要在forward里重复调用get_sinusoid_encoding_table——这是典型性能陷阱

为什么不能直接用torch.nn.Parameter注册预计算好的PE?

可以,但要注意冻结。如果把PE表声明为nn.Parameter(pe_table, requires_grad=False),它会被自动加入model.parameters(),导致优化器尝试更新(即使requires_grad=False,某些旧版PyTorch仍会报warning)。更干净的做法是注册为buffer:

self.register_buffer('pe_table', pe_table)

这样它会随model.to(device)自动迁移,且不出现在parameters()中,也不会被optimizer.step()触碰。如果你看到训练中PE值意外变化,大概率是误用了Parameter而非buffer。

真正容易被忽略的是设备同步和切片时机——PE表一旦生成就别再改动,所有动态逻辑(长度适配、设备搬运)都应在forward中完成,且必须发生在与词嵌入相加之前。

Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

2023.07.20

1671

4

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

2023.07.25

4164

7

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

2023.07.31

1669

3

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

2023.08.03

24137

23

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2947

5

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2987

5

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

1163

5

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.10

596

4

python是前端还是后端
python是前端还是后端

Python属于前端也属于后端,其灵活性和丰富的生态系统使得开发人员能够在不同的领域中灵活运用。本专题为大家提供python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

2303

5

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程