共计 1756 个字符,预计需要花费 5 分钟才能阅读完成。
从 Amdahl 定律看 token 级算力分析的意义
在分布式计算中,Amdahl 定律告诉我们:系统加速比受限于串行部分的比例。对大模型推理而言,每个 token 的处理正是这样的串行单元。当处理序列长度 $L$ 时,总耗时 $T$ 可表示为:
$$T = L \cdot (t_{\text{compute}} + t_{\text{memory}})$$
其中 $t_{\text{compute}}$ 是计算耗时,$t_{\text{memory}}$ 是内存访问耗时。优化单个 token 的算力成本,能直接带来端到端的线性加速。
硬件架构的指令周期对比
CPU 流水线(以 Intel Ice Lake 为例)
- 标量运算需要 4 - 5 个时钟周期
- SIMD 指令(AVX-512)可并行处理 16 个 fp32 操作
- 分支预测错误导致 20-25 周期惩罚

GPU 流水线(NVIDIA Ampere 架构)
- Warp 内 32 线程同步执行
- 每个 SM 每周期调度 4 条指令
- Tensor Core 处理矩阵乘仅需 1 周期(fp16 累加到 fp32)
TPU 流水线(v4 架构)
- 脉动阵列实现确定性的计算延迟
- 向量单元与标量单元解耦
- 权重固定式数据流架构
PyTorch 性能测量实战
FLOPs 计数器实现
import torch
from torch.utils.flop_counter import FlopCounterMode
model = ... # 初始化你的模型
input = torch.randn(1, seq_len, hidden_dim).cuda()
with FlopCounterMode(model) as flop_counter:
output = model(input)
print(f"FLOPs per token: {flop_counter.get_total_flops()/seq_len:.0f}")
计算强度测量
定义计算强度 $I$ 为:
$$I = \frac{\text{FLOPs}}{\text{字节访问量}}$$
def arithmetic_intensity(A, B, C):
# A@B= C 的矩阵乘法
flops = 2 * A.size(0) * A.size(1) * B.size(1)
bytes_accessed = (A.numel() + B.numel() + C.numel()) * A.element_size()
return flops / bytes_accessed
DRAM 延迟测试
使用 Nsight Compute 进行跟踪:
nv-nsight-cu-cli --metrics dram__bytes.sum ./inference_script.py
关键优化策略
精度选择的影响
| 精度 | 算力需求(TOPS) | 内存占用(GB) |
|---|---|---|
| fp32 | 16 | 4.0 |
| fp16 | 32 | 2.0 |
| int8 | 64 | 1.0 |
测试环境:A100-80GB-SXM4,batch_size=8
KV 缓存优化
- 分页注意力:将 KV 缓存分解为固定大小的块
- 动态稀疏化:根据 attention score 剪枝低权重条目
算子融合示例
def fused_attention(Q, K, V):
scale = 1.0 / math.sqrt(Q.size(-1))
attn = torch.matmul(Q, K.transpose(-2, -1)) * scale
attn = torch.softmax(attn, dim=-1)
return torch.matmul(attn, V)
# 使用 torch.jit.script 编译成单一内核
开放性问题
-
稀疏化计算:当引入 $k\%$ 稀疏度时,理论算力需求变为:
$$C_{\text{sparse}} = C_{\text{dense}} \cdot (1 – k/100) + E_{\text{overhead}}$$
其中 $E_{\text{overhead}}$ 是稀疏格式的解码开销 -
存算一体架构:三星 HBM-PIM 实测显示,在内存内执行 GEMV 操作可减少 90% 的数据搬运,但受限于工艺精度(当前仅支持 int8)
性能测试建议
- 使用
torch.backends.cuda.sdp_kernel()启用 Flash Attention - 监控
nvidia-smi dmon -s u观察 SM 利用率 - 对于长序列,优先测试
memory_attention的 P99 延迟
所有测试数据均基于 PyTorch 2.1+ 和 CUDA 11.7 环境,完整复现脚本见GitHub 仓库
正文完
发表至: 未分类
近三天内
