AI Infra 全栈技术栈 二

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 不一致。可用 SyncBNGroupNorm

  • 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

调优顺序

  1. 先跑通:BF16 + DDP/FSDP。

  2. 解决显存:激活重计算 → ZeRO → Offload。

  3. 提升吞吐:梯度累积 → Sequence Packing → 通信重叠。

  4. 极致性能:FP8 → MoE 调度 → 动态批处理。

  5. 稳定性: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

量化方法

  1. PTQ(训练后量化)

原理:用校准数据统计激活分布,确定缩放因子。

优点:无需重新训练。

缺点:精度损失,需校准。

  1. QAT(量化感知训练)

原理:训练时模拟量化误差。

优点:精度高。

缺点:需重新训练。

  1. 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

  1. 训练完整模型

  2. 评估权重重要性(L1/L2/梯度)

  3. 剪枝(置零或删除)

  4. 微调恢复精度

  5. 重复 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

  1. 算子融合(FlashAttention、TensorRT)

  2. 量化(INT8/FP8)

  3. CUDA Graph

  4. Speculative Decoding

  5. Prefix Caching

吞吐优化

text

  1. Continuous Batching

  2. PagedAttention

  3. 张量并行

  4. 量化(减少显存,增大 batch)

  5. 多副本 + 负载均衡

显存优化

text

  1. KV Cache 量化/压缩

  2. GQA/MQA

  3. 模型量化(INT4/INT8)

  4. PagedAttention

  5. 分层加载(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

性能调优

  1. 减少 Graph Break:避免不支持的操作。

  2. 增大编译区域:减少回退 Eager。

  3. 自动调优max-autotune 模式。

  4. CUDA Graph:消除 Launch 开销。

  5. 混合精度: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、代码生成流程,才能根据场景选择和调优。

相关推荐
夜雪一千1 小时前
如何使用Python做舆情分析
人工智能·python·数据分析
azhou的代码园1 小时前
景区游船预约服务系统
人工智能·spring boot·后端
进击的横打1 小时前
【人工智能】人与AI协作的四象限
人工智能
yyk333241 小时前
计算机视觉OpenCV中的FisherFace人脸识别
人工智能·opencv·计算机视觉
澳鹏Appen2 小时前
AppenTalk | 当AI推理不再只是模型的事:Dan Roth谈智能体的真正边界
人工智能·大语言模型·智能体·大模型推理
Mr数据杨2 小时前
二手车价格预测实战解析 从 Kaggle 回归赛题到可落地估价方案
人工智能·数据分析·kaggle竞赛
Zzj_tju2 小时前
Embodied Agent 小环境:用状态日志解释成功、绕路与超时
人工智能·深度学习·机器学习·自然语言处理
微财经观圈2 小时前
AI 3D生成角色怎样套用预设动作?自动绑骨与动作测试步骤
人工智能·3d
是翎2 小时前
图解大语言模型部署
人工智能·驱动开发·深度学习·开源协议·imagen