3. AI 框架与性能调优:效率的引擎
训练优化是在有限显存和算力下逼近模型性能上限的工程艺术。核心思路:用精度换速度、用计算换显存、用调度换吞吐。下面按六个维度展开。
一、混合精度(FP16 / BF16 / FP8)
精度格式对比
| 格式 | 位宽 | 指数位 | 尾数位 | 动态范围 | 精度 | 硬件支持 |
|---|---|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | 大 | 高 | 通用 |
| TF32 | 19 | 8 | 10 | 大 | 中 | Ampere+ Tensor Core |
| BF16 | 16 | 8 | 7 | 大(同 FP32) | 低 | Ampere+ / TPU |
| FP16 | 16 | 5 | 10 | 小(易溢出) | 中 | Volta+ Tensor Core |
| FP8 (E4M3) | 8 | 4 | 3 | 小 | 低 | Hopper+ |
| FP8 (E5M2) | 8 | 5 | 2 | 中 | 更低 | Hopper+ |
核心机制
FP16 训练:
-
前向/反向用 FP16,主权重用 FP32(Master Weights)。
-
Loss Scaling:防止梯度下溢。静态或动态缩放。
-
问题:动态范围小,易溢出/下溢。
BF16 训练:
-
动态范围与 FP32 相同,无需 Loss Scaling。
-
精度略低,但训练稳定性好。
-
推荐优先使用 BF16。
FP8 训练:
-
Hopper 架构原生支持,Transformer Engine 自动管理。
-
两种格式:E4M3 (前向,精度高)、E5M2(反向,范围大)。
-
Per-Tensor / Per-Channel Scaling:动态调整缩放因子。
-
需配合 FP16/BF16 主权重和 FP32 累加。
混合精度实现
python
# PyTorch AMP(自动混合精度)
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler() # FP16 需要,BF16 可不用
for data, target in loader:
optimizer.zero_grad()
with autocast(dtype=torch.bfloat16): # 或 float16
output = model(data)
loss = loss_fn(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
python
# Transformer Engine(FP8)
import transformer_engine.pytorch as te
from transformer_engine.common import recipe
fp8_recipe = recipe.DelayedScaling(
margin=0, interval=1, fp8_format=recipe.Format.HYBRID)
model = te.Linear(768, 768, params_dtype=torch.bfloat16)
with te.fp8_autocast(enabled=True, fp8_recipe=fp8_recipe):
output = model(input)
选择建议
| 场景 | 推荐精度 |
|---|---|
| 通用训练 | BF16 + FP32 主权重 |
| 老硬件(V100) | FP16 + Loss Scaling |
| Hopper+ 追求极致 | FP8 + BF16 主权重 |
| 推理 | FP16/INT8/FP8 |
| 数值敏感层 | FP32(LayerNorm、Softmax) |
二、梯度累积
原理
小 batch 分多次前向反向,累积梯度后再更新,等效大 batch。
text
实际 batch = micro_batch × grad_accum_steps × data_parallel_size
实现
python
accum_steps = 4
for i, (data, target) in enumerate(loader):
with autocast(dtype=torch.bfloat16):
output = model(data)
loss = loss_fn(output, target) / accum_steps # 关键:除以累积步数
loss.backward()
if (i + 1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad()
关键点
-
Loss 归一化:除以累积步数,保证梯度尺度一致。
-
BN 问题 :BatchNorm 统计量按 micro-batch 计算,等效 batch 不一致。可用 SyncBN 或 GroupNorm。
-
DDP 配合 :梯度累积期间不触发 All-Reduce,累积完再同步。PyTorch 需用
no_sync():
python
for i, data in enumerate(loader):
with model.no_sync() if (i+1) % accum_steps != 0 else contextlib.nullcontext():
loss = model(data)
loss.backward()
if (i+1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad()
梯度累积 vs 大 Batch
| 维度 | 梯度累积 | 真大 Batch |
|---|---|---|
| 显存 | 低 | 高 |
| 吞吐 | 略低(多次前向) | 高 |
| BN 统计 | 不准 | 准确 |
| 学习率 | 需调 | 需调 |
| 适用 | 显存受限 | 显存充足 |
三、激活重计算(Gradient Checkpointing)
原理
前向不保存中间激活,反向时重新计算。用计算换显存。
text
常规:前向保存所有激活 → 反向直接用
重计算:前向只保存边界激活 → 反向重新前向计算中间激活
显存与计算权衡
-
显存:从 O(n) 降到 O(√n)(每层都 checkpoint)或 O(n/k)(每 k 层)。
-
计算 :增加约 30%~40% 前向计算量。
-
策略:
-
Full:每层都重计算,显存最省。
-
Selective:只对显存大户(Attention)重计算。
-
Offload:激活卸载到 CPU,反向时取回。
-
实现
python
# PyTorch 原生
from torch.utils.checkpoint import checkpoint
class TransformerBlock(nn.Module):
def forward(self, x):
return checkpoint(self._forward, x, use_reentrant=False)
def _forward(self, x):
x = self.attn(x)
x = self.mlp(x)
return x
python
# DeepSpeed 激活卸载
{
"activation_checkpointing": {
"partition_activations": true,
"cpu_checkpointing": true,
"contiguous_memory_optimization": true,
"number_checkpoints": 4
}
}
选择策略
| 方法 | 显存节省 | 计算开销 | 适用 |
|---|---|---|---|
| 不重计算 | 0 | 0 | 小模型 |
| Selective | 中 | 低 | 大模型推荐 |
| Full | 高 | 中 | 显存极紧 |
| CPU Offload | 极高 | 高(PCIe) | 超大模型 |
四、Checkpointing 优化
训练 Checkpoint 的挑战
-
大模型 Checkpoint 可达 TB 级。
-
写入耗时占训练 10%~30%。
-
故障恢复需快速加载。
优化手段
| 手段 | 说明 | 收益 |
|---|---|---|
| 异步写入 | 训练继续,后台写盘 | 消除阻塞 |
| 分层存储 | 先写 NVMe,再转分布式 | 降低延迟 |
| 增量 Checkpoint | 只写变化参数 | 减少数据量 |
| 分片存储 | 每卡写自己的分片 | 并行 I/O |
| 压缩 | FP16/BF16、zstd | 减少体积 |
| 零拷贝 | GPU Direct Storage | 绕过 CPU |
| P2P 恢复 | 节点间直接传输 | 加速恢复 |
| 格式优化 | Safetensors、DCP | 加载快 |
实现
python
# PyTorch Distributed Checkpoint
from torch.distributed.checkpoint import save, load
from torch.distributed.checkpoint.state_dict import get_state_dict
state_dict = get_state_dict(model, optimizer)
save(state_dict, checkpoint_id="/ckpt/step_1000")
# 加载
load(state_dict, checkpoint_id="/ckpt/step_1000")
python
# DeepSpeed Checkpoint
model_engine.save_checkpoint("/ckpt", tag="step_1000")
model_engine.load_checkpoint("/ckpt", tag="step_1000")
最佳实践
-
周期平衡:太频繁影响吞吐,太少故障损失大。
-
分层:内存 → NVMe → 分布式存储。
-
异步:Checkpoint 与训练重叠。
-
验证:定期验证 Checkpoint 可恢复。
-
清理:保留最近 N 个 + 最优 N 个。
五、动态批处理
原理
根据序列长度或显存动态调整 batch size,最大化 GPU 利用率。
类型
| 类型 | 说明 | 场景 |
|---|---|---|
| Token-based | 按总 token 数组 batch | LLM 训练/推理 |
| Length-based | 相似长度组 batch | 减少 padding |
| Continuous Batching | 迭代级动态加入/退出 | 推理服务 |
| Sequence Packing | 多序列拼接,无 padding | 训练加速 |
Sequence Packing
python
# 将多个短序列拼成一个长序列,避免 padding 浪费
# 例:3 个序列 [100, 200, 150] → 拼接成 450,而非 pad 到 200×3=600
def pack_sequences(sequences, max_len):
packed = []
current = []
current_len = 0
for seq in sequences:
if current_len + len(seq) > max_len:
packed.append(torch.cat(current))
current, current_len = [], 0
current.append(seq)
current_len += len(seq)
if current:
packed.append(torch.cat(current))
return packed
Continuous Batching(推理)
text
传统 Batching:等所有请求完成才释放 batch
Continuous Batching:请求完成即退出,新请求立即加入
→ GPU 利用率从 30% 提升到 80%+
动态批处理框架
-
Triton Inference Server:动态 batching、序列 batching。
-
vLLM:PagedAttention + Continuous Batching。
-
TensorRT-LLM:In-flight Batching。
-
DeepSpeed-MII:动态批处理。
六、MoE 调度
MoE 核心挑战
-
负载不均:热门专家过载,冷门专家空闲。
-
All-to-All 通信:token 分发与回收开销大。
-
容量因子:每专家处理 token 上限,超限丢弃。
-
训练不稳定:路由震荡,专家退化。
调度优化
1. 负载均衡
python
# 负载均衡损失(Switch Transformer)
def load_balance_loss(gate_logits, expert_indices, num_experts):
# 每个专家的 token 比例
tokens_per_expert = torch.histc(expert_indices.float(), bins=num_experts)
fraction = tokens_per_expert / expert_indices.numel()
# 每个专家的平均门控概率
gate_prob = F.softmax(gate_logits, dim=-1).mean(dim=0)
# 辅助损失
return num_experts * (fraction * gate_prob).sum()
2. 容量因子
python
capacity = int((tokens_per_batch / num_experts) * capacity_factor)
# capacity_factor: 1.0~2.0,越大越不易丢弃,但显存增加
3. 专家并行调度
text
Token → Gate → Top-K 专家 → All-to-All 分发 → 专家计算 → All-to-All 回收 → 加权
4. 通信优化
-
通信重叠:All-to-All 与计算流水线化。
-
专家分组:同节点专家优先,减少跨节点。
-
Token 排列:按专家排序,合并通信。
-
FP8 通信:减少通信量。
5. 路由优化
| 方法 | 说明 |
|---|---|
| Top-K | 每 token 选 K 个专家 |
| Expert Choice | 每专家选 Top-K token,天然均衡 |
| Sinkhorn | 最优传输,均衡路由 |
| Soft MoE | 软加权,无离散路由 |
| DeepSeek-MoE | 细粒度专家 + 共享专家 |
DeepSeek-MoE 调度
-
细粒度专家:更多小专家,组合更灵活。
-
共享专家:所有 token 都经过,保证基础能力。
-
设备受限路由:控制跨节点通信。
-
Token 丢弃:容量因子控制。
MoE 训练配置示例
python
# DeepSpeed-MoE
{
"moe": {
"num_experts": 64,
"top_k": 2,
"capacity_factor": 1.25,
"drop_tokens": true,
"use_tutel": true,
"load_balance_loss_coeff": 0.01
},
"zero_optimization": {"stage": 1}
}
七、综合优化实践
显存优化组合
text
1. 混合精度(BF16/FP8) → 参数/激活减半
2. 激活重计算 → 激活显存降 √n
3. ZeRO-3 / FSDP → 参数/梯度/优化器分片
4. CPU/NVMe Offload → 进一步卸载
5. 梯度累积 → 小 batch 等效大 batch
吞吐优化组合
text
1. 大 batch + 梯度累积
2. Sequence Packing
3. 通信与计算重叠
4. FlashAttention
5. 动态批处理
6. MoE 负载均衡
典型配置对照
| 模型规模 | 精度 | 并行 | 优化 |
|---|---|---|---|
| 1B | BF16 | DDP | AMP + 梯度累积 |
| 10B | BF16 | FSDP | + 激活重计算 |
| 100B | BF16/FP8 | TP+PP+DP | + ZeRO-1 + Selective AC |
| 500B | FP8 | TP+PP+DP+EP | + Offload + 异步 CKPT |
| 1T+ | FP8 | 全维度 | + MoE 调度 + GDS |
调优顺序
-
先跑通:BF16 + DDP/FSDP。
-
解决显存:激活重计算 → ZeRO → Offload。
-
提升吞吐:梯度累积 → Sequence Packing → 通信重叠。
-
极致性能:FP8 → MoE 调度 → 动态批处理。
-
稳定性:Checkpoint 优化 → 容错恢复。
八、常见问题与对策
| 问题 | 原因 | 对策 |
|---|---|---|
| Loss 溢出 | FP16 动态范围小 | 用 BF16 或调 Loss Scale |
| 训练不收敛 | 精度不足/学习率 | BF16 + 调 LR + warmup |
| 显存 OOM | 激活/参数太大 | 重计算 + ZeRO + 梯度累积 |
| 吞吐低 | 通信/气泡/小 batch | 重叠 + 大 batch + Packing |
| MoE 专家退化 | 路由崩塌 | 负载均衡损失 + 容量因子 |
| Checkpoint 慢 | 同步写 + 数据大 | 异步 + 分片 + 压缩 |
| 恢复失败 | 格式/版本不兼容 | DCP + 定期验证 |
总结
训练优化的本质是在显存、计算、通信、精度之间做多维权衡:
-
混合精度:BF16 通用,FP8 极致。
-
梯度累积:小显存等效大 batch。
-
激活重计算:计算换显存。
-
Checkpointing:异步、分片、分层。
-
动态批处理:最大化利用率。
-
MoE 调度:负载均衡 + 通信优化。
实际训练中,这些技术组合使用 :FP8 + ZeRO-3 + Selective AC + 异步 CKPT + Sequence Packing + MoE 负载均衡,才能训练万亿模型。关键是根据模型结构、硬件配置、任务目标选择合适组合,并用 Profiler 持续定位瓶颈。
推理优化是在延迟、吞吐、成本之间寻找最优平衡的工程体系。核心思路:减少计算量、减少内存访问、提高硬件利用率。下面按推理引擎、KV Cache、算子融合、量化、压缩、批处理、服务化展开。
一、推理引擎对比
引擎 定位 核心优势 适用场景
vLLM LLM 推理引擎 PagedAttention、Continuous Batching、高吞吐 LLM 在线服务
TensorRT / TensorRT-LLM NVIDIA 推理优化 极致性能、算子融合、FP8/INT4 NVIDIA GPU 生产部署
ONNX Runtime 跨平台推理 跨框架、跨硬件、图优化 通用模型部署
Triton Inference Server 推理服务框架 多框架、多模型、动态批处理 生产级服务化
vLLM
核心创新:PagedAttention,将 KV Cache 分页管理,消除内存碎片。
Continuous Batching:请求完成即退出,新请求立即加入。
优势:吞吐比 HuggingFace 高 10~24 倍,支持张量并行、量化。
架构:
text
请求 → Scheduler → Block Manager(PagedAttention)
→ Worker(GPU 计算)→ Output
TensorRT-LLM
核心:针对 LLM 的 TensorRT 扩展,支持 In-flight Batching、FP8、INT4/INT8 Weight-Only。
流程:模型转换 → 编译引擎 → 部署。
优势:NVIDIA 官方,性能天花板,支持多 GPU/多节点。
劣势:编译时间长,灵活性低。
ONNX Runtime
核心:ONNX 模型跨平台执行,图优化、算子融合、多后端(CPU/GPU/TPU)。
优势:跨框架(PyTorch/TF)、跨硬件、生态好。
适用:非 LLM 模型、边缘部署、多硬件兼容。
Triton Inference Server
核心:多框架(TensorRT、ONNX、PyTorch、vLLM)、多模型、动态批处理。
特性:模型仓库、版本管理、并发执行、指标监控。
适用:生产级推理服务,统一管理多种模型。
选型建议
场景 推荐
LLM 高吞吐服务 vLLM
NVIDIA 极致性能 TensorRT-LLM
跨平台通用模型 ONNX Runtime
多模型统一服务 Triton + 各后端
边缘设备 ONNX Runtime / TensorRT
二、KV Cache 优化
为什么需要 KV Cache
自回归生成中,每步都要计算所有历史 token 的 Key/Value。缓存后,每步只计算新 token,避免重复计算。
text
无 Cache:O(n²) 计算 → 有 Cache:O(n) 计算
KV Cache 显存占用
text
KV Cache = 2 × layers × heads × head_dim × seq_len × batch × dtype_bytes
例:Llama-7B, 32层, 32头, 128维度, 4096序列, FP16
= 2 × 32 × 32 × 128 × 4096 × 2 = 4GB(单序列)
PagedAttention(vLLM)
原理:将 KV Cache 分成固定大小的 Block(如 16 token),非连续存储,按需分配。
优势:
消除内存碎片,利用率从 20~40% 提升到 95%+。
支持 Copy-on-Write,共享前缀(如 System Prompt)。
支持 Beam Search 高效内存共享。
text
传统:连续内存,预分配最大长度 → 浪费
Paged:Block 化管理,按需分配 → 高效
KV Cache 优化手段
手段 说明 收益
PagedAttention 分页管理 消除碎片
Prefix Caching 共享前缀 KV 减少重复计算
KV Cache 量化 INT8/FP8 存储 显存减半
KV Cache 压缩 淘汰不重要 token 减少显存
Multi-Query Attention 多查询头共享 KV KV 减少 N 倍
Grouped-Query Attention 分组共享 KV 平衡质量与显存
Sliding Window 只保留最近窗口 固定显存
StreamingLLM 保留 Attention Sink 支持无限长度
GQA / MQA
text
MHA: Q头=K头=V头=32 → KV Cache 最大
GQA: Q头=32, K/V头=8 → KV Cache 减少 4 倍
MQA: Q头=32, K/V头=1 → KV Cache 减少 32 倍
三、算子融合
原理
将多个小算子合并为一个大算子,减少 Kernel Launch 开销和内存访问。
常见融合
融合模式 说明
Conv + BN + ReLU 推理时 BN 折叠进 Conv
Linear + Bias + Activation 合并为单 Kernel
LayerNorm + Residual 减少内存往返
Attention 融合 QKV 投影 + Attention + 输出投影
FlashAttention 融合 Attention 全流程,IO 感知
Softmax + Dropout 推理无 Dropout,融合 Softmax
FlashAttention
核心:分块计算 Attention,避免存储完整 Attention 矩阵。
IO 感知:减少 HBM 读写,从 O(n²) 降到 O(n²/M)。
版本:FlashAttention-1/2/3,支持 FP16/BF16/FP8。
收益:速度提升 2~4 倍,显存降低 5~20 倍。
python
PyTorch 使用
from flash_attn import flash_attn_func
output = flash_attn_func(q, k, v, causal=True)
TensorRT 融合
Layer Fusion:垂直/水平融合。
Kernel Auto-Tuning:自动选择最优 Kernel。
精度校准:INT8/FP8 校准。
融合收益
text
未融合:Kernel1 → 写HBM → Kernel2 → 读HBM → Kernel3
融合后:Kernel1+2+3 → 一次读写
→ 减少内存带宽压力,提升吞吐
四、模型量化
量化类型
类型 说明 精度 压缩比
FP32 基准 最高 1x
FP16/BF16 半精度 高 2x
FP8 Hopper 原生 中高 4x
INT8 8位整数 中 4x
INT4 4位整数 中低 8x
GPTQ/AWQ 训练后量化 中 4~8x
GGUF llama.cpp 格式 可调 2~8x
量化方法
- PTQ(训练后量化)
原理:用校准数据统计激活分布,确定缩放因子。
优点:无需重新训练。
缺点:精度损失,需校准。
- QAT(量化感知训练)
原理:训练时模拟量化误差。
优点:精度高。
缺点:需重新训练。
- Weight-Only 量化
原理:只量化权重,激活保持 FP16。
代表:GPTQ、AWQ、INT4 Weight-Only。
优点:显存大幅降低,精度损失小。
适用:LLM 推理。
GPTQ vs AWQ
维度 GPTQ AWQ
原理 逐层量化,最小化重建误差 激活感知,保护重要权重
精度 好 更好
速度 快 快
支持 广泛 vLLM/TensorRT-LLM
推荐 通用 LLM 首选
INT8 量化
python
PyTorch 动态量化
import torch.quantization as quant
model_int8 = quant.quantize_dynamic(model, {nn.Linear}, dtype=torch.qint8)
TensorRT INT8 校准
from pytorch_quantization import calib
用校准数据统计分布,生成校准表
FP8 量化
python
Transformer Engine FP8
import transformer_engine.pytorch as te
fp8_recipe = recipe.DelayedScaling(
margin=0, interval=1, fp8_format=recipe.Format.E4M3)
with te.fp8_autocast(enabled=True, fp8_recipe=fp8_recipe):
output = model(input)
量化选择建议
场景 推荐
显存充足,追求精度 FP16/BF16
显存受限,精度要求高 INT8 / FP8
显存极紧,可接受损失 INT4 (GPTQ/AWQ)
边缘设备 INT8 / INT4
NVIDIA Hopper FP8
五、剪枝
类型
类型 说明 粒度
非结构化剪枝 随机置零权重 单权重
结构化剪枝 剪整个通道/头/层 通道/头
半结构化 N:M 稀疏(如 2:4) 块
2:4 稀疏
原理:每 4 个权重保留 2 个,硬件加速。
支持:Ampere+ Tensor Core。
收益:理论 2 倍加速,精度损失小。
python
PyTorch 2:4 稀疏
from torch.sparse import to_sparse_semi_structured
model = to_sparse_semi_structured(model)
剪枝流程
text
-
训练完整模型
-
评估权重重要性(L1/L2/梯度)
-
剪枝(置零或删除)
-
微调恢复精度
-
重复 2-4 直到目标稀疏度
剪枝收益
显存:减少参数存储。
计算:稀疏计算加速(需硬件支持)。
精度:结构化剪枝损失较大,需微调。
六、知识蒸馏
原理
用大模型(Teacher)指导小模型(Student)训练,迁移知识。
text
Loss = α × CE(student, label) + β × KL(student, teacher)
蒸馏类型
类型 说明
Logit 蒸馏 Student 拟合 Teacher 输出分布
Feature 蒸馏 拟合中间层特征
Attention 蒸馏 拟合 Attention 矩阵
Chain-of-Thought 蒸馏 迁移推理过程
LLM 蒸馏实践
白盒蒸馏:访问 Teacher 的 logits。
黑盒蒸馏:只用 Teacher 的输出文本(如 GPT-4 → 小模型)。
代表:DistilBERT、TinyLlama、Alpaca、Vicuna。
蒸馏 vs 量化 vs 剪枝
方法 原理 精度 压缩比 需训练
量化 降低精度 中 4~8x 否/少量
剪枝 删除权重 中 2~10x 微调
蒸馏 大教小 高 可定制 是
七、动态批处理
类型
类型 说明 引擎
Static Batching 固定 batch,等所有请求 传统
Dynamic Batching 时间窗口内聚合请求 Triton
Continuous Batching 迭代级动态加入/退出 vLLM、TensorRT-LLM
In-flight Batching 同 Continuous TensorRT-LLM
Continuous Batching 原理
text
传统:Batch = A, B, C,等 C 完成才释放 → A/B 等待浪费
Continuous:A 完成即退出,D 立即加入 → GPU 持续满载
收益
GPU 利用率从 30% → 80%+。
吞吐提升 2~10 倍。
延迟略增(等待聚合),但可调。
Triton 动态批处理配置
protobuf
dynamic_batching {
preferred_batch_size: 4, 8, 16
max_queue_delay_microseconds: 1000
}
vLLM Continuous Batching
python
from vllm import LLM, SamplingParams
llm = LLM(model="meta-llama/Llama-2-7b-hf",
max_num_seqs=256, # 最大并发序列
max_num_batched_tokens=8192)
八、服务化部署
架构分层
text
客户端 → Load Balancer / Ingress
↓
API Gateway(鉴权、限流)
↓
推理服务(Triton / vLLM / TGI)
↓
模型后端(TensorRT / ONNX / PyTorch)
↓
GPU / CPU 资源
关键组件
组件 作用
KServe K8s 原生推理服务,自动扩缩容
Triton 多模型服务,动态批处理
vLLM LLM 高吞吐服务
TGI HuggingFace 推理服务
Ray Serve 分布式推理,组合模型
BentoML 模型打包与服务
部署模式
模式 说明 适用
单模型单服务 一模型一 Deployment 简单
多模型单服务 Triton 多模型仓库 资源复用
模型组合 Pipeline 多模型串联 复杂任务
A/B 测试 多版本流量切分 灰度发布
边缘部署 ONNX/TensorRT + 轻量服务 低延迟
自动扩缩容
yaml
apiVersion: serving.kserve.io/v1beta1
kind: InferenceService
metadata:
name: llm
spec:
predictor:
minReplicas: 1
maxReplicas: 10
scaleTarget: 70
scaleMetric: gpu
model:
modelFormat: {name: vllm}
storageUri: s3://models/llama-2-7b
监控指标
延迟:P50/P90/P99、TTFT(首 token 延迟)、TPOT(每 token 延迟)。
吞吐:QPS、Tokens/s。
资源:GPU 利用率、显存、KV Cache 占用。
队列:等待请求数、队列延迟。
优化实践
预热:模型加载、CUDA Graph 捕获。
CUDA Graph:消除 Kernel Launch 开销。
Prefix Caching:共享 System Prompt。
Speculative Decoding:小模型草稿 + 大模型验证。
多副本:负载均衡 + 会话保持。
九、综合优化实践
延迟优化
text
-
算子融合(FlashAttention、TensorRT)
-
量化(INT8/FP8)
-
CUDA Graph
-
Speculative Decoding
-
Prefix Caching
吞吐优化
text
-
Continuous Batching
-
PagedAttention
-
张量并行
-
量化(减少显存,增大 batch)
-
多副本 + 负载均衡
显存优化
text
-
KV Cache 量化/压缩
-
GQA/MQA
-
模型量化(INT4/INT8)
-
PagedAttention
-
分层加载(Offload)
典型配置对照
场景 引擎 量化 批处理 优化
LLM 在线服务 vLLM AWQ/FP8 Continuous PagedAttention + Prefix Cache
NVIDIA 极致 TensorRT-LLM FP8/INT4 In-flight CUDA Graph + 融合
多模型服务 Triton INT8/FP16 Dynamic 多后端 + 版本管理
边缘部署 ONNX Runtime INT8 Static 图优化 + 量化
通用 CV TensorRT FP16/INT8 Dynamic 层融合 + 校准
调优顺序
先跑通:FP16 + 原生推理。
降显存:量化(INT8/INT4)→ KV Cache 优化。
提吞吐:Continuous Batching → PagedAttention → 张量并行。
降延迟:算子融合 → CUDA Graph → Speculative Decoding。
服务化:Triton/KServe + 自动扩缩容 + 监控。
十、常见问题与对策
问题 原因 对策
显存 OOM KV Cache 太大 PagedAttention + 量化 + GQA
吞吐低 静态批处理 Continuous Batching
延迟高 Kernel Launch 多 CUDA Graph + 算子融合
首 token 慢 Prefill 计算大 Prefix Caching + 分块 Prefill
精度下降 量化过度 混合精度 + AWQ + 校准
负载不均 请求长度差异大 Continuous Batching + 调度
扩展难 单卡显存限制 张量并行 + 流水线并行
总结
推理优化的本质是在延迟、吞吐、显存、精度之间做多维权衡:
vLLM:PagedAttention + Continuous Batching,LLM 高吞吐首选。
TensorRT:算子融合 + FP8/INT4,NVIDIA 极致性能。
ONNX Runtime:跨平台通用部署。
Triton:多模型统一服务,动态批处理。
KV Cache:PagedAttention、量化、GQA。
算子融合:FlashAttention、层融合。
量化:INT8/INT4/FP8,AWQ/GPTQ。
剪枝/蒸馏:模型压缩,精度与体积权衡。
动态批处理:Continuous Batching 提升利用率。
服务化:KServe/Triton + 扩缩容 + 监控。
实际部署中,这些技术组合使用:vLLM + AWQ + PagedAttention + Continuous Batching + Prefix Caching + 张量并行,才能在有限 GPU 上服务大模型。关键是根据模型规模、硬件配置、SLA 要求选择合适组合,并用监控持续调优。
编译与图优化是连接模型与硬件的桥梁 :训练框架产出计算图,编译器将其优化并映射到不同硬件。核心目标:减少计算量、减少内存访问、提高硬件利用率。下面按整体格局、主流编译器、核心优化技术、实践展开。
一、为什么需要编译器
框架执行的痛点
text
Eager 模式:Python 逐算子执行
├── Kernel Launch 开销大(每个算子一次)
├── 无法跨算子优化
├── 内存分配频繁
└── 难以适配多种硬件
编译器的价值
text
计算图 → 图优化 → 算子融合 → 内存规划 → 代码生成 → 硬件执行
| 收益 | 说明 |
|---|---|
| 算子融合 | 减少 Kernel Launch 和内存往返 |
| 内存复用 | 静态规划,减少分配和碎片 |
| 常量折叠 | 编译期计算不变表达式 |
| 布局优化 | 选择最优数据布局 |
| 硬件适配 | 一套图,多后端代码生成 |
| 自动调优 | 搜索最优算子实现 |
二、主流编译器对比
| 编译器 | 定位 | 输入 | 后端 | 核心优势 |
|---|---|---|---|---|
| TorchDynamo | PyTorch 图捕获 | PyTorch 代码 | 多后端 | 动态图捕获,无缝集成 |
| TorchInductor | PyTorch 默认后端 | FX Graph | Triton/C++ | 生成 Triton Kernel |
| TVM | 端到端深度学习编译 | Relay/Te | LLVM/CUDA/... | 自动调优,跨硬件 |
| MLIR | 编译器基础设施 | 多种 Dialect | LLVM/GPU | 多层 IR,可复用 |
| XLA | TensorFlow/JAX 编译 | HLO | TPU/GPU/CPU | TPU 原生,JIT |
层次关系
text
应用层: PyTorch / TensorFlow / JAX
↓
图捕获: TorchDynamo / tf.function / jit
↓
中间表示: FX Graph / HLO / Relay / StableHLO
↓
编译器: TorchInductor / XLA / TVM / MLIR
↓
代码生成: Triton / LLVM / CUDA / TPU
↓
硬件: GPU / TPU / CPU / NPU
三、TorchDynamo + TorchInductor
TorchDynamo
-
定位:PyTorch 2.0 的图捕获前端。
-
原理 :字节码分析,在 Python 运行时捕获计算图,不修改用户代码。
-
优势:
-
支持动态控制流(if/for)。
-
Graph Break:遇到不支持的操作,回退 Eager,保证正确性。
-
无缝集成,
torch.compile一行启用。
-
python
import torch
model = MyModel()
compiled_model = torch.compile(model, mode="max-autotune")
output = compiled_model(input)
TorchInductor
-
定位:PyTorch 默认编译后端。
-
输入:TorchDynamo 捕获的 FX Graph。
-
输出 :Triton Kernel (GPU)或 C++/OpenMP(CPU)。
-
核心优化:
-
算子融合:自动融合 Element-wise、Reduction。
-
内存规划:静态分配,复用缓冲区。
-
Triton 代码生成:生成高效 GPU Kernel。
-
自动调优:搜索最优配置。
-
编译流程
text
Python 代码
→ TorchDynamo(字节码分析)→ FX Graph
→ AOTAutograd(前向+反向图)→ 分解算子
→ TorchInductor(融合+调度)→ Triton/C++ 代码
→ 编译执行
模式
| 模式 | 说明 |
|---|---|
default |
平衡编译时间和性能 |
reduce-overhead |
减少开销,适合小模型 |
max-autotune |
最大性能,编译时间长 |
优势与局限
| 优势 | 局限 |
|---|---|
| 一行启用,无侵入 | Graph Break 影响性能 |
| 动态图支持好 | 复杂控制流回退 Eager |
| Triton 融合高效 | 部分算子不支持 |
| 生态活跃 | 编译时间较长 |
四、TVM
定位
端到端深度学习编译器,从模型到多硬件代码生成。
架构
text
模型 (PyTorch/ONNX/TF)
→ Relay(高层 IR,图优化)
→ TIR(底层 IR,循环优化)
→ AutoTVM / Ansor(自动调优)
→ 代码生成 (LLVM/CUDA/OpenCL/...)
→ 硬件
核心组件
| 组件 | 作用 |
|---|---|
| Relay | 高层计算图 IR,算子融合、常量折叠 |
| TIR | 底层循环 IR,循环变换、向量化 |
| AutoTVM | 基于模板的自动调优 |
| Ansor | 无模板自动调优,生成搜索空间 |
| Relax | 新一代 IR,支持动态形状 |
| BYOC | Bring Your Own Codegen,对接外部后端 |
自动调优
python
# Ansor 自动调优示例
import tvm
from tvm import auto_scheduler
@auto_scheduler.register_workload
def matmul(M, N, K):
A = te.placeholder((M, K), name="A")
B = te.placeholder((K, N), name="B")
k = te.reduce_axis((0, K), name="k")
C = te.compute((M, N), lambda i, j: te.sum(A[i, k] * B[k, j], axis=k))
return [A, B, C]
task = auto_scheduler.create_task(matmul, (1024, 1024, 1024))
tuner = auto_scheduler.TaskScheduler([task])
tuner.tune(trials=1000)
优势与局限
| 优势 | 局限 |
|---|---|
| 跨硬件支持广 | 学习曲线陡 |
| 自动调优强 | 编译时间长 |
| 可定制 | 生态不如 PyTorch |
| 支持边缘设备 | 动态形状支持弱 |
五、MLIR
定位
编译器基础设施 ,不是完整编译器,而是构建编译器的框架。
核心概念
-
Dialect:多层 IR,每层有自己的操作和类型。
-
Operation:IR 的基本单位。
-
Pass:IR 变换。
-
Progressive Lowering:从高层逐步降低到低层。
Dialect 层次
text
高层: StableHLO / TOSA / MHLO(机器学习)
↓
中层: Linalg(线性代数) / Tensor(张量)
↓
低层: Affine / SCF(循环) / Vector(向量)
↓
硬件: GPU / LLVM / SPIR-V
典型流程
text
PyTorch → Torch-MLIR → StableHLO → Linalg → Affine → LLVM IR → 机器码
优势
-
可复用:各层 Dialect 可组合。
-
可扩展:自定义 Dialect。
-
多前端多后端:统一基础设施。
-
被广泛采用:TensorFlow、JAX、PyTorch、TVM 都在用。
代表项目
| 项目 | 说明 |
|---|---|
| Torch-MLIR | PyTorch → MLIR |
| StableHLO | 机器学习高层 IR 标准 |
| IREE | 基于 MLIR 的端到端编译器 |
| Triton | 基于 MLIR 的 GPU DSL |
六、XLA
定位
TensorFlow/JAX 的编译器,TPU 原生支持。
架构
text
TensorFlow/JAX
→ HLO(High-Level Optimizer IR)
→ XLA 优化(融合、内存规划)
→ 后端(TPU/GPU/CPU)
核心优化
| 优化 | 说明 |
|---|---|
| 算子融合 | 融合 Element-wise、Reduction |
| 内存规划 | 静态分配,复用缓冲区 |
| 常量折叠 | 编译期计算 |
| 布局优化 | 选择最优数据布局 |
| 并行化 | 自动并行到多设备 |
JAX + XLA
python
import jax
import jax.numpy as jnp
@jax.jit # JIT 编译,触发 XLA
def matmul(a, b):
return jnp.dot(a, b)
# 自动微分 + XLA
grad_fn = jax.grad(loss_fn)
优势与局限
| 优势 | 局限 |
|---|---|
| TPU 原生 | 主要绑定 TF/JAX |
| JIT 编译 | 动态形状支持弱 |
| 内存规划强 | 编译时间长 |
| 多设备并行 | 调试困难 |
七、核心优化技术
1. 计算图优化
| 优化 | 说明 | 示例 |
|---|---|---|
| 常量折叠 | 编译期计算常量表达式 | 2+3 → 5 |
| 死代码消除 | 删除无用节点 | 未使用的输出 |
| 公共子表达式消除 | 复用重复计算 | 相同子图只算一次 |
| 代数化简 | 数学等价变换 | x*1 → x |
| 布局转换 | 选择最优布局 | NCHW ↔ NHWC |
| 算子替换 | 用高效算子替代 | Conv+BN → Conv |
2. 算子融合
融合类型:
text
垂直融合: producer → consumer 合并
例:Conv → BN → ReLU 合并为一个 Kernel
水平融合: 并行算子合并
例:多个独立 Element-wise 合并
跨层融合: 跨多个层融合
例:FlashAttention 融合整个 Attention
融合收益:
text
未融合:K1 → 写HBM → K2 → 读HBM → K3 → 写HBM
融合后:K1+2+3 → 一次读写
→ 减少内存带宽压力,提升吞吐
融合模式:
| 模式 | 说明 |
|---|---|
| Element-wise 融合 | 逐元素操作合并 |
| Reduction 融合 | 归约操作合并 |
| Producer-Consumer | 生产者消费者合并 |
| Attention 融合 | FlashAttention |
| GEMM 融合 | GEMM + Bias + Activation |
3. 内存复用
内存规划:
text
静态分配:编译期确定所有缓冲区大小和生命周期
内存池:预分配大块,按需切分
缓冲区复用:生命周期不重叠的缓冲区共享内存
原地操作:输出复用输入内存
技术:
| 技术 | 说明 |
|---|---|
| 生命周期分析 | 确定缓冲区存活区间 |
| 内存池 | 预分配,减少 malloc |
| 缓冲区复用 | 不重叠的缓冲区共享 |
| 原地更新 | in-place 操作 |
| 内存压缩 | 量化、稀疏 |
| Offload | 卸载到 CPU/NVMe |
示例:
text
Layer1 输出 → 缓冲区A(生命周期:Layer1~Layer2)
Layer2 输出 → 缓冲区B(生命周期:Layer2~Layer3)
Layer3 输出 → 可复用缓冲区A(Layer1 已结束)
4. 循环优化(TIR/Affine)
| 优化 | 说明 |
|---|---|
| 循环展开 | 减少循环开销 |
| 循环分块 | 提高缓存命中 |
| 循环重排 | 优化访存模式 |
| 向量化 | SIMD 指令 |
| 并行化 | 多线程/多核 |
| 软件流水 | 重叠计算与访存 |
八、实践与选型
选型建议
| 场景 | 推荐 |
|---|---|
| PyTorch 训练加速 | TorchDynamo + TorchInductor |
| PyTorch 推理 | torch.compile + TensorRT |
| TensorFlow/JAX | XLA |
| 跨硬件部署 | TVM |
| 自定义编译器 | MLIR |
| 边缘设备 | TVM / ONNX Runtime |
| TPU | XLA |
PyTorch 编译实践
python
import torch
# 基础编译
model = torch.compile(model)
# 最大性能
model = torch.compile(model, mode="max-autotune")
# 指定后端
model = torch.compile(model, backend="inductor")
# 动态形状
model = torch.compile(model, dynamic=True)
# 调试
torch._dynamo.config.verbose = True
torch._dynamo.config.suppress_errors = True
性能调优
-
减少 Graph Break:避免不支持的操作。
-
增大编译区域:减少回退 Eager。
-
自动调优 :
max-autotune模式。 -
CUDA Graph:消除 Launch 开销。
-
混合精度:FP16/BF16/FP8。
常见问题
| 问题 | 原因 | 对策 |
|---|---|---|
| Graph Break 多 | 动态控制流 | 改写代码,减少 break |
| 编译时间长 | 自动调优 | 降低调优级别 |
| 精度下降 | 融合/量化 | 检查数值稳定性 |
| 性能不升 | 未命中融合 | Profiler 分析 |
| 动态形状慢 | 重新编译 | dynamic=True 或标记 |
九、编译器对比总结
| 维度 | TorchDynamo | TVM | MLIR | XLA |
|---|---|---|---|---|
| 定位 | PyTorch 前端 | 端到端编译器 | 基础设施 | TF/JAX 编译器 |
| 输入 | PyTorch | Relay/ONNX | 多 Dialect | HLO |
| 后端 | 多后端 | 多硬件 | 多硬件 | TPU/GPU/CPU |
| 自动调优 | Inductor | AutoTVM/Ansor | 需自建 | 有限 |
| 动态形状 | 好 | 弱 | 好 | 弱 |
| 生态 | PyTorch | 独立 | 广泛 | TF/JAX |
| 学习曲线 | 低 | 高 | 高 | 中 |
| 生产成熟度 | 高 | 中 | 中 | 高 |
十、总结
编译与图优化的本质是将计算图高效映射到硬件:
-
TorchDynamo :PyTorch 图捕获,无缝集成,
torch.compile一行启用。 -
TorchInductor:生成 Triton Kernel,算子融合 + 内存规划。
-
TVM:端到端编译器,自动调优,跨硬件。
-
MLIR:编译器基础设施,多层 IR,可复用可扩展。
-
XLA:TF/JAX 编译器,TPU 原生,JIT 编译。
核心优化技术:
-
计算图优化:常量折叠、死代码消除、代数化简。
-
算子融合:垂直/水平/跨层融合,减少内存往返。
-
内存复用:生命周期分析、内存池、缓冲区复用。
-
循环优化:分块、向量化、并行化。
实际工作中,PyTorch 训练用 TorchDynamo + Inductor,推理用 TensorRT/vLLM,跨硬件用 TVM,TPU 用 XLA 。理解这些编译器的IR 层次、优化 Pass、代码生成流程,才能根据场景选择和调优。