Skip to content

推理加速算法

大模型推理的延迟主要来自两个方面: Prefill 阶段(处理输入 prompt )和 Decode 阶段(逐个生成token)。由于Decode阶段是内存带宽受限的,每生成一个token都需要加载完整的模型参数和KV Cache。本节介绍通过投机执行来突破这一瓶颈的系列算法。

7.2.1 Speculative Decoding (投机解码)

问题背景

在自回归生成中,每个token的生成都需要: 1. 加载整个模型到GPU 2. 执行一次完整的前向传播 3. 采样得到下一个token

这个过程是顺序的,无法并行化。然而,语言模型生成的 token往往具有较强的可预测性,特别是常见短语和模板化内容。

核心思想

投机解码采用小模型起草,大模型验证的策略:

  1. 起草阶段(Draft):使用轻量级小模型(draft model)快速生成 \(\gamma\) 个候选token
  2. 验证阶段(Verify):大模型(target model)并行验证这些候选token
  3. 接受/拒绝:根据概率分布决定是否接受候选token

数学原理

设:

  • \(M_p\):大模型( target model )的概率分布 \(p(x)\)
  • \(M_q\):小模型( draftmodel)的概率分布 \(q(x)\)

对于候选token\(x\),接受概率为:

\(P(\text{accept}) = \min\left(1, \frac{p(x)}{q(x)}\right)\)

修正采样

如果候选token\(x\) 被拒绝,从修正分布中采样:

\(p'(x) = \frac{p(x) - q(x)}{1 - \sum_{x'} \min(p(x'), q(x'))}\)

或者使用更简单的形式:

\(p'(x) = \text{normalize}(\max(0, p(x) - q(x)))\)

算法伪代码

python
def speculative_decoding(prompt, M_p, M_q, gamma, max_tokens):
   """
   投机解码算法
   M_p: 大模型(target model)
   M_q: 小模型(draft model)
   gamma: 每次起草的token数
   """
   tokens = tokenize(prompt)

   while len(tokens) < max_tokens:
         # === 起草阶段 ===
         draft_tokens = []
         draft_probs = []

         for _ in range(gamma):
              q_dist = M_q(tokens + draft_tokens)
              x = sample(q_dist)
              draft_tokens.append(x)
              draft_probs.append(q_dist)
          
         # ===验证阶段 ===
         # 大模型并行计算所有候选token的概率
         all_tokens = tokens + draft_tokens
         p_dists = M_p(all_tokens, return_all_probs=True)

         # ===    接受/拒绝 ===
         accepted = 0
         for i, (x, q_dist) in enumerate(zip(draft_tokens, draft_probs)):
              p_dist = p_dists[len(tokens) + i]
              p_x = p_dist[x]
              q_x = q_dist[x]

              #   计算接受概率
              accept_prob = min(1.0, p_x / q_x)

              if random() < accept_prob:
                   #   接受token
                   tokens.append(x)
                   accepted += 1
              else:
                   #   拒绝token,从修正分布采样
                   adjusted_dist = normalize(max(0, p_dist - q_dist))
                   tokens.append(sample(adjusted_dist))
                   break

         #   如果全部接受,从p采样一个额外token
         if accepted == gamma:
              tokens.append(sample(p_dists[-1]))
         return detokenize(tokens)

加速比分析

理论加速比取决于: 1. 接受率 \(\alpha\):候选 token 被接受的比例 2. 起草成本比 \(\beta\):小模型与大模型的推理时间比 3. 起草数量 \(\gamma\)

理想加速比(假设 \(\beta \ll 1\)):

\(\text{Speedup} \approx \frac{1}{1 - \alpha^{\gamma+1}} \cdot \frac{\gamma + 1}{\gamma \beta + 1}\)

典型场景( \(\alpha = 0.8, \gamma = 4, \beta = 0.1\) ): \(\text{Speedup} \approx 2.5-3.0x\)

实际应用与限制

优势:

  • 无损加速:保持大模型的输出分布
  • 通用性强:适用于任何自回归模型
  • 实现相对简单

限制:

  1. 小模型选择:需要与大模型行为相似的小模型
  2. 内存开销:需要同时加载两个模型
  3. 接受率依赖:在低可预测性内容上效果差
  4. 通信开销:如果大小模型在不同设备上

典型配置:

  • LLaMA-70B + LLaMA-7B作为draft model
  • \(\gamma = 4-8\)
  • 实测加速比:1.5-2.5x

7.2.2 Medusa

问题背景

投机解码需要维护两个独立的模型,带来额外的内存和工程复杂度。Medusa提出了一种无需辅助模型的投机解码方法,通过在原始模型上添加轻量级头来实现。

核心思想

Medusa的核心创新: 1. 多头解码:在原始模型的顶层添加多个解码头 2. 每个头预测未来不同位置的token:头1预测下一个token,头2预测下下个token,依此类推 3. 树状注意力:高效组织多个候选序列的验证

多头结构

每个Medusa Head是一个轻量级MLP:

\(\text{Medusa}_i(h) = \text{Softmax}(W_i \cdot \text{MLP}_i(h))\)

训练方法

训练数据

使用原始模型生成的数据作为监督信号: 1. 从训练集中采样prompt 2. 用原始模型生成完整序列 3. 将生成的token作为每个头的目标

损失函数

\(\mathcal{L} = \mathcal{L}_{\text{base}} + \lambda \sum_{i=1}^{k}\mathcal{L}_{\text{Medusa}_i}\)

其中 \(\mathcal{L}_{\text{base}}\) 是原始语言模型头的交叉熵损失, \(\mathcal{L}_{\text{Medusa}_i}\) 是第 \(i\) 个Medusa头的损失。

训练策略

  1. 冻结主干:保持原始模型参数不变
  2. 只训练Medusa Heads:通常只有几百万参数
  3. 典型训练时间:1-2个epoch,数小时到一天

树状注意力验证

Medusa提出树状注意力来高效验证多个候选序列

latex
候选序列树:
      t+1
     / | \
   t+2a t+2b t+2c
   /   |     |
 t+3a ...   ...

通过树状注意力,可以在一次前向传播中验证所有候选路径。Medusa 伪代码

python
class MedusaModel(nn.Module):
   def __init__(self, base_model, num_heads, num_candidates):
       super().__init__()
       self.base_model = base_model
       self.num_heads = num_heads
       self.num_candidates = num_candidates

       # Medusa Heads:   每个头预测未来第i个token
       self.medusa_heads = nn.ModuleList([
             MedusaHead(base_model.config.hidden_size, base_model.vocab_size)  # 补全右括号
             for _ in range(num_heads)
       ])

   def forward(self, input_ids):
       # 基础模型前向
       outputs = self.base_model(input_ids, output_hidden_states=True)
       hidden = outputs.hidden_states[-1]

       # 基础预测
       base_logits = outputs.logits[:, -1, :]

       # Medusa 预测
       medusa_logits = []
       for head in self.medusa_heads:
             medusa_logits.append(head(hidden[:, -1, :]))

       return base_logits, medusa_logits

   def generate(self, prompt, max_new_tokens):
       tokens = tokenize(prompt)

       for _ in range(max_new_tokens // (self.num_heads + 1)):
             #   获取所有预测
             base_logits, medusa_logits = self.forward(tokens)

             #   构建候选序列树
             candidates = build_tree(base_logits, medusa_logits, self.num_candidates)  

             #   树状注意力验证
             verified_tokens = self.verify_tree(tokens, candidates)

             tokens.extend(verified_tokens)

       return detokenize(tokens)


   def verify_tree(self, prefix, candidates):
       """使用树状注意力验证候选序列"""
       # 构建树状注意力掩码
       tree_mask = build_tree_mask(candidates)

      #   一次前向验证所有候选
      all_sequences = [prefix + cand for cand in candidates]
      logits = self.base_model(all_sequences, attention_mask=tree_mask).logits  


      #   选择接受的token序列
      accepted = []
      for i, cand in enumerate(candidates):
           if verify_acceptance(logits[i], cand):
              accepted.extend(cand)
           else:
              #    从修正分布采样
              accepted.append(sample_adjusted(logits[i]))
              break


      return accepted

与投机解码的对比

特性Speculative DecodingMedusa
辅助模型需要独立的小模型无需辅助模型
内存开销2x模型大小1.01-1.05x模型大小
训练需求无需训练需要训练Medusa Heads
加速比1.5-2.5x1.8-2.8x
实现复杂度中等较高

性能分析

在Vicuna-7B上的实测结果:

  • 平均加速比:2.0-2.5x
  • Medusa Heads数量:4-8个
  • 训练时间:约4-8小时
  • 内存增加:<5%

7.2.3 EAGLE

问题背景

Medusa在特征空间进行预测,但直接预测未来 token的分布可能面临信息不足的问题。EAGLE ( Extrapolation Algorithm for Greater Language-model Efficiency )提出在特征级别进行投机,利用更丰富的上下文信息。

核心思想

EAGLE的核心创新: 1. 特征级投机:预测未来token的特征(而非token本身) 2. 自回归头Auto-regressive Head):使用轻量级自回归模型预测特征序列 3. 特征到 token映射:通过LM头将预测特征转换为token分布

LM Head(Language Model Head)是位于主干网络之后的一个线性层(Linear Layer),它的核心作用是:将模型的隐藏状态(Hidden State)映射到词汇表大小的概率分布上

python
输入 Token → Embedding → Transformer Layers → [LM Head] → Logits → Softmax → 概率分布

                                            Linear(768, vocab_size)

特征级预测的优势

相比token级预测,特征级预测:

  • 包含更丰富的语义信息
  • 更容易捕获长距离依赖
  • 预测任务更简单(连续空间vs离散空间)

自回归头架构

EAGLE的自回归头是一个小型Transformer:

python
输入: [h_t, h_t-1, ..., h_t-k]  (历史特征序列)
  |
  v
+-----------------------+
| 小型Transformer      | (2-4层)
+-----------------------+
  |
  v
输出: [ĥ_t+1, ĥ_t+2, ..., ĥ_t+m] (预测的未来特征)

自回归头的规模通常只有原模型的1-10%。

算法流程

阶段1: 特征投机

  • 使用自回归头预测未来m个特征
  • Ĥ = AR_Head([h_t, h_t-1, ...])

阶段2: Token生成

  • 通过LM头将特征转换为token分布
  • 对每个ĥ_t+i: p_t+i = Softmax(W_lm · ĥ_t+i)

阶段3: 大模型验证

  • 大模型并行验证候选token
  • 接受/拒绝机制同投机解码

EAGLE 伪代码

python
class EAGLEModel(nn.Module):
   def __init__(self, base_model, ar_layers=2, ar_hidden=None):
       super().__init__()
       self.base_model = base_model
       hidden_size = base_model.config.hidden_size

       # 自回归头:小型Transformer
       self.ar_head = nn.TransformerDecoder(
             nn.TransformerDecoderLayer(
                  d_model=hidden_size,
                  nhead=hidden_size // 64,
                  dim_feedforward=ar_hidden or hidden_size * 2,
                  batch_first=True
             ),
             num_layers=ar_layers
       )

       #   特征投影
       self.feature_proj = nn.Linear(hidden_size, hidden_size)


   def forward(self, input_ids):
       #   获取基础模型特征
       outputs = self.base_model(input_ids, output_hidden_states=True)
       features = outputs.hidden_states[-1]      # [batch, seq, hidden]
       
       return features


   def speculate_features(self, features, num_steps):
       """ 使用自回归头投机未来特征"""
       speculated = []
       current = features


       for _ in range(num_steps):
             #  自回归预测下一个特征
             next_feature = self.ar_head(current, current)
             next_feature = self.feature_proj(next_feature[:, -1:, :])
             speculated.append(next_feature)
             current = torch.cat([current, next_feature], dim=1)

       return torch.cat(speculated, dim=1)


   def features_to_tokens(self, features):
       """ 将特征转换为token分布"""
       logits = self.base_model.lm_head(features)
       return F.softmax(logits, dim=-1)


   def generate(self, prompt, max_tokens, gamma=5):
       tokens = tokenize(prompt)

       while len(tokens) < max_tokens:
               #   获取当前特征
               features = self.forward(tokens)

               #   投机未来特征
               spec_features = self.speculate_features(features[:, -4:, :], ga

               #   转换为token分布
               draft_dists = self.features_to_tokens(spec_features)
               draft_tokens = [sample(dist) for dist in draft_dists[0]]

               #   大模型验证
               verified = self.verify(tokens, draft_tokens, draft_dists[0])
               tokens.extend(verified)

         return detokenize(tokens)


     def verify(self, prefix, draft_tokens, draft_dists):
         """验证候选token"""
         all_tokens = prefix + draft_tokens
         p_dists = self.base_model(all_tokens, return_all_probs=True)

         accepted = []
         for i, (x, q_dist) in enumerate(zip(draft_tokens, draft_dists)):
               p_dist = p_dists[len(prefix) + i]
               p_x = p_dist[x]
               q_x = q_dist[x]
             
               if random() < min(1.0, p_x / q_x):
                    accepted.append(x)
               else:
                    adjusted = normalize(max(0, p_dist - q_dist))
                    accepted.append(sample(adjusted))
                    break

         return accepted

训练方法

EAGLE的训练分为两个阶段

  1. 特征收集阶段:
    • 用基础模型处理训练数据
    • 保存每层的隐藏状态
  2. 自回归头训练:
    • 输入:历史特征序列 \([h_{t-k}, ..., h_t]\)
    • 目标:未来特征序列 \([h_{t+1}, ..., h_{t+m}]\)
    • 损失:MSE损失 \(\mathcal{L} = \sum_{i=1}^{m} ||\hat{h}_{t+i} - h_{t+i}||^2\).

性能分析

EAGLE相比 Medusa的优势:

  • 更高的接受率(特征级预测更准确)
  • 更好的长距离依赖建模 - 在复杂任务上表现更稳定

实测加速比:

  • LLaMA-2-7B: 2.5-3.0x
  • LLaMA-2-70B: 2.0-2.5x

7.2.4 Lookahead Decoding

问题背景

投机解码、Medusa和EAGLE都需要额外的模型或训练。Lookahead Decoding提出了一种无需任何额外模型或训练的并行解码方法。

核心思想

Lookahead Decoding基于 Jacobi迭代的思想: 1. 打破顺序依赖:将自回归生成视为求解方程组 2. 并行猜测:基于n-gram匹配生成多个候选token 3. 并行验证:一次前向传播验证所有猜测

  • n-gram:长度为 n 的连续子序列(子串或子词)。
  • n-gram 匹配:在两个序列中寻找共同的 n-gram,并以此衡量它们的相似度或重叠程度。

Jacobi迭代视角

自回归生成可以表示为:

\(x_{t+1} = f(x_1, x_2, ..., x_t)\)

Jacobi 迭代同时更新所有位置:

\(x_i^{(k+1)} = f(x_1^{(k)}, x_2^{(k)}, ..., x_{i-1}^{(k)})\)

对于语言模型,这意味着可以并行猜测多个未来token。

\(x_1\), \(x_2\), …, \(x_t\):已经生成的 token id 序列(或更严格地说,是这些 token 对应的离散变量 / 隐表示)\(x_i\)就是“第 \(i\) 个位置上的 token”。

这里的 \(k\)猜测/修正轮次(iteration index),Jacobi 做法是“用上一轮的全部值,同时算下一轮的全部值”。

n-gram匹配策略

Lookahead Decoding维护一个n-gram窗口:

latex
已生成序列: [The, cat, sat, on, the, ...]
                  |___| (2-gram: "cat sat")

在窗口中查找匹配的n-gram:
 位置3: "cat sat" → 下一个token是 "on"
 位置10: "cat sat" → 下一个token是 "under"
 
 候选token: ["on", "under", ...]

Lookahead Decoding 伪代码

python
class LookaheadDecoding:
   def __init__(self, model, window_size=5, ngram_size=3, num_candidates=1
       self.model = model
       self.window_size = window_size
       self.ngram_size = ngram_size
       self.num_candidates = num_candidates
        self.ngram_cache = {}     # n-gram ->    后续token列表

   def update_cache(self, tokens):
       """ 更新n-gram缓存"""
       for i in range(len(tokens) - self.ngram_size):
             ngram = tuple(tokens[i:i + self.ngram_size])
             next_token = tokens[i + self.ngram_size]

             if ngram not in self.ngram_cache:
                self.ngram_cache[ngram] = []
             if next_token not in self.ngram_cache[ngram]:
                self.ngram_cache[ngram].append(next_token)


   def generate_candidates(self, tokens):
       """ 基于n-gram匹配生成候选token"""
       candidates = []

       #   使用最近的n-gram查找
       for n in range(self.ngram_size, 1, -1):
             if len(tokens) < n:
                continue

             ngram = tuple(tokens[-n:])
             if ngram in self.ngram_cache:
                candidates.extend(self.ngram_cache[ngram])

       #   去重并限制数量
       candidates = list(dict.fromkeys(candidates))[:self.num_candidates]
       return candidates

   def verify_candidates(self, tokens, candidates):
       """ 并行验证候选token"""
       if not candidates:
             return [self.model.generate_next(tokens)]

       #   构建验证批次
       sequences = [tokens + [cand] for cand in candidates]

       #   并行前向
       logits = self.model.batch_forward(sequences)

       #   检查哪些猜测是正确的
       accepted = []
       for i, cand in enumerate(candidates):
            #   如果模型预测的下一个token与猜测一致
            predicted = logits[i].argmax(dim=-1)
            if predicted == cand:
                 accepted.append(cand)
            else:
                 #  添加正确的token并停止
                 accepted.append(predicted.item())
                 break

         return accepted


     def generate(self, prompt, max_tokens):
         tokens = tokenize(prompt)

         while len(tokens) < max_tokens:
            #   生成候选token
            candidates = self.generate_candidates(tokens)

            #   验证候选
            verified = self.verify_candidates(tokens, candidates)

            #   更新序列和缓存
            tokens.extend(verified)
            self.update_cache(tokens)

         return detokenize(tokens)

复杂度分析

操作复杂度说明
n-gram查找\(O(1)\)哈希表实现
候选生成\(O(k)\)k为候选数
并行验证\(O(1)\)单次batch前向
缓存更新\(O(w)\)w为窗口大小

实际应用与限制

优势:

  • 无需额外模型或训练
  • 实现简单
  • 适用于任何自回归模型

限制: 1. n-gram稀疏性:长n-gram匹配率低 2. 内容依赖:在模板化内容上效果好,在创造性内容上效果差 3. 缓存开销:需要维护n-gram缓存

典型加速比:

  • 代码生成:1.5-2.0x
  • 对话:1.2-1.5x
  • 创意写作:<1.2x

用心记录,持续成长