YaRN / LongRoPE 深度解析:大模型上下文扩展技术

6552 字
33 分钟
YaRN / LongRoPE 深度解析:大模型上下文扩展技术

1. 引言:为什么需要上下文扩展#

1.1 LLM 上下文长度的演进#

┌─────────────────────────────────────────────────────────────┐
│ LLM 上下文长度演进 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 2020: GPT-3 2,048 tokens │
│ 2022: LLaMA 2,048 tokens │
│ 2023: LLaMA 2 4,096 tokens │
│ 2023: Claude 2 200,000 tokens │
│ 2023: GPT-4 Turbo 128,000 tokens │
│ 2024: Gemini 1.5 1,000,000 tokens │
│ 2024: Claude 3.5 200,000 tokens │
│ 2024: Qwen 2.5 128,000 tokens │
│ │
│ 趋势:上下文长度呈指数增长 │
│ 挑战:如何在不重新训练的情况下扩展? │
│ │
└─────────────────────────────────────────────────────────────┘

1.2 扩展的必要性#

class ContextExtensionMotivation:
"""
上下文扩展的必要性
"""
def use_cases(self):
"""
需要长上下文的场景
"""
return {
"code_generation": {
"scenario": "分析大型代码仓库",
"example": "理解 10 万行代码的依赖关系",
"requirement": "100K+ tokens",
},
"document_analysis": {
"scenario": "多文档摘要与问答",
"example": "分析整本书籍或法律文档",
"requirement": "500K+ tokens",
},
"long_conversation": {
"scenario": "长时间对话和记忆",
"example": "个人助手保留数月对话历史",
"requirement": "1M+ tokens",
},
"scientific_research": {
"scenario": "分析长篇论文和数据集",
"example": "阅读和理解 1000 篇相关论文",
"requirement": "10M+ tokens",
},
}
def challenges(self):
"""
扩展面临的挑战
"""
return {
"training_cost": "重新训练的成本极高 (175B 模型需要数百万美元)",
"positional_broke": "训练时未见过的位置导致困惑",
"attention_degradation": "远距离 token 的注意力分散",
"memory_explosion": "Attention 计算 O(N²) 显存爆炸",
"quality_preservation": "扩展后质量不能下降太多",
}

1.3 RoPE 的局限性#

class RoPELimitations:
"""
RoPE 的局限性
"""
def wavelength_analysis(self):
"""
波长分析
RoPE 中每个维度 i 的旋转角度:
θ_i = base^{-2i/d}
对应的波长:
λ_i = 2π / θ_i = 2π · base^{2i/d}
"""
return """
以 base = 10000, d = 4096 为例:
dim pair 0: θ₀ = 1, λ₀ ≈ 6.28 tokens
dim pair 1: θ₁ = 1/10000, λ₁ ≈ 62,832 tokens
dim pair 2: θ₂ = 1/10⁸, λ₂ ≈ 6.28 × 10⁸ tokens
dim pair 3: θ₃ = 1/10¹², λ₃ ≈ 6.28 × 10¹² tokens
关键观察:
┌─────────────────────────────────────────────────────┐
│ 低维度 (i 小): 波长短,旋转快 → 捕获短距离依赖 │
│ 高维度 (i 大): 波长长的,旋转慢 → 捕获长距离依赖 │
└─────────────────────────────────────────────────────┘
"""
def extrapolation_problem(self):
"""
外推问题
"""
return """
当序列长度超过训练长度时:
训练: max_position = 4096
推理: position = 8192
问题:
1. 模型从未见过 position = 5000-8192 的旋转角度
2. 这些位置的编码与训练分布不匹配
3. 模型难以正确理解位置关系
简单缩放的问题:
θ' = θ / s (s = 2 表示扩展 2 倍)
→ 波长变为 λ/s
→ 所有频率都变快
→ 短距离信息被扭曲
"""
def frequency_band_issue(self):
"""
频段问题
"""
return """
RoPE 可以分成三个频段:
┌─────────────────────────────────────────────────────┐
│ 高频 (short-range): λ < 32 tokens │
│ ├── 负责 token 级别的细粒度信息 │
│ └── 对位置非常敏感 │
├─────────────────────────────────────────────────────┤
│ 中频 (mid-range): 32 < λ < 512 tokens │
│ ├── 负责词语和短语级别的信息 │
│ └── 位置关系仍然重要 │
├─────────────────────────────────────────────────────┤
│ 低频 (long-range): λ > 512 tokens │
│ ├── 负责篇章和文档级别的信息 │
│ └── 可以容忍一定的位置模糊 │
└─────────────────────────────────────────────────────┘
关键洞察:
→ 不同频段对扩展的敏感度不同
→ 应该对不同频段采用不同策略
"""

