在第 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 章(推理框架部署)的核心议题。