人人都会AI编程

8.6 长上下文训练技术:位置插值、注意力优化方案

更新时间:2026-07-09

在第 3 章中,我们详细解析了 Transformer 的位置编码机制,尤其是 RoPE(旋转位置编码)如何为模型注入序列顺序信息;在第 6 章中,我们讨论了上下文学习(In-Context Learning)的底层机制。当进入真实的预训练或继续预训练(Continual Pre-training)环节时,一个无法回避的工程难题浮现出来:绝大多数基座模型的初始预训练上下文长度只有 2k–8k Token,而业务场景(长文档问答、代码仓库级理解、多轮对话记忆)往往要求 128k 甚至 1M 级别的长窗口。

直接将短窗口模型用于长序列会遭遇两个硬约束:一是位置编码的外推崩溃(模型见到训练时未见过的远距离位置,注意力分数失控);二是显存与计算的二次方爆炸(Self-Attention 的 O(n²) 复杂度与 KV-Cache 的线性增长)。本节从这两个维度出发,给出当前工业界验证过的长上下文训练技术路线。


一、位置编码外推:为什么直接“拉长”不行?

以 RoPE 为例(当前主流开源模型如 Llama、Qwen、Baichuan 均采用),位置信息通过旋转矩阵注入 Query/Key 向量。预训练时,模型只见过位置索引 0–4095 的旋转角。当推理或训练直接扩展到 32768 时,这些三角函数值会进入分布外(Out-of-Distribution)区域,导致:

  • 注意力分数分布偏移:远距离位置的点积异常放大或缩小;
  • 温度系数失配:Softmax 输出趋于极端(极度尖锐或极度平坦);
  • 困惑度(PPL)断崖式上升:模型在长文本后半段完全“失忆”。

因此,长上下文训练的核心第一问是:如何让模型平滑地从短位置过渡到长位置?


二、位置插值(Position Interpolation, PI)及其演进

1. 线性位置插值(Linear PI)

最朴素的思路:既然模型只认识 0–L 的位置,那把实际长度 L'(如 32768)等比例压缩回 0–L 的范围内。

对于 RoPE,其旋转角 θ_m = m · θ_base(m 为位置索引)。线性 PI 将位置索引做缩放:

m' = m · (L / L')

其中 L 是原训练长度,L' 是目标长度。例如,将 4096 扩展到 32768,所有位置索引乘以 1/8。

效果与代价

  • 优点:零参数修改,仅需少量长文本继续预训练(通常几百步到几千步)即可恢复性能。
  • 缺点:所有位置被均匀压缩,短距离相对位置精度损失明显,会导致短文本任务轻微下降。

2. NTK-aware 缩放(Non-linear Interpolation)

线性 PI 对所有频率一视同仁,但 RoPE 的不同维度对应不同波长(低频对应长距离,高频对应短距离)。NTK(Neural Tangent Kernel)理论启发的改进认为:高频分量需要更少插值,低频分量需要更多插值

实现上,通过修改 RoPE 的基底频率 base:

base' = base · (L' / L)^(d / (d-2))

或采用“NTK-by-parts”策略,在不同维度区间应用不同的缩放系数。这使得短距离细节得以保留,长距离关系也能建模。

实践建议:Llama 系模型在社区外推中广泛使用 NTK-aware,无需修改模型权重,仅需在推理或训练时替换 RoPE 实现,配合 1%–10% 原数据量的长文本微调即可。

3. YaRN(Yet another RoPE extensioN)

当前工业界长上下文继续预训练的首选方案之一。YaRN 在 NTK 基础上进一步引入注意力温度系数(Temperature Scaling)

