unilm RetNet(Retentive Network)技术指南:安装、快速上手与 YOCO 中 Gated RetNet 的源码解析
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
RetNet(Retentive Network)是 unilm 研究体系中面向大语言模型的 Transformer 后继架构,本文基于 retnet/README.md 完整梳理其版本演进、安装方式与快速上手流程,并结合仓库内 YOCO 子目录中 Gated RetNet(RetNet-3)的真实源码实现,讲透其注意力计算路径与 Triton 内核细节。读完后你既能按文档部署 RetNet 环境并构建模型,也能理解该架构从并行、分块循环到单步循环的三种计算表示,以及门控 RetNet 在推理加速中的落地方式。
RetNet 是什么,在 unilm 中的定位
retnet/README.md 开篇即给出了项目的核心命题:Retentive Network 是面向大语言模型的 Transformer 后继方案(The Successor to Transformer for Large Language Models)。文档中记载了两个关键事实:
- 2023 年 7 月:预印本《Retentive Network: A Successor to Transformer for Large Language Models》(arXiv 编号 2307.08621)发布,作者为 Yutao Sun、Li Dong、Shaohan Huang 等;
- 2024 年 5 月:**Gated RetNet(即 RetNet-3)**作为 YOCO(You Only Cache Once: Decoder-Decoder Architectures for Language Models,arXiv 编号 2405.05254)的组成部分发布。
在代码层面,README 明确指向:RetNet 的官方实现托管在microsoft/torchscale仓库中,采用 MIT License(README 中同时附有 PyPI 的 torchscale 版本徽章)。需要说明的是:unilm 仓库本身并不内嵌原始 RetNet 的完整源码,retnet/目录下的 README 承担的是“索引 + 快速上手”的角色;而它的后继版本 RetNet-3(Gated RetNet)则在同一个 unilm 仓库的 YOCO/yoco/models/decoder/gate_retention.py 中有完整实现,这也是下文源码解析的依据。
从源码结构看,RetNet 的一个标志性特征是提供多种等价的计算表示。在 YOCO 的 Triton 内核文件 YOCO/yoco/models/decoder/kernel/gate_recurrent.py 中可以看到三个核心函数,分别对应三种路径:
parallel_gate_retention(L230-L240):并行表示。其核心公式实现为对衰减掩码做指数运算——mask = g[..., None] - g[..., None, :] + causal_mask,然后attn = q @ k.transpose(-1, -2)、attn = attn * exp(mask)、o = attn @ v,即保留(retention)机制中“按相对位置累积衰减的加权求和”;chunk_gate_retention(L180-L192):分块循环表示,将序列切块后,块内走注意力(inner_chunk)、块间走线性递推(cross_chunk);recurrent_gate_retention(L217-L228):逐步循环表示,仅维护一个K^T·V矩阵状态,用于解码期逐 token 更新。
这一“训练用并行/分块、推理用循环”的多表示设计,正是 RetNet 系列区别于标准 Transformer 的关键工程价值。
安装
README 的 Installation 一节给出了两种安装方式,此处完整继承并补充说明:
方式一:直接从 PyPI 安装
pip install torchscaletorchscale 包内包含RetNetConfig与RetNetDecoder等入口,适合只需要构建模型、跑实验的场景。
方式二:本地开发模式
如果你要阅读或修改 RetNet 源码(例如对照本仓库的 changelog 中的历史修复提交),README 建议克隆 microsoft/torchscale 开源仓库后进入目录执行开发模式安装:
cd torchscale pip install -e .(克隆仓库地址以 torchscale 官方仓库页面为准,本文不重复外链。)开发模式安装的额外好处是:本地对torchscale.architecture下源码的改动会即时生效,便于复现 changelog 中提到的各类稳定性修复。
快速上手:几行代码创建 RetNet 模型
README 的 Getting Started 强调 “It takes only several lines of code to create a RetNet model”,官方示例代码如下(完整继承,未作缩略):
# Creating a RetNet model >>> import torch >>> from torchscale.architecture.config import RetNetConfig >>> from torchscale.architecture.retnet import RetNetDecoder >>> config = RetNetConfig(vocab_size=64000) >>> retnet = RetNetDecoder(config) >>> print(retnet)要点解析:
RetNetConfig位于torchscale.architecture.config,示例中只覆盖了vocab_size=64000一个参数,其余超参(层数、头数、隐藏维度等)走配置类的默认值,说明该配置接口对使用者是“按需覆盖”的;RetNetDecoder位于torchscale.architecture.retnet,是纯解码器(decoder-only)结构,与大语言模型预训练范式一致;print(retnet)会输出模块树,可用于快速核对层数、注意力与归一化子模块的装配情况,适合在调试配置时验证。
Changelog:以稳定性为核心的一条演进线
README 的 Changelog 一节完整记录了 RetNet 实现自 2023 年 8 月至 2023 年 11 月的稳定性修复史。这些信息对复现训练者非常重要——它解释了为什么原始代码与后续版本的训练行为不同。完整继承如下:
| 时间 | 变更内容 | 技术解读 |
|---|---|---|
| 2023-08-04 | 修复分块循环表示(chunkwise recurrent representation)中的一个 bug | 对应三表示之一,早期分块路径的数值实现修正 |
| 2023-08-04 | 提升循环表示的数值精度(针对社区 issue #47 的建议) | 循环递推对累积误差敏感,精度修复直接影响长序列训练 |
| 2023-10 | ① 全面改用RMSNorm,以消除 LN_eps 的影响;② LayerNorm 的 eps 从1e-6 调整为 1e-5;③ 集成DeepNet 提出的初始化原则,因此不应再添加--subln或--deepnorm参数;④去掉各层 bias也提升了训练稳定性 | 这是训练稳定性的集中整改:归一化方式、超参默认值、初始化和残差 bias 一次性统一,且明确提示了命令行参数的废弃 |
| 2023-11 | 通过更好的初始化(better initialization)进一步提升稳定性 | 对权重初始化方案继续打磨 |
| 2023-11 | 修复retention 归一化(retention normalization)问题 | 对保留机制中数值尺度问题的修正 |
其中 2023 年 10 月的变更值得特别强调:文档原文明确写道 “So the arguments--subln or --deepnormshould not be added”,即如果你沿用旧版训练脚本,必须删除--subln/--deepnorm开关,否则会与内置的 DeepNet 初始化冲突。这是典型的“文档级坑点”,迁移脚本时应逐条比对。
仓库内纵深:YOCO 中的 Gated RetNet(RetNet-3)实现
retnet/README.md 首条 changelog 指出 Gated RetNet(RetNet-3)“as part of YOCO”。YOCO 是 unilm 仓库中的 decoder-decoder 架构(“You Only Cache Once”),它把 Gated RetNet 用作自身 self-decoder 的混洗器(mixer)。以下实现细节均来自当前仓库,可作为理解 RetNet 系架构工程落地的范本。
DecoderLayer 如何选择 RetNet 混洗器
在 YOCO/yoco/models/decoder/yoco.py 中,DecoderLayer.__init__按优先级选择 mixer(L54-L59):
if is_cross_layer: self.mixer = CrossAttention(args) elif args.sliding_window is not None: self.mixer = SlidingWindowAttention(args) else: self.mixer = GateRetention(args)也就是说:非交叉层且未配置滑动窗口时,默认混洗器就是 Gated RetNet(GateRetention);每层前后各接一个RMSNorm(mixer_layer_norm与final_layer_norm,均使用args.norm_eps,YOCOArgs默认 1e-5——与 retnet README changelog 中“eps 由 1e-6 调为 1e-5”的结论一脉相承)。
GateRetention 模块的内部结构
GateRetention 的参数装配(L30-L38):
- 五个投影:
q_proj、k_proj、v_proj、g_proj(SwiGLU 门控用)、gt_proj(每头一个标量的门控衰减率),其中gt_proj的输出维度只有args.n_self_heads,即门控是按头逐标量的; out_proj为行并行线性层;subln是一个elementwise_affine=False的RMSNorm,作用于每个注意力头维度(head_dim)——对应 changelog 中“RMSNorm 提升稳定性”的延续;- 所有投影均无 bias,且使用统一的
qkvg_init_method初始化。
前向流程(gate_retention.py L42-L87)中两个关键设计:
- RoPE 相对位置:
qr、kr经apply_rotary_emb(..., interleaved=True)施加旋转位置编码(L62-L63),说明 RetNet-3 在保留 retention 衰减机制的同时引入了旋转位置编码; - 对数域门控(L64):
gt = (F.logsigmoid(gt) / self.gate_logit_normalizer) # gate_logit_normalizer 默认 16门控经过logsigmoid后落在负半轴,再除以归一化常数 16,使衰减率更小、更平滑——内核中所有衰减都通过对g累加后再exp()来计算,即门控以对数衰减率形式参与运算,这与 retnet README 中 2023-11 “fix retention normalization”所关注的数值尺度问题是同类工程关注点。
三种计算路径与增量状态
GateRetention.forward根据是否处于解码阶段分派(L66-L82),与上文“多表示”特征完全对应:
- 逐 token 解码(
incremental_state is not None and not is_prefilling):走recurrent_gate_retention。它在 kernel/gate_recurrent.py L217-L228 中把当前 token 的K^T·V(k *= key_dim ** -0.5缩放后)与上一状态按门控衰减相加:kv += prev_kv * g,再输出o = q @ kv。KV 状态只有一份常数大小的矩阵,这正是 RetNet 类架构在自回归解码中显存/访存优势的实现形态; - Prefill(分块循环):走
chunk_gate_retention,固定chunk_size=256。值得注意的是SelfDecoder中self.block_size = 256(yoco.py L127),且 prefill 时若序列长度不是 256 的倍数会先做零填充(yoco.py L151-L153),保证分块整除; - Prefill 同时为解码期播种状态(L69-L79):当存在
incremental_state时,会用index_mask屏蔽 padding,按gt_sum = gt 的累加和计算整体衰减系数exp(gt_sum),把上一段历史状态缩放后合并进新的last_hidden_state并写回incremental_state,从而让 prefill 与后续逐 token 循环无缝衔接; - 并行表示:
parallel_gate_retention(L230-L240)实现了标准的 retention 并行公式,可作为数值参照或短序列路径。
Triton 内核与自校验测试
分块循环的内核由两层组成(kernel/gate_recurrent.py):
cross_chunk(L165-L170):块间传播。先按块尾衰减exp(-g + g[..., -1])加权构造每块的K^T·V,再经chunk_gate_recurrent(ChunkGateRecurrent.apply,Triton 前向_fwd_recurrence/反向_bwd_recurrence,L10-L98)做跨块线性递推;inner_chunk(L172-L178,带@torch.compile):块内因果注意力,掩码同样由对数门控差g[..., None] - g[..., None, :]加因果掩码指数化得到;chunk_gate_retention(L180-L192)把两者相加:o = cross + inner;hier_chunk_gate_retention(L195-L215):面向长序列并行(long sequence parallelism)的层次化版本,先以hier_chunk_size=16384做外层块间传播,再在外层块内递归调用chunk_gate_retention,两层 cross 相加——从源码结构看,这是为超长上下文训练(如 YOCO 的 1M 序列目标)提供的通信/并行友好路径。
更值得参考的是该文件自带的正确性校验:main()(L256-L298)用一个朴素 Python 循环naive_kv_recurrent(L242-L252,kv_state = kv_state * cross_decay + kv逐步累积)与 Triton 内核做前向/反向对拍,torch.allclose(..., atol=1e-3)校验输出与全部梯度。这意味着你可以直接把该文件当作 RetNet 类递推内核的参考实现 + 单测模板来学习。
引用
若你在研究中使用了本仓库的 RetNet 相关内容,README 建议引用原始论文:
@article{retnet, author={Yutao Sun and Li Dong and Shaohan Huang and Shuming Ma and Yuqing Xia and Jilong Xue and Jianyong Wang and Furu Wei}, title = {Retentive Network: A Successor to {Transformer} for Large Language Models}, journal = {ArXiv}, volume = {abs/2307.08621}, year = {2023} }小结与延伸阅读路径
retnet/README.md 以最小篇幅串起了 RetNet 的完整信息链:论文出处(2307.08621)、代码归属(torchscale,MIT)、安装与建模 API、以及一份以“训练稳定性”为主线的 changelog。结合本仓库的源码,可以按以下路径继续深入:
- retnet/README.md:本指南所依据的原始文档;
- YOCO/yoco/models/decoder/gate_retention.py:Gated RetNet 模块定义与三路径分派;
- YOCO/yoco/models/decoder/kernel/gate_recurrent.py:Triton 递推内核、长序列并行版本与朴素实现对照测试;
- YOCO/yoco/models/decoder/yoco.py:RetNet 混洗器在 decoder-decoder 整体结构中的装配位置;
- YOCO/README.md:Gated RetNet 所在的 YOCO 架构的训练与评测脚本。
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考