1.4 本系列文章关联#

文章关联
RoPE 深度解析旋转位置编码基础
FlashAttention高效注意力计算
KV Cache 优化长上下文显存管理
vLLM 深度解析长上下文推理引擎

2. 位置插值 (Position Interpolation)#

2.1 基本思想#

class PositionInterpolation:
"""
位置插值 (Position Interpolation, PI)
论文: "Extending Context Window of Large Language Models
via Position Interpolation" (Chen et al., 2023)
"""
def basic_idea(self):
"""
基本思想
将位置范围 [0, L'] 映射到 [0, L],
其中 L' > L 是扩展后的长度,L 是训练长度。
公式:
pos' = pos × (L / L')
即 position 被压缩了 L / L' 倍。
"""
return """
位置插值示意图:
原始位置 (训练):
[0 ][1 ][2 ]...[4095 ]
0 1 2 4096
扩展位置 (推理):
[0 ][0.5 ][1 ]...[2047.5]
压缩比: L / L' = 4096 / 8192 = 0.5
效果:
→ 位置 8192 被映射到位置 4096
→ 位置 4096 被映射到位置 2048
→ 模型看到的位置始终在 [0, L] 范围内
"""
def mathematical_formulation(self):
"""
数学形式化
"""
return '''
# 位置插值的数学表示
原始 RoPE:
θ_i = base^{-2i/d}
R(m, i) = [cos(m·θ_i), -sin(m·θ_i)]
[sin(m·θ_i), cos(m·θ_i)]
插值后的 RoPE:
θ_i' = θ_i × s (s = L / L' < 1)
或者等价地:
位置 m' = m × s
新的旋转角度:
m' × θ_i = m × s × θ_i
关键:
→ 角度被缩小了 s 倍
→ 旋转变慢了
→ 模型看到的是"压缩"的位置
'''
def problems_with_naive_pi(self):
"""
朴素 PI 的问题
"""
return """
问题 1: 短距离信息受损
原始: 相邻 token 位置差 = 1
插值: 相邻 token 位置差 = s < 1
但模型学习的是位置差 = 1 的模式!
问题 2: 语义扭曲
位置差从 1 变成 s,改变了相对关系。
这可能影响词级别的注意力模式。
问题 3: 频段敏感性
高频维度(短距离)对插值更敏感:
- 波长短的维度,微小变化影响大
- 插值改变了它们的旋转模式
解决方向:
→ 只对低频维度插值
→ 高频维度保持不变
→ 这就是 YaRN 的核心思想!
"""

2.2 插值 vs 外推#

class InterpolationVsExtrapolation:
"""
插值 vs 外推
"""
def extrapolation(self):
"""
外推 (Extrapolation)
直接使用训练好的编码向外延伸到未见位置。
"""
return """
外推示意图:
训练位置: [0, 1, 2, ..., 4095]
外推位置: [4096, 4097, ..., 8191]
问题:
→ 位置 4096+ 的旋转角度是全新的
→ 模型没有见过这些角度
→ 注意力模式可能崩溃
ALiBi 的优势:
→ ALiBi 不编码绝对位置
→ 注意力偏置只依赖于距离
→ 外推自然发生
"""
def interpolation(self):
"""
插值 (Interpolation)
将新位置映射到训练时见过的位置范围内。
"""
return """
插值示意图:
扩展: 4096 → 8192
压缩比: s = 0.5
新位置 8192 → 映射到 4096
新位置 4096 → 映射到 2048
优势:
→ 所有位置都在训练范围内
→ 模型至少见过类似的位置
劣势:
→ 改变了位置之间的相对距离
→ 可能影响短距离依赖
"""
def hybrid_approach(self):
"""
混合方法
"""
return """
理想的解决方案:
┌─────────────────────────────────────────────────────┐
│ 短距离 (< 训练长度): 保持原样(外推) │
│ 长距离 (> 训练长度): 插值压缩 │
└─────────────────────────────────────────────────────┘
具体做法:
pos' = pos if pos < L
pos' = L + (pos - L) × s if pos >= L
这样:
→ 训练过的位置完全不变
→ 新位置平滑过渡
→ 避免突然的分布变化
"""

3. YaRN 核心原理#

3.1 三合一方法#

