AI模型压缩与推理优化实战指南:从剪枝到蒸馏

前言:为什么需要模型压缩

训练时模型可以很大(GPU 多、不在乎延迟),但部署时受限于显存、延迟、能耗。模型压缩的目标:在精度损失可控的前提下,让模型更小、更快、更省

核心压缩技术:

  • 剪枝:去掉不重要的参数
  • 量化:降低参数精度
  • 蒸馏:大模型教小模型
  • 架构搜索:找到高效结构
  • 推理优化:KV Cache、Flash Attention

一、模型压缩全景

技术减小体积加速推理精度损失实现难度
剪枝
量化
蒸馏
NAS
TensorRT

二、模型剪枝

2.1 剪枝类型

类型说明硬件友好
非结构化单个权重置零否(稀疏计算支持有限)
结构化整个通道/层剪掉是(直接变小的稠密矩阵)
半结构化N:M 稀疏部分(A100 支持 2:4)

2.2 结构化剪枝实现

import torch
import torch.nn as nn
import torch.nn.utils.prune as prune

# 1. 局部剪枝(单层)
model = resnet18()
module = model.conv1
prune.l1_unstructured(module, name='weight', amount=0.3)  # 剪掉 30%

# 2. 全局剪枝(跨层)
parameters_to_prune = [
    (model.conv1, 'weight'),
    (model.layer1[0].conv1, 'weight'),
    (model.layer1[0].conv2, 'weight'),
]

prune.global_unstructured(
    parameters_to_prune,
    pruning_method=prune.L1Unstructured,
    amount=0.2
)

# 3. 通道剪枝(更实用)
def channel_prune(model, amount=0.3):
    """按通道重要性剪枝"""
    for name, module in model.named_modules():
        if isinstance(module, nn.Conv2d):
            # 计算 BN 层的 gamma 作为重要性
            bn = find_corresponding_bn(model, name)
            importance = bn.weight.data.abs()
            threshold = torch.quantile(importance, amount)

            mask = importance > threshold
            module.weight.data = module.weight.data[mask]
            module.out_channels = mask.sum().item()

2.3 剪枝的训练流程

# 1. 训练原始模型
model = train_full_model()

# 2. 剪枝
prune_model(model, amount=0.3)

# 3. 微调恢复精度
finetune_model(model, epochs=10)

# 4. 重复(渐进式剪枝)
for iteration in range(5):
    prune_model(model, amount=0.1)
    finetune_model(model, epochs=5)

2.4 剪枝的坑

坑一:剪完不微调

剪枝后精度会掉,必须微调恢复。

坑二:非结构化剪枝没加速

稀疏矩阵硬件支持有限。部署优先用结构化剪枝。

坑三:剪过头

剪超过 50% 通常精度掉得厉害。渐进式剪枝,每步 10%。

三、模型量化

3.1 量化精度对比

精度显存占用推理速度精度损失
FP32100%基准
FP16/BF1650%1.5-2x几乎无
INT825%2-4x小(需校准)
INT412.5%3-6x明显

3.2 训练后量化(PTQ)

import torch
from torch.quantization import quantize_dynamic, quantize_fx

# 1. 动态量化(最简单,只量化权重)
quantized_model = quantize_dynamic(
    model,
    {nn.Linear},
    dtype=torch.qint8
)

# 2. 静态量化(权重 + 激活,需要校准数据)
from torch.quantization import prepare, convert

model.eval()
model.qconfig = torch.quantization.get_default_qconfig('fbgemm')

# 插入观察者
prepared_model = prepare(model)

# 用校准数据跑前向
for batch in calibration_data:
    prepared_model(batch)

# 转换为量化模型
quantized_model = convert(prepared_model)

3.3 量化感知训练(QAT)

from torch.quantization import prepare_qat

# 训练时就模拟量化
model.train()
model.qconfig = torch.quantization.get_default_qat_qconfig('fbgemm')
qat_model = prepare_qat(model)

# 正常训练
for epoch in range(num_epochs):
    for batch in train_loader:
        output = qat_model(batch)
        loss = criterion(output, target)
        loss.backward()
        optimizer.step()

# 转换为真正的量化模型
quantized_model = convert(qat_model.eval())

3.4 LLM 量化(bitsandbytes)

from transformers import AutoModelForCausalLM, BitsAndBytesConfig

# 4-bit 量化(QLoRA 用的)
bnb_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.bfloat16,
    bnb_4bit_use_double_quant=True
)

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-2-7b",
    quantization_config=bnb_config,
    device_map="auto"
)

3.5 GGUF 量化(llama.cpp)

# 转换为 GGUF
python convert-hf-to-gguf.py model/ --outtype f16