Attention Score = (Q·K^T) / (√d · scale)
scale ≈ 0.1 · ln(L'/L) + 1

当序列变长时,RoPE 内积天然有放大趋势,YaRN 通过显式温度缩放将其压回训练时的分布区间。同时,YaRN 对高频和低频分量采用不同的插值策略,兼顾短距精度与长距外推。

实验级结论:在 Llama 2 上,YaRN 支持将 4k 模型扩展到 128k,仅需在数百亿 Token 的长文本上继续预训练,PPL 即可收敛到与短文本相当水平。


三、注意力计算优化:从 O(n²) 到工程可接受

即使位置编码问题解决,长序列训练的第二大瓶颈是显存与计算量。设序列长度为 n,隐藏层维度为 d,则:

  • 计算复杂度:Self-Attention 为 O(n²·d),FFN 为 O(n·d²);当 n 很大时,注意力成为瓶颈。
  • 显存占用:KV-Cache 大小为 2 · n · 层数 · 头数 · 头维度 · 精度字节数,与 n 线性正相关。在 128k 长度下,即使 7B 模型,KV-Cache 也可能吃掉数十 GB 显存。

以下是在训练阶段(区别于推理阶段)常用的优化方案:

1. FlashAttention-2 / FlashAttention-3:IO 感知的精确注意力

这是目前长上下文训练的基线标配。核心思想不是近似计算,而是通过分块(Tiling)和重计算(Recomputation),将注意力计算适配到 GPU SRAM 的容量层级,减少对 HBM(高带宽显存)的读写次数。

  • FlashAttention-2:减少了非矩阵乘法运算的瓶颈,优化了 Warp 级并行,训练速度相比标准 Attention 提升 2–4 倍。
  • FlashAttention-3:针对 Hopper 架构(H100)的异步拷贝和 Tensor Core 特性进一步优化,支持 FP8 低精度,单卡吞吐提升显著。

对长上下文的意义:它不降低理论计算量(仍是 O(n²)),但将常数因子压到极低,使得在 A100/H100 上训练 32k–128k 序列从“不可能”变为“昂贵但可行”。在长上下文训练中,没有 FlashAttention 的方案基本不具备工程落地价值

2. 序列并行(Sequence Parallelism):跨卡切分长序列

当单卡显存放不下完整的 n × d 注意力矩阵时,需要将序列维度本身切分到多张 GPU 上。

  • Ring Attention:将序列切成环状块,每个 GPU 负责一块 Q,通过环状通信依次获取其他块的 K/V,在本地计算局部 Softmax 并同步全局统计量(Online Softmax)。这使得理论上可以训练无限长的序列(只要 GPU 数量足够)。
  • Megatron-LM / DeepSpeed Ulysses:将序列在数据并行维度上切分(All-to-All 通信交换序列和注意力头的维度),配合张量并行,实现超长长度的稳定训练。

实用边界:序列并行引入了额外的通信开销,通常在千卡以上集群、目标长度 ≥ 64k 时才值得启用。对于 8–32k 的扩展,单卡/单机配合 FlashAttention 通常足够。

3. 稀疏与局部注意力:选择性丢弃全连接

在长文本中,并非所有 Token 之间都需要直接交互。训练阶段的稀疏策略包括:

  • 滑动窗口注意力(Sliding Window / Local Attention):每个 Token 只 attend 到左右固定窗口(如 4k)内的 Token。这直接将 O(n²) 降为 O(n·w),但对需要全局依赖的任务(如文档级指代消解)会损失能力。
  • Dilated Attention / Longformer 变体:在窗口内引入空洞(Dilation),或对不同层交替使用局部和全局注意力,平衡感受野与计算量。
  • 混合策略:底层网络使用全注意力,上层使用局部注意力——因为底层负责语法和局部语义,上层负责长程逻辑。

取舍建议:稀疏注意力在推理端应用更广泛(如 Mistral 的 Sliding Window Attention)。在训练端,如果算力允许,尽量保留全注意力 + FlashAttention + 序列并行的组合,避免过早引入稀疏性导致模型能力天花板降低。


四、继续预训练(Continual Pre-training)的工程要点

将 4k 模型扩展到 128k,通常不是从头预训练,而是长上下文继续预训练。以下是经过验证的流程:

1. 数据准备

  • 来源:书籍、技术文档、代码仓库(GitHub commits with long dependencies)、合成数据(将多个短文档用特定分隔符拼接)。
  • Packing 策略:将多个短样本打包成长序列(如 32k/64k/128k),用 attention mask 隔离不同样本,提升训练效率。注意避免边界处的注意力渗透。
  • 数据配比:长文本训练步数通常只占原预训练的 1%–5%,但学习率需要重新热身(Warmup)并采用余弦退火,防止破坏短上下文能力。

2. 训练配置

  • 学习率:通常比原预训练低 1/10 到 1/5,例如原预训练 3e-4,继续预训练用 1e-4 或 5e-5。
  • 并行策略:长序列下,ZeRO-3 的显存优化配合张量并行(TP)必备;若序列长度 ≥ 64k,启用序列并行。
  • 长度渐进:实践中常采用多阶段扩展:4k → 16k → 32k → 128k,每阶段训练数百步到数 k 步。这比直接拉到 128k 更稳定。

3. 效果验证:Needle in a Haystack

不要只看 PPL。长上下文模型的标准压力测试是大海捞针(Needle in a Haystack):在极长文本的不同深度(前 10%、25%、50%、75%、90%)插入一个关键事实,然后提问。如果模型能准确召回,说明位置插值与注意力优化真正生效。


五、小结

长上下文训练是预训练工程中最具挑战性的子领域之一,其核心矛盾是位置编码的外推分布偏移注意力计算的二次方复杂度

关键 Takeaway:

  • 位置编码侧:优先采用 YaRN 或 NTK-aware 缩放,避免粗暴线性插值;RoPE 基底的非线性调整是短文本能力不崩塌的关键。
  • 计算侧:FlashAttention-2/3 是训练 32k+ 序列的入场券;当长度超过单卡承载极限时,引入 Ring Attention 或 Megatron 序列并行。
  • 训练策略:采用渐进式长度扩展 + 低学习率继续预训练,配合 Needle in a Haystack 验证长程信息保持能力。
  • 成本意识:每扩大 2 倍上下文,训练显存占用接近翻倍(若不用稀疏策略),需在业务收益与算力投入之间做硬性权衡。

完成长上下文扩展后,模型便具备了处理整本书籍、大型代码库、长篇报告的能力。但如何将这种能力稳定地接入生产环境,如何在推理阶段进一步压缩 KV-Cache、降低延迟,则是第 22 章(推理优化技术)与第 23 章(推理框架部署)的核心议题。