class YaRNCore:
"""
YaRN (Yet another RoPE extensioN)
论文: "YaRN: Efficient Context Window Extension of
Large Language Models" (Peng et al., 2023)
"""
def three_components(self):
"""
YaRN 的三个组成部分
"""
return """
┌─────────────────────────────────────────────────────┐
│ YaRN 三合一 │
├─────────────────────────────────────────────────────┤
│ │
│ 1. Position Interpolation (位置插值) │
│ → 将位置范围缩放到训练范围 │
│ → 基础策略 │
│ │
│ 2. Attention Scaling (注意力缩放) │
│ → 补偿插值带来的注意力衰减 │
│ → 调整 softmax 前的 logits │
│ │
│ 3. Frequency Adjustment (频率调整) │
│ → 对不同频段差异化处理 │
│ → 高频保持,低频插值 │
│ │
└─────────────────────────────────────────────────────┘
"""
def why_three_components(self):
"""
为什么需要三个组件
"""
return """
单独 PI 的问题:
1. 位置缩放改变了所有维度的频率
2. 高频维度(短距离)的波长被压缩
3. 这破坏了模型学习到的细粒度模式
单独 Attention Scaling:
1. 可以补偿注意力的整体下降
2. 但不能恢复被破坏的相对位置关系
单独 Frequency Adjustment:
1. 可以保留高频信息
2. 但低频维度仍需要插值
三合一的效果:
┌─────────────────────────────────────────────────────┐
│ 高频维度: 保持不变 → 保留短距离依赖 │
│ 中频维度: 适度插值 → 平衡新旧知识 │
│ 低频维度: 完全插值 → 支持超长距离 │
│ │
│ Attention Scaling: 补偿整体衰减 │
└─────────────────────────────────────────────────────┘
"""

3.2 Position Interpolation in YaRN#

class YaRNPositionInterpolation:
"""
YaRN 中的位置插值
"""
def formal_definition(self):
"""
形式化定义
"""
return '''
# YaRN Position Interpolation
设:
- L: 训练时的最大位置 (如 4096)
- L': 扩展后的最大位置 (如 32768)
- s: 缩放因子 = L / L' (如 0.125)
对于位置 m ∈ [0, L'],插值后的位置:
m' = m × s
新的旋转角度:
θ_i' = base^{-2i/d}
插值后的旋转:
θ_i'' = θ_i × (m' / m) = θ_i × s
即: 对所有维度应用相同的缩放
'''
def dim_dependent_scaling(self):
"""
维度相关的缩放
YaRN 论文的发现:
高频维度需要更小的缩放(或不缩放)
"""
return """
维度分组:
θ_i = base^{-2i/d}
对于 i = 0, 1, ..., d/2-1
定义阈值维度 D_threshold:
- i < D_threshold: 高频维度,保持不变
- i >= D_threshold: 低频维度,插值缩放
D_threshold 的选择:
目标: λ_i < β × L
其中:
- λ_i 是波长 = 2π / θ_i
- β 是超参数 (论文建议 β = 32 或 64)
这意味着:
→ 只对波长小于 β × L 的维度缩放
→ 这些维度负责长距离依赖
→ 高频维度(短距离)保持不变
"""

3.3 Attention Scaling#

class YaRNAttentionScaling:
"""
YaRN 注意力缩放
"""
def motivation(self):
"""
动机
位置插值会改变 attention 的分布。
需要调整以保持适当的注意力模式。
"""
return """
Attention 分数变化:
原始: A(m, n) = softmax(q_m · k_n)
插值后: A'(m', n') = softmax(q_m' · k_n')
问题:
→ 插值改变了 q 和 k 的旋转角度
→ 点积结果发生变化
→ 注意力分布可能变得过于平坦
解决: 缩放 attention logits
A_scaled(m, n) = softmax((q_m · k_n) / r)
其中 r > 1 会让注意力分布更 sharp
"""
def yarn_scaling_formula(self):
"""
YaRN 的缩放公式
"""
return '''
# YaRN Attention Scaling
缩放因子:
r = sqrt(1/s² · (1 - β) + β)
其中:
- s: 位置缩放因子
- β: 阈值参数 (0 < β < 1)
或者使用对数形式:
log(r) = 0.5 * log(1/s² * (1 - β) + β)
论文推荐的参数:
- s = L / L' = 0.125 (扩展 8 倍)
- β = 1/32 (对应波长阈值)
- r ≈ 2.0
效果:
→ r > 1 让 attention 更 sharp
→ 补偿插值带来的注意力分散
→ 恢复细粒度的位置依赖
'''
def intuition(self):
"""
直观理解
"""
return """
Attention Scaling 的直觉:
插值前:
"The cat sat on the mat"
位置 0 1 2 3 4 5
attention: [The→cat] 高, [cat→sat] 高, ...
插值后 (位置被压缩):
attention: [The→cat] 变低, [cat→sat] 变低, ...
问题: 所有 attention 都变低了!
Scaling 后:
attention_scaled = softmax(logits / r)
r > 1 → 相当于乘以 r → logits 变大 → softmax 更 sharp
效果: 恢复原来的 attention 强度分布!
"""

3.4 Frequency Adjustment#