# 量化为 Q4_K_M(推荐)
./quantize model-f16.gguf model-q4_k_m.gguf Q4_K_M

# 量化级别选择
# Q4_K_S: 最激进,4-bit
# Q4_K_M: 平衡(推荐)
# Q5_K_M: 更高精度
# Q8_0:   接近 FP16

四、知识蒸馏

4.1 蒸馏原理

graph LR A[Teacher 大模型] -->|soft labels| C[Student 小模型] B[Hard Labels 真实标签] --> C C -->|学习| D[蒸馏后小模型]

核心思想: Teacher 输出的概率分布比 hard label 信息更丰富(暗知识)。

4.2 标准蒸馏

class DistillationTrainer:
    def __init__(self, teacher, student, temperature=4.0, alpha=0.5):
        self.teacher = teacher.eval()  # 冻结
        self.student = student
        self.temperature = temperature
        self.alpha = alpha

    def compute_loss(self, x, y):
        # Teacher 输出(不计算梯度)
        with torch.no_grad():
            teacher_logits = self.teacher(x)

        # Student 输出
        student_logits = self.student(x)

        # 硬标签损失
        hard_loss = F.cross_entropy(student_logits, y)

        # 软标签损失(KL 散度)
        soft_loss = F.kl_div(
            F.log_softmax(student_logits / self.temperature, dim=-1),
            F.softmax(teacher_logits / self.temperature, dim=-1),
            reduction='batchmean'
        ) * (self.temperature ** 2)

        # 加权组合
        return self.alpha * hard_loss + (1 - self.alpha) * soft_loss

4.3 特征蒸馏

class FeatureDistillationModel(nn.Module):
    def __init__(self, teacher, student):
        super().__init__()
        self.teacher = teacher
        self.student = student
        # 适配层(对齐维度)
        self.adapter = nn.Linear(student_dim, teacher_dim)

    def forward(self, x):
        # Teacher 中间层特征
        with torch.no_grad():
            teacher_features = self.teacher.get_intermediate(x)

        # Student 中间层特征
        student_features = self.student.get_intermediate(x)
        adapted = self.adapter(student_features)

        # 特征蒸馏损失
        feature_loss = F.mse_loss(adapted, teacher_features)
        return feature_loss

4.4 LLM 蒸馏

# GPT-4 蒸馏到小模型
def generate_distillation_data(teacher, prompts):
    """用 Teacher 生成训练数据"""
    dataset = []
    for prompt in prompts:
        # Teacher 生成高质量回答
        response = teacher.generate(prompt)
        dataset.append({'instruction': prompt, 'output': response})
    return dataset

# 然后用这些数据 SFT 训练 Student

4.5 蒸馏的坑

坑一:Teacher 不够强

Teacher 准确率 80%,Student 上限就是 80%。Teacher 要远强于 Student。

坑二:温度太高/太低

  • T 太低(1):接近 hard label,失去暗知识
  • T 太高(20):分布太平,信号弱
  • 推荐:4-8

坑三:只蒸馏最后输出

中间层特征也很有价值。结合 logit 蒸馏 + 特征蒸馏。

五、推理优化

5.1 KV Cache

# Transformer 推理的 KV Cache
# 不用每次重新计算历史 token 的 K、V

class CachedAttention(nn.Module):
    def forward(self, x, past_kv=None):
        q, k, v = self.qkv(x).chunk(3, dim=-1)

        if past_kv is not None:
            past_k, past_v = past_kv
            k = torch.cat([past_k, k], dim=-2)
            v = torch.cat([past_v, v], dim=-2)

        new_kv = (k, v)

        # 只对当前 q 和累积的 k/v 计算 attention
        attn = attention(q, k, v)
        return attn, new_kv

效果: 推理速度提升 5-10 倍。

5.2 Flash Attention

# 标准 Attention:IO 密集(频繁读写 HBM)
# Flash Attention:分块计算,减少 HBM 读写

from flash_attn import flash_attn_func

attn_output = flash_attn_func(q, k, v, causal=True)

效果: 长序列训练/推理加速 2-4 倍。

5.3 Continuous Batching(vLLM)

# vLLM 自动实现连续批处理
python -m vllm.entrypoints.api_server \
    --model meta-llama/Llama-2-7b \
    --enable-chunked-prefill \
    --max-num-batched-tokens 4096

效果: 吞吐量提升 4-5 倍。

5.4 推理框架对比

框架适用特点
Transformers通用易用但慢
vLLMLLM 推理PagedAttention、高吞吐
TensorRT-LLMNVIDIA GPU极致性能
llama.cppCPU/边缘GGUF 量化
ONNX Runtime跨平台通用优化
OpenVINOIntel CPUIntel 优化

