外观
A2 系统:先解释瓶颈,再优化计算
同一个数学表达式可以对应非常不同的执行计划。A2 要求构建 profiling/benchmark 工具、实现 Triton FlashAttention-2 内核、数据并行训练和优化器状态分片。本章将这四部分连成一条路线。CPU 实验只验证在线注意力的数学等价性;没有 GPU 性能数据,也没有把 Python 分块循环称为高性能 FlashAttention 内核。
测量必须先定义边界
GPU 调用通常异步返回,Python 计时可能只测到提交任务的时间。先做 warmup,排除编译与首次分配,再在计时边界同步,或使用正确放置的 CUDA events。记录前向、前向加反向、完整 optimizer step 三种范围,不能拿一种范围的加速比宣传另一种。固定形状、dtype、mask、设备和批大小,重复测量并报告中位数及波动。
python
# GPU 环境的测量结构示意;这里没有运行 CUDA
for _ in range(warmup):
run_step()
torch.cuda.synchronize()
start = time.perf_counter()
for _ in range(repeats):
run_step()
torch.cuda.synchronize()
seconds_per_step = (time.perf_counter() - start) / repeatsprofile 进一步问:CPU 是否在等待数据?有多少 kernel launch?算子之间是否频繁搬运中间结果?峰值内存是否来自 attention、logits、激活还是 optimizer 状态?如果 attention 只占总时间 20%,即使把它加速到零成本,总速度上限也只有
FLOPs 与字节必须一起看
算术强度是计算量除以搬运字节数。一个算子的执行时间下界近似受
限制。矩阵乘法通过分块复用数据,提高算术强度;逐元素算子可能更受带宽和 launch 开销限制。混合精度能减少存储并使用专用矩阵硬件,但稳定归一化与归约通常需要更高精度累加。性能优化后必须再检查数值误差,而不只检查是否能运行。
标准注意力显式产生
在线 softmax 的推导
对一个 query,假设已处理一些 key,保存最大分数
结束时 -inf - (-inf)。本地 online_attention 对每个 query 只遍历允许的 key,并与稠密实现逐元素比较。
从公式到 Triton 程序
真实内核让一个 program 负责某个 batch、head、query 行块,在寄存器/片上缓存保留
text
program(b, h, query_block):
Q = masked_load(Q_pointer, query_rows, head_columns)
m = -infinity; l = 0; acc = 0
for key_block intersecting allowed causal region:
K, V = masked_load(...)
S = dot(Q, transpose(K)) / sqrt(head_dim)
S = where(valid_rows & valid_keys & causal, S, -infinity)
update m, l, acc with online-softmax equations
masked_store(output, acc / l)落到 Triton 时要显式传 strides,不能假设任意输入都连续;tl.arange 的块边界与真实 T 的尾块要分开;tl.dot 与累加器 dtype 要明确。若用 exp2 代替自然指数,分数先乘
反向传播可保存每行 log-sum-exp 和输出,重算局部 probabilities,再使用上一章的梯度公式,减少激活存储。多个 query 块会同时贡献
DDP 与通信重叠
数据并行每个 rank 有完整模型,处理不同数据。各 rank 计算本地平均梯度,再 all-reduce 求和并除 world size;所有 rank 用相同参数与优化器状态更新。若本地有效 token 数不同,简单平均 rank 均值会偏置,应按 token 总数加权。
flowchart LR A[rank 0 数据] --> G0[本地梯度] B[rank 1 数据] --> G1[本地梯度] G0 --> R[all-reduce 并正确归一化] G1 --> R R --> U0[相同更新] R --> U1[相同更新]
查看流程图文本
flowchart LR A[rank 0 数据] --> G0[本地梯度] B[rank 1 数据] --> G1[本地梯度] G0 --> R[all-reduce 并正确归一化] G1 --> R R --> U0[相同更新] R --> U1[相同更新]
朴素实现等所有反向完成再通信;bucketed DDP 可在一组梯度就绪后异步通信,与后续层反向重叠。更新前必须等待所有相关通信完成,bucket 次序需在 rank 间一致,否则可能死锁。小 bucket 便于重叠却增加延迟次数;大 bucket 减少调用却可能拖到反向末尾。基准要检查 global batch 是否随 GPU 数改变,否则速度与训练目标同时变化。
优化器状态分片解决哪一块内存
Adam 的一二阶矩本身占两个参数规模;若 FP32,每个参数需要 8 bytes 的这两项状态。把参数按归属 rank 划分,归属者维护该片参数的 optimizer state 并更新,随后广播更新后的参数,可降低每 rank 的状态内存。它不自动消除完整参数副本、全部激活或所有梯度;不要把 optimizer sharding 直接等同于完全分片训练。
验收时先比较单卡与多 rank 在相同 global batch 上的一步更新,再测多步 loss 与峰值内存。多 rank 随机种子既要保证相同初始化,又要避免每个 rank 抽到完全相同的数据。故障排查从数据划分、归一化、同步时机开始,比盲目调整通信环境变量有效。
自测
FlashAttention 没有保存 T×T 矩阵,为什么还能精确求结果?
输出只需要 softmax 的分母与概率加权 value 之和;在线递推保留了这些充分统计量,并通过最大值重标定保持数值稳定。没有丢弃任何允许的 key,也没有对概率做低秩近似。浮点求和顺序变化仍会产生小误差。
两张 GPU 各用 batch 8,是否应该与单 GPU batch 8 比较训练轨迹?
要验证梯度等价,应与单 GPU global batch 16 对照,使用相同 16 个样本,控制归一化、随机层和精度。单卡 batch 8 是另一个优化设置;可以比较资源效率,但不能期待逐步参数相等。
来源:A2 2025 固定讲义、FlashAttention、FlashAttention-2、Triton fused attention 教程、ZeRO。下一章:Scaling laws。