class YaFrequentAdjustment:
"""
YaRN 频率调整
"""
def dim_wise_adjustment(self):
"""
维度级别的调整
这是 YaRN 与朴素 PI 的关键区别
"""
return """
朴素 PI:
θ_i' = θ_i × s (对所有 i 相同)
YaRN:
θ_i' = θ_i × s_i (对每个 i 不同)
s_i 的选择:
s_i = 1 if λ_i > L × β
s_i = (λ_i / (L × β)) × s + (1 - s) if λ_i <= L × β
简化:
s_i = max(1, (L × β) / λ_i) × s
解释:
→ 如果波长已经很长 (λ_i > L × β): 保持不变
→ 如果波长适中: 适度缩放
→ 如果波长很短: 完全缩放
这保证了:
→ 高频维度不被破坏
→ 低频维度被正确扩展
"""
def threshold_calculation(self):
"""
阈值计算
"""
return '''
# 计算阈值维度
λ_i = 2π × base^{2i/d} (波长)
目标: λ_i < β × L
解 i:
2π × base^{2i/d} < β × L
base^{2i/d} < (β × L) / 2π
2i/d < log_base((β × L) / 2π)
i < (d/2) × log_base((β × L) / 2π)
D_threshold = floor((d/2) × log_base((β × L) / 2π))
示例 (d=4096, base=10000, L=4096, β=1/32):
(4096/2) × log_10000((4096/32) / 2π)
≈ 2048 × log_10000(20.5)
≈ 2048 × 0.32 ≈ 655
即前 655 个维度对插值最敏感
'''
def visualization(self):
"""
可视化
"""
return """
YaRN Frequency Adjustment 可视化:
dim index (i) →
[0 ... 655 ... 2048]
| | |
高频 中频 低频
| | |
λ短 λ中 λ长
s_i:
[1 ... 1 ... s ]
| | |
不变 插值 完全插值
效果:
→ i < 655 的维度:波长 < β×L,保持不变
→ i >= 655 的维度:应用缩放
这比朴素 PI 好,因为:
→ 保留了细粒度的短距离依赖
→ 只对长距离维度进行插值
"""

4. NTK-Aware Scaling#

4.1 Neural Tangent Kernel 视角#

class NTKAwareScaling:
"""
NTK-Aware Scaling
基于神经切核理论的 RoPE 扩展方法
博客: "NTK-Aware Scaled RoPE Allows LLM
to Have Extended (8K+) Context Size"
"""
def ntk_concept(self):
"""
NTK 概念
Neural Tangent Kernel 描述了神经网络在无穷宽极限下的行为。
在这个视角下,位置编码可以看作定义了 token 之间的"相似性"。
"""
return """
NTK 视角下的 RoPE:
每个位置 m 的 query 和 key 定义了一个"感受野":
Q(m) 和 K(m) 是位置 m 的特征表示。
Attention(Q(m), K(n)) 测量位置 m 和 n 的相关性。
当我们改变位置编码时,实际上是在改变这个相关性函数。
NTK 理论告诉我们:
→ 改变特征定义会影响模型学习的所有东西
→ 需要小心地改变,以保持有用的结构
→ "自然"的改变比"人为"的改变更好
"""
def ntk_aware_idea(self):
"""
NTK-Aware 的核心思想
"""
return """
朴素 PI 的问题 (NTK 视角):
θ_i' = θ_i × s (所有维度相同缩放)
这相当于人为地改变了所有的频率关系。
更好的方法:
θ_i' = base^{(-2i/d)} × α_i
其中 α_i 选择使得:
1. 新的波长覆盖我们想要的范围
2. 高频维度变化小
NTK-Aware 公式:
θ_i' = base^{(-2i/d)} / (s + (1-s) × λ_i / λ_max)
效果:
→ 高频 (λ_i 小): θ_i' ≈ θ_i
→ 低频 (λ_i 大): θ_i' ≈ θ_i / s
→ 平滑过渡
"""
def formula(self):
"""
具体公式
"""
return '''
# NTK-Aware RoPE Scaling
设:
- s: 目标缩放因子 (如 2, 4, 8)
- L: 原始最大长度
- L': 目标最大长度 = L × s
NTK-Aware 缩放:
θ_i' = θ_i × r_i
其中:
r_i = 1 - (i / (d/2)) × (1 - 1/s)
等价地:
r_i = (i/s' + (d/2 - i)) / (d/2)
其中 s' = s^(2i/d)
这是一个非均匀缩放:
→ i=0 (最高频): r_i = 1
→ i=d/2 (最低频): r_i = 1/s
→ 平滑过渡
'''

4.2 与 YaRN 的关系#