5.5 LLM 推理优化技巧

# 1. 温度和采样参数
generation_config = GenerationConfig(
    temperature=0.7,
    top_p=0.9,
    top_k=50,
    do_sample=True,
    max_new_tokens=512,
    num_beams=1,        # 贪婪解码更快
    repetition_penalty=1.1
)

# 2. Speculative Decoding(推测解码)
# 用小模型先生成 draft,大模型验证
from transformers import AutoModelForCausalLM

draft_model = AutoModelForCausalLM.from_pretrained("small-model")
target_model = AutoModelForCausalLM.from_pretrained("large-model")

# 3. 模型合并(LoRA 合并)
merged_model = model.merge_and_unload()
merged_model.save_pretrained("./merged")

六、Logits 处理器

6.1 什么是 Logits 处理器

在生成时修改 logits,控制输出。

from transformers import LogitsProcessor, LogitsProcessorList

class NoBadWordsProcessor(LogitsProcessor):
    """禁止某些词"""
    def __init__(self, bad_words_ids):
        self.bad_words_ids = bad_words_ids

    def __call__(self, input_ids, scores):
        for bad_word_id in self.bad_words_ids:
            scores[:, bad_word_id] = -float('inf')
        return scores

class TemperatureLogitsProcessor(LogitsProcessor):
    """温度调节"""
    def __init__(self, temperature):
        self.temperature = temperature

    def __call__(self, input_ids, scores):
        return scores / self.temperature

class TopKLogitsProcessor(LogitsProcessor):
    """Top-K 采样"""
    def __init__(self, top_k):
        self.top_k = top_k

    def __call__(self, input_ids, scores):
        top_k = min(self.top_k, scores.size(-1))
        indices_to_remove = scores < torch.topk(scores, top_k)[0][..., -1, None]
        scores[indices_to_remove] = -float('inf')
        return scores

# 使用
logits_processor = LogitsProcessorList([
    NoBadWordsProcessor([bad_word_id_1, bad_word_id_2]),
    TemperatureLogitsProcessor(0.7),
    TopKLogitsProcessor(50),
])

output = model.generate(
    input_ids,
    logits_processor=logits_processor,
    max_new_tokens=100
)

7.1 为什么需要混合搜索

纯向量搜索的问题:

  • 具体术语(错误代码、版本号)召回差
  • 长尾查询不精准

纯关键词搜索的问题:

  • 语义理解差
  • 同义词匹配不上

混合搜索 = 向量搜索 + BM25 关键词搜索

7.2 实现

from rank_bm25 import BM25Okapi
from sentence_transformers import SentenceTransformer
import faiss

class HybridSearch:
    def __init__(self, documents):
        self.documents = documents

        # BM25 索引
        tokenized_docs = [doc.split() for doc in documents]
        self.bm25 = BM25Okapi(tokenized_docs)

        # 向量索引
        self.embedder = SentenceTransformer('paraphrase-multilingual-MiniLM-L12-v2')
        embeddings = self.embedder.encode(documents)
        self.vector_index = faiss.IndexFlatIP(embeddings.shape[1])
        self.vector_index.add(embeddings.astype('float32'))

    def search(self, query, top_k=5, alpha=0.5):
        # BM25 搜索
        bm25_scores = self.bm25.get_scores(query.split())

        # 向量搜索
        query_emb = self.embedder.encode([query]).astype('float32')
        _, vector_indices = self.vector_index.search(query_emb, top_k * 2)
        vector_scores = np.zeros(len(self.documents))
        for rank, idx in enumerate(vector_indices[0]):
            vector_scores[idx] = 1.0 / (rank + 1)  # 倒数排名

        # 归一化
        bm25_norm = (bm25_scores - bm25_scores.min()) / (bm25_scores.max() - bm25_scores.min() + 1e-8)
        vector_norm = (vector_scores - vector_scores.min()) / (vector_scores.max() - vector_scores.min() + 1e-8)

        # 加权融合
        final_scores = alpha * vector_norm + (1 - alpha) * bm25_norm
        top_indices = np.argsort(final_scores)[::-1][:top_k]

        return [(self.documents[i], final_scores[i]) for i in top_indices]

7.3 重排序(Reranking)

from sentence_transformers import CrossEncoder

# 用 Cross Encoder 重排序(更准但更慢)
reranker = CrossEncoder('BAAI/bge-reranker-base')

def hybrid_search_with_rerank(query, top_k=5):
    # 1. 初步召回(快)
    candidates = hybrid_search.search(query, top_k=top_k * 4)

    # 2. 重排序(准)
    pairs = [(query, doc) for doc, _ in candidates]
    rerank_scores = reranker.predict(pairs)

    # 3. 返回 top_k
    ranked = sorted(zip(candidates, rerank_scores), key=lambda x: x[1], reverse=True)
    return ranked[:top_k]

