SDPA backend 与 FlashAttention
Self-attention 的 O(n²) 内存和算力随序列长度平方涨——长序列要算 FLOPs/activation 预算,FlashAttention 和线性近似是工程选项不是免费 lunch。

1. 序列从 512 拉到 4096,OOM 在 attention
文档分类 ViT 改 patch 后 seq len 4096,forward 在 attention 层 OOM——不是 embedding 太大,是标准 self-attention 的 activation 随 涨。profile 显示 attention 占 forward CUDA time 60%+,这才值得换 kernel;若 conv backbone 主导,先别折腾 FlashAttention。OOM 栈指向 softmax(QK^T) 时,应优先算 seq len 与 batch 的平方律预算,而不是先减 hidden dim。产品要 long context 时,硬件选型阶段就要算 activation 上限。
2. 复杂度预算
标准 attention: FLOPs,中间 score matrix 占显存。seq len 加倍,time/activation 约 4×——固定 batch,只改 seq len,用 profiler 看 attention op 是否 ~4×。
长序列选项:FlashAttention-2(IO-aware 分块)、PyTorch SDPA、Swin window 、Linformer/Performer(有表达力代价)。选型前先写一行 、batch、head_dim 的 activation 估算,避免盲目上 H100 仍 OOM。
3. PyTorch SDPA
import torch.nn.functional as F
out = F.scaled_dot_product_attention(
q, k, v, dropout_p=0.0, is_causal=True)PyTorch 2.0+ 自动选 flash / mem-efficient / math kernel。Flash 需要 head_dim 整除、sm80+ 等——不满足 fallback math,仍 显存。训练 is_causal=True 与 inference 全 attention 图不同,profile 分场景。bf16 训练时确认 Flash 路径支持当前 dtype,否则 silent fallback。
4. 何时换架构而非只换 kernel
profile 确认 attention >40% forward 再投入。Swin window_size=7 把全局 变 。减 seq(patch merge、滑动窗口)往往比换 linear attention 更稳。盲目上 Performer 掉点后再 profile,是双倍浪费。LLM 预fill 与 decode 瓶颈不同,分别优化。
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CUDA],
) as prof:
model(x)
print(prof.key_averages().table(sort_by="cuda_time_total"))5. 训练与推理:KV cache
训练要存 dropout 与 backward 中间态;推理 KV cache 把复杂度摊到 decode,但 cache 占显存 。long context 部署要分 prefill 与 decode 两阶段测 latency。batch decode 与 batch prefill 的瓶颈不同,应分别 profile。服务 SLA 常绑 P99 decode latency,不是单次 prefill benchmark。
6. activation checkpoint 与 attention
OOM 时 gradient_checkpointing_enable() 换 compute 省 activation——与 FlashAttention 可叠加。先 profile 确认 bottleneck 在 attention 还是 FFN;FFN 主导时 checkpoint FFN 层更划算。checkpoint 增加 20–30% step time,验收要 metric parity + wall-clock 一起看。不要为了省显存 checkpoint 到 loss 计算图断裂。
7. 失败模式
| 现象 | 原因 |
|---|---|
| SDPA 没加速 | 不满足 flash 条件,fallback math |
| 4× seq 非 4× time | bottleneck 不在 attention |
| linear attention 掉点 | 近似损失 metric |
8. 验收
seq len ×2 测 attention CUDA time 是否 ~4×。换 SDPA/Flash 后 metric parity + latency 增益同报。固定 seed 单变量;val best → test 一次。上线用生产 seq len 分布测 P99 latency。
9. 案例:Flash 未生效
某模型 head_dim=80 不满足 Flash 对齐,SDPA silent fallback math,seq 2048 OOM。padding head_dim 到 96 后 Flash 生效,同 batch 训通且 step time 降 35%,metric 无差。教训:profile 里看 attention kernel 名,确认不是 aten::matmul 大包代替 Flash。
10. 与 FFN、embedding 的占比
profile 若 FFN 占 50%+,减 hidden 或 SwiGLU 中间维比 FlashAttention 更直接。embedding 层 seq×dim 也大,long context 时考虑 ALiBi、RoPE 与 cache 布局,不只是 attention kernel。报告里分开 attention / FFN / embedding 三栏 time,避免「上了 Flash 仍 OOM」却未查 FFN activation。
11. 工程预算模板
立项 long context 前先填:目标 seq n、batch b、head h、dim d,估算 activation 与 KV cache 字节数。超预算则选 window、sparse 或减 n,而不是先买卡。上线 SLA 绑 P99 decode latency,profile 用生产 trace 的 n 分布,不是实验室固定 n=512。多模态模型若 vision token 与 text token 同 attention,先数清各模态 n 再优化——只 optimize text 侧 Flash 可能仍 OOM 在 image patch。训练 gradient checkpoint 与 inference KV cache 是不同问题,别混谈「省显存技巧」。speculative decoding 减 decode 步数,与 FlashAttention 正交——先 profile 逐步时间,再选组合优化。长 context 预fill OOM 时,先减 batch 或 gradient checkpoint,再考虑 kernel;batch=1 仍 OOM 才是 kernel/架构问题。Multi-head attention 与 GQA/MQA 减 KV cache 是架构层省显存,与 SDPA kernel 正交——long context 服务应架构与 kernel 两条线分别评估。训练时 flash 不可用不要 silent 以为已优化,log backend 名进 step 0。Paper 报 long-context 结果应附 n 与 batch,否则他人无法判断是 kernel 还是规模贡献。
12. 案例:Flash 未生效 OOM
LLM prefill n=8k,SDPA fallback math OOM;padding head 96 后 Flash 生效,step time −38%。monitor 首 step log attention backend;不满足 Flash 条件时优先降 n 或开 checkpoint,而不是加卡忽视 profile。GQA 减 KV 与 Flash 可叠加,long context 服务两条线各做 ablation 再组合。Speculative decoding 与 Flash 无关,profile 勿把 prefill 加速算进 decode SLA。MoE 路由 overhead 有时大于 attention,profile 看 expert dispatch 占比。Context parallel 与 Flash 是不同维度优化,long context 训练先算 activation 账再选并行策略。Inference 批大小增大时 attention 仍 O(n²) 主导,batch 不是 linear 省 time。Prefill 与 decode 分开 profile 再报 SLA;训练 step time 不能代替 serving P99 decode。Ring attention 与 Flash 选型前先确认 model 已支持对应并行,别在 unsupported arch 上 profile 空转。Training 用 flash 而 inference 用 math 时,latency 验收要在 deploy EP 上重做,不能抄训练 profile。Long context 容量规划应 publish prefill 与 decode 分项 SLA,用生产 trace 的 n 分布做依据。Paper 对比 Flash 与 math 时应固定 hardware 与 driver,否则 kernel 对比无意义。Serving 团队与训练团队共用 profile 模板字段(n、batch、EP、dtype),避免口头传参。Capacity planning spreadsheet 应含 n² activation 项与 KV 字节项,别只填 FLOPs。Team wiki 维护「long context 优化决策树」:先 profile 再选 kernel/架构/并行,避免跳步。
Attention 贵是序列长度的平方律——先 profile 再选 kernel 或架构。
相关
也可以看看
- ·3 分钟阅读
PyTorch SDPA 的两个静默坑:布尔 mask 与 eval dropout
拆清 scaled_dot_product_attention 的 keep-mask 语义、广播维度和 dropout_p 行为,用最小张量测试阻止注意力泄漏。
- ·5 分钟阅读
Pre-LN 与 RMSNorm eps
LayerNorm 在 channel/ token 维归一化,不依赖 batch 统计——Transformer 和小 batch 场景比 BatchNorm 稳,但和 BN 的 inductive bias 不同。
- ·16 分钟阅读
KV cache 分页与前缀复用:吞吐上去后正确性怎么验
沿 block table、前缀哈希与抢占回收拆开缓存一致性风险,说明错页复用为何表现为「偶发胡话」;用确定性前缀集、强制抢占和逐 token 对照验收。
johan's blog