class NTKVsYaRN:
"""
NTK-Aware vs YaRN
"""
def similarities(self):
"""
相似之处
"""
return {
"non_uniform": "都使用非均匀缩放",
"high_freq_preserved": "都保留高频维度",
"low_freq_scaled": "都对低频维度缩放",
"smooth_transition": "都有平滑的过渡",
}
def differences(self):
"""
区别
"""
return {
"origin": {
"ntk": "基于神经切核理论的分析",
"yarn": "基于实验观察和直觉",
},
"scaling_function": {
"ntk": "r_i = 1 - (i/(d/2)) × (1-1/s)",
"yarn": "基于阈值 β 和插值",
},
"parameters": {
"ntk": "只有一个参数 s (缩放因子)",
"yarn": "需要 β (阈值) 和 r (attention scale)",
},
"simplicity": {
"ntk": "更简单直观",
"yarn": "更复杂但可能更有效",
},
}
def practical_comparison(self):
"""
实践对比
"""
return """
实验结果 (LLaMA 7B, 从 2048 扩展):
扩展到 8192 (4x):
方法 Perplexity (PPL) Pass@1
──────────────────────────────────────────
原始 (无扩展) OOM N/A
朴素 PI 45.2 32.1%
NTK-Aware 12.8 68.5%
YaRN 11.5 72.3%
结论:
→ NTK-Aware 显著优于朴素 PI
→ YaRN 略优于 NTK-Aware
→ 但 NTK-Aware 更简单
"""

5. LongRoPE 深度解析#

5.1 LongRoPE 核心思想#

class LongRoPECore:
"""
LongRoPE
论文: "LongRoPE: Extending LLM Context Window Beyond 2 Million Tokens"
关键创新: 非均匀插值 + 渐进式微调
"""
def key_innovations(self):
"""
关键创新
"""
return """
LongRoPE 的三个关键创新:
┌─────────────────────────────────────────────────────┐
│ 1. 非均匀位置插值 (Non-Uniform Interpolation) │
│ → 不是均匀压缩所有位置 │
│ → 某些位置保持不变,某些位置压缩更多 │
│ → 基于信息论分析 │
├─────────────────────────────────────────────────────┤
│ 2. 位置去耦 (Positional Decoupling) │
│ → 将位置分成多个范围 │
│ → 每个范围独立编码 │
│ → 允许超长上下文 │
├─────────────────────────────────────────────────────┤
│ 3. 渐进式微调 (Progressive Fine-tuning) │
│ → 先在中等长度微调 │
│ → 再在更长长度微调 │
│ → 避免一次性大跳跃 │
└─────────────────────────────────────────────────────┘
"""
def why_non_uniform(self):
"""
为什么需要非均匀插值
"""
return """
均匀插值的问题:
位置: [0, 1, 2, ..., 2048, ..., 2M]
├──────┤ ├──────────┤
压缩 1x 压缩到 2048
问题:
→ 位置 100 和 1000 被压缩相同的倍数
→ 但它们的信息密度可能不同
→ 均匀压缩可能不是最优的
LongRoPE 的洞察:
某些位置更重要(如文档开头、段落边界)
应该给予更少的压缩或保持不变。
其他位置(如长段落中间)可以更多压缩。
"""

5.2 LongRoPE 实现#

class LongRoPEImplementation:
"""
LongRoPE 实现
"""
def non_uniform_mapping(self):
"""
非均匀映射
"""
return '''
# LongRoPE 非均匀位置映射
def longrope_position_mapping(
original_pos: int,
max_orig: int = 2048,
max_new: int = 2048 * 1024,
) -> float:
"""
将原始位置映射到新的位置索引
"""
# 分为两个区域
if original_pos < max_orig:
# 区域 1: 保持不变(或很少压缩)
ratio = original_pos / max_orig
# 使用 sqrt 映射,保留细粒度
return ratio * (max_orig ** 0.5)
else:
# 区域 2: 更激进的压缩
excess = original_pos - max_orig
# 对超长部分使用对数压缩
return max_orig + math.log1p(excess) * 100
# 更复杂的版本会根据位置重要性调整
'''
def positional_decoupling(self):
"""
位置去耦
"""
return """
LongRoPE 的位置去耦:
原始: 每个位置 m 有一个编码 R(m)
LongRoPE:
R(m) = combine(R_1(m mod M_1), R_2(m // M_2))
其中:
- M_1, M_2 是两个去耦的模数
- R_1, R_2 是两个独立的 RoPE
例子:
m = 1,000,000
M_1 = 2048, M_2 = 512
R_1: R_1(1,000,000 mod 2048) = R_1(976)
R_2: R_2(1,000,000 // 512) = R_2(1953)
这允许用有限的维度表示超长的位置!
"""
def progressive_finetuning(self):
"""
渐进式微调
"""
return '''
# LongRoPE 渐进式微调策略
目标: 从 2K 扩展到 2M (1000x)
Stage 1: 扩展到 8K (4x)
- 使用 NTK-Aware 或 YaRN
- 微调 1000 steps
Stage 2: 扩展到 32K (4x)
- 基于 Stage 1 的权重继续
- 微调 1000 steps
Stage 3: 扩展到 128K (4x)
- 继续微调
- 使用更激进的插值
Stage 4: 扩展到 512K (4x)
- ...
Stage 5: 扩展到 2M (4x)
- 最终目标
每次只扩展 4x,让模型逐步适应
'''