八、模型服务性能优化

8.1 批处理

# 动态批处理:等待一小段时间收集请求
import asyncio
from collections import deque

class DynamicBatcher:
    def __init__(self, model, max_batch=32, max_wait=0.05):
        self.model = model
        self.max_batch = max_batch
        self.max_wait = max_wait
        self.queue = deque()

    async def predict(self, input_data):
        future = asyncio.Future()
        self.queue.append((input_data, future))
        return await future

    async def batch_worker(self):
        while True:
            if len(self.queue) >= self.max_batch:
                batch = [self.queue.popleft() for _ in range(self.max_batch)]
            elif self.queue:
                await asyncio.sleep(self.max_wait)
                batch = list(self.queue)
                self.queue.clear()
            else:
                await asyncio.sleep(0.01)
                continue

            # 批量推理
            inputs = [item[0] for item in batch]
            outputs = self.model(inputs)

            # 返回结果
            for (_, future), output in zip(batch, outputs):
                future.set_result(output)

8.2 异步推理

# FastAPI 异步推理
from fastapi import FastAPI
from fastapi.concurrency import run_in_threadpool

app = FastAPI()
model = load_model()

@app.post("/predict")
async def predict(input_data: dict):
    # 模型推理放到线程池,不阻塞事件循环
    result = await run_in_threadpool(model.predict, input_data)
    return {"result": result}

8.3 缓存

from functools import lru_cache
import hashlib

# 结果缓存
@lru_cache(maxsize=1000)
def cached_predict(input_hash):
    return model.predict(input_hash)

def predict(input_data):
    input_hash = hashlib.md5(str(input_data).encode()).hexdigest()
    return cached_predict(input_hash)

# 语义缓存(相似输入复用)
class SemanticCache:
    def __init__(self, threshold=0.95):
        self.embedder = SentenceTransformer(...)
        self.cache = []

    def get(self, query):
        query_emb = self.embedder.encode([query])
        for cached_query, cached_emb, result in self.cache:
            sim = cosine_similarity(query_emb, cached_emb)
            if sim > self.threshold:
                return result
        return None

九、踩坑总结

坑一:量化后精度掉太多

解决: 用 QAT 代替 PTQ,或用更高精度(INT8 代替 INT4)。

坑二:蒸馏后 Student 学不到 Teacher 水平

解决:

  • 检查 Teacher 是否足够强
  • 调温度参数(4-8)
  • 加特征蒸馏
  • 用更多训练数据

坑三:剪枝后推理没变快

非结构化剪枝稀疏矩阵硬件不支持。用结构化剪枝。

坑四:ONNX 导出失败

某些自定义算子 ONNX 不支持。torch.onnx.exportopset_version 设高一些。

坑五:TensorRT 编译时间长

TensorRT 编译 engine 需要几分钟。编译一次保存 engine 文件,下次直接加载。

十、技术选型建议

按场景选压缩技术

场景推荐方案
LLM 部署QLoRA 4-bit + vLLM
移动端INT8 量化 + 剪枝
CPU 推理ONNX Runtime + INT8
GPU 推理TensorRT + FP16
边缘设备蒸馏到超小模型 + INT4

按需求选推理框架

需求推荐框架
高吞吐 LLMvLLM
极致 GPU 性能TensorRT-LLM
跨平台ONNX Runtime
CPU/边缘llama.cpp
快速原型HuggingFace Transformers

十一、写在最后

模型压缩是一个精度、速度、大小的三角权衡

几条核心原则:

  1. 先量化,再考虑其他:量化是最简单、最有效的优化
  2. 蒸馏适合结构性压缩:大模型 → 小模型
  3. 剪枝要配合微调:剪完必须恢复训练
  4. 推理优化比模型优化更立竿见影:vLLM、Flash Attention
  5. 端到端优化:数据加载、预处理、推理、后处理都要看
  6. 监控推理性能:延迟、吞吐、显存

压缩不是一次性的。随着模型迭代、数据变化、业务需求调整,需要持续优化。


本文整合了 7 篇模型压缩与推理优化相关文章,涵盖剪枝、量化、知识蒸馏、Logits 处理、混合搜索、推理框架选型、模型服务性能优化等核心技术。

版权声明: 本文首发于 指尖魔法屋-AI模型压缩与推理优化实战指南:从剪枝到蒸馏https://blog.thinkmoon.cn/post/ai-model-compression-comprehensive-guide/) 转载或引用必须申明原指尖魔法屋来源及源地址!