Skip to content

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) / repeats

profile 进一步问:CPU 是否在等待数据?有多少 kernel launch?算子之间是否频繁搬运中间结果?峰值内存是否来自 attention、logits、激活还是 optimizer 状态?如果 attention 只占总时间 20%,即使把它加速到零成本,总速度上限也只有 1/0.8=1.25 倍。这是 Amdahl 定律给实验选题的直接约束。

FLOPs 与字节必须一起看 ​

算术强度是计算量除以搬运字节数。一个算子的执行时间下界近似受

t≳max(FLOPs算力,传输字节带宽)

限制。矩阵乘法通过分块复用数据,提高算术强度;逐元素算子可能更受带宽和 launch 开销限制。混合精度能减少存储并使用专用矩阵硬件,但稳定归一化与归约通常需要更高精度累加。性能优化后必须再检查数值误差,而不只检查是否能运行。

标准注意力显式产生 T×T 的 scores 和 probabilities,并在 HBM 写出、读回。对长上下文,这些中间量很昂贵。FlashAttention 的关键是改变 IO 计划,用片上块与在线 softmax 避免把完整注意力矩阵写到 HBM;它仍计算精确稠密注意力,通常仍有 O(T2dh) 算术复杂度。

在线 softmax 的推导 ​

对一个 query,假设已处理一些 key,保存最大分数 m、分母 l=∑jesj−m、未归一化输出 a=∑jesj−mvj。读入新块分数 s 后:

m′=max(m,maxs),α=em−m′,p=es−m′,l′=αl+∑p,a′=αa+pVblock.

结束时 o=a/l。乘 α 是为了把旧累积量改写到新最大值下;省掉它会让分块结果错误。初值为 m=−∞,l=0,a=0,首个有效块有 α=0。因果计算应跳过全被屏蔽的块,避免 -inf - (-inf)。本地 online_attention 对每个 query 只遍历允许的 key,并与稠密实现逐元素比较。

从公式到 Triton 程序 ​

真实内核让一个 program 负责某个 batch、head、query 行块,在寄存器/片上缓存保留 Q 块、每行 m,l 和输出累积量,循环加载 K/V 块。以下是算法骨架,不是可以直接运行的 kernel:

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 代替自然指数,分数先乘 log2⁡e,保存的 log-sum-exp 也要保持同一底数约定。块越大不一定越快:寄存器/共享内存压力可能降低 occupancy。FlashAttention-2 在工作划分上进一步降低非矩阵运算和 warp 间通信,并让 query 块提供更多并行机会;不能把“分块”三个字等同于这些调度改进。

反向传播可保存每行 log-sum-exp 和输出,重算局部 probabilities,再使用上一章的梯度公式,减少激活存储。多个 query 块会同时贡献 dK,dV,并行实现必须解决归约/写冲突。前向通过测试并不代表训练可用;需验证 backward、因果 mask、尾块、非连续输入和多种长度。

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。