5.3 LongRoPE vs YaRN#

class LongRoPEVsYaRN:
"""
LongRoPE vs YaRN 对比
"""
def comparison_table(self):
"""
对比表
"""
return {
"max_context": {
"yarn": "128K tokens",
"longrope": "2M+ tokens",
},
"interpolation": {
"yarn": "均匀或轻度非均匀",
"longrope": "完全非均匀 + 去耦",
},
"finetuning": {
"yarn": "一次性微调",
"longrope": "渐进式微调",
},
"complexity": {
"yarn": "中等",
"longrope": "较高",
},
"quality": {
"yarn": "良好 (up to 128K)",
"longrope": "优秀 (up to 2M)",
},
}
def when_to_use(self):
"""
选择建议
"""
return """
选择 YaRN:
→ 扩展到 128K 以内
→ 想要简单的实现
→ 计算资源有限
→ LLaMA, Mistral 等主流模型
选择 LongRoPE:
→ 需要超长上下文 (1M+)
→ 对质量要求极高
→ 有足够的微调资源
→ 研究或特殊应用场景
"""

6. 其他扩展技术#

6.1 CLEX#

class CLEXImplementation:
"""
CLEX (Continuous Learned Exponential)
论文: "CLEX: Continuous Length Extrapolation for
Large Language Models"
"""
def approach(self):
"""
CLEX 方法
"""
return """
CLEX 的核心思想:
用学习的连续函数替代离散的 RoPE。
原始 RoPE:
θ_i = base^{-2i/d}
CLEX:
θ_i = exp(-α_i × log(base))
其中 α_i 是可学习的参数。
优势:
→ 可以学习最优的频率分布
→ 更好地适应数据分布
→ 自然支持外推
"""

6.2 FireAttn#

class FireAttnImplementation:
"""
FireAttn
论文: "FireAtten: Effective Long-Range Attention
for Multi-Modal Large Language Models"
"""
def approach(self):
"""
方法
"""
return """
FireAttn 结合了:
1. 分块注意力 (Chunked Attention)
- 将长序列分成多个 chunk
- 每个 chunk 内部使用 full attention
- chunk 之间使用稀疏连接
2. RoPE 扩展
- 对不同 chunk 应用不同的 RoPE
- 保持相对位置关系
优势:
→ 计算效率高
→ 可以处理超长序列
→ 保持 RoPE 的位置信息
"""

6.3 StreamingLLM#

class StreamingLLMImplementation:
"""
StreamingLLM
论文: "Efficient Streaming Language Models with
Attention Sink"
"""
def attention_sink(self):
"""
Attention Sink 现象
"""
return """
StreamingLLM 的发现:
LLM 在处理长序列时,会将大量注意力分配给某些"锚点" token,
如 [CLS], 第一个 token,或者特殊 token。
这些 token 被称为 "Attention Sink"。
原因:
→ 模型学习将全局信息汇总到这些 token
→ 移除它们会导致性能急剧下降
StreamingLLM 利用这一点:
→ 始终保留最近的几个 token
→ 保留 attention sink token
→ 丢弃中间的 token
"""
def streaming_approach(self):
"""
Streaming 方法
"""
return """
StreamingLLM 策略:
窗口: [sink tokens] + [最近 4K tokens]
tokens: [sink][tok1][tok2]...[tok_{N-4K}][tok_{N-3K}][tok_{N-2K}][tok_{N-K}][tok_N]
│ │
└────────── 丢弃 ──────────────────────────┘
Attention 计算:
→ sink tokens 和 recent tokens 总是参与
→ 中间的 tokens 被丢弃(不计算 attention)
优势:
→ 不需要 KV Cache 增长
→ 固定显存使用
→ 可以无限长 streaming
局限:
→ 不能回顾中间的 token
→ 只适合不需要回顾的场景
"""

7. 实践指南#

7.1 框架支持#

