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
其中 是训练长度, 是目标长度。
8.2 YaRN Attention Scaling
其中 是缩放因子, 是阈值参数。
8.3 NTK-Aware Scaling
8.4 LongRoPE 非均匀映射
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. 理论理解 → 为什么有效 → 最优的扩展策略 → 外推的理论保证推荐阅读
- YaRN Paper (Peng et al., 2023) —— 原始论文
- NTK-Aware 博客 —— 实现细节
- LongRoPE Paper —— 超长上下文
- StreamingLLM Paper —— 流式推理
参考资料
- Peng, B., et al. (2023). “YaRN: Efficient Context Window Extension of Large Language Models.” arXiv<2309>2309>.00071.
- Chen, S., et al. (2023). “Extending Context Window of Large Language Models via Position Interpolation.” arXiv<2306>2306>.15595.
- Chen, S. (2023). “NTK-Aware Scaled RoPE Allows LLaMA Models to Have Extended (8K+) Context Size.” Blog Post.
- Ding, Y., et al. (2024). “LongRoPE: Extending LLM Context Window Beyond 2 Million Tokens.” arXiv.
- Xiao, G., et al. (2023). “StreamingLLM: Efficient Streaming Language Models with Attention Sinks.” arXiv.
- Dao, T., et al. (2022). “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.” NeurIPS.
- Kwon, W., et al. (2023). “Efficient Memory Management for Large Language Model Serving with PagedAttention.” SOSP.
- Touvron, H., et al. (2023). “LLaMA 2: Open Foundation and Fine-Tuned Chat Models.” Meta Research.
- Team, Q. (2024). “Qwen2 Technical Report.” arXiv.
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
YaRN / LongRoPE 深度解析:大模型上下文扩展技术
https://aiattnstudio.link/posts/yarn-longrope/
