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 量化精度对比
| 精度 | 显存占用 | 推理速度 | 精度损失 |
|---|---|---|---|
| FP32 | 100% | 基准 | 无 |
| FP16/BF16 | 50% | 1.5-2x | 几乎无 |
| INT8 | 25% | 2-4x | 小(需校准) |
| INT4 | 12.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 蒸馏原理
核心思想: 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 | 通用 | 易用但慢 |
| vLLM | LLM 推理 | PagedAttention、高吞吐 |
| TensorRT-LLM | NVIDIA GPU | 极致性能 |
| llama.cpp | CPU/边缘 | GGUF 量化 |
| ONNX Runtime | 跨平台 | 通用优化 |
| OpenVINO | Intel CPU | Intel 优化 |
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
)
七、混合搜索(Hybrid Search)
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.export 的 opset_version 设高一些。
坑五:TensorRT 编译时间长
TensorRT 编译 engine 需要几分钟。编译一次保存 engine 文件,下次直接加载。
十、技术选型建议
按场景选压缩技术
| 场景 | 推荐方案 |
|---|---|
| LLM 部署 | QLoRA 4-bit + vLLM |
| 移动端 | INT8 量化 + 剪枝 |
| CPU 推理 | ONNX Runtime + INT8 |
| GPU 推理 | TensorRT + FP16 |
| 边缘设备 | 蒸馏到超小模型 + INT4 |
按需求选推理框架
| 需求 | 推荐框架 |
|---|---|
| 高吞吐 LLM | vLLM |
| 极致 GPU 性能 | TensorRT-LLM |
| 跨平台 | ONNX Runtime |
| CPU/边缘 | llama.cpp |
| 快速原型 | HuggingFace Transformers |
十一、写在最后
模型压缩是一个精度、速度、大小的三角权衡。
几条核心原则:
- 先量化,再考虑其他:量化是最简单、最有效的优化
- 蒸馏适合结构性压缩:大模型 → 小模型
- 剪枝要配合微调:剪完必须恢复训练
- 推理优化比模型优化更立竿见影:vLLM、Flash Attention
- 端到端优化:数据加载、预处理、推理、后处理都要看
- 监控推理性能:延迟、吞吐、显存
压缩不是一次性的。随着模型迭代、数据变化、业务需求调整,需要持续优化。
本文整合了 7 篇模型压缩与推理优化相关文章,涵盖剪枝、量化、知识蒸馏、Logits 处理、混合搜索、推理框架选型、模型服务性能优化等核心技术。
版权声明: 本文首发于 指尖魔法屋-AI模型压缩与推理优化实战指南:从剪枝到蒸馏(https://blog.thinkmoon.cn/post/ai-model-compression-comprehensive-guide/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。