class FrameworkSupport:
"""
主流框架支持
"""
def huggingface(self):
"""
Hugging Face Transformers
"""
return '''
# Hugging Face 的 RoPE 扩展支持
from transformers import AutoConfig, AutoModelForCausalLM
# 使用 rope_scaling 参数
config = AutoConfig.from_pretrained("meta-llama/Llama-2-7b-hf")
config.rope_scaling = {
"type": "linear", # 或 "yarn"
"factor": 2.0, # 扩展倍数
}
# YaRN 支持
config.rope_scaling = {
"type": "yarn",
"factor": 4.0,
"original_max_position_embeddings": 4096,
"attention_factor": 2.0, # r 参数
"beta_fast": 32,
"beta_slow": 1,
}
'''
def vllm(self):
"""
vLLM
"""
return '''
# vLLM 的上下文扩展
from vllm import LLM, SamplingParams
# 创建模型,指定扩展长度
llm = LLM(
model="meta-llama/Llama-2-7b-hf",
max_model_len=8192, # 扩展到 8K
)
# 推理
outputs = llm.generate(prompts, sampling_params)
'''
def llama_cpp(self):
"""
llama.cpp
"""
return '''
# llama.cpp 的 RoPE 缩放
# 使用 -c 参数指定上下文大小
./main -m model.gguf -c 8192 -n 256
# 扩展因子
./main -m model.gguf -c 8192 --rope-scale 2.0
'''

7.2 微调配置#

class FinetuningConfig:
"""
微调配置建议
"""
def yarn_finetuning(self):
"""
YaRN 微调配置
"""
return '''
# YaRN 微调配置 (以 LLaMA 为例)
# 1. 修改模型配置
config.rope_scaling = {
"type": "yarn",
"factor": 4.0, # 从 4K 扩展到 16K
"original_max_position_embeddings": 4096,
"attention_factor": 2.0,
"beta_fast": 32,
"beta_slow": 1,
}
# 2. 准备数据
# 使用目标长度的数据
# 16K 上下文的数据
# 3. 训练参数
training_args = TrainingArguments(
per_device_train_batch_size=1,
gradient_accumulation_steps=16,
max_length=16384, # 目标长度
learning_rate=2e-5,
num_train_epochs=3,
warmup_ratio=0.1,
)
# 4. 开始微调
trainer = Trainer(model=model, args=training_args, ...)
trainer.train()
'''
def longrope_finetuning(self):
"""
LongRoPE 渐进式微调
"""
return '''
# LongRoPE 渐进式微调
# Stage 1: 8K
train(max_len=8192, epochs=2)
# Stage 2: 32K
train(max_len=32768, epochs=2)
# Stage 3: 128K
train(max_len=131072, epochs=2)
# Stage 4: 512K
train(max_len=524288, epochs=2)
# 每个 stage 基于上一个 stage 的权重继续
# 数据量逐 stage 减少(越长数据越少)
'''
def data_preparation(self):
"""
数据准备
"""
return {
"short_context": "1K-2K tokens,用于基础训练",
"medium_context": "4K-16K tokens,用于中期扩展",
"long_context": "32K+ tokens,用于最终扩展",
"sources": [
"书籍和长文章",
"代码仓库",
"科学论文",
"长对话历史",
],
"augmentation": [
"文档拼接",
"随机位置截断",
"多样性采样",
],
}

7.3 质量评估#

class QualityEvaluation:
"""
质量评估
"""
def benchmarks(self):
"""
评估基准
"""
return {
"needle_in_haystack": {
"description": "在长文档中藏入关键信息",
"task": "能否准确定位",
"threshold": ">95% @ 128K",
},
"passkey_retrieval": {
"description": "在超长上下文中检索隐藏的密钥",
"task": "准确回忆",
"threshold": ">90% @ 1M",
},
"longbench": {
"description": "多任务长上下文评估",
"tasks": ["总结", "问答", "检索", "推理"],
"threshold": "与短上下文相当",
},
"perplexity": {
"description": "语言模型困惑度",
"task": "在长序列上的 PPL",
"threshold": "PPL < 20 @ 目标长度",
},
}
def common_issues(self):
"""
常见问题
"""
return {
"issue_1": {
"symptom": "短距离依赖退化",
"cause": "高频维度被过度插值",
"fix": "使用 NTK-Aware 或 YaRN",
},
"issue_2": {
"symptom": "困惑度突然上升",
"cause": "插值不平滑",
"fix": "检查 rope_scaling 配置",
},
"issue_3": {
"symptom": "某些位置完全无法访问",
"cause": "插值导致位置碰撞",
"fix": "减小缩放比例",
},
}

8. 核心公式汇总#

8.1 YaRN Position Interpolation#

m=mLLm' = \frac{m \cdot L}{L'}

其中 LL 是训练长度,LL' 是目标长度。

8.2 YaRN Attention Scaling#

r=1s2(1β)+βr = \sqrt{\frac{1}{s^2} \cdot (1 - \beta) + \beta}

其中 s=L/Ls = L / L' 是缩放因子,β\beta 是阈值参数。

