Skip to content

高效推理论文精选

大语言模型的推理成本是实际部署的关键瓶颈。从量化、蒸馏到稀疏化和架构创新,高效推理研究致力于在保持模型性能的同时大幅降低计算和内存开销。

高效推理技术

推理瓶颈分析

LLM 推理的主要瓶颈:

  1. 内存带宽:参数加载是瓶颈,而非计算
  2. KV 缓存:上下文越长,KV 缓存越大
  3. 批处理效率:自回归生成难以有效批处理
  4. 延迟:首 Token 延迟和每 Token 延迟

量化技术

训练后量化(PTQ)

无需重新训练,直接量化已训练模型:

python
import torch

class LLMInt8Quantizer:
    """LLM.int8() 量化"""
    def __init__(self, threshold=6.0):
        self.threshold = threshold

    def quantize_weight(self, weight):
        """量化权重到 int8"""
        # 检测离群值维度
        outlier_mask = torch.any(weight.abs() > self.threshold, dim=0)

        # 正常维度:int8 量化
        normal_weight = weight[:, ~outlier_mask]
        scale = normal_weight.abs().max(dim=0).values / 127.0
        quantized = (normal_weight / scale).round().to(torch.int8)

        # 离群维度:保持 fp16
        outlier_weight = weight[:, outlier_mask]

        return quantized, scale, outlier_weight, outlier_mask

    def dequantize_and_matmul(self, quantized, scale, outlier_weight, outlier_mask, x):
        """反量化并计算矩阵乘法"""
        B, T, D = x.shape
        output = torch.zeros(B, T, quantized.shape[0], dtype=torch.float16, device=x.device)

        # 正常维度:反量化后计算
        x_normal = x[:, :, ~outlier_mask]
        dequantized = quantized.float() * scale.unsqueeze(0)
        output[:, :, ~outlier_mask] = x_normal @ dequantized.T

        # 离群维度:fp16 计算
        x_outlier = x[:, :, outlier_mask]
        output[:, :, outlier_mask] = x_outlier @ outlier_weight.T

        return output

量化方案对比

方案精度方法精度损失
FP1616-bit基线0%
LLM.int8()8-bit混合精度<1%
GPTQ4-bit逐层最优量化~1%
AWQ4-bit注意力加权量化<1%
GGUF2-8-bit多种量化级别可变

知识蒸馏

将大模型的知识转移到小模型:

白盒蒸馏

使用大模型的 logits 或中间表示训练小模型:

  • MiniLLM:逆向 KL 蒸馏,生成能力更强
  • GKD:自蒸馏,用模型自身的生成进行蒸馏

黑盒蒸馏

仅使用大模型的输入输出对:

  • Alpaca:GPT-3.5 输出微调 LLaMA
  • Vicuna:ShareGPT 数据微调

蒸馏的关键

蒸馏的效果取决于两个因素:数据质量和教师模型的输出分布。白盒蒸馏利用 logits 中的"暗知识"(类别间的相似度),通常优于纯黑盒蒸馏。

稀疏化

结构化剪枝

移除整个神经元/注意力头/层:

  • LLM-Pruner:发现并剪除不重要的结构
  • SliceGPT:矩阵切片式剪枝
  • ShortGPT:跳过冗余层

非结构化剪枝

将个别权重置零:

  • 理论压缩率高但硬件不友好
  • 通常需要专用稀疏计算库

推理优化系统

推理引擎

引擎特点加速比
vLLMPagedAttention2-4x
TensorRT-LLMNVIDIA 优化2-5x
LMDeploy量化+推理2-4x
SGLang编程式推理2-6x

投机解码

用小模型快速生成候选,大模型并行验证:

python
def speculative_decoding(draft_model, target_model, prompt, max_tokens=100, gamma=4):
    """投机解码"""
    generated = prompt
    for _ in range(max_tokens // gamma):
        # 小模型快速生成 gamma 个候选 Token
        draft_tokens = draft_model.generate(generated, max_new_tokens=gamma)

        # 大模型并行验证
        target_probs = target_model.get_probs(generated + draft_tokens)

        # 接受/拒绝每个 Token
        accepted = 0
        for i, token in enumerate(draft_tokens):
            if should_accept(token, target_probs[i], draft_probs[i]):
                accepted += 1
            else:
                # 从大模型分布重新采样
                token = sample_from(target_probs[i])
                accepted += 1
                break

        generated += draft_tokens[:accepted]
    return generated

相关资源

最近更新