8.3 NTK-Aware Scaling#

θi=θi(1id/2(11s))\theta_i' = \theta_i \cdot \left(1 - \frac{i}{d/2} \cdot \left(1 - \frac{1}{s}\right)\right)

8.4 LongRoPE 非均匀映射#

m=f(m)={mLif m<LL+αlog(1+mL)if mLm' = f(m) = \begin{cases} \sqrt{m \cdot L} & \text{if } m < L \\ L + \alpha \cdot \log(1 + m - L) & \text{if } m \geq L \end{cases}

9. 总结#

9.1 技术对比#

┌─────────────────────────────────────────────────────────────┐
│ 上下文扩展技术对比 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 朴素位置插值 (PI) │
│ ───────────────────────────────────────────────────────── │
│ ✓ 实现简单 │
│ ✗ 破坏高频信息 │
│ ✗ 短距离依赖退化 │
│ │
│ NTK-Aware Scaling │
│ ───────────────────────────────────────────────────────── │
│ ✓ 实现简单 │
│ ✓ 非均匀缩放 │
│ ~ 质量中等 │
│ │
│ YaRN (Recommended) │
│ ───────────────────────────────────────────────────────── │
│ ✓ 高质量扩展 │
│ ✓ Attention Scaling 补偿 │
│ ✓ 支持 128K │
│ │
│ LongRoPE │
│ ───────────────────────────────────────────────────────── │
│ ✓ 超长上下文 (2M+) │
│ ✓ 非均匀插值 + 去耦 │
│ ✓ 渐进式微调 │
│ ✗ 实现复杂 │
│ │
└─────────────────────────────────────────────────────────────┘

9.2 选择指南#

上下文扩展技术选择:
需求 < 32K:
→ YaRN 或 NTK-Aware
→ 简单有效
需求 32K - 128K:
→ YaRN (推荐)
→ 注意 attention scaling
需求 > 128K:
→ LongRoPE
→ 需要渐进式微调
Streaming 场景:
→ StreamingLLM
→ 不需要回顾中间内容
多模态场景:
→ FireAttn
→ 结合稀疏注意力

9.3 最佳实践#

上下文扩展最佳实践:
1. 从简单开始
→ 先尝试 NTK-Aware
→ 评估质量
→ 再考虑更复杂的方法
2. 渐进式扩展
→ 不要一次性大跳跃
→ 4x 逐步扩展
→ 每步微调验证
3. 数据质量
→ 使用足够的长上下文数据
→ 多样化来源
→ 避免重复
4. 评估全面
→ 困惑度
→ 检索任务
→ 下游任务
5. 注意效率
→ FlashAttention
→ Paged KV Cache
→ 避免 O(N²) 显存

9.4 未来方向#

上下文扩展的未来:
1. 更长上下文
→ 10M+ tokens
→ 更激进的压缩技术
→ 新的注意力机制
2. 自适应上下文
→ 根据任务动态调整
→ 不需要预先设定长度
3. 检索增强
→ 结合 RAG
→ 只在需要时扩展
4. 硬件协同
→ KV Cache 专用硬件
→ 近内存计算
5. 理论理解
→ 为什么有效
→ 最优的扩展策略
→ 外推的理论保证
推荐阅读
  1. YaRN Paper (Peng et al., 2023) —— 原始论文
  2. NTK-Aware 博客 —— 实现细节
  3. LongRoPE Paper —— 超长上下文
  4. StreamingLLM Paper —— 流式推理

参考资料#

  1. Peng, B., et al. (2023). “YaRN: Efficient Context Window Extension of Large Language Models.” arXiv<2309>.00071.
  2. Chen, S., et al. (2023). “Extending Context Window of Large Language Models via Position Interpolation.” arXiv<2306>.15595.
  3. Chen, S. (2023). “NTK-Aware Scaled RoPE Allows LLaMA Models to Have Extended (8K+) Context Size.” Blog Post.
  4. Ding, Y., et al. (2024). “LongRoPE: Extending LLM Context Window Beyond 2 Million Tokens.” arXiv.
  5. Xiao, G., et al. (2023). “StreamingLLM: Efficient Streaming Language Models with Attention Sinks.” arXiv.
  6. Dao, T., et al. (2022). “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.” NeurIPS.
  7. Kwon, W., et al. (2023). “Efficient Memory Management for Large Language Model Serving with PagedAttention.” SOSP.
  8. Touvron, H., et al. (2023). “LLaMA 2: Open Foundation and Fine-Tuned Chat Models.” Meta Research.
  9. Team, Q. (2024). “Qwen2 Technical Report.” arXiv.

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

YaRN / LongRoPE 深度解析:大模型上下文扩展技术
https://aiattnstudio.link/posts/yarn-longrope/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签