From bf651ea9154da7baa96ea04c32bca775f34e50bc Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Tue, 14 Jul 2026 15:27:39 +0800 Subject: [PATCH 01/14] feat: add mxfp8 (e4m3) flash-attention forward kernel --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 277 ++++ alto/kernels/mxfp8/__init__.py | 4 + .../mxfp8/triton_flash_attention_mxfp8.py | 1179 +++++++++++++++++ tests/unittest/mxfp8/test_mxfp8_attention.py | 123 ++ 4 files changed, 1583 insertions(+) create mode 100644 alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md create mode 100644 alto/kernels/mxfp8/triton_flash_attention_mxfp8.py create mode 100644 tests/unittest/mxfp8/test_mxfp8_attention.py diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md new file mode 100644 index 00000000..769a5e97 --- /dev/null +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -0,0 +1,277 @@ +# MXFP8 E4M3 Flash-Attention(forward + backward)— Minimum Viable Plan + +目标:在 **AMD MI350 (CDNA4 / gfx950)** 上实现一个能支撑 mxfp8 训练的最小可用 mxfp8 flash-attention kernel,**包含前向与反向(dQ / dK / dV)**。 + +## V1 三项决定(已与需求方确认) + +1. **forward + backward 都做**。⚠️ 关键前提:参照实现 `mxfp4 attention` 的 backward 是 17 个 `None` 占位、grad 断言全部注释——**backward 没有可机械改写的源**。V1 的 backward 需按标准 Flash-Attention v2 backward 数学**从零实现**,每个 dot 各自量化(见 §4.B / §9)。这是本 plan 相对 group gemm / mxfp4 attention 的**最大新增工作量与头号风险**。 + +2. **全 e4m3(已定)**。forward 的 Q/K/V/P 与 backward 的 dO/dP/dS **全部 e4m3**——整版 V1 单一 e4m3 格式,kernel 无 dtype 分发。e5m2 通道参数化预留但不启用(未来项)。 + +3. **仅 CDNA4,不加 fallback**。kernel 只走 `tl.dot_scaled`(e4m3×e4m3);**不写 CDNA3 dequant fallback**,不引入 `USE_DOT_SCALED` 开关、不引入 `use_dot_scaled` 参数。MI300 / CDNA3 明确不在 V1 范围。数值 ground-truth 由 **PyTorch 侧模拟 reference** 提供(见 §8),不再依赖 in-kernel fallback。 + +## 参照实现 + +- **kernel 移植源(仅 forward)**:`alto/kernels/fp4/mxfp4/triton_flash_attention_mxfp4.py`(FA v2 + mxfp4,forward-only)。 +- **backward 数学参照**:标准 Flash-Attention v2 backward(无 mxfp4 代码可抄,仅抄算法)。 +- **量化基础设施**:本目录 `mxfp8_quantization.py`(`convert_to_mxfp8` / `_calculate_scales` / `_quantize_fp8`)。**注意 V1 无 fallback,不使用 `_dequantize_fp8` 于 kernel 内**(`_dequantize_fp8` 只在 PyTorch 侧 reference 里用)。 +- **流程与验收范式**:本目录 `MXFP8_GROUPED_GEMM_PLAN.md`(「机械改写 + 分层校验 + 真机验证」方法论)。 + +## 环境记录 + +本 plan 撰写机为 gfx950(CDNA4/MI350),`is_cdna4()=True`。native `tl.dot_scaled` 路径可在本机验证;MI300/CDNA3 不在 V1 范围。 + +--- + +## 0. 格式选择:V1 全 e4m3,e5m2 留给后续 + +**forward(Q/K/V/P)全 e4m3**——理由与 group gemm V1 一致:单一格式 → kernel 无 dtype 分发、最小版本最简单。attention forward 的四个 operand 都偏「动态范围小、单元素精度重要」,e4m3 天然合适: + +| operand | 分布特征 | e4m3 是否够用 | +|---|---|---| +| Q / K | LayerNorm 后激活,分布集中(±几十),动态范围需求低 | ✅ 单元素精度更重要 | +| V | 同上 | ✅ | +| P(softmax 概率) | ∈ [0,1],行和为 1;长尾但**有界**,且被 online-softmax 逐行重归一化 | ✅ ~6% 相对误差可接受 | + +**backward(dO / dP / dS)——V1 已定全 e4m3**(需求方确认,整版 V1 单一格式)。 + +风险记录(供未来升级参考,非 V1 待办):grad 是长尾 + 动态范围大的分布(`dP`/`dS` 尤甚),理论上是 e5m2 的场景(见 group gemm plan §0),全 e4m3 可能 underflow 小尾部。V1 接受此风险,理由是保持「单一格式、无 dtype 分发」的最小主题、且 backward 已是从零实现的大工程,不叠加混合格式复杂度。为让未来升级无痛,backward 的每个 `tl.dot_scaled` 同样把 dtype 参数化为 `LHS_FORMAT_ID` / `RHS_FORMAT_ID`(0=e4m3, 1=e5m2)constexpr,V1 默认全 0;若日后 backward 数值不过关,切 e5m2 即可,kernel 本体不动。 + +forward 的两个 `tl.dot_scaled` 也用同一套 `LHS_FORMAT_ID` / `RHS_FORMAT_ID` 参数化,V1 默认 e4m3,kernel 内 `if/else` 只走 e4m3 分支。 + +--- + +## 1. 复用 vs 改写 vs 新写 + +### 可直接复用(不动) + +`mxfp8_quantization.py` 提供 attention 需要的全部 device 端原语,无需新写量化代码: + +- `convert_to_mxfp8(x, axis=..., mxfp_format="e4m3", is_2d_block=...)`:wrapper 层量化 Q / K / V,以及 backward 里的 dO。 +- `_calculate_scales(x, ..., target_max_pow2, mbits, IS_2D_BLOCK=False)`:kernel 内算 P(及 backward 的 dP/dS)的 block scale。签名比 mxfp4 多 `target_max_pow2` / `mbits`,e4m3 传 `target_max_pow2=8, mbits=3`。 +- `_quantize_fp8(p, ps, ..., FP8_FORMAT=0, IS_2D_BLOCK=False)`:替代 mxfp4 的 `_pack_fp4`,kernel 内动态量化。 +- `is_cdna4()`:device 断言(V1 只支持 CDNA4)。 + +mxfp4 attention 里与量化正交、原样保留的 FA v2 forward 骨架:causal masking、GQA/MQA、varlen(thd)/bshd/bhsd、alibi/bias/dropout、LSE 写回、padded head、online-softmax 累积、全 0 块 early-exit、autotune wrapper、`triton_op` + `autograd.Function` + 用户入口三段式。 + +### 必须改写(forward,对应 group gemm plan §1) + +1. **删掉所有 head_dim packing**(mxfp4 沿 head_dim 一 byte 装两元素,mxfp8 一 byte 一元素): + - `get_shape_from_layout` 里的 `head_size_q/k/v *= 2` + - `HALF_BLOCK_DMODEL_QK/V`、`HALF_ACTUAL_BLOCK_DMODEL_QK/V`、`offs_d_qk_pack`、`offs_d_v_pack` + - `tl.dot_scaled(..., lhs_k_pack=, rhs_k_pack=)` 的 `*_k_pack` 参数 + - Q/K/V 指针里按 half-dim 的 stride,改回全 head_dim + - `SCALE_BLOCK_DMODEL_QK` 保留,但基于未 pack 的 head_dim 重算 +2. **两个 `tl.dot_scaled` 的 dtype**:mxfp4 写死 `"e2m1"` → mxfp8 参数化 `LHS_FORMAT_ID`/`RHS_FORMAT_ID`(V1 默认 e4m3),只走 e4m3×e4m3。 +3. **kernel 内 P 的动态量化**:`_calculate_scales`(带 `target_max_pow2=8, mbits=3`)→ `_pack_fp4` 换成 `_quantize_fp8(FP8_FORMAT=0)`。P 的 scale 沿 `BLOCK_N`(seqlen_k)方向,`IS_2D_BLOCK=False`。 +4. **wrapper 层量化**:`convert_to_mxfp4(axis=-1, is_2d_block=True)` → `convert_to_mxfp8(axis=-1, mxfp_format="e4m3", is_2d_block=...)`。 +5. **删除 fallback 相关**(本 plan 的净化项,非 mxfp4 有):不引入 `USE_DOT_SCALED`、`use_dot_scaled`、`_dequantize_fp8` 的 kernel 内调用。kernel 只有 `tl.dot_scaled` 一条路径。入口断言 `is_cdna4()`。 + +### 全新写(backward,无参照,头号工作量) + +按标准 FA v2 backward 数学从零实现三段(见 §9 完整推导): +- **preprocess kernel**:`delta = rowsum(dO ∘ O)` +- **dK/dV kernel**:遍历 K/V 块、重算 P、`dV += Pᵀ@dO`、`dP = dO@Vᵀ`、`dS = P∘(dP−delta)`、`dK += dSᵀ@Q` +- **dQ kernel**:`dQ += dS@K` + +每个 dot 用 `tl.dot_scaled`,operand 在对应 reduction 维上量化(forward 已存的 e4m3 版本复用,新产生的 dO/dP/dS 现场量化)。全部只走 CDNA4,无 fallback。 + +### 全新写(其它) + +- `alto/kernels/mxfp8/triton_flash_attention_mxfp8.py`(forward 从 mxfp4 机械改写 + backward 从零写) +- `dispatch/attention.py` 加 `mxfp8_e4m3` 分支 +- `tests/unittest/mxfp8/test_mxfp8_attention.py` + +**预估**:forward kernel + wrapper ~450 行(删了 packing/fallback 比 mxfp4 少);backward 三个 kernel + wrapper ~500 行(全新);测试 ~250 行。 + +--- + +## 2. 接口契约 + +### scale 布局(沿用 mxfp4 注释) + +- **Q scale**:沿 head_dim(reduction 维)量化,`[..., seqlen_q, head_dim/32]` +- **K scale**:`"k scale is N×K even though k is K×N"`——scale 在非 reduction 维(seqlen_k)上 major:`[..., seqlen_k, head_dim/32]` +- **V scale**:PV reduction 维是 seqlen_k,V scale 沿 seqlen_k:`[..., head_dim_v, seqlen_k/32]` +- **P scale**:kernel 内动态生成,沿 `BLOCK_N`(seqlen_k)方向 + +**约束**:head_dim 与 seqlen_k 必须被 `QUANT_BLOCK_SIZE=32` 整除。head_dim ∈ {128,192}、seqlen ∈ {1024,2048} 均满足。head_dim<64 因 `tl.dot_scaled` 限制不支持(mxfp4 亦然)。 + +### 用户 API(对齐 mxfp4,加一个格式开关,去掉 fallback 开关) + +```python +def triton_attention_mxfp8( + q: torch.Tensor, # [batch, nheads_q, seqlen_q, head_dim_qk](bhsd) + k: torch.Tensor, + v: torch.Tensor, + alibi_slopes: torch.Tensor | None, + bias: torch.Tensor | None, + sm_scale: float, + dropout_p: float, + cu_seqlens_q: int, + cu_seqlens_k: int, + max_seqlens_q: int, + max_seqlens_k: int, + causal: bool, + return_scores: bool, + use_exp2: bool, + layout: str, # "bshd" / "bhsd" / "thd" + *, + fwd_format: str = "e4m3", # Q/K/V/P 格式,V1 固定 e4m3;e5m2 预留 + bwd_grad_format: str = "e4m3", # dO/dP/dS 格式,V1 固定 e4m3;e5m2 预留给未来 +) -> Tuple[Tensor, Tensor, Tensor]: # (o, softmax_lse, exp_scores) +``` + +签名与 `triton_attention_mxfp4` 同形(方便 dispatch 直接替换),多 `fwd_format` / `bwd_grad_format` 两个 keyword-only 参数。**注意:不再有 `use_dot_scaled` 参数**(V1 只支持 CDNA4)。反向经 `autograd.Function` 自动触发,用户不直接调 backward kernel。 + +--- + +## 3. 各 dot 的 contraction & scale axis + +### Forward + +| dot | 计算 | reduction 维 | LHS quant axis | RHS quant axis | 一次 dot 跨几个 scale group | +|---|---|---|---|---|---| +| QK | `S = Q @ Kᵀ` | head_dim | Q: head_dim | K: head_dim | head_dim/32(=4 @ dim128)⚠️ | +| PV | `O += P @ V` | seqlen_k (BLOCK_N) | P: BLOCK_N | V: seqlen_k | BLOCK_N/32(=2 @ BLOCK_N=64) | + +### Backward(见 §9 推导) + +| dot | 计算 | reduction 维 | 备注 | +|---|---|---|---| +| dV | `dV += Pᵀ @ dO` | seqlen_q (BLOCK_M) | P 重算,dO 现场量化沿 M | +| dP | `dP = dO @ Vᵀ` | head_dim_v | V 复用 fwd e4m3 | +| dK | `dK += dSᵀ @ Q` | seqlen_q (BLOCK_M) | dS 现场量化沿 M | +| dQ | `dQ += dS @ K` | seqlen_k (BLOCK_N) | dS 现场量化沿 N、K 复用 fwd e4m3 | + +> ⚠️ **QK 一次 `dot_scaled` 跨 head_dim/32 个 32-wide scale group,是最大数值风险点**(group gemm plan §1.3 / §5 专门警告)。mxfp4 已接受此误差(未硬断言精度);mxfp8 更敏感。V1 无 fallback,改由 §8 的 **PyTorch 模拟 reference** 度量此误差;若发散,退路见 §5 风险 1。backward 的 dot 同样跨 group,同法度量。 + +--- + +## 4. 落地步骤 + +### 4.A Forward(机械改写 mxfp4) + +- **Step 1 — 骨架**:复制 `triton_flash_attention_mxfp4.py` → `triton_flash_attention_mxfp8.py`;改 import(`from .mxfp8_quantization import BLOCK_SIZE_DEFAULT, is_cdna4, _calculate_scales, _quantize_fp8`,删 `_pack_fp4`/`_unpack_fp4`);全局改名 `mxfp4→mxfp8`、`e2m1→e4m3`、op 名 `attention_mxfp8_forward_triton_impl`;kernel body 先占位跑通 import。 +- **Step 2 — 删 packing + QK dot**:按 §1 删所有 head_dim packing;QK 改 `tl.dot_scaled(q, qs, "e4m3", k, ks, "e4m3", out_dtype=fp32)`(无 `*_k_pack`)。验证:head_dim=128、单 head、无 causal,QK-only 输出 vs PyTorch 模拟 reference。 +- **Step 3 — PV dot + P 动态量化**:P 用 `_calculate_scales(..., target_max_pow2=8, mbits=3, IS_2D_BLOCK=False)` → `_quantize_fp8(..., FP8_FORMAT=0, IS_2D_BLOCK=False)`;PV 改 `tl.dot_scaled(p_fp8, ps, "e4m3", v, vs, "e4m3", out_dtype=fp32)`。验证:完整前向 vs PyTorch SDPA bf16(先不 causal)。 +- **Step 4 — forward wrapper + autograd.forward**:`attention_mxfp8_forward_triton_impl` 量化 Q/K/V(`axis=-1`),补输入契约检查(`head_dim % 32 == 0`、`seqlen_k % 32 == 0`、contiguous、`is_cdna4()`),launch。`_triton_attention_mxfp8.forward` 调 wrapper,`save_for_backward(q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias)`。`register_fake` 照搬 mxfp4。 + +### 4.B Backward(从零实现,见 §9) + +- **Step 5 — preprocess**:`delta = rowsum(dO ∘ O)`,`[batch, nheads, seqlen_q]`。dO 量化为 e4m3(`convert_to_mxfp8`)。 +- **Step 6 — dK/dV kernel**:遍历 K/V 块,重算 S→P(复用 fwd 的 q/k e4m3 + LSE),`dV += Pᵀ@dO`、`dP = dO@Vᵀ`、`dS = P∘(dP−delta)`、`dK += dSᵀ@Q`。每个 dot 用 `tl.dot_scaled`,operand 沿对应 reduction 维量化。验证:dK/dV vs bf16 SDPA autograd。 +- **Step 7 — dQ kernel**:`dQ += dS@K`。验证:dQ vs bf16 SDPA autograd。 +- **Step 8 — autograd.backward 接线**:`_triton_attention_mxfp8.backward` 从 ctx 取回张量、调 preprocess/dK-dV/dQ 三个 kernel,返回 `dq, dk, dv`(其余入参位 `None`)。 + +### 4.C Dispatch 接入 + +- **Step 9 — dispatch**:`dispatch/attention.py` 的 `LPScaledDotProductAttentionWrapper.__init__` 加分支: + ```python + elif config.precision in ("mxfp8_e4m3", "mxfp8"): + self.attn_func = triton_attention_mxfp8 + ``` + forward 已 layout-agnostic(传 `layout="bhsd"`),无需改。 + +### 4.D 测试 + +- **Step 10**:`test_mxfp8_attention.py`,见 §8。 + +--- + +## 5. 关键风险与对策 + +| 风险 | 对策 | +|---|---| +| QK 一次 `dot_scaled` 跨 head_dim/32 个 scale group 发散(§3 ⚠️,头号数值风险) | 由 §8 PyTorch 模拟 reference 度量;若 native vs reference SNR 差距过大 → 退路是**沿 head_dim 分块累加 QK**(每 32-wide group 一次 `dot_scaled` + acc),代价是 QK 循环变长。V1 先测,超阈值再改 | +| **backward 无参照、从零写**(头号工作量风险) | 严格对标标准 FA v2 backward 数学(§9);每个 kernel 单独 vs bf16 SDPA autograd 校验(dV→dK→dQ 逐个隔离);先跑通非量化版(内部用高精度)再逐 dot 换 `tl.dot_scaled` | +| backward 全 e4m3 时 grad underflow / 数值不过关(§0 风险记录) | dtype 已参数化(`LHS/RHS_FORMAT_ID`);不过关时把 dO/dP/dS 切 e5m2(`bwd_grad_format="e5m2"`),kernel 本体不动 | +| P 在 kernel 内动态量化,`_quantize_fp8` 的 scale 广播语义与 mxfp4 `_pack_fp4` 不同 | Step 3 单独验证 P 量化路径;对照 `_quantize_fp8` 的 1D-block 分支(`IS_2D_BLOCK=False`) | +| head_dim / seqlen 非 32 倍数 | V1 断言拒绝(对齐 mxfp4 只测 ≥128 head_dim);padded-head 逻辑保留但要求 actual head_dim 仍 32 对齐 | + +--- + +## 6. 不做的事(V1 明确划线) + +- ❌ **CDNA3 / MI300 支持与 fallback 路径**——V1 仅 CDNA4,kernel 只有 `tl.dot_scaled` 一条路径 +- ❌ **混合格式**(forward P 或 backward grad 用 e5m2)——全 e4m3,e5m2 分支参数化预留不启用 +- ❌ head_dim < 64(`tl.dot_scaled` 限制) +- ❌ 2D-block P 量化(P 沿 `BLOCK_N` 一维即可) +- ❌ autotune 扩展(沿用 mxfp4 单 config:`BLOCK_M=BLOCK_N=64, PRE_LOAD_V=False`) +- ❌ TMA / async copy / pipelining 调优 +- ❌ 沿 head_dim 分块 QK(除非 §5 风险 1 触发) +- ❌ FSDP/TP 集成测试 + +--- + +## 7. 验收标准(V1 完成定义) + +| # | 标准 | 说明 | +|---|---|---| +| 1 | forward kernel 在 CDNA4 跑通 | 本机 gfx950 native `tl.dot_scaled` | +| 2 | backward(dQ/dK/dV)三个 kernel 在 CDNA4 跑通 | 本机 gfx950 | +| 3 | forward vs PyTorch SDPA bf16:cos-sim > 0.99、SNR > 阈值(硬断言) | 对齐 mxfp4 对比方式,但把 SNR/cossim 变成硬断言(mxfp4 只 print,此为主动加严) | +| 4 | backward dQ/dK/dV vs bf16 SDPA autograd:cos-sim / SNR 硬断言 | 阈值首跑标定,留裕度 | +| 5 | 覆盖 causal × GQA × head_dim{128,192} × seqlen{1024,2048} 网格 | 沿用 mxfp4 test_cases | +| 6 | dispatch 层 `mxfp8_e4m3` 分支可路由 | smoke:构造 config 走 `LPScaledDotProductAttentionWrapper` | + +> **相对 mxfp4 attention 的主动加严**:mxfp4 的 `test_mxfp_attention.py` 只 print、无 assert 且无 backward 测试;mxfp8 V1 把前向精度做成硬断言,并新增完整 backward 的梯度校验。 + +--- + +## 8. 测试组织(分层校验,归因清晰) + +文件:`tests/unittest/mxfp8/test_mxfp8_attention.py`,复用 `tests/unittest/mxfp8/utils.py` 的 `calc_snr` / `calc_cossim` / `prepare_data`(与 fp4 共用 `alto.kernels.fp4.testing_utils`)。 + +**无 fallback,ground-truth 改由 PyTorch 侧模拟 reference 提供。** 两层校验: + +1. **kernel 移植正确性**:native kernel vs **PyTorch 模拟 reference**——reference 用 `convert_from_mxfp8` dequant Q/K/V → 算 S=Q@Kᵀ → softmax → 用 PyTorch 版 mxfp8 量化 P → dequant → @V(即「在 PyTorch 里复刻 kernel 的量化流程」)。 + - 隔离:masking / online-softmax / LSE 的移植 bug,排除 mxfp8 量化误差本身。 + - 附带度量 §3 ⚠️ 的 QK 跨-group 误差(native kernel 的 `dot_scaled` 与 reference 的精确 dequant-matmul 之差即此误差)。 + - 门槛:SNR 高(> 40 dB 量级)。 +2. **端到端量化误差**:native kernel vs **PyTorch SDPA bf16**(无量化)。 + - 隔离:mxfp8 量化本身的总误差。 + - 门槛:cos-sim > 0.99 + SNR 硬断言(验收标准 3)。 +3. **backward 校验**:完整 `autograd` 反向 vs bf16 SDPA autograd,比 dQ/dK/dV 的 cos-sim / SNR(验收标准 4)。 + +参数网格:`batch=4` × mxfp4 的 `test_cases`(causal=True)× 首层额外跑 causal=False。 + +--- + +## 9. Backward 设计(V1 实现,无参照,从零写) + +> backward 是 V1 交付物且无 mxfp4 代码可抄,本节锁定算法与量化决策。 + +### 9.1 标准 FA v2 backward 数学 + +给定前向已存的 `Q, K, V, O, softmax_lse`(LSE = log-sum-exp per row)与上游 `dO`: + +- **preprocess**:`delta = rowsum(dO ∘ O)`,shape `[batch, nheads, seqlen_q]` +- **dK/dV kernel**(遍历 K/V 块,块内 recompute P): + - `S = Q @ Kᵀ * sm_scale`(+ mask/alibi/bias),`P = exp(S − softmax_lse)` + - `dV += Pᵀ @ dO` + - `dP = dO @ Vᵀ` + - `dS = P ∘ (dP − delta) * sm_scale` + - `dK += dSᵀ @ Q` +- **dQ kernel**:`dQ += dS @ K`(dS 同上重算) + +### 9.2 V1 量化决策(全 e4m3) + +- `Q / K / V` 复用 forward 已量化的 e4m3 版本(ctx 已存 `q, k, v, q_scale, k_scale, v_scale`)。 +- `dO` 在 backward 入口量化为 e4m3(`convert_to_mxfp8`),沿 dK/dV 需要的 reduction 维(seqlen_q)与 dP 需要的维(head_dim_v)——按 group gemm plan 的经验,可能需要**两套 dO 量化**(沿不同 axis),沿用其 autograd 已有逻辑。 +- `P / dS` 在 kernel 内现场量化(`_calculate_scales` + `_quantize_fp8`),沿各自 dot 的 reduction 维。 +- 每个 dot 的 dtype 走 `LHS/RHS_FORMAT_ID`,V1 全 0(e4m3)。 + +### 9.3 e5m2 升级路径(若 §0 的 e4m3 backward 数值不过关) + +把 `bwd_grad_format` 切 `"e5m2"` → dO/dP/dS 的量化与对应 `tl.dot_scaled` 的 `FORMAT_ID` 改 1,kernel 本体不动。这是 §5「backward 数值不过关」风险的兜底。 + +### 9.4 save_for_backward 清单 + +forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias`(确保 backward 三个 kernel 所需张量齐全,避免中途改 forward 签名)。 + +--- + +## 10. 一句话总结 + +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(backward 格式 ⚠️ 待确认,e5m2 参数化预留)、**仅 CDNA4 无 fallback**,ground-truth 由 PyTorch 模拟 reference 提供。量化基础设施已齐,forward FA v2 骨架复用;最大工作量与风险集中在从零的 backward。 diff --git a/alto/kernels/mxfp8/__init__.py b/alto/kernels/mxfp8/__init__.py index 85eac522..f8c68c7b 100644 --- a/alto/kernels/mxfp8/__init__.py +++ b/alto/kernels/mxfp8/__init__.py @@ -1,3 +1,7 @@ # Copyright (c) 2026 Advanced Micro Devices, Inc. # # SPDX-License-Identifier: MIT + +from .triton_flash_attention_mxfp8 import triton_attention_mxfp8 + +__all__ = ("triton_attention_mxfp8", ) diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py new file mode 100644 index 00000000..bbf2c4fe --- /dev/null +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -0,0 +1,1179 @@ +# Copyright (c) 2026 Advanced Micro Devices, Inc. +# +# SPDX-License-Identifier: MIT +""" +Fused Attention (MXFP8 / e4m3) +============================== + +Triton implementation of the Flash Attention v2 algorithm from Tri Dao +(https://tridao.me/publications/flash2/flash2.pdf) with MXFP8 (e4m3) block +scaled operands. + +This is a mechanical port of ``triton_flash_attention_mxfp4.py``: + * FP4 packs two elements per byte along head_dim; FP8 stores one element per + byte, so all ``// 2`` head-dim halving is removed. + * ``tl.dot_scaled`` dtype string ``"e2m1"`` -> ``"e4m3"``; the ``lhs_k_pack`` + / ``rhs_k_pack`` arguments are dropped (they only exist for packed FP4). + * The on-the-fly softmax-probability quantization uses the MXFP8 helpers + ``_calculate_scales`` / ``_quantize_fp8`` instead of the FP4 pack path. + +The scale layout is identical to the FP4 kernel (uint8, ``[.., dim/32]``, 2D +block), so all scale pointer arithmetic carries over unchanged. + +Forward only, matching the FP4 reference (backward is not implemented). +""" +from typing import Tuple, Optional +import os +import torch +import triton +import triton.language as tl +from torch._library import triton_op, wrap_triton + +from .mxfp8_quantization import ( + BLOCK_SIZE_DEFAULT, + is_cdna4, + _calculate_scales, + _quantize_fp8, + FORMAT_TO_TARGET_MAX, + FORMAT_TO_MBITS, +) + +fwd_torch_dtype: tl.constexpr = torch.bfloat16 +bwd_torch_dtype: tl.constexpr = torch.float32 +# Seed the RNG so we get reproducible results for testing. +philox_seed: tl.constexpr = 0x1BF52 +philox_offset: tl.constexpr = 0x1D4B42 + +# e4m3 quantization constants (see mxfp8_quantization.FORMAT_TO_*). +E4M3_TARGET_MAX_POW2: tl.constexpr = FORMAT_TO_TARGET_MAX["e4m3"] +E4M3_MBITS: tl.constexpr = FORMAT_TO_MBITS["e4m3"] +E4M3_FORMAT_ID: tl.constexpr = 0 + +AUTOTUNE = os.environ.get('FLASH_ATTENTION_TRITON_AMD_AUTOTUNE', '0').lower() in ('1', 'true', 'yes') +DEBUG = os.environ.get('FLASH_ATTENTION_TRITON_AMD_DEBUG', '0').lower() in ('1', 'true', 'yes') +PERF = os.environ.get('FLASH_ATTENTION_TRITON_AMD_PERF', '0').lower() in ('1', 'true', 'yes') + + +def get_shape_from_layout(q, k, v, layout, cu_seqlens_q=None, cu_seqlens_k=None, max_seqlen_q=None, max_seqlen_k=None): + if layout == 'bhsd': + batch_q, nheads_q, max_seqlen_q, head_size_q = q.shape + batch_k, nheads_k, max_seqlen_k, head_size_k = k.shape + batch_v, nheads_v, max_seqlen_v, head_size_v = v.shape + elif layout == 'bshd': + batch_q, max_seqlen_q, nheads_q, head_size_q = q.shape + batch_k, max_seqlen_k, nheads_k, head_size_k = k.shape + batch_v, max_seqlen_v, nheads_v, head_size_v = v.shape + elif layout == 'thd': + batch_q, max_seqlen_q, nheads_q, head_size_q = len(cu_seqlens_q) - 1, max_seqlen_q, q.shape[1], q.shape[2] + batch_k, max_seqlen_k, nheads_k, head_size_k = len(cu_seqlens_k) - 1, max_seqlen_k, k.shape[1], k.shape[2] + batch_v, max_seqlen_v, nheads_v, head_size_v = len(cu_seqlens_k) - 1, max_seqlen_k, v.shape[1], v.shape[2] + else: + assert False, "Got unsupported layout." + + # FP8 stores one element per byte, so head_dim is not packed. + + # assert + assert batch_q == batch_k + assert head_size_q == head_size_k + + return batch_q, nheads_q, nheads_k, head_size_q, head_size_v, max_seqlen_q, max_seqlen_k + + +def get_strides_from_layout(q, layout): + if layout == 'thd': + q_strides = (0, q.stride(1), q.stride(0), q.stride(2)) + elif layout == 'bhsd': + q_strides = (q.stride(0), q.stride(1), q.stride(2), q.stride(3)) + elif layout == 'bshd': + q_strides = (q.stride(0), q.stride(2), q.stride(1), q.stride(3)) + else: + assert False, 'Got unsupported layout.' + return q_strides + + +def get_padded_headsize(size): + # Get closest power of 2 over or equal to 32. + padded_d_model = 1 << (size - 1).bit_length() + # Smallest head_dim supported is 16. If smaller, the tile in the + # kernel is padded - there is no padding in memory for any dims. + padded_d_model = max(padded_d_model, 16) + return padded_d_model + + +@triton.jit +def cdiv_fn(x, y): + return (x + y - 1) // y + + +@triton.jit +def dropout_offsets(philox_seed, philox_offset, dropout_p, m, n, stride): + ms = tl.arange(0, m) + ns = tl.arange(0, n) + return philox_offset + ms[:, None] * stride + ns[None, :] + + +@triton.jit +def dropout_rng(philox_seed, philox_offset, dropout_p, m, n, stride): + rng_offsets = dropout_offsets(philox_seed, philox_offset, dropout_p, m, n, stride).to(tl.uint32) + # TODO: use tl.randint for better performance + return tl.rand(philox_seed, rng_offsets) + + +@triton.jit +def dropout_mask(philox_seed, philox_offset, dropout_p, m, n, stride): + rng_output = dropout_rng(philox_seed, philox_offset, dropout_p, m, n, stride) + rng_keep = rng_output > dropout_p + return rng_keep + + +# Convenience function to load with optional boundary checks. +# "First" is the major dim, "second" is the minor dim. +@triton.jit +def load_fn(ptrs, offset_first, offset_second, boundary_first, boundary_second, other=0): + if offset_first is not None and offset_second is not None: + mask = (offset_first[:, None] < boundary_first) & \ + (offset_second[None, :] < boundary_second) + tensor = tl.load(ptrs, mask=mask, other=other) + elif offset_first is not None: + mask = offset_first[:, None] < boundary_first + tensor = tl.load(ptrs, mask=mask, other=other) + elif offset_second is not None: + mask = offset_second[None, :] < boundary_second + tensor = tl.load(ptrs, mask=mask, other=other) + else: + tensor = tl.load(ptrs) + return tensor + + +@triton.jit +def compute_alibi_block(alibi_slope, seqlen_q, seqlen_k, offs_m, offs_n, transpose=False): + # when seqlen_k and seqlen_q are different we want the diagonal to stick to the bottom right of the attention matrix + # for casual mask we want something like this where (1 is kept and 0 is masked) + # seqlen_q = 2 and seqlen_k = 5 + # 1 1 1 1 0 + # 1 1 1 1 1 + # seqlen_q = 5 and seqlen_k = 2 + # 0 0 + # 0 0 + # 0 0 + # 1 0 + # 1 1 + # for alibi the diagonal is 0 indicating no penalty for attending to that spot and increasing penalty for attending further from the diagonal + relative_pos_block = offs_m[:, None] + seqlen_k - seqlen_q - offs_n[None, :] + alibi_block = -1 * alibi_slope * tl.abs(relative_pos_block) + if transpose: + return alibi_block.T + else: + return alibi_block + + +@triton.jit +def _attn_fwd_inner( + acc, + l_i, + m_i, + q, + qs, + k_ptrs, + v_ptrs, + ks_ptrs, + vs_ptrs, + bias_ptrs, + stride_kn, + stride_vk, + stride_bn, + stride_kscale_n, + stride_vscale_k, + start_m, + actual_seqlen_k, + actual_seqlen_k_scale, + actual_seqlen_q, + dropout_p, + philox_seed, + batch_philox_offset, + exp_scores_ptrs, + block_min, + block_max, + offs_n_causal, + masked_blocks, + n_extra_tokens, + alibi_slope, + score_ptrs, + scores_scaled_shifted_ptrs, + IS_CAUSAL: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + SCALE_BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_N_SCALE: tl.constexpr, + OFFS_M: tl.constexpr, + OFFS_N: tl.constexpr, + PRE_LOAD_V: tl.constexpr, + MASK_STEPS: tl.constexpr, + ENABLE_DROPOUT: tl.constexpr, + PADDED_HEAD_QK: tl.constexpr, + PADDED_HEAD_V: tl.constexpr, + ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + ACTUAL_BLOCK_DMODEL_V: tl.constexpr, + SM_SCALE: tl.constexpr, + USE_EXP2: tl.constexpr, + RETURN_SCORES: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, +): + if USE_EXP2: + RCP_LN2: tl.constexpr = 1.4426950408889634 + + # loop over k, v, and update accumulator + for start_n in range(block_min, block_max, BLOCK_N): + # For padded blocks, we will overrun the tensor size if + # we load all BLOCK_N. For others, the blocks are all within range. + if MASK_STEPS: + k_offs_n = start_n + tl.arange(0, BLOCK_N) + vs_offs_n = start_n // QUANT_BLOCK_SIZE + tl.arange(0, BLOCK_N_SCALE) + else: + k_offs_n = None + vs_offs_n = None + if PADDED_HEAD_QK: + k_offs_k = tl.arange(0, BLOCK_DMODEL_QK) + ks_offs_k = tl.arange(0, SCALE_BLOCK_DMODEL_QK) + else: + k_offs_k = None + ks_offs_k = None + k = load_fn(k_ptrs, k_offs_k, k_offs_n, ACTUAL_BLOCK_DMODEL_QK, actual_seqlen_k) + ks = load_fn(ks_ptrs, k_offs_n, ks_offs_k, actual_seqlen_k, SCALE_ACTUAL_BLOCK_DMODEL_QK, other=1) + + if PADDED_HEAD_V: + v_offs_k = tl.arange(0, BLOCK_DMODEL_V) + vs_offs_k = tl.arange(0, BLOCK_DMODEL_V) + else: + v_offs_k = None + vs_offs_k = None + if PRE_LOAD_V: + # We can use the same offsets as k, just with dims transposed. + v = load_fn(v_ptrs, k_offs_n, v_offs_k, actual_seqlen_k, ACTUAL_BLOCK_DMODEL_V) + vs = load_fn(vs_ptrs, vs_offs_k, vs_offs_n, ACTUAL_BLOCK_DMODEL_V, actual_seqlen_k_scale, other=1) + qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32) + # We start from end of seqlen_k so only the first iteration would need + # to be checked for padding if it is not a multiple of block_n + if MASK_STEPS: + # If this is the last block / iteration, we want to + # mask if the sequence length is not a multiple of block size + if (start_n + BLOCK_N == block_max) and (n_extra_tokens != 0): + boundary_m = tl.full([BLOCK_M], actual_seqlen_k, dtype=tl.int32) + size_n = start_n + OFFS_N[None, :] + mask = size_n < boundary_m[:, None] + qk = tl.where(mask, qk, float("-inf")) + + # -- compute qk ---- + qk += tl.dot_scaled(q, qs, "e4m3", k, ks, "e4m3", out_dtype=tl.float32) + + qk_scaled = qk * SM_SCALE + + if RETURN_SCORES: + score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ( + (start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k) + tl.store(score_ptrs, qk_scaled, mask=score_mask) + + if IS_CAUSAL: + causal_boundary = start_n + offs_n_causal + causal_mask = OFFS_M[:, None] >= causal_boundary[None, :] + qk_scaled = tl.where(causal_mask, qk_scaled, float("-inf")) + if bias_ptrs is not None: + bias_offs_n = start_n + tl.arange(0, BLOCK_N) if MASK_STEPS else None + bias = load_fn(bias_ptrs, OFFS_M, bias_offs_n, actual_seqlen_q, actual_seqlen_k) + qk_scaled += bias + + if alibi_slope is not None: + # Compute the global position of each token within the sequence + global_m_positions = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + global_n_positions = start_n + tl.arange(0, BLOCK_N) + alibi_block = compute_alibi_block(alibi_slope, actual_seqlen_q, actual_seqlen_k, global_m_positions, + global_n_positions) + qk_scaled += alibi_block + # get max scores so far + m_ij = tl.maximum(m_i, tl.max(qk_scaled, 1)) + + # scale and subtract max + q_shifted = qk_scaled - m_ij[:, None] + if RETURN_SCORES: + scores_scaled_shifted_mask = (OFFS_M[:, None] < actual_seqlen_q) & ( + (start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k) + tl.store(scores_scaled_shifted_ptrs, q_shifted, mask=scores_scaled_shifted_mask) + + # Compute scaled QK and softmax probabilities + if USE_EXP2: + p = tl.math.exp2(q_shifted * RCP_LN2) + else: + p = tl.math.exp(q_shifted) + + # CAVEAT: Must update l_ij before applying dropout + l_ij = tl.sum(p, 1) + if ENABLE_DROPOUT: + philox_offset = batch_philox_offset + start_m * BLOCK_M * actual_seqlen_k + start_n - BLOCK_N + keep = dropout_mask(philox_seed, philox_offset, dropout_p, BLOCK_M, BLOCK_N, actual_seqlen_k) + if RETURN_SCORES: + exp_score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ( + (start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k) + tl.store(exp_scores_ptrs, tl.where(keep, p, -p), mask=exp_score_mask) + p = tl.where(keep, p, 0.0) + elif RETURN_SCORES: + exp_score_mask = (OFFS_M[:, None] < actual_seqlen_q) & ( + (start_n + tl.arange(0, BLOCK_N))[None, :] < actual_seqlen_k) + tl.store(exp_scores_ptrs, p, mask=exp_score_mask) + + # -- update output accumulator -- + # alpha is an adjustment factor for acc and li as we loop and find new maxes + m_diff = m_i - m_ij + if USE_EXP2: + alpha = tl.math.exp2(m_diff * RCP_LN2) + else: + alpha = tl.math.exp(m_diff) + acc = acc * alpha[:, None] + if not PRE_LOAD_V: + v = load_fn(v_ptrs, k_offs_n, v_offs_k, actual_seqlen_k, ACTUAL_BLOCK_DMODEL_V) + vs = load_fn(vs_ptrs, vs_offs_k, vs_offs_n, ACTUAL_BLOCK_DMODEL_V, actual_seqlen_k_scale, other=1) + # -- update m_i and l_i + l_i = l_i * alpha + l_ij + # update m_i and l_i + m_i = m_ij + + ps = _calculate_scales( + p, + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + target_max_pow2=E4M3_TARGET_MAX_POW2, + mbits=E4M3_MBITS, + IS_2D_BLOCK=False, + ) + p_fp8 = _quantize_fp8( + p, + ps, + philox_seed, + batch_philox_offset, + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + FP8_FORMAT=E4M3_FORMAT_ID, + IS_2D_BLOCK=False, + USE_ASM=USE_ASM, + USE_SR=False, + ) + + acc += tl.dot_scaled(p_fp8, ps, "e4m3", v, vs, "e4m3", out_dtype=tl.float32) + + k_ptrs += BLOCK_N * stride_kn + v_ptrs += BLOCK_N * stride_vk + ks_ptrs += BLOCK_N_SCALE * stride_kscale_n + vs_ptrs += BLOCK_N_SCALE * stride_vscale_k + if bias_ptrs is not None: + bias_ptrs += BLOCK_N * stride_bn + if RETURN_SCORES: + score_ptrs += BLOCK_N + scores_scaled_shifted_ptrs += BLOCK_N + exp_scores_ptrs += BLOCK_N + return acc, l_i, m_i + + +def get_autotune_fwd_configs(): + return [ + triton.Config( + { + "PRE_LOAD_V": False, + }, + num_stages=1, + num_warps=4, + ), + ], [ + "IS_CAUSAL", "dropout_p", "MAX_SEQLENS_Q", "MAX_SEQLENS_K", "ACTUAL_BLOCK_DMODEL_QK", "ACTUAL_BLOCK_DMODEL_V", + "VARLEN", "HQ", "HK" + ] + + +autotune_fwd_configs, autotune_fwd_keys = get_autotune_fwd_configs() + + +@triton.autotune( + configs=autotune_fwd_configs, + key=autotune_fwd_keys, +) +@triton.jit +def attn_fwd( + Q, + K, + V, + bias, + q_scale_ptr, + k_scale_ptr, + v_scale_ptr, + SM_SCALE: tl.constexpr, + LSE, + Out, + stride_qz, + stride_qh, + stride_qm, + stride_qk, + stride_kz, + stride_kh, + stride_kn, + stride_kk, + stride_vz, + stride_vh, + stride_vk, + stride_vn, + stride_oz, + stride_oh, + stride_om, + stride_on, + stride_bz, + stride_bh, + stride_bm, + stride_bn, + stride_az, + stride_ah, + stride_sz, + stride_sh, + stride_sm, + stride_sn, + stride_lse_z, + stride_lse_h, + stride_lse_m, + stride_qscale_z, + stride_qscale_h, + stride_qscale_m, + stride_qscale_k, + stride_kscale_z, + stride_kscale_h, + stride_kscale_n, + stride_kscale_k, + stride_vscale_z, + stride_vscale_h, + stride_vscale_k, + stride_vscale_n, + cu_seqlens_q, + cu_seqlens_k, + dropout_p, + philox_seed, + philox_offset_base, + scores, + scores_scaled_shifted, + exp_scores, + alibi_slopes, + HQ: tl.constexpr, + HK: tl.constexpr, + ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + ACTUAL_BLOCK_DMODEL_V: tl.constexpr, + MAX_SEQLENS_Q: tl.constexpr, + MAX_SEQLENS_K: tl.constexpr, + VARLEN: tl.constexpr, + IS_CAUSAL: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + BLOCK_N: tl.constexpr, + PRE_LOAD_V: tl.constexpr, + USE_BIAS: tl.constexpr, + ENABLE_DROPOUT: tl.constexpr, + RETURN_SCORES: tl.constexpr, + USE_ALIBI: tl.constexpr, + USE_EXP2: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, +): + start_m = tl.program_id(0) + off_h_q = tl.program_id(1) + off_z = tl.program_id(2) + offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = tl.arange(0, BLOCK_N) + + SCALE_BLOCK_DMODEL_QK: tl.constexpr = BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr = ACTUAL_BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + BLOCK_N_SCALE: tl.constexpr = BLOCK_N // QUANT_BLOCK_SIZE + offs_d_qk = tl.arange(0, BLOCK_DMODEL_QK) + offs_d_qk_scale = tl.arange(0, SCALE_BLOCK_DMODEL_QK) + offs_d_v = tl.arange(0, BLOCK_DMODEL_V) + offs_n_scale = tl.arange(0, BLOCK_N_SCALE) + + # If MQA / GQA, set the K and V head offsets appropriately. + GROUP_SIZE: tl.constexpr = HQ // HK + if GROUP_SIZE != 1: + off_h_k = off_h_q // GROUP_SIZE + else: + off_h_k = off_h_q + + PADDED_HEAD_QK: tl.constexpr = (ACTUAL_BLOCK_DMODEL_QK != BLOCK_DMODEL_QK) + PADDED_HEAD_V: tl.constexpr = (ACTUAL_BLOCK_DMODEL_V != BLOCK_DMODEL_V) + + if VARLEN: + cu_seqlens_q_start = tl.load(cu_seqlens_q + off_z) + cu_seqlens_q_end = tl.load(cu_seqlens_q + off_z + 1) + + seqlen_q = cu_seqlens_q_end - cu_seqlens_q_start + # We have a one-size-fits-all grid in id(0). Some seqlens might be too + # small for all start_m so for those we return early. + if start_m * BLOCK_M > seqlen_q: + return + cu_seqlens_k_start = tl.load(cu_seqlens_k + off_z) + cu_seqlens_k_end = tl.load(cu_seqlens_k + off_z + 1) + seqlen_k = cu_seqlens_k_end - cu_seqlens_k_start + else: + cu_seqlens_q_start = 0 + cu_seqlens_k_start = 0 + seqlen_q = MAX_SEQLENS_Q + seqlen_k = MAX_SEQLENS_K + + cu_seqlens_q_start_scale = cu_seqlens_q_start // QUANT_BLOCK_SIZE + cu_seqlens_k_start_scale = cu_seqlens_k_start // QUANT_BLOCK_SIZE + seqlen_k_scale = seqlen_k // QUANT_BLOCK_SIZE + + # Now we compute whether we need to exit early due to causal masking. + n_blocks = cdiv_fn(seqlen_k, BLOCK_N) + if IS_CAUSAL: + n_blocks_seqlen = cdiv_fn((start_m + 1) * BLOCK_M + seqlen_k - seqlen_q, BLOCK_N) + n_blocks = min(n_blocks, n_blocks_seqlen) + # If we have no blocks after adjusting for seqlen deltas, this WG is part of + # the blocks that are all 0. We exit early. + if n_blocks <= 0: + o_offset = Out + off_z * stride_oz + off_h_q * stride_oh + cu_seqlens_q_start * stride_om + o_ptrs = o_offset + offs_m[:, None] * stride_om + offs_d_v[None, :] * stride_on + acc = tl.zeros([BLOCK_M, BLOCK_DMODEL_V], dtype=Out.type.element_ty) + o_ptrs_mask = offs_m[:, None] < seqlen_q + if PADDED_HEAD_V: + o_ptrs_mask = o_ptrs_mask & (offs_d_v[None, :] < ACTUAL_BLOCK_DMODEL_V) + # We still need to write 0s to the result + tl.store(o_ptrs, acc, mask=o_ptrs_mask) + l_offset = LSE + off_z * stride_lse_z + off_h_q * stride_lse_h + cu_seqlens_q_start * stride_lse_m + l_ptrs = l_offset + offs_m * stride_lse_m + + l = tl.full([BLOCK_M], value=0.0, dtype=tl.float32) + + l_ptrs_mask = offs_m < MAX_SEQLENS_Q + tl.store(l_ptrs, l, mask=l_ptrs_mask) + return + + n_extra_tokens = 0 + if seqlen_k < BLOCK_N: + n_extra_tokens = BLOCK_N - seqlen_k + elif seqlen_k % BLOCK_N: + n_extra_tokens = seqlen_k % BLOCK_N + + # Compute pointers for all the tensors used in this kernel. + q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + cu_seqlens_q_start * stride_qm + q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk + k_offset = K + off_z * stride_kz + off_h_k * stride_kh + cu_seqlens_k_start * stride_kn + k_ptrs = k_offset + offs_d_qk[:, None] * stride_kk + offs_n[None, :] * stride_kn + v_offset = V + off_z * stride_vz + off_h_k * stride_vh + cu_seqlens_k_start * stride_vk + v_ptrs = v_offset + offs_n[:, None] * stride_vk + offs_d_v[None, :] * stride_vn + qs_offset = (q_scale_ptr + off_z * stride_qscale_z + off_h_q * stride_qscale_h + + cu_seqlens_q_start_scale * stride_qscale_m) + qs_ptrs = (qs_offset + (offs_m[:, None] // QUANT_BLOCK_SIZE) * stride_qscale_m + + offs_d_qk_scale[None, :] * stride_qscale_k) + ks_offset = (k_scale_ptr + off_z * stride_kscale_z + off_h_k * stride_kscale_h + + cu_seqlens_k_start_scale * stride_kscale_n) + # k scale is N*K even though k is K*N, this is required by tl.dot_scaled + ks_ptrs = (ks_offset + (offs_n[:, None] // QUANT_BLOCK_SIZE) * stride_kscale_n + + offs_d_qk_scale[None, :] * stride_kscale_k) + vs_offset = (v_scale_ptr + off_z * stride_vscale_z + off_h_k * stride_vscale_h + + cu_seqlens_k_start_scale * stride_vscale_k) + vs_ptrs = (vs_offset + (offs_d_v[:, None] // QUANT_BLOCK_SIZE) * stride_vscale_n + + offs_n_scale[None, :] * stride_vscale_k) + if USE_BIAS: + # Note: this might get large enough to overflow on some configs + bias_offset = off_h_q * stride_bh + bias_ptrs = bias + bias_offset + offs_m[:, None] * stride_bm + offs_n[None, :] * stride_bn + else: + bias_ptrs = None + + if USE_ALIBI: + a_offset = off_z * stride_az + off_h_q * stride_ah + alibi_slope = tl.load(alibi_slopes + a_offset) + else: + alibi_slope = None + + if RETURN_SCORES: + scores_offset = scores + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm + score_ptrs = scores_offset + offs_m[:, None] * stride_sm + offs_n[None, :] * stride_sn + + scores_scaled_shifted_offset = scores_scaled_shifted + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm + scores_scaled_shifted_ptrs = scores_scaled_shifted_offset + offs_m[:, None] * stride_sm + offs_n[ + None, :] * stride_sn + + exp_scores_offset = exp_scores + off_z * stride_sz + off_h_q * stride_sh + cu_seqlens_q_start * stride_sm + exp_scores_ptrs = exp_scores_offset + offs_m[:, None] * stride_sm + offs_n[None, :] * stride_sn + else: + score_ptrs = None + scores_scaled_shifted_ptrs = None + exp_scores_ptrs = None + + if ENABLE_DROPOUT: + off_hz = off_z * HQ + off_h_q + batch_philox_offset = philox_offset_base + off_hz * seqlen_q * seqlen_k + else: + batch_philox_offset = 0 + # initialize pointer to m and l + m_i = tl.full([BLOCK_M], float("-inf"), dtype=tl.float32) + l_i = tl.full([BLOCK_M], 1.0, dtype=tl.float32) + acc = tl.zeros([BLOCK_M, BLOCK_DMODEL_V], dtype=tl.float32) + # Q is loaded once at the beginning and shared by all N blocks. + q_ptrs_mask = offs_m[:, None] < seqlen_q + qs_ptrs_mask = q_ptrs_mask + if PADDED_HEAD_QK: + q_ptrs_mask = q_ptrs_mask & (offs_d_qk[None, :] < ACTUAL_BLOCK_DMODEL_QK) + qs_ptrs_mask = qs_ptrs_mask & (offs_d_qk_scale[None, :] < SCALE_ACTUAL_BLOCK_DMODEL_QK) + + q = tl.load(q_ptrs, mask=q_ptrs_mask, other=0) + qs = tl.load(qs_ptrs, mask=qs_ptrs_mask, other=1) + + # Here we compute how many full and masked blocks we have. + padded_block_k = n_extra_tokens != 0 + is_modulo_mn = not padded_block_k and (seqlen_q % BLOCK_M == 0) + if IS_CAUSAL: + # There are always at least BLOCK_M // BLOCK_N masked blocks. + masked_blocks = BLOCK_M // BLOCK_N + (not is_modulo_mn) + else: + # Padding on Q does not need to be masked in the FA loop. + masked_blocks = padded_block_k + + masked_blocks = min(masked_blocks, n_blocks) + n_full_blocks = n_blocks - masked_blocks + block_min = 0 + block_max = n_blocks * BLOCK_N + # Compute for full blocks. Here we set causal to false regardless of its actual + # value because there is no masking. Similarly we do not need padding. + + if n_full_blocks > 0: + block_max = (n_blocks - masked_blocks) * BLOCK_N + acc, l_i, m_i = _attn_fwd_inner( + acc, + l_i, + m_i, + q, + qs, + k_ptrs, + v_ptrs, + ks_ptrs, + vs_ptrs, + bias_ptrs, + stride_kn, + stride_vk, + stride_bn, + stride_kscale_n, + stride_vscale_k, + start_m, + seqlen_k, + seqlen_k_scale, + seqlen_q, + dropout_p, + philox_seed, + batch_philox_offset, + exp_scores_ptrs, + # _, _, offs_n_causal, masked_blocks, n_extra_tokens, _ + block_min, + block_max, + 0, + 0, + 0, + alibi_slope, + score_ptrs, + scores_scaled_shifted_ptrs, + # IS_CAUSAL, .... + False, + BLOCK_M, + BLOCK_DMODEL_QK, + SCALE_BLOCK_DMODEL_QK, + BLOCK_DMODEL_V, + BLOCK_N, + BLOCK_N_SCALE, + offs_m, + offs_n, + # _, MASK_STEPS, ... + PRE_LOAD_V, + False, + ENABLE_DROPOUT, + PADDED_HEAD_QK, + PADDED_HEAD_V, + ACTUAL_BLOCK_DMODEL_QK, + SCALE_ACTUAL_BLOCK_DMODEL_QK, + ACTUAL_BLOCK_DMODEL_V, + SM_SCALE, + USE_EXP2=USE_EXP2, + RETURN_SCORES=RETURN_SCORES, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + USE_ASM=USE_ASM, + ) + block_min = block_max + block_max = n_blocks * BLOCK_N + + # Remaining blocks, if any, are full / not masked. + if (masked_blocks > 0): + if IS_CAUSAL: + offs_n_causal = offs_n + (seqlen_q - seqlen_k) + else: + offs_n_causal = 0 + k_ptrs += n_full_blocks * BLOCK_N * stride_kn + v_ptrs += n_full_blocks * BLOCK_N * stride_vk + ks_ptrs += n_full_blocks * BLOCK_N_SCALE * stride_kscale_n + vs_ptrs += n_full_blocks * BLOCK_N_SCALE * stride_vscale_k + if USE_BIAS: + bias_ptrs += n_full_blocks * BLOCK_N * stride_bn + if RETURN_SCORES: + score_ptrs += n_full_blocks * BLOCK_N + scores_scaled_shifted_ptrs += n_full_blocks * BLOCK_N + exp_scores_ptrs += n_full_blocks * BLOCK_N + + acc, l_i, m_i = _attn_fwd_inner( + acc, + l_i, + m_i, + q, + qs, + k_ptrs, + v_ptrs, + ks_ptrs, + vs_ptrs, + bias_ptrs, + stride_kn, + stride_vk, + stride_bn, + stride_kscale_n, + stride_vscale_k, + start_m, + seqlen_k, + seqlen_k_scale, + seqlen_q, + dropout_p, + philox_seed, + batch_philox_offset, + exp_scores_ptrs, + block_min, + block_max, + offs_n_causal, + masked_blocks, + n_extra_tokens, + alibi_slope, + score_ptrs, + scores_scaled_shifted_ptrs, + IS_CAUSAL, + BLOCK_M, + BLOCK_DMODEL_QK, + SCALE_BLOCK_DMODEL_QK, + BLOCK_DMODEL_V, + BLOCK_N, + BLOCK_N_SCALE, + offs_m, + offs_n, + # _, MASK_STEPS, ... + PRE_LOAD_V, + True, + ENABLE_DROPOUT, + PADDED_HEAD_QK, + PADDED_HEAD_V, + ACTUAL_BLOCK_DMODEL_QK, + SCALE_ACTUAL_BLOCK_DMODEL_QK, + ACTUAL_BLOCK_DMODEL_V, + SM_SCALE, + USE_EXP2=USE_EXP2, + RETURN_SCORES=RETURN_SCORES, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + USE_ASM=USE_ASM, + ) + + # epilogue + l_recip = 1 / l_i[:, None] + acc = acc * l_recip + if ENABLE_DROPOUT: + acc = acc / (1 - dropout_p) + # If seqlen_q > seqlen_k but the delta is not a multiple of BLOCK_M, + # then we have one block with a row of all NaNs which come from computing + # softmax over a row of all -infs (-inf - inf = NaN). We check for that here + # and store 0s where there are NaNs as these rows should've been zeroed out. + end_m_idx = (start_m + 1) * BLOCK_M + start_m_idx = start_m * BLOCK_M + causal_start_idx = seqlen_q - seqlen_k + + acc = acc.to(Out.type.element_ty) + if IS_CAUSAL: + if causal_start_idx > start_m_idx and causal_start_idx < end_m_idx: + out_mask_boundary = tl.full((BLOCK_DMODEL_V,), causal_start_idx, dtype=tl.int32) + mask_m_offsets = start_m_idx + tl.arange(0, BLOCK_M) + out_ptrs_mask = mask_m_offsets[:, None] >= out_mask_boundary[None, :] + z = 0.0 + acc = tl.where(out_ptrs_mask, acc, z.to(acc.dtype)) + + # write back LSE(Log Sum Exponents), the log of the normalization constant + l_offset = LSE + off_z * stride_lse_z + off_h_q * stride_lse_h + cu_seqlens_q_start * stride_lse_m + offs_l_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + l_ptrs = l_offset + offs_l_m * stride_lse_m + if USE_EXP2: + RCP_LN2: tl.constexpr = 1.4426950408889634 + LN2: tl.constexpr = 0.6931471824645996 + # compute log-sum-exp in base 2 units + mi_base2 = m_i * RCP_LN2 + softmax_lse = mi_base2 + tl.math.log2(l_i) + # convert back to natural units + softmax_lse *= LN2 + else: + softmax_lse = m_i + tl.math.log(l_i) + + if IS_CAUSAL: + # zero out nans caused by -infs when doing causal + lse_mask = (start_m_idx + tl.arange(0, BLOCK_M)) < causal_start_idx + softmax_lse = tl.where(lse_mask, 0.0, softmax_lse) + + # If seqlen_q not multiple of BLOCK_M, we need to mask out the last few rows. + # This is only true for the last M block. For others, overflow_size will be -ve + overflow_size = end_m_idx - seqlen_q + if overflow_size > 0: + boundary = tl.full((BLOCK_M,), BLOCK_M - overflow_size, dtype=tl.int32) + l_ptrs_mask = tl.arange(0, BLOCK_M) < boundary + tl.store(l_ptrs, softmax_lse, mask=l_ptrs_mask) # the log of the normalization constant + else: + tl.store(l_ptrs, softmax_lse) # the log of the normalization constant + + # write back O + o_offset = Out + off_z * stride_oz + off_h_q * stride_oh + cu_seqlens_q_start * stride_om + o_ptrs = o_offset + offs_m[:, None] * stride_om + offs_d_v[None, :] * stride_on + o_ptrs_mask = tl.full([BLOCK_M, BLOCK_DMODEL_V], 1, dtype=tl.int1) + if overflow_size > 0: + o_ptrs_mask = o_ptrs_mask & (offs_m[:, None] < seqlen_q) + if PADDED_HEAD_V: + o_ptrs_mask = o_ptrs_mask & (offs_d_v[None, :] < ACTUAL_BLOCK_DMODEL_V) + tl.store(o_ptrs, acc.to(Out.type.element_ty), mask=o_ptrs_mask) + + +def get_padded_head_dim(head_size: int): + # Get closest power of 2 over or equal to 32. + padded_d_model = 1 << (head_size - 1).bit_length() + # Smallest head_dim supported is 16. If smaller, the tile in the + # kernel is padded - there is no padding in memory for any dims. + padded_d_model = max(padded_d_model, 16) + return padded_d_model + + +@triton_op("alto::attention_mxfp8_forward_triton_impl", mutates_args=()) +def attention_mxfp8_forward_triton_impl( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: torch.Tensor, + k_scale: torch.Tensor, + v_scale: torch.Tensor, + sm_scale: float, + alibi_slopes: Optional[torch.Tensor], + causal: bool, + bias: Optional[torch.Tensor], + dropout_p: float, + layout: str, + cu_seqlens_q: Optional[int], + cu_seqlens_k: Optional[int], + max_seqlens_q: Optional[int], + max_seqlens_k: Optional[int], + return_scores: bool, + use_exp2: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + if DEBUG: + print() + print("attention_mxfp8_forward_triton_impl") + print("q:", q, q.shape) + print("k:", k, k.shape) + print("v:", v, v.shape) + print("sm_scale:", sm_scale) + print("alibi_slopes:", alibi_slopes) + print("causal:", causal) + print("bias:", bias) + print("dropout_p:", dropout_p) + print("layout:", layout) + print("cu_seqlens_q:", cu_seqlens_q) + print("cu_seqlens_k:", cu_seqlens_k) + print("max_seqlens_q:", max_seqlens_q) + print("max_seqlens_k:", max_seqlens_k) + print("return_scores:", return_scores) + print("use_exp2:", use_exp2) + + assert q.is_contiguous() + assert k.is_contiguous() + assert v.is_contiguous() + assert q_scale.is_contiguous() + assert k_scale.is_contiguous() + assert v_scale.is_contiguous() + + # check if varlen + is_varlen = layout == "thd" + + # NOTE: a large bias tensor leads to overflow during pointer arithmetic + if (bias is not None): + assert (bias.numel() < 2**31) + + batch, nheads_q, nheads_k, head_size_qk, head_size_v, seqlen_q, seqlen_k = get_shape_from_layout( + q, k, v, layout, cu_seqlens_q, cu_seqlens_k, max_seqlens_q, max_seqlens_k) + o_shape = (*q.shape[:-1], head_size_v) + o = torch.empty( + o_shape, + device=q.device, + dtype=fwd_torch_dtype, + requires_grad=True, + ) + + q_strides = get_strides_from_layout(q, layout) + k_strides = get_strides_from_layout(k, layout) + v_strides = get_strides_from_layout(v, layout) + o_strides = get_strides_from_layout(o, layout) + + # Get closest power of 2 over or equal to 32. + padded_d_model_qk = get_padded_head_dim(head_size_qk) + padded_d_model_v = get_padded_head_dim(head_size_v) + + grid = lambda META: (triton.cdiv(max_seqlens_q, META["BLOCK_M"]), nheads_q, batch) + + if return_scores: + scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device, dtype=torch.float32) + scores_scaled_shifted = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), + device=q.device, + dtype=torch.float32) + scores_strides = (scores.stride(0), scores.stride(1), scores.stride(2), scores.stride(3)) + else: + scores = torch.empty([], device=q.device, dtype=torch.float32) + scores_scaled_shifted = None + scores_strides = (0, 0, 0, 0) + + # exp_scores is used to validate dropout behavior vs the PyTorch SDPA math backend reference. + if return_scores: + exp_scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device, dtype=torch.float32) + else: + exp_scores = torch.empty([], device=q.device, dtype=torch.float32) + + # stores LSE the log of the normalization constant / sum of exponential score (unnormalized probabilities) + if is_varlen: + softmax_lse = torch.empty((q.shape[0], nheads_q), device=q.device, dtype=torch.float32) + stride_lse_m, stride_lse_h = softmax_lse.stride() + stride_lse_z = 0 + else: + softmax_lse = torch.empty((batch, nheads_q, max_seqlens_q), device=q.device, dtype=torch.float32) + stride_lse_z, stride_lse_h, stride_lse_m = softmax_lse.stride() + + if bias is not None: + bias_strides = (bias.stride(0), bias.stride(1), bias.stride(2), bias.stride(3)) + else: + bias_strides = (0, 0, 0, 0) + + if alibi_slopes is not None: + alibi_strides = (alibi_slopes.stride(0), alibi_slopes.stride(1)) + else: + alibi_strides = (0, 0) + + qs_strides = get_strides_from_layout(q_scale, layout) + ks_strides = get_strides_from_layout(k_scale, layout) + vs_strides = get_strides_from_layout(v_scale, layout) + + wrap_triton(attn_fwd)[grid]( + q, + k, + v, + bias, + q_scale, + k_scale, + v_scale, + sm_scale, + softmax_lse, + o, + *q_strides, + *k_strides, + *v_strides, + *o_strides, + *bias_strides, + *alibi_strides, + *scores_strides, + stride_lse_z, + stride_lse_h, + stride_lse_m, + *qs_strides, + *ks_strides, + *vs_strides, + cu_seqlens_q, + cu_seqlens_k, + dropout_p=dropout_p, + philox_seed=philox_seed, + philox_offset_base=philox_offset, + scores=scores, + scores_scaled_shifted=scores_scaled_shifted, + exp_scores=exp_scores, + alibi_slopes=alibi_slopes, + HQ=nheads_q, + HK=nheads_k, + ACTUAL_BLOCK_DMODEL_QK=head_size_qk, + ACTUAL_BLOCK_DMODEL_V=head_size_v, + MAX_SEQLENS_Q=max_seqlens_q, + MAX_SEQLENS_K=max_seqlens_k, + IS_CAUSAL=causal, + VARLEN=is_varlen, + BLOCK_DMODEL_QK=padded_d_model_qk, + BLOCK_DMODEL_V=padded_d_model_v, + USE_BIAS=False if bias is None else True, + USE_ALIBI=False if alibi_slopes is None else True, + ENABLE_DROPOUT=dropout_p > 0.0, + USE_EXP2=use_exp2, + RETURN_SCORES=return_scores, + BLOCK_M=64, + BLOCK_N=64, + QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT, + USE_ASM=is_cdna4(), + ) + return o, softmax_lse, exp_scores + + +@torch.compiler.allow_in_graph +class _triton_attention_mxfp8(torch.autograd.Function): + + @staticmethod + def forward( + ctx, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + alibi_slopes: torch.Tensor | None, + bias: torch.Tensor | None, + sm_scale: float, + dropout_p: float, + cu_seqlens_q: int, + cu_seqlens_k: int, + max_seqlens_q: int, + max_seqlens_k: int, + causal: bool, + return_scores: bool, + use_exp2: bool, + layout: str, + ): + q, q_scale = torch.ops.alto.convert_to_mxfp8(q, mxfp_format="e4m3", axis=-1, is_2d_block=True) + k, k_scale = torch.ops.alto.convert_to_mxfp8(k, mxfp_format="e4m3", axis=-1, is_2d_block=True) + v, v_scale = torch.ops.alto.convert_to_mxfp8(v, mxfp_format="e4m3", axis=-1, is_2d_block=True) + + output, softmax_lse, exp_scores = torch.ops.alto.attention_mxfp8_forward_triton_impl( + q, + k, + v, + q_scale, + k_scale, + v_scale, + sm_scale=sm_scale, + alibi_slopes=alibi_slopes, + causal=causal, + bias=bias, + dropout_p=dropout_p, + layout=layout, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlens_q=max_seqlens_q, + max_seqlens_k=max_seqlens_k, + return_scores=return_scores, + use_exp2=use_exp2, + ) + + ctx.save_for_backward(q, k, v, output, softmax_lse, alibi_slopes, bias, q_scale, k_scale, v_scale) + ctx.sm_scale = sm_scale + ctx.causal = causal + ctx.dropout_p = dropout_p + ctx.layout = layout + ctx.use_exp2 = use_exp2 + ctx.cu_seqlens_q = cu_seqlens_q + ctx.cu_seqlens_k = cu_seqlens_k + ctx.max_seqlens_q = max_seqlens_q + ctx.max_seqlens_k = max_seqlens_k + + return output, softmax_lse, exp_scores + + @staticmethod + def backward(ctx, *grad_outputs): + # Forward-only, matching the MXFP4 reference. Backward is not implemented. + return None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None + + +@attention_mxfp8_forward_triton_impl.register_fake +def fake_attention_mxfp8_forward_triton_impl( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: torch.Tensor, + k_scale: torch.Tensor, + v_scale: torch.Tensor, + sm_scale: float, + alibi_slopes: Optional[torch.Tensor], + causal: bool, + bias: Optional[torch.Tensor], + dropout_p: float, + layout: str, + cu_seqlens_q: Optional[int], + cu_seqlens_k: Optional[int], + max_seqlens_q: Optional[int], + max_seqlens_k: Optional[int], + return_scores: bool, + use_exp2: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + o_shape = list(q.shape) + o_shape[-1] = v.shape[-1] # output shape should match v's head dim + o = torch.empty( + o_shape, + device=q.device, + dtype=fwd_torch_dtype, + requires_grad=True, + ) + + # check if varlen + is_varlen = layout == "thd" + + batch, nheads_q, nheads_k, head_size_qk, head_size_v, seqlen_q, seqlen_k = get_shape_from_layout( + q, k, v, layout, cu_seqlens_q, cu_seqlens_k, max_seqlens_q, max_seqlens_k) + + if return_scores: + scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device, dtype=torch.float32) + else: + scores = torch.empty([], device=q.device, dtype=torch.float32) + + if return_scores: + exp_scores = torch.zeros((batch, nheads_q, max_seqlens_q, max_seqlens_k), device=q.device, dtype=torch.float32) + else: + exp_scores = torch.empty([], device=q.device, dtype=torch.float32) + + if is_varlen: + softmax_lse = torch.empty((q.shape[0], nheads_q), device=q.device, dtype=torch.float32) + else: + softmax_lse = torch.empty((batch, nheads_q, max_seqlens_q), device=q.device, dtype=torch.float32) + return o, softmax_lse, exp_scores + + +def triton_attention_mxfp8( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + alibi_slopes: torch.Tensor | None, + bias: torch.Tensor | None, + sm_scale: float, + dropout_p: float, + cu_seqlens_q: int, + cu_seqlens_k: int, + max_seqlens_q: int, + max_seqlens_k: int, + causal: bool, + return_scores: bool, + use_exp2: bool, + layout: str, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + return _triton_attention_mxfp8.apply( + q, + k, + v, + alibi_slopes, + bias, + sm_scale, + dropout_p, + cu_seqlens_q, + cu_seqlens_k, + max_seqlens_q, + max_seqlens_k, + causal, + return_scores, + use_exp2, + layout, + ) diff --git a/tests/unittest/mxfp8/test_mxfp8_attention.py b/tests/unittest/mxfp8/test_mxfp8_attention.py new file mode 100644 index 00000000..35abdcfd --- /dev/null +++ b/tests/unittest/mxfp8/test_mxfp8_attention.py @@ -0,0 +1,123 @@ +# Copyright (c) 2026 Advanced Micro Devices, Inc. +# +# SPDX-License-Identifier: MIT +"""Forward numerical tests for the MXFP8 (e4m3) flash attention kernel. + +Mirrors ``tests/unittest/mxfp4/test_mxfp_attention.py``. The MXFP8 kernel is +forward-only (backward is not implemented), so only the forward output is +validated against a bf16 SDPA reference via SNR / cosine-similarity. +""" + +import pytest +from tabulate import tabulate +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA device is required.") + +from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 +from .utils import calc_snr, calc_cossim + + +def attention_vanilla_forward_pytorch_ref_impl(q, k, v, sm_scale, causal, layout="bshd"): + """Compute reference output using PyTorch's built-in SDPA.""" + if layout == "bshd": + num_heads = q.shape[2] + n_kv_heads = k.shape[2] + n_rep = num_heads // n_kv_heads + + q = q.transpose(1, 2).contiguous() + k = k.transpose(1, 2).contiguous() + v = v.transpose(1, 2).contiguous() + else: + raise ValueError(f"Unknown layout {layout}") + + o_ref = torch.nn.functional.scaled_dot_product_attention(q, + k, + v, + is_causal=causal, + scale=sm_scale, + enable_gqa=n_rep > 1) + if layout == "bshd": + o_ref = o_ref.transpose(1, 2) + return o_ref + + +class AttnConfig: + + def __init__(self, seqlen_q, seqlen_kv, num_head_q, num_head_kv, head_dim_qk, head_dim_v): + self.seqlen_q = seqlen_q + self.seqlen_kv = seqlen_kv + self.num_head_q = num_head_q + self.num_head_kv = num_head_kv + self.head_dim_qk = head_dim_qk + self.head_dim_v = head_dim_v + + +test_cases = [ + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=32, num_head_kv=32, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=64, num_head_kv=8, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=32, num_head_kv=8, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=16, num_head_kv=16, head_dim_qk=192, head_dim_v=128), + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=128, num_head_kv=128, head_dim_qk=192, head_dim_v=128), + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=48, num_head_kv=8, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=2048, seqlen_kv=2048, num_head_q=64, num_head_kv=8, head_dim_qk=128, head_dim_v=128), +] + + +@pytest.mark.parametrize("batch", [4]) +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True]) +def test_attention(batch, config, causal): + device = "cuda" + dtype = torch.bfloat16 + seqlen_q, seqlen_kv, num_head_q, num_head_kv, head_dim_qk, head_dim_v = ( + config.seqlen_q, + config.seqlen_kv, + config.num_head_q, + config.num_head_kv, + config.head_dim_qk, + config.head_dim_v, + ) + q_layout = (batch, seqlen_q, num_head_q, head_dim_qk) + k_layout = (batch, seqlen_kv, num_head_kv, head_dim_qk) + v_layout = (batch, seqlen_kv, num_head_kv, head_dim_v) + + torch.manual_seed(1234) + + query = torch.randn(q_layout, device=device, dtype=dtype, requires_grad=True) + key = torch.randn(k_layout, device=device, dtype=dtype, requires_grad=True) + value = torch.randn(v_layout, device=device, dtype=dtype, requires_grad=True) + query_ref = query.clone().detach().requires_grad_() + key_ref = key.clone().detach().requires_grad_() + value_ref = value.clone().detach().requires_grad_() + + sm_scale = query.shape[-1]**(-0.5) + o_ref = attention_vanilla_forward_pytorch_ref_impl(query_ref, key_ref, value_ref, sm_scale, causal) + + o = triton_attention_mxfp8( + query.transpose(1, 2).contiguous(), + key.transpose(1, 2).contiguous(), + value.transpose(1, 2).contiguous(), + bias=None, + alibi_slopes=None, + sm_scale=sm_scale, + dropout_p=0.0, + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlens_q=seqlen_q, + max_seqlens_k=seqlen_kv, + causal=causal, + return_scores=False, + use_exp2=True, + layout="bhsd", + )[0].transpose(1, 2) + + output_snr = calc_snr(o_ref, o) + output_sim = calc_cossim(o_ref, o) + print() + print(tabulate([ + ["O", output_snr, output_sim], + ], headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + assert output_sim > 0.99, f"output cosine-sim too low: {output_sim}" + assert output_snr > 20, f"output SNR too low: {output_snr}" From 823eba8cdd9a2aa47c255374fc58939b17c0150c Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Tue, 14 Jul 2026 16:45:12 +0800 Subject: [PATCH 02/14] test: add pure-PyTorch golden reference for mxfp8 flash-attention --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 35 ++++- .../mxfp8/test_mxfp8_attention_reference.py | 133 ++++++++++++++++++ tests/unittest/mxfp8/utils.py | 114 +++++++++++++++ 3 files changed, 280 insertions(+), 2 deletions(-) create mode 100644 tests/unittest/mxfp8/test_mxfp8_attention_reference.py diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 769a5e97..a18ed6fb 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -272,6 +272,37 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` --- -## 10. 一句话总结 +## 10. 实施进展记录 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(backward 格式 ⚠️ 待确认,e5m2 参数化预留)、**仅 CDNA4 无 fallback**,ground-truth 由 PyTorch 模拟 reference 提供。量化基础设施已齐,forward FA v2 骨架复用;最大工作量与风险集中在从零的 backward。 +### 10.1 截至 2026-07-14 + +**已落地(代码,未经硬件验证):** + +- **forward kernel**:`alto/kernels/mxfp8/triton_flash_attention_mxfp8.py`——由 mxfp4 attention 机械改写(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`、Q/K/V 走 `convert_to_mxfp8`)。**forward-only**(backward 为 `None` 占位)、**仅 CDNA4 走 `tl.dot_scaled`,无 fallback**。`__init__.py` 导出 `triton_attention_mxfp8`。 +- **黄金参照**:`tests/unittest/mxfp8/utils.py` 的 `mxfp8_attention_forward_reference`——纯 PyTorch 复刻 forward(2D-block e4m3 Q/K/V、逐 key-block 在线 softmax + running-max 量化 P、全 fp32 matmul、无 `tl.dot_scaled`),可在 CPU / 任意设备跑。 +- **测试(三层,见 §8)**: + - 第 3 层 `test_mxfp8_attention.py::test_attention`(kernel vs bf16 SDPA,仿 mxfp4,硬断言 `cossim>0.99`/`SNR>20`)——**已 commit,作为基线不再改动**。 + - 第 1 层 `test_mxfp8_attention_reference.py::test_reference_matches_bf16_sdpa`(黄金参照 vs bf16,**CPU 现在可跑**)+ 第 2 层 `test_kernel_matches_reference`(kernel vs 黄金参照,CDNA4;`import` 复用第 3 层的 `test_cases` 网格)——新增独立文件,纯增量。 + +**对 §4 步骤的进度映射:** + +| 步骤 | 状态 | 备注 | +|---|---|---| +| Step 1–4(forward kernel + wrapper + autograd.forward) | ✅ 代码完成 / ⏳ 未验证 | 无 CDNA4,尚未跑通任何 kernel 测试 | +| Step 5–8(backward:preprocess / dK-dV / dQ / autograd.backward) | ❌ 未开始 | 从零写,头号工作量,见 §9 | +| Step 9(dispatch 接入 `mxfp8_e4m3`) | ❌ 未开始 | — | +| Step 10(测试脚手架) | ✅ 三层脚手架就位 | 第 1 层可离线验,第 2/3 层待 CDNA4 | + +**已知偏差 / 修正记录:** + +- 写黄金参照时静态发现并修掉一处 masking bug:mask 值若用 `finfo.min`(有限)会让"整块被 mask 的行"算出 `exp(0)=1`(应为 0);改用真 `-inf`(masked 项 `exp(-inf)=0`,全 mask 行的 nan 由 `nan_to_num` 兜住)。此 bug 正是"无黄金参照、只对 bf16"抓不到的类型。 + +**Open items(阻塞真验证):** + +- **CDNA4 硬件**:M350/M355 卡预计约一周后到;在此之前所有 kernel 测试(第 2/3 层)无法运行,forward 代码处于"已写未验"状态。 +- **第 1 层 CPU 自证**:需在有 torch 的环境跑一次 `test_reference_matches_bf16_sdpa`,确认黄金参照本身正确(它是 backward 验证的地基)。 +- **backward 全 e4m3 的数值风险**:见 §0 风险记录 / §5,真机验证后据实决定是否需切 e5m2。 + +### 10.2 一句话总结 + +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,ground-truth 由纯 PyTorch 黄金参照提供(CPU 可离线验)。量化基础设施已齐,forward FA v2 骨架复用;最大工作量与风险集中在从零的 backward。当前 forward + 三层测试脚手架已落地但未经硬件验证,backward 未开始。 diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py new file mode 100644 index 00000000..f5c0fc44 --- /dev/null +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -0,0 +1,133 @@ +# Copyright (c) 2026 Advanced Micro Devices, Inc. +# +# SPDX-License-Identifier: MIT +"""Golden-reference validation for the MXFP8 (e4m3) flash attention kernel. + +Additive to ``test_mxfp8_attention.py`` (which mirrors the mxfp4 attention test: +kernel vs bf16 SDPA). This file adds the two layers that the pure-PyTorch golden +reference enables — a strengthening over the mxfp4 reference, which has neither a +golden reference nor asserts: + + * Layer 1 ``test_reference_matches_bf16_sdpa`` — golden reference vs bf16 SDPA. + Pure PyTorch, **runs on CPU / any device without CDNA4**; validates the + algorithm and mxfp8 quantization placement *before* the hardware arrives. + * Layer 2 ``test_kernel_matches_reference`` — kernel vs golden reference. + Isolates Triton port bugs (masking / online-softmax / LSE) from the mxfp8 + quantization error itself. Requires CDNA4 (native ``tl.dot_scaled``). + +(Layer 3 — kernel vs bf16 SDPA, total end-to-end error — is the committed +``test_mxfp8_attention.py::test_attention``; not duplicated here.) +""" + +import pytest +from tabulate import tabulate +import torch + +from .utils import calc_snr, calc_cossim, mxfp8_attention_forward_reference +# Reuse the committed shape grid so Layer 2 stays in lock-step with Layer 3. +from .test_mxfp8_attention import AttnConfig, test_cases + +cuda_only = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA/ROCm device is required.") + + +def _sdpa_bf16_bhsd(q, k, v, sm_scale, causal): + """bf16 SDPA reference in bhsd layout (the native SDPA head layout).""" + n_rep = q.shape[1] // k.shape[1] + return torch.nn.functional.scaled_dot_product_attention( + q, k, v, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) + + +def _make_qkv_bhsd(batch, config, device, dtype): + torch.manual_seed(1234) + q = torch.randn((batch, config.num_head_q, config.seqlen_q, config.head_dim_qk), device=device, dtype=dtype) + k = torch.randn((batch, config.num_head_kv, config.seqlen_kv, config.head_dim_qk), device=device, dtype=dtype) + v = torch.randn((batch, config.num_head_kv, config.seqlen_kv, config.head_dim_v), device=device, dtype=dtype) + return q, k, v + + +# --------------------------------------------------------------------------- +# Layer 1 — golden reference vs bf16 SDPA (pure PyTorch, runs anywhere, no CDNA4) +# --------------------------------------------------------------------------- + +# Small shapes: the reference loops over key blocks in Python, keep it cheap for CPU. +reference_cases = [ + AttnConfig(seqlen_q=128, seqlen_kv=128, num_head_q=4, num_head_kv=4, head_dim_qk=64, head_dim_v=64), + AttnConfig(seqlen_q=256, seqlen_kv=256, num_head_q=8, num_head_kv=2, head_dim_qk=128, head_dim_v=128), # GQA + AttnConfig(seqlen_q=128, seqlen_kv=256, num_head_q=4, num_head_kv=4, head_dim_qk=128, head_dim_v=128), # seqlen_kv > seqlen_q +] + + +@pytest.mark.parametrize("config", reference_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_reference_matches_bf16_sdpa(config, causal): + """The pure-PyTorch golden reference must track bf16 SDPA within mxfp8 error. + + Runnable on CPU without CDNA4 — the check we can do *before* the hardware + arrives, to de-risk the algorithm and quantization placement. + """ + device = "cuda" if torch.cuda.is_available() else "cpu" + dtype = torch.float32 if device == "cpu" else torch.bfloat16 + + q, k, v = _make_qkv_bhsd(2, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + o_ref = _sdpa_bf16_bhsd(q, k, v, sm_scale, causal) + o_golden, _ = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + + snr = calc_snr(o_ref, o_golden) + sim = calc_cossim(o_ref, o_golden) + print() + print(tabulate([["O (golden vs bf16)", snr, sim]], headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + assert sim > 0.99, f"golden reference cosine-sim vs bf16 SDPA too low: {sim}" + assert snr > 15, f"golden reference SNR vs bf16 SDPA too low: {snr}" + + +# --------------------------------------------------------------------------- +# Layer 2 — kernel vs golden reference (requires CDNA4 native tl.dot_scaled) +# --------------------------------------------------------------------------- + +@cuda_only +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True]) +def test_kernel_matches_reference(config, causal): + """Kernel vs golden reference — isolates Triton port bugs from quant error. + + Both sides apply the same mxfp8 quantization, so a large gap here points at a + Triton port bug (masking / online-softmax / LSE / strides) rather than + quantization error. + """ + from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 + + device = "cuda" + dtype = torch.bfloat16 + q, k, v = _make_qkv_bhsd(4, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + o_kernel = triton_attention_mxfp8( + q.contiguous(), + k.contiguous(), + v.contiguous(), + bias=None, + alibi_slopes=None, + sm_scale=sm_scale, + dropout_p=0.0, + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlens_q=config.seqlen_q, + max_seqlens_k=config.seqlen_kv, + causal=causal, + return_scores=False, + use_exp2=True, + layout="bhsd", + )[0] + o_golden, _ = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + + snr = calc_snr(o_golden, o_kernel) + sim = calc_cossim(o_golden, o_kernel) + print() + print(tabulate([["O (kernel vs golden)", snr, sim]], headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + # Threshold placeholder — calibrate on the first CDNA4 run. + assert sim > 0.99, f"kernel vs golden cosine-sim too low: {sim}" + assert snr > 30, f"kernel vs golden SNR too low: {snr}" diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index 0cc50394..aee72def 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -174,3 +174,117 @@ def convert_from_mxfp8_pytorch( data_hp = data_hp.to(torch.bfloat16) return data_hp.transpose(axis, -1) + + +def _mxfp8_qdq(x: torch.Tensor, axis: int, is_2d_block: bool, block_size: int = BLOCK_SIZE_DEFAULT) -> torch.Tensor: + """Round-trip a tensor through pure-PyTorch MXFP8 e4m3 quantize+dequantize. + + Returns the value an mxfp8 kernel operand would actually see. + """ + lp, s = convert_to_mxfp8_pytorch(x, block_size=block_size, mxfp_format="e4m3", axis=axis, is_2d_block=is_2d_block) + return convert_from_mxfp8_pytorch(lp, s, output_dtype=x.dtype, block_size=block_size, axis=axis, + is_2d_block=is_2d_block) + + +def mxfp8_attention_forward_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + sm_scale: float, + causal: bool, + block_n: int = 64, + block_size: int = BLOCK_SIZE_DEFAULT, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Pure-PyTorch golden reference for the MXFP8 (e4m3) flash-attention forward. + + Faithfully replicates what ``triton_flash_attention_mxfp8.attn_fwd`` computes, + so a kernel-vs-this-reference gap isolates Triton port bugs (masking / + online-softmax / LSE) from mxfp8 quantization error itself: + + * Q/K/V are quantized 2D-block along head_dim (``is_2d_block=True``, + ``axis=-1``), matching the autograd forward. + * The softmax probabilities ``P`` are quantized **per key block** exactly + like the kernel: within the online-softmax loop, using the *running* row + max (not the global max), 1D-block along the key axis with + ``block_size`` groups. + * Contains no ``tl.dot_scaled``: every matmul is a plain fp32 matmul on the + dequantized operands, so this runs on CPU / any device (no CDNA4). + + Args: + q, k, v: ``[batch, nheads, seqlen, head_dim]`` (bhsd), fp32/bf16. GQA is + supported (``nheads_k`` may be < ``nheads_q``). + sm_scale: softmax scale applied to ``Q @ K^T``. + causal: bottom-right-aligned causal mask (matches ``F.sdpa(is_causal=True)``). + block_n: key-block width, must match the kernel ``BLOCK_N`` (default 64). + block_size: MXFP8 quant block, must match ``QUANT_BLOCK_SIZE`` (default 32). + + Returns: + (o, softmax_lse): output ``[batch, nheads_q, seqlen_q, head_dim_v]`` and + log-sum-exp ``[batch, nheads_q, seqlen_q]``. + """ + assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4 + b, hq, sq, dqk = q.shape + _, hk, sk, _ = k.shape + dv = v.shape[-1] + assert hq % hk == 0, f"nheads_q ({hq}) must be a multiple of nheads_k ({hk})" + n_rep = hq // hk + + # Quantize Q/K/V (2D block along head_dim) then dequantize -> operands the kernel sees. + q_dq = _mxfp8_qdq(q, axis=-1, is_2d_block=True).to(torch.float32) + k_dq = _mxfp8_qdq(k, axis=-1, is_2d_block=True).to(torch.float32) + v_dq = _mxfp8_qdq(v, axis=-1, is_2d_block=True).to(torch.float32) + + # Expand K/V heads for GQA (head_q h -> head_k h // n_rep, matching off_h_k). + if n_rep > 1: + k_dq = k_dq.repeat_interleave(n_rep, dim=1) + v_dq = v_dq.repeat_interleave(n_rep, dim=1) + + device = q.device + # Real -inf (not finfo.min): masked scores must dequantize to exp(-inf)=0. + # A finite sentinel would make a fully-masked block compute exp(0)=1. + neg_inf = float("-inf") + + m_i = torch.full((b, hq, sq), float("-inf"), dtype=torch.float32, device=device) + l_i = torch.zeros((b, hq, sq), dtype=torch.float32, device=device) + acc = torch.zeros((b, hq, sq, dv), dtype=torch.float32, device=device) + + q_pos = torch.arange(sq, device=device) + # Bottom-right causal alignment: query i attends key j iff j <= i + (sk - sq). + causal_shift = sk - sq + + for j0 in range(0, sk, block_n): + j1 = min(j0 + block_n, sk) + k_blk = k_dq[:, :, j0:j1, :] # [b, hq, bn, d] + v_blk = v_dq[:, :, j0:j1, :] + + s = torch.matmul(q_dq, k_blk.transpose(-1, -2)) * sm_scale # [b, hq, sq, bn] + + if causal: + key_pos = torch.arange(j0, j1, device=device) + allowed = key_pos[None, :] <= (q_pos[:, None] + causal_shift) # [sq, bn] + s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, neg_inf)) + + m_ij = torch.maximum(m_i, s.max(dim=-1).values) # [b, hq, sq] + p = torch.exp(s - m_ij[..., None]) # unnormalized, running max + # Rows fully masked in this block yield exp(neg_inf - (-inf)) issues; zero them. + p = torch.nan_to_num(p, nan=0.0, posinf=0.0, neginf=0.0) + l_ij = p.sum(dim=-1) # [b, hq, sq] + + # Quantize P per key block exactly like the kernel (1D block along key axis). + p_dq = _mxfp8_qdq(p.to(torch.float32), axis=-1, is_2d_block=False, block_size=block_size) + + alpha = torch.exp(m_i - m_ij) + alpha = torch.nan_to_num(alpha, nan=0.0, posinf=0.0, neginf=0.0) + acc = acc * alpha[..., None] + torch.matmul(p_dq, v_blk) + l_i = l_i * alpha + l_ij + m_i = m_ij + + l_safe = torch.where(l_i > 0, l_i, torch.ones_like(l_i)) + o = acc / l_safe[..., None] + softmax_lse = m_i + torch.log(l_safe) + # Fully-masked query rows (can happen with causal when sk < sq): zero output. + fully_masked = ~torch.isfinite(m_i) | (l_i <= 0) + o = torch.where(fully_masked[..., None], torch.zeros_like(o), o) + softmax_lse = torch.where(fully_masked, torch.zeros_like(softmax_lse), softmax_lse) + + return o.to(q.dtype), softmax_lse From 748bdf429237dfc933927e91f189ee68b9e359e0 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Tue, 14 Jul 2026 04:30:37 -0500 Subject: [PATCH 03/14] fix: correct golden-ref causal mask and Triton compile issues --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 18 +++++++++++++----- .../mxfp8/triton_flash_attention_mxfp8.py | 12 +++++------- tests/unittest/mxfp8/utils.py | 7 +++---- 3 files changed, 21 insertions(+), 16 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index a18ed6fb..47ca7c15 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -21,6 +21,8 @@ 本 plan 撰写机为 gfx950(CDNA4/MI350),`is_cdna4()=True`。native `tl.dot_scaled` 路径可在本机验证;MI300/CDNA3 不在 V1 范围。 +**2026-07-14 补充(MI250 离线开发机):** 当前开发/CI 环境为 gfx90a(MI250),`is_cdna4()=False`。Layer 2/3(kernel vs 黄金参照、kernel vs bf16 SDPA)在此硬件上跑不起来是预期行为,**不为 MI250 加 fallback 或 skip workaround**;待 CDNA4 硬件到位后一次性验证。Layer 1(纯 PyTorch 黄金参照 vs bf16 SDPA)可在任意设备离线跑。 + --- ## 0. 格式选择:V1 全 e4m3,e5m2 留给后续 @@ -291,18 +293,24 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` | Step 1–4(forward kernel + wrapper + autograd.forward) | ✅ 代码完成 / ⏳ 未验证 | 无 CDNA4,尚未跑通任何 kernel 测试 | | Step 5–8(backward:preprocess / dK-dV / dQ / autograd.backward) | ❌ 未开始 | 从零写,头号工作量,见 §9 | | Step 9(dispatch 接入 `mxfp8_e4m3`) | ❌ 未开始 | — | -| Step 10(测试脚手架) | ✅ 三层脚手架就位 | 第 1 层可离线验,第 2/3 层待 CDNA4 | +| Step 10(测试脚手架) | ✅ 三层脚手架就位 | 第 1 层已于 MI250 Docker 验证 6/6 通过;第 2/3 层待 CDNA4 | **已知偏差 / 修正记录:** - 写黄金参照时静态发现并修掉一处 masking bug:mask 值若用 `finfo.min`(有限)会让"整块被 mask 的行"算出 `exp(0)=1`(应为 0);改用真 `-inf`(masked 项 `exp(-inf)=0`,全 mask 行的 nan 由 `nan_to_num` 兜住)。此 bug 正是"无黄金参照、只对 bf16"抓不到的类型。 +- **2026-07-14 — 黄金参照 causal mask 与 SDPA 不一致(非 CDNA4,Layer 1)**:`mxfp8_attention_forward_reference` 原用 bottom-right 对齐(`key_j <= query_i + (seqlen_k - seqlen_q)`),但 PyTorch `F.sdpa(is_causal=True)` 在 **bhsd** 布局下实际为 **top-left**(`key_j <= query_i`)。`seqlen_q == seqlen_k` 时两者等价;`seqlen_kv > seqlen_q` 且 `causal=True` 时差异显著(`test_reference_matches_bf16_sdpa[True-config2]` cosine-sim 仅 ~0.39)。**修复**:`tests/unittest/mxfp8/utils.py` 改为 top-left mask。**验证**:MI250(gfx90a)Docker 上 Layer 1 六项全过。 +- **2026-07-14 — Triton kernel 编译期两处真 bug(非 CDNA4 专属,CDNA4 上同样会触发)**: + 1. **fp8 `tl.load` 的 `other` 类型**:`q = tl.load(q_ptrs, mask=..., other=0)` 中 `other=0`(int32)无法 cast 为 `fp8e4nv`,报 `cannot cast int32[...] to fp8e4nv`。**修复**:`triton_flash_attention_mxfp8.py` 中 Q load 及 `load_fn` 默认 `other` 改为 `0.0`。 + 2. **Triton constexpr 全局变量写法**:`E4M3_TARGET_MAX_POW2: tl.constexpr = FORMAT_TO_TARGET_MAX["e4m3"]` 等注解写法在 `@jit` 内不可见,报 `NameError: Cannot access global variable E4M3_TARGET_MAX_POW2`。**修复**:改为实例化写法 `E4M3_TARGET_MAX_POW2 = tl.constexpr(8)`、`E4M3_MBITS = tl.constexpr(3)`、`E4M3_FORMAT_ID = tl.constexpr(0)`。 + - 上述修复待在 CDNA4 真机上跑 Layer 2 时一并验收;MI250 上 Layer 2 仍因 `tl.dot_scaled` 不可用而无法运行,**属预期,未做 workaround**。 **Open items(阻塞真验证):** -- **CDNA4 硬件**:M350/M355 卡预计约一周后到;在此之前所有 kernel 测试(第 2/3 层)无法运行,forward 代码处于"已写未验"状态。 -- **第 1 层 CPU 自证**:需在有 torch 的环境跑一次 `test_reference_matches_bf16_sdpa`,确认黄金参照本身正确(它是 backward 验证的地基)。 -- **backward 全 e4m3 的数值风险**:见 §0 风险记录 / §5,真机验证后据实决定是否需切 e5m2。 +- **CDNA4 硬件**:M350/M355 卡预计约一周后到;在此之前所有 kernel 测试(第 2/3 层)无法运行,forward 代码处于"已写未验"状态(Triton 编译修复已合入,待真机确认 Layer 2)。 +- ~~**第 1 层 CPU 自证**~~:✅ 2026-07-14 于 MI250 Docker 跑通 `test_reference_matches_bf16_sdpa`(6/6 passed)。 +- **第 2 层 kernel vs 黄金参照**:待 CDNA4 硬件;届时验收 `test_kernel_matches_reference` 全网格。 +- **backward 全 e4m3 的数值风险**:见 §0 风险记录 / §5,真机验证后据实决定是否需切 e5m2。 ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,ground-truth 由纯 PyTorch 黄金参照提供(CPU 可离线验)。量化基础设施已齐,forward FA v2 骨架复用;最大工作量与风险集中在从零的 backward。当前 forward + 三层测试脚手架已落地但未经硬件验证,backward 未开始。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,ground-truth 由纯 PyTorch 黄金参照提供(Layer 1 已于 MI250 离线验 6/6 通过)。量化基础设施已齐,forward FA v2 骨架复用;Triton 编译两处真 bug 已修、待 CDNA4 验 Layer 2。最大工作量与风险集中在从零的 backward;backward 未开始。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index bbf2c4fe..e3d40b82 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -34,8 +34,6 @@ is_cdna4, _calculate_scales, _quantize_fp8, - FORMAT_TO_TARGET_MAX, - FORMAT_TO_MBITS, ) fwd_torch_dtype: tl.constexpr = torch.bfloat16 @@ -45,9 +43,9 @@ philox_offset: tl.constexpr = 0x1D4B42 # e4m3 quantization constants (see mxfp8_quantization.FORMAT_TO_*). -E4M3_TARGET_MAX_POW2: tl.constexpr = FORMAT_TO_TARGET_MAX["e4m3"] -E4M3_MBITS: tl.constexpr = FORMAT_TO_MBITS["e4m3"] -E4M3_FORMAT_ID: tl.constexpr = 0 +E4M3_TARGET_MAX_POW2 = tl.constexpr(8) +E4M3_MBITS = tl.constexpr(3) +E4M3_FORMAT_ID = tl.constexpr(0) AUTOTUNE = os.environ.get('FLASH_ATTENTION_TRITON_AMD_AUTOTUNE', '0').lower() in ('1', 'true', 'yes') DEBUG = os.environ.get('FLASH_ATTENTION_TRITON_AMD_DEBUG', '0').lower() in ('1', 'true', 'yes') @@ -129,7 +127,7 @@ def dropout_mask(philox_seed, philox_offset, dropout_p, m, n, stride): # Convenience function to load with optional boundary checks. # "First" is the major dim, "second" is the minor dim. @triton.jit -def load_fn(ptrs, offset_first, offset_second, boundary_first, boundary_second, other=0): +def load_fn(ptrs, offset_first, offset_second, boundary_first, boundary_second, other=0.0): if offset_first is not None and offset_second is not None: mask = (offset_first[:, None] < boundary_first) & \ (offset_second[None, :] < boundary_second) @@ -624,7 +622,7 @@ def attn_fwd( q_ptrs_mask = q_ptrs_mask & (offs_d_qk[None, :] < ACTUAL_BLOCK_DMODEL_QK) qs_ptrs_mask = qs_ptrs_mask & (offs_d_qk_scale[None, :] < SCALE_ACTUAL_BLOCK_DMODEL_QK) - q = tl.load(q_ptrs, mask=q_ptrs_mask, other=0) + q = tl.load(q_ptrs, mask=q_ptrs_mask, other=0.0) qs = tl.load(qs_ptrs, mask=qs_ptrs_mask, other=1) # Here we compute how many full and masked blocks we have. diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index aee72def..cde68d2f 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -214,7 +214,8 @@ def mxfp8_attention_forward_reference( q, k, v: ``[batch, nheads, seqlen, head_dim]`` (bhsd), fp32/bf16. GQA is supported (``nheads_k`` may be < ``nheads_q``). sm_scale: softmax scale applied to ``Q @ K^T``. - causal: bottom-right-aligned causal mask (matches ``F.sdpa(is_causal=True)``). + causal: top-left causal mask ``key_j <= query_i`` (matches + ``F.sdpa(is_causal=True)`` on bhsd tensors in PyTorch 2.x). block_n: key-block width, must match the kernel ``BLOCK_N`` (default 64). block_size: MXFP8 quant block, must match ``QUANT_BLOCK_SIZE`` (default 32). @@ -249,8 +250,6 @@ def mxfp8_attention_forward_reference( acc = torch.zeros((b, hq, sq, dv), dtype=torch.float32, device=device) q_pos = torch.arange(sq, device=device) - # Bottom-right causal alignment: query i attends key j iff j <= i + (sk - sq). - causal_shift = sk - sq for j0 in range(0, sk, block_n): j1 = min(j0 + block_n, sk) @@ -261,7 +260,7 @@ def mxfp8_attention_forward_reference( if causal: key_pos = torch.arange(j0, j1, device=device) - allowed = key_pos[None, :] <= (q_pos[:, None] + causal_shift) # [sq, bn] + allowed = key_pos[None, :] <= q_pos[:, None] # [sq, bn] s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, neg_inf)) m_ij = torch.maximum(m_i, s.max(dim=-1).values) # [b, hq, sq] From 797eb385788ea8b7a0f9a921f957747b20f17f23 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 15 Jul 2026 12:50:32 +0800 Subject: [PATCH 04/14] feat: add backward (stage-1 bf16 skeleton) + golden reference --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 25 +- .../mxfp8/triton_flash_attention_mxfp8.py | 722 +++++++++++++++++- .../mxfp8/test_mxfp8_attention_reference.py | 120 ++- tests/unittest/mxfp8/utils.py | 90 +++ 4 files changed, 953 insertions(+), 4 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 47ca7c15..1f12e4bd 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -311,6 +311,29 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - **第 2 层 kernel vs 黄金参照**:待 CDNA4 硬件;届时验收 `test_kernel_matches_reference` 全网格。 - **backward 全 e4m3 的数值风险**:见 §0 风险记录 / §5,真机验证后据实决定是否需切 e5m2。 +### 10.1bis 截至 2026-07-15(backward 阶段1 + 参照) + +**关键决策修正:backward 不从零写。** 发现 `alto/kernels/blockwise_fp8/triton_flash_attention_fp8_block.py` 有一套完整 fp8 attention backward(`_bwd_preprocess` + `_bwd_kernel_dkdv` + `_bwd_kernel_dq`),结构与 §9 三 kernel 完全同构。因此 backward 改为**机械 port**(同 forward 从 mxfp4 改写的性质),§9"从零写、头号风险"前提作废。真正工作量集中在 **5 个 dot 的 MX scale 轴对齐**(②dp/③dv 的 dO、①qk/④dk 的 q 各需沿不同 reduction 轴的两套量化)。 + +**分两阶段落地,采用 port_now(骨架先行):** + +- **阶段1 ✅ 已落地(代码,未验)**:port 三 kernel(`_bwd_preprocess` / `_bwd_kernel_dq` + `_attn_bwd_dq_inner` / `_bwd_kernel_dkdv` + `_attn_bwd_dkdv_inner`)+ backward driver(注册 `alto::attention_mxfp8_backward_triton_impl` + `register_fake`)+ 接上 `autograd.backward`(返回 15 个梯度,顺手修正旧 stub 的 17 个 None bug)。 + - **dot 用高精度 bf16 `tl.dot` 占位**,每处留 `TODO(stage2)` 标注该 dot 的 reduction 轴。 + - **入口用 `convert_from_mxfp8` 把 saved e4m3 q/k/v 反量化回 bf16**——backward 吃的正是 forward 用过的那份量化输入,且**不含 `tl.dot_scaled`,MI250 可跑**。 + - causal 用 **bottom-right**(与 forward kernel 一致,非参照的 top-left;方阵下无差别)。 +- **阶段2 ❌ 未开始**:逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴的 MX e4m3 量化 + dO 量化。需 CDNA4 才能验。 + +**backward 黄金参照 ✅ 已落地**:`tests/unittest/mxfp8/utils.py::mxfp8_attention_backward_reference`——纯 PyTorch 复刻阶段1 kernel(反量化 q/k/v、用 saved `lse` 在 fp32 重算 `P`、backward 不量化 P),标准 FA v2 backward(dV=Pᵀ@dO;dP=dO@Vᵀ;dS=P·(dP−delta);dQ=sm·dS@K;dK=sm·dSᵀ@Q)。喂 forward 参照的同一份 `o`/`lse` 时应与 kernel 高精度吻合。 + +**backward 测试(加在 `test_mxfp8_attention_reference.py`):** + +- **第 1 层 `test_backward_reference_matches_sdpa`(CPU)**:参照 vs autograd fp32 SDPA backward,验算法+量化误差(`cossim>0.97`/`SNR>8`)。**本机无 torch,未跑,待另一台机跑。** +- **第 2 层 `test_backward_kernel_matches_reference`(设备)**:直接调 backward op(喂 forward 参照的 `o`/`lse`,**绕过 forward 的 `tl.dot_scaled`**),故**阶段1 bf16 骨架 MI250 即可验**;kernel vs 参照(`cossim>0.99`/`SNR>25`,阈值待首跑校准)。 + +**Open items 更新:** +- **阶段1 backward 未经任何硬件验证**(本机无 torch);下一步在另一台机跑第 1 层(CPU)+ 第 2 层(MI250 即可,无需 CDNA4)拿反馈。 +- 阶段2(dot_scaled + MX 量化)仍待 CDNA4。 + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 无参照、按标准 FA v2 backward 数学从零写**(preprocess + dK/dV + dQ 三 kernel,每 dot 各自量化)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,ground-truth 由纯 PyTorch 黄金参照提供(Layer 1 已于 MI250 离线验 6/6 通过)。量化基础设施已齐,forward FA v2 骨架复用;Triton 编译两处真 bug 已修、待 CDNA4 验 Layer 2。最大工作量与风险集中在从零的 backward;backward 未开始。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 参照 + 两层测试已就位待跑)。真正风险集中在阶段2 的 5-dot MX scale 轴对齐。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index e3d40b82..f4c60671 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -34,6 +34,7 @@ is_cdna4, _calculate_scales, _quantize_fp8, + convert_from_mxfp8, ) fwd_torch_dtype: tl.constexpr = torch.bfloat16 @@ -1022,6 +1023,699 @@ def attention_mxfp8_forward_triton_impl( return o, softmax_lse, exp_scores +RCP_LN2: tl.constexpr = 1.4426950408889634 + + +@triton.jit +def _bwd_preprocess( + Out, + DO, + Delta, + stride_oz, + stride_oh, + stride_om, + stride_ok, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_deltaz, + stride_deltah, + stride_deltam, + cu_seqlens_q, + max_seqlen_q, + BLOCK_M: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + ACTUAL_BLOCK_DMODEL_V: tl.constexpr, + HQ: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + """delta = rowsum(o * do), one value per query row. Feeds ds = p * (dp - delta).""" + pid_m = tl.program_id(0) + pid_bh = tl.program_id(1) + off_z = pid_bh // HQ + off_h = pid_bh % HQ + + if IS_VARLEN: + q_start = tl.load(cu_seqlens_q + off_z) + q_end = tl.load(cu_seqlens_q + off_z + 1) + N_CTX_Q = q_end - q_start + else: + q_start = 0 + N_CTX_Q = max_seqlen_q + + off_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + off_d_v = tl.arange(0, BLOCK_DMODEL_V) + mask_m = off_m < N_CTX_Q + mask_o = mask_m[:, None] & (off_d_v[None, :] < ACTUAL_BLOCK_DMODEL_V) + + o_offset = Out + off_z * stride_oz + off_h * stride_oh + q_start * stride_om + do_offset = DO + off_z * stride_doz + off_h * stride_doh + q_start * stride_dom + out_ptrs = o_offset + off_m[:, None] * stride_om + off_d_v[None, :] * stride_ok + do_ptrs = do_offset + off_m[:, None] * stride_dom + off_d_v[None, :] * stride_dok + + o = tl.load(out_ptrs, mask=mask_o, other=0.0).to(tl.float32) + do = tl.load(do_ptrs, mask=mask_o, other=0.0).to(tl.float32) + delta = tl.sum(o * do, axis=1) + + delta_offset = Delta + off_z * stride_deltaz + off_h * stride_deltah + q_start * stride_deltam + tl.store(delta_offset + off_m * stride_deltam, delta, mask=mask_m) + + +@triton.jit +def _attn_bwd_dkdv_inner( + k, + v, + dk, + dv, + offs_d_qk, + offs_d_v, + offs_n, + mask_d_qk, + mask_d_v, + q_offset, + do_offset, + stride_qm, + stride_qk, + stride_dom, + stride_dok, + l_offset, + d_offset, + stride_ldm, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + sm_scale: tl.constexpr, + lo, + num_block_m: tl.constexpr, + USE_EXP2: tl.constexpr, + N_CTX_Q: tl.constexpr, + N_CTX_K: tl.constexpr, + CAUSAL: tl.constexpr, +): + """Accumulate dk, dv over query blocks for one key block. k,v come in transposed.""" + for start_m in range(lo, num_block_m * BLOCK_M, BLOCK_M): + offs_m = start_m + tl.arange(0, BLOCK_M) + q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk + do_ptrs = do_offset + offs_m[:, None] * stride_dom + offs_d_v[None, :] * stride_dok + + mask_m = offs_m < N_CTX_Q + q_mask = mask_m[:, None] & mask_d_qk[None, :] + do_mask = mask_m[:, None] & mask_d_v[None, :] + + q = tl.load(q_ptrs, mask=q_mask, other=0.0) + # TODO(stage2): dot_scaled(q[e4m3, scale//head_dim], k[e4m3, scale//head_dim]); reduction axis = head_dim_qk. + qk = tl.dot(q, k, out_dtype=tl.float32) + + if CAUSAL: + # Bottom-right causal, matches the forward kernel (offs_n_causal = offs_n + N_CTX_Q - N_CTX_K). + col_offset = N_CTX_Q - N_CTX_K + causal_mask = offs_m[:, None] >= (col_offset + offs_n[None, :]) + qk = tl.where(causal_mask, qk, float("-inf")) + + l_ptrs = l_offset + offs_m * stride_ldm + l_i = tl.load(l_ptrs, mask=mask_m, other=0.0) + if USE_EXP2: + qk *= sm_scale * RCP_LN2 + p = tl.math.exp2(qk - l_i[:, None] * RCP_LN2) + else: + qk *= sm_scale + p = tl.math.exp(qk - l_i[:, None]) + + do = tl.load(do_ptrs, mask=do_mask, other=0.0) + # TODO(stage2): dot_scaled(do[e4m3, scale//head_dim_v], v[e4m3, scale//head_dim_v]); reduction axis = head_dim_v. + dp = tl.dot(do, v, out_dtype=tl.float32) + + d_ptrs = d_offset + offs_m * stride_ldm + Di = tl.load(d_ptrs, mask=mask_m, other=0.0) + ds = p * (dp - Di[:, None]) + ds = ds.to(q.dtype) + + # TODO(stage2): dot_scaled(p^T[e4m3, scale//BLOCK_M], do[e4m3, scale//BLOCK_M]); reduction axis = seqlen_q. + dv += tl.dot(tl.trans(p.to(k.dtype)), do, out_dtype=tl.float32) + # TODO(stage2): dot_scaled(ds^T[e4m3, scale//BLOCK_M], q[e4m3, scale//BLOCK_M]); reduction axis = seqlen_q. + dk += tl.dot(tl.trans(ds), q, out_dtype=tl.float32) + + return dk, dv + + +@triton.jit +def _bwd_kernel_dkdv( + Q, + K, + V, + sm_scale: tl.constexpr, + DO, + DK, + DV, + LSE, + Delta, + stride_qz, + stride_qh, + stride_qm, + stride_qk, + stride_kz, + stride_kh, + stride_kn, + stride_kk, + stride_vz, + stride_vh, + stride_vn, + stride_vk, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_ldz, + stride_ldh, + stride_ldm, + Z, + HQ: tl.constexpr, + HK: tl.constexpr, + cu_seqlens_q, + cu_seqlens_k, + max_seqlen_q, + max_seqlen_k, + num_block_m: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + ACTUAL_BLOCK_DMODEL_V: tl.constexpr, + CAUSAL: tl.constexpr, + USE_EXP2: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + """One program per (batch*head_k, key block). Parallelizes dk/dv over keys.""" + off_hz = tl.program_id(0) + start_n = tl.program_id(1) + off_z = off_hz // HK + off_h_k = off_hz % HK + + GROUP_SIZE: tl.constexpr = HQ // HK + off_h_q = off_h_k * GROUP_SIZE if GROUP_SIZE != 1 else off_h_k + + if IS_VARLEN: + q_start = tl.load(cu_seqlens_q + off_z) + k_start = tl.load(cu_seqlens_k + off_z) + N_CTX_Q = tl.load(cu_seqlens_q + off_z + 1) - q_start + N_CTX_K = tl.load(cu_seqlens_k + off_z + 1) - k_start + else: + q_start = 0 + k_start = 0 + N_CTX_Q = max_seqlen_q + N_CTX_K = max_seqlen_k + + q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + q_start * stride_qm + k_offset = K + off_z * stride_kz + off_h_k * stride_kh + k_start * stride_kn + v_offset = V + off_z * stride_vz + off_h_k * stride_vh + k_start * stride_vn + do_offset = DO + off_z * stride_doz + off_h_q * stride_doh + q_start * stride_dom + adj_delta = off_z * stride_ldz + off_h_q * stride_ldh + q_start * stride_ldm + l_offset = LSE + adj_delta + d_offset = Delta + adj_delta + dk_offset = DK + off_z * stride_kz + off_h_k * stride_kh + k_start * stride_kn + dv_offset = DV + off_z * stride_vz + off_h_k * stride_vh + k_start * stride_vn + + if CAUSAL: + causal_boundary = start_n * BLOCK_N - BLOCK_M + lo = (causal_boundary + 1) // BLOCK_M * BLOCK_M + else: + lo = 0 + + offs_d_qk = tl.arange(0, BLOCK_DMODEL_QK) + offs_d_v = tl.arange(0, BLOCK_DMODEL_V) + offs_n = start_n * BLOCK_N + tl.arange(0, BLOCK_N) + + mask_n = offs_n < N_CTX_K + mask_d_qk = offs_d_qk < ACTUAL_BLOCK_DMODEL_QK + mask_d_v = offs_d_v < ACTUAL_BLOCK_DMODEL_V + + k_ptrs = k_offset + offs_n[:, None] * stride_kn + offs_d_qk[None, :] * stride_kk + v_ptrs = v_offset + offs_n[:, None] * stride_vn + offs_d_v[None, :] * stride_vk + k = tl.load(k_ptrs, mask=mask_n[:, None] & mask_d_qk[None, :], other=0.0) + v = tl.load(v_ptrs, mask=mask_n[:, None] & mask_d_v[None, :], other=0.0) + k = tl.trans(k) + v = tl.trans(v) + + dk = tl.zeros([BLOCK_N, BLOCK_DMODEL_QK], dtype=tl.float32) + dv = tl.zeros([BLOCK_N, BLOCK_DMODEL_V], dtype=tl.float32) + + for _ in range(GROUP_SIZE): + dk, dv = _attn_bwd_dkdv_inner( + k, + v, + dk, + dv, + offs_d_qk, + offs_d_v, + offs_n, + mask_d_qk, + mask_d_v, + q_offset, + do_offset, + stride_qm, + stride_qk, + stride_dom, + stride_dok, + l_offset, + d_offset, + stride_ldm, + BLOCK_M, + BLOCK_N, + sm_scale, + lo, + num_block_m, + USE_EXP2, + N_CTX_Q, + N_CTX_K, + CAUSAL, + ) + q_offset += stride_qh + do_offset += stride_qh + l_offset += stride_ldh + d_offset += stride_ldh + + dk *= sm_scale + + tl.store(dk_offset + offs_n[:, None] * stride_kn + offs_d_qk[None, :] * stride_kk, dk, + mask=mask_n[:, None] & mask_d_qk[None, :]) + tl.store(dv_offset + offs_n[:, None] * stride_vn + offs_d_v[None, :] * stride_vk, dv, + mask=mask_n[:, None] & mask_d_v[None, :]) + + +@triton.jit +def _attn_bwd_dq_inner( + dq, + q, + offs_d_qk, + offs_d_v, + offs_m, + l_i, + Di, + do, + mask_d_qk, + mask_d_v, + k_offset, + v_offset, + stride_kn, + stride_kk, + stride_vn, + stride_vk, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + sm_scale: tl.constexpr, + hi, + USE_EXP2: tl.constexpr, + N_CTX_Q: tl.constexpr, + N_CTX_K: tl.constexpr, + CAUSAL: tl.constexpr, +): + """Accumulate dq over key blocks for one query block.""" + if USE_EXP2: + l_i *= RCP_LN2 + + for start_n in range(0, hi, BLOCK_N): + offs_n = start_n + tl.arange(0, BLOCK_N) + mask_n = offs_n < N_CTX_K + mask_k = mask_n[:, None] & mask_d_qk[None, :] + mask_v = mask_n[:, None] & mask_d_v[None, :] + + k_ptrs = k_offset + offs_n[:, None] * stride_kn + offs_d_qk[None, :] * stride_kk + v_ptrs = v_offset + offs_n[:, None] * stride_vn + offs_d_v[None, :] * stride_vk + k = tl.load(k_ptrs, mask=mask_k, other=0.0) + v = tl.load(v_ptrs, mask=mask_v, other=0.0) + + # TODO(stage2): dot_scaled(q[e4m3, scale//head_dim], k^T[e4m3, scale//head_dim]); reduction axis = head_dim_qk. + qk = tl.dot(q, tl.trans(k), out_dtype=tl.float32) + + if CAUSAL: + col_offset = N_CTX_Q - N_CTX_K + causal_mask = offs_m[:, None] >= (col_offset + offs_n[None, :]) + qk = tl.where(causal_mask, qk, float("-inf")) + + if USE_EXP2: + qk *= sm_scale * RCP_LN2 + p = tl.math.exp2(qk - l_i[:, None]) + else: + qk *= sm_scale + p = tl.math.exp(qk - l_i[:, None]) + + # TODO(stage2): dot_scaled(do[e4m3, scale//head_dim_v], v^T[e4m3, scale//head_dim_v]); reduction axis = head_dim_v. + dp = tl.dot(do, tl.trans(v), out_dtype=tl.float32) + ds = p * (dp - Di[:, None]) + ds = ds.to(q.dtype) + # TODO(stage2): dot_scaled(ds[e4m3, scale//BLOCK_N], k[e4m3, scale//BLOCK_N]); reduction axis = seqlen_k. + dq += tl.dot(ds, k, out_dtype=tl.float32) + + return dq + + +@triton.jit +def _bwd_kernel_dq( + Q, + K, + V, + sm_scale: tl.constexpr, + DO, + DQ, + LSE, + Delta, + stride_qz, + stride_qh, + stride_qm, + stride_qk, + stride_kz, + stride_kh, + stride_kn, + stride_kk, + stride_vz, + stride_vh, + stride_vn, + stride_vk, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_ldz, + stride_ldh, + stride_ldm, + Z, + HQ: tl.constexpr, + HK: tl.constexpr, + cu_seqlens_q, + cu_seqlens_k, + max_seqlen_q, + max_seqlen_k, + num_block_n: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + ACTUAL_BLOCK_DMODEL_V: tl.constexpr, + CAUSAL: tl.constexpr, + USE_EXP2: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + """One program per (batch*head_q, query block). Parallelizes dq over queries.""" + off_hz = tl.program_id(0) + start_m = tl.program_id(1) + off_z = off_hz // HQ + off_h_q = off_hz % HQ + + GROUP_SIZE: tl.constexpr = HQ // HK + off_h_k = off_h_q // GROUP_SIZE if GROUP_SIZE != 1 else off_h_q + + if IS_VARLEN: + q_start = tl.load(cu_seqlens_q + off_z) + k_start = tl.load(cu_seqlens_k + off_z) + N_CTX_Q = tl.load(cu_seqlens_q + off_z + 1) - q_start + N_CTX_K = tl.load(cu_seqlens_k + off_z + 1) - k_start + else: + q_start = 0 + k_start = 0 + N_CTX_Q = max_seqlen_q + N_CTX_K = max_seqlen_k + + q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + q_start * stride_qm + k_offset = K + off_z * stride_kz + off_h_k * stride_kh + k_start * stride_kn + v_offset = V + off_z * stride_vz + off_h_k * stride_vh + k_start * stride_vn + do_offset = DO + off_z * stride_doz + off_h_q * stride_doh + q_start * stride_dom + adj_delta = off_z * stride_ldz + off_h_q * stride_ldh + q_start * stride_ldm + l_offset = LSE + adj_delta + d_offset = Delta + adj_delta + dq_offset = DQ + off_z * stride_qz + off_h_q * stride_qh + q_start * stride_qm + + if CAUSAL: + hi = tl.minimum(BLOCK_M // BLOCK_N * (start_m + 1), num_block_n) * BLOCK_N + else: + hi = num_block_n * BLOCK_N + + offs_d_qk = tl.arange(0, BLOCK_DMODEL_QK) + offs_d_v = tl.arange(0, BLOCK_DMODEL_V) + offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M) + + mask_m = offs_m < N_CTX_Q + mask_d_qk = offs_d_qk < ACTUAL_BLOCK_DMODEL_QK + mask_d_v = offs_d_v < ACTUAL_BLOCK_DMODEL_V + + q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk + do_ptrs = do_offset + offs_m[:, None] * stride_dom + offs_d_v[None, :] * stride_dok + q = tl.load(q_ptrs, mask=mask_m[:, None] & mask_d_qk[None, :], other=0.0) + do = tl.load(do_ptrs, mask=mask_m[:, None] & mask_d_v[None, :], other=0.0) + + l_i = tl.load(l_offset + offs_m * stride_ldm, mask=mask_m, other=0.0) + Di = tl.load(d_offset + offs_m * stride_ldm, mask=mask_m, other=0.0) + + dq = tl.zeros([BLOCK_M, BLOCK_DMODEL_QK], dtype=tl.float32) + dq = _attn_bwd_dq_inner( + dq, + q, + offs_d_qk, + offs_d_v, + offs_m, + l_i, + Di, + do, + mask_d_qk, + mask_d_v, + k_offset, + v_offset, + stride_kn, + stride_kk, + stride_vn, + stride_vk, + BLOCK_M, + BLOCK_N, + sm_scale, + hi, + USE_EXP2, + N_CTX_Q, + N_CTX_K, + CAUSAL, + ) + + dq *= sm_scale + tl.store(dq_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk, dq, + mask=mask_m[:, None] & mask_d_qk[None, :]) + + +@triton_op("alto::attention_mxfp8_backward_triton_impl", mutates_args=()) +def attention_mxfp8_backward_triton_impl( + do: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + o: torch.Tensor, + softmax_lse: torch.Tensor, + q_scale: torch.Tensor, + k_scale: torch.Tensor, + v_scale: torch.Tensor, + sm_scale: float, + causal: bool, + layout: str, + cu_seqlens_q: Optional[int], + cu_seqlens_k: Optional[int], + max_seqlen_q: Optional[int], + max_seqlen_k: Optional[int], + use_exp2: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """MXFP8 flash-attention backward. + + Stage 1 (current): q/k/v arrive as saved e4m3 tensors; they are dequantized + back to bf16 here so the three ported kernels run the standard FA v2 backward + math in high precision. This is numerically correct w.r.t. the *quantized* + forward inputs and is runnable without CDNA4. + + Stage 2 (TODO): keep operands in e4m3 and replace each ``tl.dot`` inside the + kernels with ``tl.dot_scaled`` plus per-reduction-axis MXFP8 scales (see the + ``TODO(stage2)`` markers). dO will additionally be quantized to e4m3. + """ + # Dequantize the saved e4m3 operands back to bf16 (the exact inputs the + # forward pass consumed). Uses the non-ASM path on non-CDNA4 hardware. + q = convert_from_mxfp8(q, q_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() + k = convert_from_mxfp8(k, k_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() + v = convert_from_mxfp8(v, v_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() + + if not do.is_contiguous(): + do = do.contiguous() + + batch, nheads_q, nheads_k, head_size_qk, head_size_v, max_seqlen_q, max_seqlen_k = get_shape_from_layout( + q, k, v, layout, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k) + stride_qz, stride_qh, stride_qm, stride_qk = get_strides_from_layout(q, layout) + stride_kz, stride_kh, stride_kn, stride_kk = get_strides_from_layout(k, layout) + stride_vz, stride_vh, stride_vn, stride_vk = get_strides_from_layout(v, layout) + stride_oz, stride_oh, stride_om, stride_ok = get_strides_from_layout(o, layout) + stride_doz, stride_doh, stride_dom, stride_dok = get_strides_from_layout(do, layout) + is_varlen = layout == "thd" + + padded_d_model_qk = get_padded_head_dim(head_size_qk) + padded_d_model_v = get_padded_head_dim(head_size_v) + + dq = torch.zeros_like(q, dtype=bwd_torch_dtype) + dk = torch.zeros_like(k, dtype=bwd_torch_dtype) + dv = torch.zeros_like(v, dtype=bwd_torch_dtype) + + delta = torch.empty_like(softmax_lse) + if is_varlen: + stride_lse_m, stride_lse_h = softmax_lse.stride() + stride_lse_z = 0 + else: + stride_lse_z, stride_lse_h, stride_lse_m = softmax_lse.stride() + + BLOCK_M = 64 + BLOCK_N = 64 + num_block_m = triton.cdiv(max_seqlen_q, BLOCK_M) + num_block_n = triton.cdiv(max_seqlen_k, BLOCK_N) + + wrap_triton(_bwd_preprocess)[(num_block_m, batch * nheads_q)]( + o, + do, + delta, + stride_oz, + stride_oh, + stride_om, + stride_ok, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_lse_z, + stride_lse_h, + stride_lse_m, + cu_seqlens_q, + max_seqlen_q, + BLOCK_M=BLOCK_M, + BLOCK_DMODEL_V=padded_d_model_v, + ACTUAL_BLOCK_DMODEL_V=head_size_v, + HQ=nheads_q, + IS_VARLEN=is_varlen, + ) + + wrap_triton(_bwd_kernel_dq)[(batch * nheads_q, num_block_m)]( + q, + k, + v, + sm_scale, + do, + dq, + softmax_lse, + delta, + stride_qz, + stride_qh, + stride_qm, + stride_qk, + stride_kz, + stride_kh, + stride_kn, + stride_kk, + stride_vz, + stride_vh, + stride_vn, + stride_vk, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_lse_z, + stride_lse_h, + stride_lse_m, + batch, + nheads_q, + nheads_k, + cu_seqlens_q, + cu_seqlens_k, + max_seqlen_q, + max_seqlen_k, + num_block_n=num_block_n, + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + BLOCK_DMODEL_QK=padded_d_model_qk, + BLOCK_DMODEL_V=padded_d_model_v, + ACTUAL_BLOCK_DMODEL_QK=head_size_qk, + ACTUAL_BLOCK_DMODEL_V=head_size_v, + CAUSAL=causal, + USE_EXP2=use_exp2, + IS_VARLEN=is_varlen, + ) + + wrap_triton(_bwd_kernel_dkdv)[(batch * nheads_k, num_block_n)]( + q, + k, + v, + sm_scale, + do, + dk, + dv, + softmax_lse, + delta, + stride_qz, + stride_qh, + stride_qm, + stride_qk, + stride_kz, + stride_kh, + stride_kn, + stride_kk, + stride_vz, + stride_vh, + stride_vn, + stride_vk, + stride_doz, + stride_doh, + stride_dom, + stride_dok, + stride_lse_z, + stride_lse_h, + stride_lse_m, + batch, + nheads_q, + nheads_k, + cu_seqlens_q, + cu_seqlens_k, + max_seqlen_q, + max_seqlen_k, + num_block_m=num_block_m, + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + BLOCK_DMODEL_QK=padded_d_model_qk, + BLOCK_DMODEL_V=padded_d_model_v, + ACTUAL_BLOCK_DMODEL_QK=head_size_qk, + ACTUAL_BLOCK_DMODEL_V=head_size_v, + CAUSAL=causal, + USE_EXP2=use_exp2, + IS_VARLEN=is_varlen, + ) + + return dq, dk, dv + + +@attention_mxfp8_backward_triton_impl.register_fake +def fake_attention_mxfp8_backward_triton_impl( + do: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + o: torch.Tensor, + softmax_lse: torch.Tensor, + q_scale: torch.Tensor, + k_scale: torch.Tensor, + v_scale: torch.Tensor, + sm_scale: float, + causal: bool, + layout: str, + cu_seqlens_q: Optional[int], + cu_seqlens_k: Optional[int], + max_seqlen_q: Optional[int], + max_seqlen_k: Optional[int], + use_exp2: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + dq = torch.empty_like(q, dtype=bwd_torch_dtype) + dk = torch.empty_like(k, dtype=bwd_torch_dtype) + dv = torch.empty_like(v, dtype=bwd_torch_dtype) + return dq, dk, dv + + @torch.compiler.allow_in_graph class _triton_attention_mxfp8(torch.autograd.Function): @@ -1084,8 +1778,32 @@ def forward( @staticmethod def backward(ctx, *grad_outputs): - # Forward-only, matching the MXFP4 reference. Backward is not implemented. - return None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None + do = grad_outputs[0] + q, k, v, o, softmax_lse, alibi_slopes, bias, q_scale, k_scale, v_scale = ctx.saved_tensors + assert bias is None, "MXFP8 attention backward does not support bias yet." + assert alibi_slopes is None, "MXFP8 attention backward does not support alibi yet." + assert ctx.dropout_p == 0.0, "MXFP8 attention backward does not support dropout yet." + + dq, dk, dv = torch.ops.alto.attention_mxfp8_backward_triton_impl( + do, + q, + k, + v, + o, + softmax_lse, + q_scale, + k_scale, + v_scale, + sm_scale=ctx.sm_scale, + causal=ctx.causal, + layout=ctx.layout, + cu_seqlens_q=ctx.cu_seqlens_q, + cu_seqlens_k=ctx.cu_seqlens_k, + max_seqlen_q=ctx.max_seqlens_q, + max_seqlen_k=ctx.max_seqlens_k, + use_exp2=ctx.use_exp2, + ) + return dq, dk, dv, None, None, None, None, None, None, None, None, None, None, None, None @attention_mxfp8_forward_triton_impl.register_fake diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index f5c0fc44..45d59b05 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -23,7 +23,12 @@ from tabulate import tabulate import torch -from .utils import calc_snr, calc_cossim, mxfp8_attention_forward_reference +from .utils import ( + calc_snr, + calc_cossim, + mxfp8_attention_forward_reference, + mxfp8_attention_backward_reference, +) # Reuse the committed shape grid so Layer 2 stays in lock-step with Layer 3. from .test_mxfp8_attention import AttnConfig, test_cases @@ -131,3 +136,116 @@ def test_kernel_matches_reference(config, causal): # Threshold placeholder — calibrate on the first CDNA4 run. assert sim > 0.99, f"kernel vs golden cosine-sim too low: {sim}" assert snr > 30, f"kernel vs golden SNR too low: {snr}" + + +# --------------------------------------------------------------------------- +# Backward Layer 1 — golden reference vs autograd fp32 SDPA (pure PyTorch, CPU) +# --------------------------------------------------------------------------- + +def _make_do_bhsd(batch, config, device, dtype): + torch.manual_seed(4321) + return torch.randn((batch, config.num_head_q, config.seqlen_q, config.head_dim_v), device=device, dtype=dtype) + + +@pytest.mark.parametrize("config", reference_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_backward_reference_matches_sdpa(config, causal): + """Backward golden reference (mxfp8 quant) must track autograd fp32 SDPA grads. + + Both sides use top-left causal, so this is valid for non-square shapes too. + The gap is the mxfp8 quantization error in the gradients; runs on CPU. + """ + device = "cuda" if torch.cuda.is_available() else "cpu" + dtype = torch.float32 if device == "cpu" else torch.bfloat16 + n_rep = config.num_head_q // config.num_head_kv + + q, k, v = _make_qkv_bhsd(2, config, device, dtype) + do = _make_do_bhsd(2, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + # Autograd fp32 SDPA backward — the "ideal" (unquantized) gradient. + qg = q.clone().requires_grad_(True) + kg = k.clone().requires_grad_(True) + vg = v.clone().requires_grad_(True) + o_auto = torch.nn.functional.scaled_dot_product_attention( + qg, kg, vg, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) + o_auto.backward(do) + + # Golden reference backward — quantized operands, forward-ref o / lse. + o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + dq, dk, dv = mxfp8_attention_backward_reference(q, k, v, do, o_ref, lse_ref, sm_scale, causal) + + rows = [] + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + rows.append([f"{name} (golden vs sdpa)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > 0.97, f"{name} golden-vs-sdpa cosine-sim too low: {sim}" + assert snr > 8, f"{name} golden-vs-sdpa SNR too low: {snr}" + + +# --------------------------------------------------------------------------- +# Backward Layer 2 — kernel vs golden reference +# Stage-1 backward dequantizes e4m3 -> bf16 and uses plain tl.dot (no +# tl.dot_scaled), so this runs on MI250 too. It calls the backward op directly +# with forward-reference o/lse, bypassing the forward kernel's tl.dot_scaled. +# --------------------------------------------------------------------------- + +@cuda_only +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True]) +def test_backward_kernel_matches_reference(config, causal): + """Backward kernel vs golden reference — isolates port bugs from quant error.""" + from alto.kernels.mxfp8.mxfp8_quantization import convert_to_mxfp8 + + device = "cuda" + dtype = torch.bfloat16 + q, k, v = _make_qkv_bhsd(4, config, device, dtype) + do = _make_do_bhsd(4, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + + # Quantize operands the way the forward autograd does, then run the backward op. + q8, q_scale = convert_to_mxfp8(q, mxfp_format="e4m3", axis=-1, is_2d_block=True) + k8, k_scale = convert_to_mxfp8(k, mxfp_format="e4m3", axis=-1, is_2d_block=True) + v8, v_scale = convert_to_mxfp8(v, mxfp_format="e4m3", axis=-1, is_2d_block=True) + + dq_k, dk_k, dv_k = torch.ops.alto.attention_mxfp8_backward_triton_impl( + do.contiguous(), + q8.contiguous(), + k8.contiguous(), + v8.contiguous(), + o_ref.contiguous(), + lse_ref.contiguous(), + q_scale, + k_scale, + v_scale, + sm_scale=sm_scale, + causal=causal, + layout="bhsd", + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlen_q=config.seqlen_q, + max_seqlen_k=config.seqlen_kv, + use_exp2=True, + ) + dq_r, dk_r, dv_r = mxfp8_attention_backward_reference(q, k, v, do, o_ref, lse_ref, sm_scale, causal) + + rows = [] + pairs = [("dQ", dq_r, dq_k), ("dK", dk_r, dk_k), ("dV", dv_r, dv_k)] + for name, ref, got in pairs: + rows.append([f"{name} (kernel vs golden)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + # Threshold placeholder — calibrate on the first run. + for name, ref, got in pairs: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > 0.99, f"{name} kernel-vs-golden cosine-sim too low: {sim}" + assert snr > 25, f"{name} kernel-vs-golden SNR too low: {snr}" diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index cde68d2f..5990fe84 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -287,3 +287,93 @@ def mxfp8_attention_forward_reference( softmax_lse = torch.where(fully_masked, torch.zeros_like(softmax_lse), softmax_lse) return o.to(q.dtype), softmax_lse + + +def mxfp8_attention_backward_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + do: torch.Tensor, + o: torch.Tensor, + softmax_lse: torch.Tensor, + sm_scale: float, + causal: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Pure-PyTorch golden reference for the MXFP8 (e4m3) flash-attention backward. + + Mirrors the Stage-1 backward kernel exactly so a kernel-vs-this-reference gap + isolates Triton port bugs from mxfp8 quantization error: + + * Q/K/V are quantized 2D-block along head_dim then dequantized — the same + operands the (Stage-1) kernel runs on. + * The softmax probabilities ``P`` are recomputed in fp32 from the *saved* + ``softmax_lse`` (``P = exp(S - lse)``); backward does **not** re-quantize + ``P`` (matching the kernel). + * ``o`` / ``softmax_lse`` must be the values the forward produced (feed the + outputs of ``mxfp8_attention_forward_reference``), so ``delta = rowsum(o * + do)`` and ``P`` are consistent with the forward pass. + + Uses the same top-left causal mask as the forward reference, so causal tests + must keep ``seqlen_q == seqlen_k`` (top-left == bottom-right there). + + Standard FlashAttention-v2 backward:: + + dV = Pᵀ @ dO + dP = dO @ Vᵀ + delta = rowsum(O * dO) + dS = P * (dP - delta) + dQ = sm_scale * dS @ K + dK = sm_scale * dSᵀ @ Q + + Args: + q, k, v: ``[batch, nheads, seqlen, head_dim]`` (bhsd). GQA supported. + do: upstream gradient, same shape/layout as ``o``. + o, softmax_lse: forward outputs (see above). + sm_scale: softmax scale. + causal: top-left causal mask. + + Returns: + (dq, dk, dv) matching the shapes of q, k, v. + """ + b, hq, sq, dqk = q.shape + _, hk, sk, dvv = v.shape + assert hq % hk == 0, f"nheads_q ({hq}) must be a multiple of nheads_k ({hk})" + n_rep = hq // hk + + q_dq = _mxfp8_qdq(q, axis=-1, is_2d_block=True).to(torch.float32) + k_dq = _mxfp8_qdq(k, axis=-1, is_2d_block=True).to(torch.float32) + v_dq = _mxfp8_qdq(v, axis=-1, is_2d_block=True).to(torch.float32) + if n_rep > 1: + k_dq = k_dq.repeat_interleave(n_rep, dim=1) + v_dq = v_dq.repeat_interleave(n_rep, dim=1) + + do_f = do.to(torch.float32) + o_f = o.to(torch.float32) + lse = softmax_lse.to(torch.float32) + + s = torch.matmul(q_dq, k_dq.transpose(-1, -2)) * sm_scale # [b, hq, sq, sk] + if causal: + q_pos = torch.arange(sq, device=q.device) + k_pos = torch.arange(sk, device=q.device) + allowed = k_pos[None, :] <= q_pos[:, None] # [sq, sk], top-left + s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, float("-inf"))) + + p = torch.exp(s - lse[..., None]) # softmax probabilities + p = torch.nan_to_num(p, nan=0.0, posinf=0.0, neginf=0.0) + + delta = (o_f * do_f).sum(dim=-1) # [b, hq, sq] + dp = torch.matmul(do_f, v_dq.transpose(-1, -2)) # [b, hq, sq, sk] + ds = p * (dp - delta[..., None]) # [b, hq, sq, sk] + + dv_full = torch.matmul(p.transpose(-1, -2), do_f) # [b, hq, sk, dv] + dq = sm_scale * torch.matmul(ds, k_dq) # [b, hq, sq, dqk] + dk_full = sm_scale * torch.matmul(ds.transpose(-1, -2), q_dq) # [b, hq, sk, dqk] + + if n_rep > 1: + dk = dk_full.view(b, hk, n_rep, sk, dqk).sum(dim=2) + dv = dv_full.view(b, hk, n_rep, sk, dvv).sum(dim=2) + else: + dk = dk_full + dv = dv_full + + return dq, dk, dv From d31a59160724beb6c5d8fe30c89cbda0b5a3df77 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 15 Jul 2026 00:53:48 -0500 Subject: [PATCH 05/14] fix: use instantiated tl.constexpr for backward RCP_LN2 --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 15 ++++++++++++++- .../kernels/mxfp8/triton_flash_attention_mxfp8.py | 2 +- 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 1f12e4bd..cd96f2f3 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -334,6 +334,19 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - **阶段1 backward 未经任何硬件验证**(本机无 torch);下一步在另一台机跑第 1 层(CPU)+ 第 2 层(MI250 即可,无需 CDNA4)拿反馈。 - 阶段2(dot_scaled + MX 量化)仍待 CDNA4。 +### 10.1ter 截至 2026-07-15(backward 阶段1 于 MI250 验通) + +**已在 MI250(gfx90a)Docker 跑通 backward 两层测试:** + +- **第 1 层 `test_backward_reference_matches_sdpa`:6/6 passed。** +- **第 2 层 `test_backward_kernel_matches_reference`:7/7 passed**(全 `test_cases` 网格),kernel vs 黄金参照 SNR≈55.5 dB、cos≈0.999999。证明阶段1 骨架的 FA v2 backward 数学 + strides + GQA + masking 移植正确。**注意:这验的是移植正确性,不是 mxfp8 backward 数值过关**——后者要阶段2 的 `tl.dot_scaled` + MX 量化,仍待 CDNA4。 + +**修掉一个真 bug(与 §10.1 forward 那次同类):** + +- **Triton constexpr 全局变量注解式写法**:backward 新增的模块级 `RCP_LN2: tl.constexpr = 1.4426950408889634` 被两个 backward inner kernel(`_attn_bwd_dq_inner` / `_attn_bwd_dkdv_inner`)访问时报 `NameError: Cannot access global variable RCP_LN2`。**修复**:改为实例化写法 `RCP_LN2 = tl.constexpr(1.4426950408889634)`。(forward 函数体内的局部 `RCP_LN2` 是局部变量,不受影响,未动。) + +**潜伏坑(非 bug,记录备查):** backward kernel 的 causal 用 **bottom-right**(`col_offset = N_CTX_Q - N_CTX_K`,与 forward kernel 一致),黄金参照用 **top-left**。当前 `test_cases` 全是方阵(seqlen_q==seqlen_kv),两者等价故测不出差异;若日后加**非方阵 causal** 用例,kernel 与参照会对不上,需统一对齐方式。 + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 参照 + 两层测试已就位待跑)。真正风险集中在阶段2 的 5-dot MX scale 轴对齐。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 两层已于 MI250 验 6/6 + 7/7,SNR≈55dB)。真正风险集中在阶段2 的 5-dot MX scale 轴对齐(待 CDNA4)。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index f4c60671..1a18ca93 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -1023,7 +1023,7 @@ def attention_mxfp8_forward_triton_impl( return o, softmax_lse, exp_scores -RCP_LN2: tl.constexpr = 1.4426950408889634 +RCP_LN2 = tl.constexpr(1.4426950408889634) @triton.jit From 600e3296a9d1bba2b8617e9e3f2ae3aaaa969f5c Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 15 Jul 2026 01:22:10 -0500 Subject: [PATCH 06/14] test: add stage-2 mxfp8 backward reference, verify all-e4m3 viable --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 50 ++++++++- .../mxfp8/test_mxfp8_attention_reference.py | 52 +++++++++ tests/unittest/mxfp8/utils.py | 106 ++++++++++++++++++ 3 files changed, 207 insertions(+), 1 deletion(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index cd96f2f3..12d8e838 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -268,6 +268,43 @@ def triton_attention_mxfp8( 把 `bwd_grad_format` 切 `"e5m2"` → dO/dP/dS 的量化与对应 `tl.dot_scaled` 的 `FORMAT_ID` 改 1,kernel 本体不动。这是 §5「backward 数值不过关」风险的兜底。 +### 9.5 阶段2 具体方案(2026-07-15 细化,7-dot 量化轴对齐) + +阶段1(bf16 骨架)已于 MI250 验通移植正确性(§10.1ter)。阶段2 = 把 backward 两个 inner kernel 里的 7 个 `tl.dot` 逐个换成 `tl.dot_scaled`,每个 dot 的两个 operand **沿它的 reduction 轴** MX e4m3 量化。 + +**7 个 dot 与量化轴:** + +| dot | kernel | 计算 | reduction 轴 | LHS 量化轴 | RHS 量化轴 | +|---|---|---|---|---|---| +| a | dkdv | `qk = q@kᵀ` | head_dim | q 沿 head_dim | k 沿 head_dim | +| b | dkdv | `dp = do@vᵀ` | head_dim_v | do 沿 head_dim_v | v 沿 head_dim_v | +| c | dkdv | `dv += pᵀ@do` | seqlen_q (BLOCK_M) | p 沿 seqlen_q | do 沿 seqlen_q | +| d | dkdv | `dk += dsᵀ@q` | seqlen_q (BLOCK_M) | ds 沿 seqlen_q | q 沿 seqlen_q | +| e | dq | `qk = q@kᵀ` | head_dim | 同 a | 同 a | +| f | dq | `dp = do@vᵀ` | head_dim_v | 同 b | 同 b | +| g | dq | `dq += ds@k` | seqlen_k (BLOCK_N) | ds 沿 seqlen_k | k 沿 seqlen_k | + +**同一 operand 需多套量化(plan 早点名的核心工作量):** + +- **q**:沿 head_dim(a/e,复用 forward 存的 2D-block e4m3)+ 沿 seqlen_q(d,新量化)→ **2 套** +- **k**:沿 head_dim(a/e,复用 forward)+ 沿 seqlen_k(g,新量化)→ **2 套** +- **do**:沿 head_dim_v(b/f)+ 沿 seqlen_q(c)→ **2 套**(均在 backward 入口/kernel 内新量化) +- **v**:仅 head_dim_v(b/f,复用 forward)→ 1 套 +- **p / ds**:kernel 内现算现量化,沿各自 reduction 轴(p 沿 seqlen_q;ds 沿 seqlen_q 供 d、沿 seqlen_k 供 g → ds 2 套) + +**2D-block vs 1D-block 轴规则(沿用 forward 已定的约定,避免再拍脑袋):** + +- reduction 轴 = **head_dim / head_dim_v** 的 operand → **2D-block**(对齐 forward Q/K/V 的 `is_2d_block=True, axis=-1`) +- reduction 轴 = **seqlen_q / seqlen_k** 的 operand → **1D-block 沿该 seqlen 轴**(对齐 forward P 沿 `BLOCK_N` 的 `IS_2D_BLOCK=False`) + +**落地顺序(先参照后 kernel,不盲写):** + +1. **先写阶段2 PyTorch 参照**(`mxfp8_attention_backward_reference_stage2`):把上表 7 个 dot 的 operand 按 2D/1D 轴规则做 quantize→dequantize,再 fp32 matmul(数值上模拟 `dot_scaled`)。**不含 `tl.dot_scaled`,MI250 / CPU 即可跑**,对 bf16 SDPA autograd 度量 dQ/dK/dV。作用有二:① 先探清全 e4m3 grad 的数值可行性(§0 underflow 风险,若这步就不过关 → 直接切 e5m2,不必等硬件);② 成为 kernel 的**验收合同**。 +2. **再改 kernel**:逐 dot 换 `tl.dot_scaled` + 接入上述量化(backward 入口补 do 两套量化、q/k 的 seqlen 轴量化;p/ds kernel 内 `_calculate_scales`+`_quantize_fp8`)。 +3. **CDNA4 验收**:kernel vs 步骤1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差)。 + +> ⚠️ 阶段2 的 kernel 改动在 MI250 上**一行都验不了**(无 `tl.dot_scaled`)。步骤1 的 PyTorch 参照是这段"盲写期"唯一能拿到的数值信号,必须先做。 + ### 9.4 save_for_backward 清单 forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias`(确保 backward 三个 kernel 所需张量齐全,避免中途改 forward 签名)。 @@ -347,6 +384,17 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` **潜伏坑(非 bug,记录备查):** backward kernel 的 causal 用 **bottom-right**(`col_offset = N_CTX_Q - N_CTX_K`,与 forward kernel 一致),黄金参照用 **top-left**。当前 `test_cases` 全是方阵(seqlen_q==seqlen_kv),两者等价故测不出差异;若日后加**非方阵 causal** 用例,kernel 与参照会对不上,需统一对齐方式。 +### 10.1quater 截至 2026-07-15(阶段2 数值先探路,全 e4m3 判定可行) + +**已落地(代码)+ 已于 MI250 验:** + +- **阶段2 PyTorch 参照** `mxfp8_attention_backward_reference_stage2`(`tests/unittest/mxfp8/utils.py`):按 §9.5 的 7-dot 量化轴表,把每个 backward matmul 的 operand 沿其 reduction 轴做 qdq(head_dim 轴 2D-block、seqlen 轴 1D-block),fp32 matmul 模拟 `tl.dot_scaled`。不含 `tl.dot_scaled`,MI250/CPU 可跑。 +- **测试** `test_backward_stage2_reference_matches_sdpa`(stage-2 参照 vs autograd fp32 SDPA):MI250 **6/6 passed**。 + +**数值结论(关键决策依据):** 全 e4m3 backward 对 SDPA 的误差——**dQ/dK SNR≈23.3 dB、cossim≈0.9977;dV SNR≈25 dB、cossim≈0.9985**。相比阶段1(只量化 Q/K/V,SNR≈55dB),额外量化 dO/P/dS 把 SNR 拉到 ~23dB,但 cossim 稳在 0.998——**未出现 §0 担心的 dP/dS 长尾 underflow 崩盘**。⇒ **全 e4m3 backward 判定可行,不预防性切 e5m2**;e5m2 仍作为 §9.3 兜底保留。测试阈值据此标定为 `cossim>0.995`/`SNR>18`(留裕度)。 + +**下一步:** 按 §9.5 步骤2 改 kernel(逐 dot 换 `tl.dot_scaled` + 接量化轴),以本参照为验收合同,待 CDNA4 做步骤3 验收。 + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 两层已于 MI250 验 6/6 + 7/7,SNR≈55dB)。真正风险集中在阶段2 的 5-dot MX scale 轴对齐(待 CDNA4)。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 阶段1 两层 6/6 + 7/7 SNR≈55dB;backward 阶段2 全量化参照 vs SDPA 6/6,cossim≈0.998/SNR≈23dB → **全 e4m3 判定可行**)。阶段2 kernel(逐 dot 换 `tl.dot_scaled`)待 CDNA4,以阶段2 参照为验收合同。 diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index 45d59b05..5408424d 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -28,6 +28,7 @@ calc_cossim, mxfp8_attention_forward_reference, mxfp8_attention_backward_reference, + mxfp8_attention_backward_reference_stage2, ) # Reuse the committed shape grid so Layer 2 stays in lock-step with Layer 3. from .test_mxfp8_attention import AttnConfig, test_cases @@ -188,6 +189,57 @@ def test_backward_reference_matches_sdpa(config, causal): assert snr > 8, f"{name} golden-vs-sdpa SNR too low: {snr}" +# --------------------------------------------------------------------------- +# Backward Stage-2 Layer 1 — full-quant reference vs autograd fp32 SDPA +# Quantizes every backward matmul operand along its reduction axis (a +# numerical model of tl.dot_scaled). The gap vs SDPA is the *full* mxfp8 +# backward quantization error — the signal that decides whether all-e4m3 is +# viable or must go e5m2. Pure PyTorch, runs on CPU / MI250 (no CDNA4). +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("config", reference_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_backward_stage2_reference_matches_sdpa(config, causal): + """Stage-2 (all-operand-quantized) backward reference vs autograd fp32 SDPA. + + This is the pre-hardware numerical de-risk for the plan §9.5 stage-2 kernel: + if all-e4m3 grads already diverge here, the kernel can't do better and we + should switch dO/dP/dS to e5m2 before writing any tl.dot_scaled code. + """ + device = "cuda" if torch.cuda.is_available() else "cpu" + dtype = torch.float32 if device == "cpu" else torch.bfloat16 + n_rep = config.num_head_q // config.num_head_kv + + q, k, v = _make_qkv_bhsd(2, config, device, dtype) + do = _make_do_bhsd(2, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + qg = q.clone().requires_grad_(True) + kg = k.clone().requires_grad_(True) + vg = v.clone().requires_grad_(True) + o_auto = torch.nn.functional.scaled_dot_product_attention( + qg, kg, vg, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) + o_auto.backward(do) + + o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + dq, dk, dv = mxfp8_attention_backward_reference_stage2(q, k, v, do, o_ref, lse_ref, sm_scale, causal) + + rows = [] + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + rows.append([f"{name} (stage2 vs sdpa)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + # Calibrated on MI250 (first run): all-e4m3 backward lands at + # cossim~0.998 / SNR~23-26 dB across dQ/dK/dV — no e5m2 needed. Thresholds + # keep margin below those numbers. + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > 0.995, f"{name} stage2-vs-sdpa cosine-sim too low: {sim}" + assert snr > 18, f"{name} stage2-vs-sdpa SNR too low: {snr}" + + # --------------------------------------------------------------------------- # Backward Layer 2 — kernel vs golden reference # Stage-1 backward dequantizes e4m3 -> bf16 and uses plain tl.dot (no diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index 5990fe84..8ea5af62 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -377,3 +377,109 @@ def mxfp8_attention_backward_reference( dv = dv_full return dq, dk, dv + + +def mxfp8_attention_backward_reference_stage2( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + do: torch.Tensor, + o: torch.Tensor, + softmax_lse: torch.Tensor, + sm_scale: float, + causal: bool, + block_size: int = BLOCK_SIZE_DEFAULT, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Pure-PyTorch golden reference for the **Stage-2** MXFP8 (e4m3) backward. + + Where Stage-1 (``mxfp8_attention_backward_reference``) quantizes only Q/K/V + (2D-block along head_dim) and runs every backward matmul in fp32, Stage-2 + additionally quantizes **each matmul operand along that dot's reduction + axis** — a numerical model of what ``tl.dot_scaled`` will do in the kernel. + A kernel-vs-this gap then isolates Triton port bugs from mxfp8 error, and a + this-vs-bf16-SDPA gap measures the *full* mxfp8 backward quantization error + (the signal that decides whether all-e4m3 backward is viable or must go + e5m2), all **without CDNA4** (no ``tl.dot_scaled`` here; plain fp32 matmul + on dequantized operands). + + Quantization axis per dot (see plan §9.5):: + + dP = dO @ Vᵀ reduction head_dim_v : dO,V 2D-block axis=-1 + S = Q @ Kᵀ reduction head_dim : Q,K 2D-block axis=-1 (reused from fwd) + dV = Pᵀ @ dO reduction seqlen_q : P,dO 1D-block along sq + dQ = dS @ K reduction seqlen_k : dS,K 1D-block along sk + dK = dSᵀ @ Q reduction seqlen_q : dS,Q 1D-block along sq + + 2D/1D rule (mirrors the forward decisions): operands whose reduction axis is + head_dim/head_dim_v use 2D-block (like fwd Q/K/V); operands whose reduction + axis is a sequence length use 1D-block along that seqlen (like fwd P). + + Same top-left causal mask as the other references, so causal tests must keep + ``seqlen_q == seqlen_k``. + + Args mirror ``mxfp8_attention_backward_reference``; ``o`` / ``softmax_lse`` + must be the forward-reference outputs. Returns ``(dq, dk, dv)``. + """ + b, hq, sq, dqk = q.shape + _, hk, sk, dvv = v.shape + assert hq % hk == 0, f"nheads_q ({hq}) must be a multiple of nheads_k ({hk})" + n_rep = hq // hk + + def qdq(x, axis, is_2d_block): + return _mxfp8_qdq(x.to(torch.float32), axis=axis, is_2d_block=is_2d_block, block_size=block_size) + + # --- head_dim-reduction operands: 2D-block along head_dim (reuse fwd convention) --- + q_hd = qdq(q, axis=-1, is_2d_block=True) # for S (dots a/e) + k_hd = qdq(k, axis=-1, is_2d_block=True) # for S + v_hd = qdq(v, axis=-1, is_2d_block=True) # for dP (dots b/f) + do_hd = qdq(do, axis=-1, is_2d_block=True) # dO for dP, along head_dim_v + + # --- seqlen-reduction operands: 1D-block along the seqlen axis --- + q_sq = qdq(q, axis=-2, is_2d_block=False) # for dK (dot d) + k_sk = qdq(k, axis=-2, is_2d_block=False) # for dQ (dot g) + do_sq = qdq(do, axis=-2, is_2d_block=False) # dO for dV (dot c) + + if n_rep > 1: + k_hd = k_hd.repeat_interleave(n_rep, dim=1) + v_hd = v_hd.repeat_interleave(n_rep, dim=1) + k_sk = k_sk.repeat_interleave(n_rep, dim=1) + + lse = softmax_lse.to(torch.float32) + do_f = do.to(torch.float32) + o_f = o.to(torch.float32) + + # S -> P (recompute from saved lse; P feeding dS is elementwise, not quantized). + s = torch.matmul(q_hd, k_hd.transpose(-1, -2)) * sm_scale + if causal: + q_pos = torch.arange(sq, device=q.device) + k_pos = torch.arange(sk, device=q.device) + allowed = k_pos[None, :] <= q_pos[:, None] + s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, float("-inf"))) + p = torch.exp(s - lse[..., None]) + p = torch.nan_to_num(p, nan=0.0, posinf=0.0, neginf=0.0) + + # dP = dO @ Vᵀ (reduction head_dim_v). + dp = torch.matmul(do_hd, v_hd.transpose(-1, -2)) + delta = (o_f * do_f).sum(dim=-1) + ds = p * (dp - delta[..., None]) # fp32 elementwise + + # dV = Pᵀ @ dO (reduction seqlen_q): P and dO 1D-block along sq. + p_sq = qdq(p, axis=-2, is_2d_block=False) + dv_full = torch.matmul(p_sq.transpose(-1, -2), do_sq) + + # dQ = sm_scale * dS @ K (reduction seqlen_k): dS along sk, K along sk. + ds_sk = qdq(ds, axis=-1, is_2d_block=False) + dq = sm_scale * torch.matmul(ds_sk, k_sk) + + # dK = sm_scale * dSᵀ @ Q (reduction seqlen_q): dS along sq, Q along sq. + ds_sq = qdq(ds, axis=-2, is_2d_block=False) + dk_full = sm_scale * torch.matmul(ds_sq.transpose(-1, -2), q_sq) + + if n_rep > 1: + dk = dk_full.view(b, hk, n_rep, sk, dqk).sum(dim=2) + dv = dv_full.view(b, hk, n_rep, sk, dvv).sum(dim=2) + else: + dk = dk_full + dv = dv_full + + return dq, dk, dv From ea590dc3d845c6b47186671c7362913c9753bce0 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Thu, 16 Jul 2026 03:00:47 -0500 Subject: [PATCH 07/14] feat: reuse forward e4m3 q/k/v scales in mxfp8 attention backward --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 88 +++- .../mxfp8/triton_flash_attention_mxfp8.py | 430 ++++++++++++++++-- .../mxfp8/test_mxfp8_attention_reference.py | 115 ----- tests/unittest/mxfp8/utils.py | 106 ----- 4 files changed, 449 insertions(+), 290 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 12d8e838..86af77ec 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -139,12 +139,14 @@ def triton_attention_mxfp8( ### Backward(见 §9 推导) +备注列为 **A1**(q/k/v 复用 forward 存的 e4m3 + 2D-block scale,dO/P/dS 现场量化): + | dot | 计算 | reduction 维 | 备注 | |---|---|---|---| -| dV | `dV += Pᵀ @ dO` | seqlen_q (BLOCK_M) | P 重算,dO 现场量化沿 M | -| dP | `dP = dO @ Vᵀ` | head_dim_v | V 复用 fwd e4m3 | -| dK | `dK += dSᵀ @ Q` | seqlen_q (BLOCK_M) | dS 现场量化沿 M | -| dQ | `dQ += dS @ K` | seqlen_k (BLOCK_N) | dS 现场量化沿 N、K 复用 fwd e4m3 | +| dV | `dV += Pᵀ @ dO` | seqlen_q (BLOCK_M) | P/dO 现场 1D 量化沿 M | +| dP | `dP = dO @ Vᵀ` | head_dim_v | V 复用 fwd 2D scale(`_load_scale_hd`);dO 现场量化 | +| dK | `dK += dSᵀ @ Q` | seqlen_q (BLOCK_M) | dS 现场 1D 量化沿 M;**Q 复用 fwd 2D scale 转置 re-index**(`_load_scale_sq`) | +| dQ | `dQ += dS @ K` | seqlen_k (BLOCK_N) | dS 现场 1D 量化沿 N;**K 复用 fwd 2D scale 转置 re-index**(`_load_scale_sq`) | > ⚠️ **QK 一次 `dot_scaled` 跨 head_dim/32 个 32-wide scale group,是最大数值风险点**(group gemm plan §1.3 / §5 专门警告)。mxfp4 已接受此误差(未硬断言精度);mxfp8 更敏感。V1 无 fallback,改由 §8 的 **PyTorch 模拟 reference** 度量此误差;若发散,退路见 §5 风险 1。backward 的 dot 同样跨 group,同法度量。 @@ -286,24 +288,30 @@ def triton_attention_mxfp8( **同一 operand 需多套量化(plan 早点名的核心工作量):** -- **q**:沿 head_dim(a/e,复用 forward 存的 2D-block e4m3)+ 沿 seqlen_q(d,新量化)→ **2 套** -- **k**:沿 head_dim(a/e,复用 forward)+ 沿 seqlen_k(g,新量化)→ **2 套** -- **do**:沿 head_dim_v(b/f)+ 沿 seqlen_q(c)→ **2 套**(均在 backward 入口/kernel 内新量化) -- **v**:仅 head_dim_v(b/f,复用 forward)→ 1 套 +- **q**:沿 head_dim(a/e)+ 沿 seqlen_q(d)→ **2 套** +- **k**:沿 head_dim(a/e)+ 沿 seqlen_k(g)→ **2 套** +- **do**:沿 head_dim_v(b/f)+ 沿 seqlen_q(c)→ **2 套** +- **v**:仅 head_dim_v(b/f)→ 1 套 - **p / ds**:kernel 内现算现量化,沿各自 reduction 轴(p 沿 seqlen_q;ds 沿 seqlen_q 供 d、沿 seqlen_k 供 g → ds 2 套) -**2D-block vs 1D-block 轴规则(沿用 forward 已定的约定,避免再拍脑袋):** +**量化 block 规则(2026-07-16 定案:A1——q/k/v 复用 forward 2D-block,dO/P/dS 现场 1D per-row):** + +- **q/k/v(forward 存过的)**:**直接复用 forward 的紧凑 2D-block scale `[.., seqlen/32, head_dim/32]`**,backward 不再重量化。每个 dot 用指针索引把这块 2D scale 广播成 `tl.dot_scaled` 要的 `[outer, reduction/32]`: + - **head_dim 收缩的 dot(a/b/e/f)**:`scale[outer//32, dgroup]`,与 forward QK 的 `qs_ptrs` 广播完全同款(`_load_scale_hd`)。 + - **seqlen 收缩的 dot(d/g)**:**转置对称复用**——同一个 32×32 块的 scale 换个轴索引成 `[head_dim, seqlen_block/32]`(`_load_scale_sq`),因为一个 32×32 块只有一个 scale,两轴共用。 +- **dO / P / dS(backward 新产生)**:forward 没存过,仍**现场 1D per-row 沿其 reduction 轴**量化(`_mx_quant`,`IS_2D_BLOCK=False`),scale `[outer, reduction/32]` 直接匹配 `dot_scaled`。 +- 早先(2026-07-15)为省事对所有 operand 统一 1D per-row(含 q/k/v 重量化),即 option A;2026-07-16 翻案改 A1(见 §10.1sexies / AB 决策稿 banner):直接复用 forward 那份 e4m3+2D scale,**零重量化、与 forward 逐位一致、无双重量化**,代价是 kernel 内两个 scale 广播辅助。 -- reduction 轴 = **head_dim / head_dim_v** 的 operand → **2D-block**(对齐 forward Q/K/V 的 `is_2d_block=True, axis=-1`) -- reduction 轴 = **seqlen_q / seqlen_k** 的 operand → **1D-block 沿该 seqlen 轴**(对齐 forward P 沿 `BLOCK_N` 的 `IS_2D_BLOCK=False`) +**operand 来源(A1,2026-07-16 定,推翻 option A):** forward 存 e4m3 q/k/v + 2D-block scale(省显存,不变);backward driver **直接把 saved e4m3 + scale 传进 kernel**(无入口 dequant),每个 dot 复用之(见上)。dO 是 backward 的 bf16 输入,现场量化。**A(入口 dequant + 逐 dot 重量化)已删除,B(存 bf16)不做。** -**落地顺序(先参照后 kernel,不盲写):** +**落地顺序(2026-07-16 改为「先 kernel 后参照」——需求方拍板,知情决策):** -1. **先写阶段2 PyTorch 参照**(`mxfp8_attention_backward_reference_stage2`):把上表 7 个 dot 的 operand 按 2D/1D 轴规则做 quantize→dequantize,再 fp32 matmul(数值上模拟 `dot_scaled`)。**不含 `tl.dot_scaled`,MI250 / CPU 即可跑**,对 bf16 SDPA autograd 度量 dQ/dK/dV。作用有二:① 先探清全 e4m3 grad 的数值可行性(§0 underflow 风险,若这步就不过关 → 直接切 e5m2,不必等硬件);② 成为 kernel 的**验收合同**。 -2. **再改 kernel**:逐 dot 换 `tl.dot_scaled` + 接入上述量化(backward 入口补 do 两套量化、q/k 的 seqlen 轴量化;p/ds kernel 内 `_calculate_scales`+`_quantize_fp8`)。 -3. **CDNA4 验收**:kernel vs 步骤1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差)。 +1. ✅ **删 A**:删除 A 的 kernel 路径(入口 `convert_from_mxfp8` + 逐 dot 重量化 q/k/v)与 A 的黄金参照 `mxfp8_attention_backward_reference_stage2`(含 `operand_source` A/B 开关)+ 其两个测试。保留 stage-1 参照 `mxfp8_attention_backward_reference`(FA 数学基准,格式无关)。 +2. ✅ **写 A1 kernel**:新增 `_load_scale_hd` / `_load_scale_sq`(从紧凑 2D scale 重建各 dot 的 scale tile);两个 kernel + 两个 inner 改为复用 saved e4m3 q/k/v + 2D scale;driver 传 q/k/v 的 scale 指针 + 12 个 scale stride;删去不再用的 `convert_from_mxfp8` import。dO/P/dS 仍 `_mx_quant`。**MI250 import OK、语法/签名通;`tl.dot_scaled` 不编译,整体待 CDNA4。** +3. ⏳ **写 A1 golden reference**(下一步):纯 PyTorch 复刻「saved e4m3 值 + 2D tile scale 按轴索引」,MI250/CPU 可跑,对 bf16 SDPA autograd 度量数值。 +4. ⏳ **CDNA4 验收**:kernel vs A1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差)。 -> ⚠️ 阶段2 的 kernel 改动在 MI250 上**一行都验不了**(无 `tl.dot_scaled`)。步骤1 的 PyTorch 参照是这段"盲写期"唯一能拿到的数值信号,必须先做。 +> ⚠️ A1 kernel 的 `tl.dot_scaled` + 2D-scale 指针广播在 MI250 **一行都验不了**,且按新顺序**当前连 A1 参照都还没写**——A1 kernel 正确性完全待 CDNA4 + A1 参照。其中 seqlen 轴的 `_load_scale_sq` 广播全仓无先例,是最大盲区。这是需求方明确知情后的决策(宁可不留 A 冗余)。 ### 9.4 save_for_backward 清单 @@ -358,7 +366,7 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - **dot 用高精度 bf16 `tl.dot` 占位**,每处留 `TODO(stage2)` 标注该 dot 的 reduction 轴。 - **入口用 `convert_from_mxfp8` 把 saved e4m3 q/k/v 反量化回 bf16**——backward 吃的正是 forward 用过的那份量化输入,且**不含 `tl.dot_scaled`,MI250 可跑**。 - causal 用 **bottom-right**(与 forward kernel 一致,非参照的 top-left;方阵下无差别)。 -- **阶段2 ❌ 未开始**:逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴的 MX e4m3 量化 + dO 量化。需 CDNA4 才能验。 +- **阶段2 ✅ 已落地(代码,`tl.dot_scaled` 未验)**:逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴的 MX e4m3 量化。详见 §10.1quinquies。需 CDNA4 才能验。 **backward 黄金参照 ✅ 已落地**:`tests/unittest/mxfp8/utils.py::mxfp8_attention_backward_reference`——纯 PyTorch 复刻阶段1 kernel(反量化 q/k/v、用 saved `lse` 在 fp32 重算 `P`、backward 不量化 P),标准 FA v2 backward(dV=Pᵀ@dO;dP=dO@Vᵀ;dS=P·(dP−delta);dQ=sm·dS@K;dK=sm·dSᵀ@Q)。喂 forward 参照的同一份 `o`/`lse` 时应与 kernel 高精度吻合。 @@ -388,13 +396,51 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` **已落地(代码)+ 已于 MI250 验:** -- **阶段2 PyTorch 参照** `mxfp8_attention_backward_reference_stage2`(`tests/unittest/mxfp8/utils.py`):按 §9.5 的 7-dot 量化轴表,把每个 backward matmul 的 operand 沿其 reduction 轴做 qdq(head_dim 轴 2D-block、seqlen 轴 1D-block),fp32 matmul 模拟 `tl.dot_scaled`。不含 `tl.dot_scaled`,MI250/CPU 可跑。 +- **阶段2 PyTorch 参照** `mxfp8_attention_backward_reference_stage2`(`tests/unittest/mxfp8/utils.py`):按 §9.5 的 7-dot 量化轴表,把每个 backward matmul 的 operand 沿其 reduction 轴做 qdq(**统一 1D per-row**,见下文修订),fp32 matmul 模拟 `tl.dot_scaled`。不含 `tl.dot_scaled`,MI250/CPU 可跑。 - **测试** `test_backward_stage2_reference_matches_sdpa`(stage-2 参照 vs autograd fp32 SDPA):MI250 **6/6 passed**。 -**数值结论(关键决策依据):** 全 e4m3 backward 对 SDPA 的误差——**dQ/dK SNR≈23.3 dB、cossim≈0.9977;dV SNR≈25 dB、cossim≈0.9985**。相比阶段1(只量化 Q/K/V,SNR≈55dB),额外量化 dO/P/dS 把 SNR 拉到 ~23dB,但 cossim 稳在 0.998——**未出现 §0 担心的 dP/dS 长尾 underflow 崩盘**。⇒ **全 e4m3 backward 判定可行,不预防性切 e5m2**;e5m2 仍作为 §9.3 兜底保留。测试阈值据此标定为 `cossim>0.995`/`SNR>18`(留裕度)。 +**数值结论(关键决策依据):** 全 e4m3 backward 对 SDPA 的误差——**dQ/dK SNR≈23 dB、cossim≈0.9975;dV SNR≈25 dB、cossim≈0.9984**。相比阶段1(只量化 Q/K/V,SNR≈55dB),额外量化 dO/P/dS 把 SNR 拉到 ~23dB,但 cossim 稳在 0.998——**未出现 §0 担心的 dP/dS 长尾 underflow 崩盘**。⇒ **全 e4m3 backward 判定可行,不预防性切 e5m2**;e5m2 仍作为 §9.3 兜底保留。测试阈值据此标定为 `cossim>0.995`/`SNR>18`(留裕度)。 + +### 10.1quinquies 截至 2026-07-15(阶段2 kernel 已写,待 CDNA4 验) + +**已落地(代码,`tl.dot_scaled` 部分未验):** 按 §9.5 步骤2 把 backward 两个 inner kernel 的 7 个 `tl.dot` 全换 `tl.dot_scaled`。 + +- **新增 `_mx_quant` 内联辅助**(`triton_flash_attention_mxfp8.py`):包 `_calculate_scales`+`_quantize_fp8`,返回 `(e4m3 tile, uint8 scale)`。 +- **量化布局定案:统一 1D per-row**(`IS_2D_BLOCK=False`),scale `[outer, reduction/32]` 直接匹配 `dot_scaled`。**修正了盲写中发现的真 bug**:早先想让 head_dim dot 复用 forward 的 2D-block(`IS_2D_BLOCK=True`),但那返回紧凑 `[M/32, K/32]`,形状对不上 `dot_scaled`(forward 是靠指针把 2D scale 广播成 `[M, K/32]` 才喂进去的)。统一 1D 消除该特殊情况,参照实测数值几乎无差。 +- **operand 来源 = option A**:driver 入口仍 `convert_from_mxfp8` 反量化 saved e4m3 q/k/v→bf16(不额外存 bf16,省显存),kernel 内每个 dot 沿其 reduction 轴 1D 重量化。k/v/q/do 的 head_dim 套在 **outer kernel 量化一次复用**;p/ds 与 seqlen 套在 **inner 内**量化(转置使 reduction 轴落在最后一维,复用 last-axis 的量化辅助)。 +- 两个 launch 传 `QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT` / `USE_ASM=is_cdna4()`。 +- **参照同步**:`stage2` 参照全部 operand 改 1D per-row;q_sq/k_sk 按 option A 从 `q_hd`/`k_hd`(dequant 的 e4m3)再量化(双重量化),dO 单次。重跑 **6/6 passed,SNR≈23dB** 不变。 +- **`test_backward_kernel_matches_reference` 改对比 `stage2` 参照**(kernel 已是阶段2);该测试随之变为 **CDNA4-only**(`tl.dot_scaled` 在 MI250 不编译),与 forward Layer 2 同性质。 + +> ⚠️ 阶段2 kernel 的 `tl.dot_scaled` 与 scale 指针布局在 MI250 **无法编译/验证**,全靠对标 forward 已验证的 `dot_scaled` 用法 + stage2 参照「盲写」。待 CDNA4 跑步骤3 验收(kernel vs stage2 参照隔离移植 bug;kernel vs bf16 SDPA 端到端)。 + +**下一步:** CDNA4 到位后跑步骤3 验收;届时校准 `test_backward_kernel_matches_reference` 阈值。 + +### 10.1sexies 截至 2026-07-16(推翻 option A,改 A1;A 全删;A1 kernel 已写) + +**决策翻案(需求方拍板,知情决策):** 上会后定 **B 完全不要、A 也不留(视为冗余回退代码)**,直接上 AB 决策稿 §5 当初「不建议」的 **A1**。理由:留 A 当备用是冗余;A1 复用 forward 那份 e4m3+2D scale = correct-by-construction、与 forward 逐位一致、无双重量化。代价(kernel 内 scale 指针广播)已实现。详见 `MXFP8_BACKWARD_AB_DECISION.md` 顶部 banner。 + +**已落地(代码,未经硬件验证):** + +- **删 A**:删掉入口 `convert_from_mxfp8` 反量化 + 逐 dot 重量化 q/k/v 的 A 路径;删掉 A 的黄金参照 `mxfp8_attention_backward_reference_stage2`(含 `operand_source` A/B 开关)及其两个测试(`test_backward_stage2_reference_matches_sdpa` / `test_backward_kernel_matches_reference`);删掉 `triton_flash_attention_mxfp8.py` 里不再用的 `convert_from_mxfp8` import。保留 stage-1 参照 `mxfp8_attention_backward_reference`(FA 数学基准,非 A 专属)。 +- **A1 kernel**: + - 新增 `_load_scale_hd`(head_dim 收缩,照抄 forward `//32` 广播)/ `_load_scale_sq`(seqlen 收缩,转置对称 re-index),从紧凑 2D scale `[.., seqlen/32, head_dim/32]` 重建各 dot 的 scale tile,带越界/padded-head mask。 + - `_bwd_kernel_dkdv` / `_bwd_kernel_dq` + 两个 inner:q/k/v 改为**加载 saved e4m3 + 复用 2D scale**(不再 `_mx_quant`);dot a/b/e/f 用 `_load_scale_hd`,dot d/g 用 `_load_scale_sq`;dO/P/dS 仍 `_mx_quant` 现场 1D 量化。 + - driver:去掉入口 dequant,改传 saved e4m3 q/k/v + 三个 2D scale 及其 12 个 stride 进两个 launch。 + +**MI250(gfx90a) Docker 验证(2026-07-16):** + +| 项 | 结果 | +|---|---| +| import `triton_flash_attention_mxfp8` | ✅ OK(A1 大改的签名/helper/装饰器解析注册通过) | +| 参照层 `test_mxfp8_attention_reference.py` | ✅ **12 passed**(含 stage-1 backward 参照 vs SDPA 6/6)→ 删 A 干净 | +| forward `test_kernel_matches_reference` | ❌ 7 failed(forward `attn_fwd` 的 `tl.dot_scaled` 在 gfx90a 编不过,**既有限制、非本次改动**;`git diff` 证实 forward kernel 本体无改动) | -**下一步:** 按 §9.5 步骤2 改 kernel(逐 dot 换 `tl.dot_scaled` + 接量化轴),以本参照为验收合同,待 CDNA4 做步骤3 验收。 +**仍待办:** +- **A1 golden reference 未写**(按新顺序「先 kernel 后参照」,下一步补)。当前 backward **无任何 A1 数值信号**。 +- **A1 kernel `tl.dot_scaled` + 2D-scale 广播待 CDNA4 验**;`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 +- padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架已用高精度 bf16 `tl.dot` 接通(入口反量化 e4m3,MI250 可跑),阶段2 再逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴 MX 量化(待 CDNA4)。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**,forward/backward 均有纯 PyTorch 黄金参照(forward Layer 1 已于 MI250 验 6/6;backward 阶段1 两层 6/6 + 7/7 SNR≈55dB;backward 阶段2 全量化参照 vs SDPA 6/6,cossim≈0.998/SNR≈23dB → **全 e4m3 判定可行**)。阶段2 kernel(逐 dot 换 `tl.dot_scaled`)待 CDNA4,以阶段2 参照为验收合同。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA 于 MI250 验 6/6(阶段2 全量化参照 vs SDPA 曾 6/6、cossim≈0.998/SNR≈23dB,已随 A 删除,A1 参照待补)。**A1 kernel 的 `tl.dot_scaled` + 2D-scale 广播在 MI250 编不了、当前无 A1 参照,正确性完全待 CDNA4 + A1 参照**(`_load_scale_sq` 为最大盲区)。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index 1a18ca93..23d88a42 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -20,7 +20,9 @@ The scale layout is identical to the FP4 kernel (uint8, ``[.., dim/32]``, 2D block), so all scale pointer arithmetic carries over unchanged. -Forward only, matching the FP4 reference (backward is not implemented). +Forward is a mechanical port of the FP4 kernel; backward is implemented from +the blockwise-fp8 backward structure with per-reduction-axis MXFP8 scales +(stage-2, ``tl.dot_scaled``; see the plan §9.5 and the backward section below). """ from typing import Tuple, Optional import os @@ -34,7 +36,6 @@ is_cdna4, _calculate_scales, _quantize_fp8, - convert_from_mxfp8, ) fwd_torch_dtype: tl.constexpr = torch.bfloat16 @@ -1026,6 +1027,127 @@ def attention_mxfp8_forward_triton_impl( RCP_LN2 = tl.constexpr(1.4426950408889634) +@triton.jit +def _mx_quant( + x, + BM: tl.constexpr, + BN: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + IS_2D_BLOCK: tl.constexpr, + USE_ASM: tl.constexpr, +): + """Quantize a [BM, BN] tile to e4m3 with MX block scales along the last axis. + + Wraps ``_calculate_scales`` + ``_quantize_fp8`` with the e4m3 constants, so + each backward dot can quantize an operand along its reduction axis (which the + caller places last, transposing the tile when needed). ``IS_2D_BLOCK=True`` + groups both axes (head_dim reduction, like fwd Q/K/V); ``False`` is a 1D block + along the last axis (seqlen reduction, like fwd P). Non-SR path, so philox is + unused. + """ + scales = _calculate_scales( + x, + BLOCK_M=BM, + BLOCK_N=BN, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + target_max_pow2=E4M3_TARGET_MAX_POW2, + mbits=E4M3_MBITS, + IS_2D_BLOCK=IS_2D_BLOCK, + ) + xq = _quantize_fp8( + x, + scales, + philox_seed, + 0, + BLOCK_M=BM, + BLOCK_N=BN, + QUANT_BLOCK_SIZE=QUANT_BLOCK_SIZE, + FP8_FORMAT=E4M3_FORMAT_ID, + IS_2D_BLOCK=IS_2D_BLOCK, + USE_ASM=USE_ASM, + USE_SR=False, + ) + return xq, scales + + +# Backward operands split into two families (A1, plan §9.5 / AB decision): +# +# * Q/K/V come from the forward pass as saved e4m3 + a compact 2D-block scale +# ``[.., seqlen/32, head_dim/32]``. A1 **reuses** that exact quantization — +# no dequant, no re-quantization. Each dot rebuilds the scale tile the +# ``tl.dot_scaled`` layout wants (``[outer, reduction/32]``) directly from the +# saved compact scale via pointer index math (the same broadcast the forward +# kernel uses for QK). Because a 32x32 block shares one scale, the same saved +# e4m3 values + tile serve both the head_dim-reduction dots (a/b/e/f) and the +# seqlen-reduction dots (d/g); only the broadcast axis differs. +# * dO/P/dS are backward-only tensors the forward never saw, so they are +# quantized fresh in-kernel with 1D per-row scales (``IS_2D_BLOCK=False``) +# along their reduction axis — the canonical layout ``tl.dot_scaled`` eats +# with no broadcast. +_MX_2D = tl.constexpr(False) + + +@triton.jit +def _load_scale_hd( + scale_base, + offs_outer, + stride_sg, + stride_dg, + N_CTX, + SCALE_D: tl.constexpr, + SCALE_ACTUAL_D: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, +): + """Rebuild a ``[len(offs_outer), SCALE_D]`` scale tile for a **head_dim**- + reduction dot from the saved compact 2D-block scale. + + ``offs_outer`` indexes the outer (token/seqlen) rows; each maps to compact + seqlen-group ``offs_outer // QUANT_BLOCK_SIZE``. ``SCALE_D = head_dim/32`` + columns index the reduction (head_dim) groups directly. This mirrors the + forward kernel's ``qs_ptrs`` broadcast (a 32-row block shares one scale). + Out-of-range rows (past ``N_CTX``) and padded head groups load a neutral + scale; they only ever multiply masked-to-zero e4m3 values. + """ + offs_dg = tl.arange(0, SCALE_D) + mask = (offs_outer[:, None] < N_CTX) & (offs_dg[None, :] < SCALE_ACTUAL_D) + ptrs = (scale_base + (offs_outer[:, None] // QUANT_BLOCK_SIZE) * stride_sg + + offs_dg[None, :] * stride_dg) + return tl.load(ptrs, mask=mask, other=127) + + +@triton.jit +def _load_scale_sq( + scale_base, + start_outer, + offs_d, + stride_sg, + stride_dg, + N_CTX, + BLOCK_OUTER: tl.constexpr, + SCALE_ACTUAL_D: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, +): + """Rebuild a ``[len(offs_d), BLOCK_OUTER/32]`` scale tile for a **seqlen**- + reduction dot from the same saved compact 2D-block scale. + + This is the transpose-symmetric reuse (AB decision A1): the reduction axis is + now the seqlen block ``[start_outer, start_outer+BLOCK_OUTER)``, grouped by 32 + into ``BLOCK_OUTER/32`` columns, while ``offs_d`` (head_dim rows) map to + compact head_dim-group ``offs_d // QUANT_BLOCK_SIZE`` and are broadcast across + the 32 rows of each group. Element ``[d, sg] = scale[start_outer/32 + sg, + d/32]`` — the same 32x32 block scale the head_dim path reads, indexed for the + other axis. + """ + n_sg: tl.constexpr = BLOCK_OUTER // QUANT_BLOCK_SIZE + offs_sg = tl.arange(0, n_sg) + seq_group = start_outer // QUANT_BLOCK_SIZE + offs_sg + mask = (offs_d[:, None] < SCALE_ACTUAL_D * QUANT_BLOCK_SIZE) & \ + ((seq_group[None, :] * QUANT_BLOCK_SIZE) < N_CTX) + ptrs = (scale_base + seq_group[None, :] * stride_sg + + (offs_d[:, None] // QUANT_BLOCK_SIZE) * stride_dg) + return tl.load(ptrs, mask=mask, other=127) + + @triton.jit def _bwd_preprocess( Out, @@ -1084,10 +1206,15 @@ def _bwd_preprocess( @triton.jit def _attn_bwd_dkdv_inner( - k, - v, + kt_fp8, + ks_hd, + vt_fp8, + vs_hd, dk, dv, + q_scale_base, + stride_qsm, + stride_qsk, offs_d_qk, offs_d_v, offs_n, @@ -1104,6 +1231,10 @@ def _attn_bwd_dkdv_inner( stride_ldm, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + SCALE_BLOCK_DMODEL_QK: tl.constexpr, + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, sm_scale: tl.constexpr, lo, num_block_m: tl.constexpr, @@ -1111,8 +1242,19 @@ def _attn_bwd_dkdv_inner( N_CTX_Q: tl.constexpr, N_CTX_K: tl.constexpr, CAUSAL: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, ): - """Accumulate dk, dv over query blocks for one key block. k,v come in transposed.""" + """Accumulate dk, dv over query blocks for one key block (A1, dot_scaled). + + ``kt_fp8`` / ``vt_fp8`` are the saved-e4m3 key/value tiles transposed to + ``[head_dim, BLOCK_N]`` (fixed across the m-loop); ``ks_hd`` / ``vs_hd`` are + their ``[BLOCK_N, head_dim/32]`` scales rebuilt from the forward's compact 2D + scale by the caller. ``q`` is likewise the **saved e4m3** and is reused with + no re-quantization: dot a reads its head_dim-grouped scale, dot d reads the + same 2D block scale re-indexed along seqlen. Only dO/P/dS are quantized fresh + (plan §9.5 dots a/b/c/d). + """ for start_m in range(lo, num_block_m * BLOCK_M, BLOCK_M): offs_m = start_m + tl.arange(0, BLOCK_M) q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk @@ -1122,9 +1264,12 @@ def _attn_bwd_dkdv_inner( q_mask = mask_m[:, None] & mask_d_qk[None, :] do_mask = mask_m[:, None] & mask_d_v[None, :] + # Saved e4m3 q — reused directly, no re-quantization (A1). q = tl.load(q_ptrs, mask=q_mask, other=0.0) - # TODO(stage2): dot_scaled(q[e4m3, scale//head_dim], k[e4m3, scale//head_dim]); reduction axis = head_dim_qk. - qk = tl.dot(q, k, out_dtype=tl.float32) + # dot a: qk = q @ kᵀ, reduction = head_dim. Reuse q's saved head_dim scale. + qs_hd = _load_scale_hd(q_scale_base, offs_m, stride_qsm, stride_qsk, N_CTX_Q, + SCALE_BLOCK_DMODEL_QK, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + qk = tl.dot_scaled(q, qs_hd, "e4m3", kt_fp8, ks_hd, "e4m3", out_dtype=tl.float32) if CAUSAL: # Bottom-right causal, matches the forward kernel (offs_n_causal = offs_n + N_CTX_Q - N_CTX_K). @@ -1142,18 +1287,26 @@ def _attn_bwd_dkdv_inner( p = tl.math.exp(qk - l_i[:, None]) do = tl.load(do_ptrs, mask=do_mask, other=0.0) - # TODO(stage2): dot_scaled(do[e4m3, scale//head_dim_v], v[e4m3, scale//head_dim_v]); reduction axis = head_dim_v. - dp = tl.dot(do, v, out_dtype=tl.float32) + # dot b: dp = do @ vᵀ, reduction = head_dim_v. dO is fresh -> quantize here. + do_fp8_hd, dos_hd = _mx_quant(do, BLOCK_M, BLOCK_DMODEL_V, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + dp = tl.dot_scaled(do_fp8_hd, dos_hd, "e4m3", vt_fp8, vs_hd, "e4m3", out_dtype=tl.float32) d_ptrs = d_offset + offs_m * stride_ldm Di = tl.load(d_ptrs, mask=mask_m, other=0.0) ds = p * (dp - Di[:, None]) - ds = ds.to(q.dtype) - # TODO(stage2): dot_scaled(p^T[e4m3, scale//BLOCK_M], do[e4m3, scale//BLOCK_M]); reduction axis = seqlen_q. - dv += tl.dot(tl.trans(p.to(k.dtype)), do, out_dtype=tl.float32) - # TODO(stage2): dot_scaled(ds^T[e4m3, scale//BLOCK_M], q[e4m3, scale//BLOCK_M]); reduction axis = seqlen_q. - dk += tl.dot(tl.trans(ds), q, out_dtype=tl.float32) + # dot c: dv += pᵀ @ do, reduction = seqlen_q (BLOCK_M). P/dO fresh -> quantize + # 1D-block along BLOCK_M (their last axis after transpose). + pt_fp8, ps_c = _mx_quant(tl.trans(p), BLOCK_N, BLOCK_M, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + do_t_fp8, dos_c = _mx_quant(tl.trans(do), BLOCK_DMODEL_V, BLOCK_M, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + dv += tl.dot_scaled(pt_fp8, ps_c, "e4m3", tl.trans(do_t_fp8), dos_c, "e4m3", out_dtype=tl.float32) + + # dot d: dk += dsᵀ @ q, reduction = seqlen_q (BLOCK_M). dS fresh (1D along + # BLOCK_M); q reuses the saved 2D block scale re-indexed along seqlen. + dst_fp8, dss_d = _mx_quant(tl.trans(ds), BLOCK_N, BLOCK_M, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + qs_sq = _load_scale_sq(q_scale_base, start_m, offs_d_qk, stride_qsm, stride_qsk, N_CTX_Q, + BLOCK_M, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + dk += tl.dot_scaled(dst_fp8, dss_d, "e4m3", q, qs_sq, "e4m3", out_dtype=tl.float32) return dk, dv @@ -1163,6 +1316,9 @@ def _bwd_kernel_dkdv( Q, K, V, + Q_scale, + K_scale, + V_scale, sm_scale: tl.constexpr, DO, DK, @@ -1185,6 +1341,18 @@ def _bwd_kernel_dkdv( stride_doh, stride_dom, stride_dok, + stride_qsz, + stride_qsh, + stride_qsm, + stride_qsk, + stride_ksz, + stride_ksh, + stride_ksn, + stride_ksk, + stride_vsz, + stride_vsh, + stride_vsn, + stride_vsk, stride_ldz, stride_ldh, stride_ldm, @@ -1205,8 +1373,20 @@ def _bwd_kernel_dkdv( CAUSAL: tl.constexpr, USE_EXP2: tl.constexpr, IS_VARLEN: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, ): - """One program per (batch*head_k, key block). Parallelizes dk/dv over keys.""" + """One program per (batch*head_k, key block). Parallelizes dk/dv over keys. + + A1: q/k/v arrive as the forward's saved e4m3 with compact 2D-block scales; + the kernel reuses them (no re-quant), rebuilding each dot's scale tile from + the saved scale. Only dO/P/dS are quantized fresh (see ``_attn_bwd_dkdv_inner``). + """ + SCALE_BLOCK_DMODEL_QK: tl.constexpr = BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + SCALE_BLOCK_DMODEL_V: tl.constexpr = BLOCK_DMODEL_V // QUANT_BLOCK_SIZE + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr = ACTUAL_BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + SCALE_ACTUAL_BLOCK_DMODEL_V: tl.constexpr = ACTUAL_BLOCK_DMODEL_V // QUANT_BLOCK_SIZE + off_hz = tl.program_id(0) start_n = tl.program_id(1) off_z = off_hz // HK @@ -1226,10 +1406,16 @@ def _bwd_kernel_dkdv( N_CTX_Q = max_seqlen_q N_CTX_K = max_seqlen_k + q_scale_start = q_start // QUANT_BLOCK_SIZE + k_scale_start = k_start // QUANT_BLOCK_SIZE + q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + q_start * stride_qm k_offset = K + off_z * stride_kz + off_h_k * stride_kh + k_start * stride_kn v_offset = V + off_z * stride_vz + off_h_k * stride_vh + k_start * stride_vn do_offset = DO + off_z * stride_doz + off_h_q * stride_doh + q_start * stride_dom + q_scale_base = Q_scale + off_z * stride_qsz + off_h_q * stride_qsh + q_scale_start * stride_qsm + k_scale_base = K_scale + off_z * stride_ksz + off_h_k * stride_ksh + k_scale_start * stride_ksn + v_scale_base = V_scale + off_z * stride_vsz + off_h_k * stride_vsh + k_scale_start * stride_vsn adj_delta = off_z * stride_ldz + off_h_q * stride_ldh + q_start * stride_ldm l_offset = LSE + adj_delta d_offset = Delta + adj_delta @@ -1252,20 +1438,31 @@ def _bwd_kernel_dkdv( k_ptrs = k_offset + offs_n[:, None] * stride_kn + offs_d_qk[None, :] * stride_kk v_ptrs = v_offset + offs_n[:, None] * stride_vn + offs_d_v[None, :] * stride_vk - k = tl.load(k_ptrs, mask=mask_n[:, None] & mask_d_qk[None, :], other=0.0) - v = tl.load(v_ptrs, mask=mask_n[:, None] & mask_d_v[None, :], other=0.0) - k = tl.trans(k) - v = tl.trans(v) + # Saved e4m3 k/v reused directly; transpose to [head_dim, BLOCK_N] for the + # QK / dP dots (a/b). Scales are rebuilt from the saved compact 2D scale. + k_fp8 = tl.load(k_ptrs, mask=mask_n[:, None] & mask_d_qk[None, :], other=0.0) + v_fp8 = tl.load(v_ptrs, mask=mask_n[:, None] & mask_d_v[None, :], other=0.0) + ks_hd = _load_scale_hd(k_scale_base, offs_n, stride_ksn, stride_ksk, N_CTX_K, + SCALE_BLOCK_DMODEL_QK, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + vs_hd = _load_scale_hd(v_scale_base, offs_n, stride_vsn, stride_vsk, N_CTX_K, + SCALE_BLOCK_DMODEL_V, SCALE_ACTUAL_BLOCK_DMODEL_V, QUANT_BLOCK_SIZE) + kt_fp8 = tl.trans(k_fp8) + vt_fp8 = tl.trans(v_fp8) dk = tl.zeros([BLOCK_N, BLOCK_DMODEL_QK], dtype=tl.float32) dv = tl.zeros([BLOCK_N, BLOCK_DMODEL_V], dtype=tl.float32) for _ in range(GROUP_SIZE): dk, dv = _attn_bwd_dkdv_inner( - k, - v, + kt_fp8, + ks_hd, + vt_fp8, + vs_hd, dk, dv, + q_scale_base, + stride_qsm, + stride_qsk, offs_d_qk, offs_d_v, offs_n, @@ -1282,6 +1479,10 @@ def _bwd_kernel_dkdv( stride_ldm, BLOCK_M, BLOCK_N, + BLOCK_DMODEL_QK, + BLOCK_DMODEL_V, + SCALE_BLOCK_DMODEL_QK, + SCALE_ACTUAL_BLOCK_DMODEL_QK, sm_scale, lo, num_block_m, @@ -1289,9 +1490,12 @@ def _bwd_kernel_dkdv( N_CTX_Q, N_CTX_K, CAUSAL, + QUANT_BLOCK_SIZE, + USE_ASM, ) q_offset += stride_qh do_offset += stride_qh + q_scale_base += stride_qsh l_offset += stride_ldh d_offset += stride_ldh @@ -1306,13 +1510,21 @@ def _bwd_kernel_dkdv( @triton.jit def _attn_bwd_dq_inner( dq, - q, + q_fp8_hd, + qs_hd, + do_fp8_hd, + dos_hd, + k_scale_base, + v_scale_base, + stride_ksn, + stride_ksk, + stride_vsn, + stride_vsk, offs_d_qk, offs_d_v, offs_m, l_i, Di, - do, mask_d_qk, mask_d_v, k_offset, @@ -1323,14 +1535,29 @@ def _attn_bwd_dq_inner( stride_vk, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, + BLOCK_DMODEL_QK: tl.constexpr, + BLOCK_DMODEL_V: tl.constexpr, + SCALE_BLOCK_DMODEL_QK: tl.constexpr, + SCALE_BLOCK_DMODEL_V: tl.constexpr, + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr, + SCALE_ACTUAL_BLOCK_DMODEL_V: tl.constexpr, sm_scale: tl.constexpr, hi, USE_EXP2: tl.constexpr, N_CTX_Q: tl.constexpr, N_CTX_K: tl.constexpr, CAUSAL: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, ): - """Accumulate dq over key blocks for one query block.""" + """Accumulate dq over key blocks for one query block (A1, dot_scaled). + + ``q_fp8_hd`` / ``qs_hd`` are the saved-e4m3 q + its head_dim scale (fixed + across the n-loop, built once by the caller); ``do_fp8_hd`` / ``dos_hd`` are + the fresh-quantized dO. k/v are the saved e4m3 reused with no re-quant: dots + e/f read their head_dim scale, dot g reads k's 2D block scale re-indexed along + seqlen. Only dS is quantized fresh (plan §9.5 dots e/f/g). + """ if USE_EXP2: l_i *= RCP_LN2 @@ -1342,11 +1569,14 @@ def _attn_bwd_dq_inner( k_ptrs = k_offset + offs_n[:, None] * stride_kn + offs_d_qk[None, :] * stride_kk v_ptrs = v_offset + offs_n[:, None] * stride_vn + offs_d_v[None, :] * stride_vk + # Saved e4m3 k/v, reused directly. k = tl.load(k_ptrs, mask=mask_k, other=0.0) v = tl.load(v_ptrs, mask=mask_v, other=0.0) - # TODO(stage2): dot_scaled(q[e4m3, scale//head_dim], k^T[e4m3, scale//head_dim]); reduction axis = head_dim_qk. - qk = tl.dot(q, tl.trans(k), out_dtype=tl.float32) + # dot e: qk = q @ kᵀ, reduction = head_dim. Reuse k's saved head_dim scale. + ks_e = _load_scale_hd(k_scale_base, offs_n, stride_ksn, stride_ksk, N_CTX_K, + SCALE_BLOCK_DMODEL_QK, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + qk = tl.dot_scaled(q_fp8_hd, qs_hd, "e4m3", tl.trans(k), ks_e, "e4m3", out_dtype=tl.float32) if CAUSAL: col_offset = N_CTX_Q - N_CTX_K @@ -1360,12 +1590,18 @@ def _attn_bwd_dq_inner( qk *= sm_scale p = tl.math.exp(qk - l_i[:, None]) - # TODO(stage2): dot_scaled(do[e4m3, scale//head_dim_v], v^T[e4m3, scale//head_dim_v]); reduction axis = head_dim_v. - dp = tl.dot(do, tl.trans(v), out_dtype=tl.float32) + # dot f: dp = do @ vᵀ, reduction = head_dim_v. Reuse v's saved head_dim scale. + vs_f = _load_scale_hd(v_scale_base, offs_n, stride_vsn, stride_vsk, N_CTX_K, + SCALE_BLOCK_DMODEL_V, SCALE_ACTUAL_BLOCK_DMODEL_V, QUANT_BLOCK_SIZE) + dp = tl.dot_scaled(do_fp8_hd, dos_hd, "e4m3", tl.trans(v), vs_f, "e4m3", out_dtype=tl.float32) ds = p * (dp - Di[:, None]) - ds = ds.to(q.dtype) - # TODO(stage2): dot_scaled(ds[e4m3, scale//BLOCK_N], k[e4m3, scale//BLOCK_N]); reduction axis = seqlen_k. - dq += tl.dot(ds, k, out_dtype=tl.float32) + + # dot g: dq += ds @ k, reduction = seqlen_k (BLOCK_N). dS fresh (1D along + # BLOCK_N); k reuses the saved 2D block scale re-indexed along seqlen. + ds_fp8, dss_g = _mx_quant(ds, BLOCK_M, BLOCK_N, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + ks_sk = _load_scale_sq(k_scale_base, start_n, offs_d_qk, stride_ksn, stride_ksk, N_CTX_K, + BLOCK_N, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + dq += tl.dot_scaled(ds_fp8, dss_g, "e4m3", k, ks_sk, "e4m3", out_dtype=tl.float32) return dq @@ -1375,6 +1611,9 @@ def _bwd_kernel_dq( Q, K, V, + Q_scale, + K_scale, + V_scale, sm_scale: tl.constexpr, DO, DQ, @@ -1396,6 +1635,18 @@ def _bwd_kernel_dq( stride_doh, stride_dom, stride_dok, + stride_qsz, + stride_qsh, + stride_qsm, + stride_qsk, + stride_ksz, + stride_ksh, + stride_ksn, + stride_ksk, + stride_vsz, + stride_vsh, + stride_vsn, + stride_vsk, stride_ldz, stride_ldh, stride_ldm, @@ -1416,8 +1667,19 @@ def _bwd_kernel_dq( CAUSAL: tl.constexpr, USE_EXP2: tl.constexpr, IS_VARLEN: tl.constexpr, + QUANT_BLOCK_SIZE: tl.constexpr, + USE_ASM: tl.constexpr, ): - """One program per (batch*head_q, query block). Parallelizes dq over queries.""" + """One program per (batch*head_q, query block). Parallelizes dq over queries. + + A1: q/k/v are the forward's saved e4m3 + compact 2D-block scales, reused with + no re-quant. Only dO (fresh input) and dS are quantized in-kernel. + """ + SCALE_BLOCK_DMODEL_QK: tl.constexpr = BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + SCALE_BLOCK_DMODEL_V: tl.constexpr = BLOCK_DMODEL_V // QUANT_BLOCK_SIZE + SCALE_ACTUAL_BLOCK_DMODEL_QK: tl.constexpr = ACTUAL_BLOCK_DMODEL_QK // QUANT_BLOCK_SIZE + SCALE_ACTUAL_BLOCK_DMODEL_V: tl.constexpr = ACTUAL_BLOCK_DMODEL_V // QUANT_BLOCK_SIZE + off_hz = tl.program_id(0) start_m = tl.program_id(1) off_z = off_hz // HQ @@ -1437,10 +1699,16 @@ def _bwd_kernel_dq( N_CTX_Q = max_seqlen_q N_CTX_K = max_seqlen_k + q_scale_start = q_start // QUANT_BLOCK_SIZE + k_scale_start = k_start // QUANT_BLOCK_SIZE + q_offset = Q + off_z * stride_qz + off_h_q * stride_qh + q_start * stride_qm k_offset = K + off_z * stride_kz + off_h_k * stride_kh + k_start * stride_kn v_offset = V + off_z * stride_vz + off_h_k * stride_vh + k_start * stride_vn do_offset = DO + off_z * stride_doz + off_h_q * stride_doh + q_start * stride_dom + q_scale_base = Q_scale + off_z * stride_qsz + off_h_q * stride_qsh + q_scale_start * stride_qsm + k_scale_base = K_scale + off_z * stride_ksz + off_h_k * stride_ksh + k_scale_start * stride_ksn + v_scale_base = V_scale + off_z * stride_vsz + off_h_k * stride_vsh + k_scale_start * stride_vsn adj_delta = off_z * stride_ldz + off_h_q * stride_ldh + q_start * stride_ldm l_offset = LSE + adj_delta d_offset = Delta + adj_delta @@ -1461,22 +1729,36 @@ def _bwd_kernel_dq( q_ptrs = q_offset + offs_m[:, None] * stride_qm + offs_d_qk[None, :] * stride_qk do_ptrs = do_offset + offs_m[:, None] * stride_dom + offs_d_v[None, :] * stride_dok - q = tl.load(q_ptrs, mask=mask_m[:, None] & mask_d_qk[None, :], other=0.0) + # Saved e4m3 q reused directly; dO is a fresh input and is quantized here. + q_fp8_hd = tl.load(q_ptrs, mask=mask_m[:, None] & mask_d_qk[None, :], other=0.0) do = tl.load(do_ptrs, mask=mask_m[:, None] & mask_d_v[None, :], other=0.0) l_i = tl.load(l_offset + offs_m * stride_ldm, mask=mask_m, other=0.0) Di = tl.load(d_offset + offs_m * stride_ldm, mask=mask_m, other=0.0) + # q's saved head_dim scale (dot e), fixed across the n-loop; dO quantized fresh. + qs_hd = _load_scale_hd(q_scale_base, offs_m, stride_qsm, stride_qsk, N_CTX_Q, + SCALE_BLOCK_DMODEL_QK, SCALE_ACTUAL_BLOCK_DMODEL_QK, QUANT_BLOCK_SIZE) + do_fp8_hd, dos_hd = _mx_quant(do, BLOCK_M, BLOCK_DMODEL_V, QUANT_BLOCK_SIZE, _MX_2D, USE_ASM) + dq = tl.zeros([BLOCK_M, BLOCK_DMODEL_QK], dtype=tl.float32) dq = _attn_bwd_dq_inner( dq, - q, + q_fp8_hd, + qs_hd, + do_fp8_hd, + dos_hd, + k_scale_base, + v_scale_base, + stride_ksn, + stride_ksk, + stride_vsn, + stride_vsk, offs_d_qk, offs_d_v, offs_m, l_i, Di, - do, mask_d_qk, mask_d_v, k_offset, @@ -1487,12 +1769,20 @@ def _bwd_kernel_dq( stride_vk, BLOCK_M, BLOCK_N, + BLOCK_DMODEL_QK, + BLOCK_DMODEL_V, + SCALE_BLOCK_DMODEL_QK, + SCALE_BLOCK_DMODEL_V, + SCALE_ACTUAL_BLOCK_DMODEL_QK, + SCALE_ACTUAL_BLOCK_DMODEL_V, sm_scale, hi, USE_EXP2, N_CTX_Q, N_CTX_K, CAUSAL, + QUANT_BLOCK_SIZE, + USE_ASM, ) dq *= sm_scale @@ -1520,22 +1810,27 @@ def attention_mxfp8_backward_triton_impl( max_seqlen_k: Optional[int], use_exp2: bool, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """MXFP8 flash-attention backward. + """MXFP8 flash-attention backward (stage-2, ``tl.dot_scaled``). - Stage 1 (current): q/k/v arrive as saved e4m3 tensors; they are dequantized - back to bf16 here so the three ported kernels run the standard FA v2 backward - math in high precision. This is numerically correct w.r.t. the *quantized* - forward inputs and is runnable without CDNA4. + The saved e4m3 q/k/v are dequantized back to bf16 at entry (the exact inputs + the forward consumed; option A — no extra fwd memory), then **each of the 7 + backward dots quantizes its two operands in-kernel with 1D per-row MX scales + along that dot's reduction axis** and runs ``tl.dot_scaled`` (plan §9.5). + This uniform per-row layout is exactly what ``tl.dot_scaled`` consumes (no + broadcast), and matches ``mxfp8_attention_backward_reference_stage2``. - Stage 2 (TODO): keep operands in e4m3 and replace each ``tl.dot`` inside the - kernels with ``tl.dot_scaled`` plus per-reduction-axis MXFP8 scales (see the - ``TODO(stage2)`` markers). dO will additionally be quantized to e4m3. + Requires CDNA4 (native ``tl.dot_scaled``). + + A1 (AB decision): the saved e4m3 q/k/v and their compact 2D-block scales are + passed straight into the kernels, which reuse them in every dot (no entry + dequant, no re-quantization of q/k/v). Only dO/P/dS are quantized in-kernel. """ - # Dequantize the saved e4m3 operands back to bf16 (the exact inputs the - # forward pass consumed). Uses the non-ASM path on non-CDNA4 hardware. - q = convert_from_mxfp8(q, q_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() - k = convert_from_mxfp8(k, k_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() - v = convert_from_mxfp8(v, v_scale, output_dtype=fwd_torch_dtype, axis=-1, is_2d_block=True).contiguous() + q = q.contiguous() + k = k.contiguous() + v = v.contiguous() + q_scale = q_scale.contiguous() + k_scale = k_scale.contiguous() + v_scale = v_scale.contiguous() if not do.is_contiguous(): do = do.contiguous() @@ -1547,6 +1842,11 @@ def attention_mxfp8_backward_triton_impl( stride_vz, stride_vh, stride_vn, stride_vk = get_strides_from_layout(v, layout) stride_oz, stride_oh, stride_om, stride_ok = get_strides_from_layout(o, layout) stride_doz, stride_doh, stride_dom, stride_dok = get_strides_from_layout(do, layout) + # Compact 2D-block scale strides ([.., seqlen/32, head_dim/32]); reused by the + # kernels for both head_dim- and seqlen-reduction dots (A1). + stride_qsz, stride_qsh, stride_qsm, stride_qsk = get_strides_from_layout(q_scale, layout) + stride_ksz, stride_ksh, stride_ksn, stride_ksk = get_strides_from_layout(k_scale, layout) + stride_vsz, stride_vsh, stride_vsn, stride_vsk = get_strides_from_layout(v_scale, layout) is_varlen = layout == "thd" padded_d_model_qk = get_padded_head_dim(head_size_qk) @@ -1596,6 +1896,9 @@ def attention_mxfp8_backward_triton_impl( q, k, v, + q_scale, + k_scale, + v_scale, sm_scale, do, dq, @@ -1617,6 +1920,18 @@ def attention_mxfp8_backward_triton_impl( stride_doh, stride_dom, stride_dok, + stride_qsz, + stride_qsh, + stride_qsm, + stride_qsk, + stride_ksz, + stride_ksh, + stride_ksn, + stride_ksk, + stride_vsz, + stride_vsh, + stride_vsn, + stride_vsk, stride_lse_z, stride_lse_h, stride_lse_m, @@ -1637,12 +1952,17 @@ def attention_mxfp8_backward_triton_impl( CAUSAL=causal, USE_EXP2=use_exp2, IS_VARLEN=is_varlen, + QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT, + USE_ASM=is_cdna4(), ) wrap_triton(_bwd_kernel_dkdv)[(batch * nheads_k, num_block_n)]( q, k, v, + q_scale, + k_scale, + v_scale, sm_scale, do, dk, @@ -1665,6 +1985,18 @@ def attention_mxfp8_backward_triton_impl( stride_doh, stride_dom, stride_dok, + stride_qsz, + stride_qsh, + stride_qsm, + stride_qsk, + stride_ksz, + stride_ksh, + stride_ksn, + stride_ksk, + stride_vsz, + stride_vsh, + stride_vsn, + stride_vsk, stride_lse_z, stride_lse_h, stride_lse_m, @@ -1685,6 +2017,8 @@ def attention_mxfp8_backward_triton_impl( CAUSAL=causal, USE_EXP2=use_exp2, IS_VARLEN=is_varlen, + QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT, + USE_ASM=is_cdna4(), ) return dq, dk, dv diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index 5408424d..741e5694 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -28,7 +28,6 @@ calc_cossim, mxfp8_attention_forward_reference, mxfp8_attention_backward_reference, - mxfp8_attention_backward_reference_stage2, ) # Reuse the committed shape grid so Layer 2 stays in lock-step with Layer 3. from .test_mxfp8_attention import AttnConfig, test_cases @@ -187,117 +186,3 @@ def test_backward_reference_matches_sdpa(config, causal): snr = calc_snr(ref, got) assert sim > 0.97, f"{name} golden-vs-sdpa cosine-sim too low: {sim}" assert snr > 8, f"{name} golden-vs-sdpa SNR too low: {snr}" - - -# --------------------------------------------------------------------------- -# Backward Stage-2 Layer 1 — full-quant reference vs autograd fp32 SDPA -# Quantizes every backward matmul operand along its reduction axis (a -# numerical model of tl.dot_scaled). The gap vs SDPA is the *full* mxfp8 -# backward quantization error — the signal that decides whether all-e4m3 is -# viable or must go e5m2. Pure PyTorch, runs on CPU / MI250 (no CDNA4). -# --------------------------------------------------------------------------- - -@pytest.mark.parametrize("config", reference_cases) -@pytest.mark.parametrize("causal", [True, False]) -def test_backward_stage2_reference_matches_sdpa(config, causal): - """Stage-2 (all-operand-quantized) backward reference vs autograd fp32 SDPA. - - This is the pre-hardware numerical de-risk for the plan §9.5 stage-2 kernel: - if all-e4m3 grads already diverge here, the kernel can't do better and we - should switch dO/dP/dS to e5m2 before writing any tl.dot_scaled code. - """ - device = "cuda" if torch.cuda.is_available() else "cpu" - dtype = torch.float32 if device == "cpu" else torch.bfloat16 - n_rep = config.num_head_q // config.num_head_kv - - q, k, v = _make_qkv_bhsd(2, config, device, dtype) - do = _make_do_bhsd(2, config, device, dtype) - sm_scale = config.head_dim_qk**(-0.5) - - qg = q.clone().requires_grad_(True) - kg = k.clone().requires_grad_(True) - vg = v.clone().requires_grad_(True) - o_auto = torch.nn.functional.scaled_dot_product_attention( - qg, kg, vg, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) - o_auto.backward(do) - - o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) - dq, dk, dv = mxfp8_attention_backward_reference_stage2(q, k, v, do, o_ref, lse_ref, sm_scale, causal) - - rows = [] - for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: - rows.append([f"{name} (stage2 vs sdpa)", calc_snr(ref, got), calc_cossim(ref, got)]) - print() - print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) - - # Calibrated on MI250 (first run): all-e4m3 backward lands at - # cossim~0.998 / SNR~23-26 dB across dQ/dK/dV — no e5m2 needed. Thresholds - # keep margin below those numbers. - for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: - sim = calc_cossim(ref, got) - snr = calc_snr(ref, got) - assert sim > 0.995, f"{name} stage2-vs-sdpa cosine-sim too low: {sim}" - assert snr > 18, f"{name} stage2-vs-sdpa SNR too low: {snr}" - - -# --------------------------------------------------------------------------- -# Backward Layer 2 — kernel vs golden reference -# Stage-1 backward dequantizes e4m3 -> bf16 and uses plain tl.dot (no -# tl.dot_scaled), so this runs on MI250 too. It calls the backward op directly -# with forward-reference o/lse, bypassing the forward kernel's tl.dot_scaled. -# --------------------------------------------------------------------------- - -@cuda_only -@pytest.mark.parametrize("config", test_cases) -@pytest.mark.parametrize("causal", [True]) -def test_backward_kernel_matches_reference(config, causal): - """Backward kernel vs golden reference — isolates port bugs from quant error.""" - from alto.kernels.mxfp8.mxfp8_quantization import convert_to_mxfp8 - - device = "cuda" - dtype = torch.bfloat16 - q, k, v = _make_qkv_bhsd(4, config, device, dtype) - do = _make_do_bhsd(4, config, device, dtype) - sm_scale = config.head_dim_qk**(-0.5) - - o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) - - # Quantize operands the way the forward autograd does, then run the backward op. - q8, q_scale = convert_to_mxfp8(q, mxfp_format="e4m3", axis=-1, is_2d_block=True) - k8, k_scale = convert_to_mxfp8(k, mxfp_format="e4m3", axis=-1, is_2d_block=True) - v8, v_scale = convert_to_mxfp8(v, mxfp_format="e4m3", axis=-1, is_2d_block=True) - - dq_k, dk_k, dv_k = torch.ops.alto.attention_mxfp8_backward_triton_impl( - do.contiguous(), - q8.contiguous(), - k8.contiguous(), - v8.contiguous(), - o_ref.contiguous(), - lse_ref.contiguous(), - q_scale, - k_scale, - v_scale, - sm_scale=sm_scale, - causal=causal, - layout="bhsd", - cu_seqlens_q=0, - cu_seqlens_k=0, - max_seqlen_q=config.seqlen_q, - max_seqlen_k=config.seqlen_kv, - use_exp2=True, - ) - dq_r, dk_r, dv_r = mxfp8_attention_backward_reference(q, k, v, do, o_ref, lse_ref, sm_scale, causal) - - rows = [] - pairs = [("dQ", dq_r, dq_k), ("dK", dk_r, dk_k), ("dV", dv_r, dv_k)] - for name, ref, got in pairs: - rows.append([f"{name} (kernel vs golden)", calc_snr(ref, got), calc_cossim(ref, got)]) - print() - print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) - - # Threshold placeholder — calibrate on the first run. - for name, ref, got in pairs: - sim = calc_cossim(ref, got) - snr = calc_snr(ref, got) - assert sim > 0.99, f"{name} kernel-vs-golden cosine-sim too low: {sim}" - assert snr > 25, f"{name} kernel-vs-golden SNR too low: {snr}" diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index 8ea5af62..5990fe84 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -377,109 +377,3 @@ def mxfp8_attention_backward_reference( dv = dv_full return dq, dk, dv - - -def mxfp8_attention_backward_reference_stage2( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - do: torch.Tensor, - o: torch.Tensor, - softmax_lse: torch.Tensor, - sm_scale: float, - causal: bool, - block_size: int = BLOCK_SIZE_DEFAULT, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Pure-PyTorch golden reference for the **Stage-2** MXFP8 (e4m3) backward. - - Where Stage-1 (``mxfp8_attention_backward_reference``) quantizes only Q/K/V - (2D-block along head_dim) and runs every backward matmul in fp32, Stage-2 - additionally quantizes **each matmul operand along that dot's reduction - axis** — a numerical model of what ``tl.dot_scaled`` will do in the kernel. - A kernel-vs-this gap then isolates Triton port bugs from mxfp8 error, and a - this-vs-bf16-SDPA gap measures the *full* mxfp8 backward quantization error - (the signal that decides whether all-e4m3 backward is viable or must go - e5m2), all **without CDNA4** (no ``tl.dot_scaled`` here; plain fp32 matmul - on dequantized operands). - - Quantization axis per dot (see plan §9.5):: - - dP = dO @ Vᵀ reduction head_dim_v : dO,V 2D-block axis=-1 - S = Q @ Kᵀ reduction head_dim : Q,K 2D-block axis=-1 (reused from fwd) - dV = Pᵀ @ dO reduction seqlen_q : P,dO 1D-block along sq - dQ = dS @ K reduction seqlen_k : dS,K 1D-block along sk - dK = dSᵀ @ Q reduction seqlen_q : dS,Q 1D-block along sq - - 2D/1D rule (mirrors the forward decisions): operands whose reduction axis is - head_dim/head_dim_v use 2D-block (like fwd Q/K/V); operands whose reduction - axis is a sequence length use 1D-block along that seqlen (like fwd P). - - Same top-left causal mask as the other references, so causal tests must keep - ``seqlen_q == seqlen_k``. - - Args mirror ``mxfp8_attention_backward_reference``; ``o`` / ``softmax_lse`` - must be the forward-reference outputs. Returns ``(dq, dk, dv)``. - """ - b, hq, sq, dqk = q.shape - _, hk, sk, dvv = v.shape - assert hq % hk == 0, f"nheads_q ({hq}) must be a multiple of nheads_k ({hk})" - n_rep = hq // hk - - def qdq(x, axis, is_2d_block): - return _mxfp8_qdq(x.to(torch.float32), axis=axis, is_2d_block=is_2d_block, block_size=block_size) - - # --- head_dim-reduction operands: 2D-block along head_dim (reuse fwd convention) --- - q_hd = qdq(q, axis=-1, is_2d_block=True) # for S (dots a/e) - k_hd = qdq(k, axis=-1, is_2d_block=True) # for S - v_hd = qdq(v, axis=-1, is_2d_block=True) # for dP (dots b/f) - do_hd = qdq(do, axis=-1, is_2d_block=True) # dO for dP, along head_dim_v - - # --- seqlen-reduction operands: 1D-block along the seqlen axis --- - q_sq = qdq(q, axis=-2, is_2d_block=False) # for dK (dot d) - k_sk = qdq(k, axis=-2, is_2d_block=False) # for dQ (dot g) - do_sq = qdq(do, axis=-2, is_2d_block=False) # dO for dV (dot c) - - if n_rep > 1: - k_hd = k_hd.repeat_interleave(n_rep, dim=1) - v_hd = v_hd.repeat_interleave(n_rep, dim=1) - k_sk = k_sk.repeat_interleave(n_rep, dim=1) - - lse = softmax_lse.to(torch.float32) - do_f = do.to(torch.float32) - o_f = o.to(torch.float32) - - # S -> P (recompute from saved lse; P feeding dS is elementwise, not quantized). - s = torch.matmul(q_hd, k_hd.transpose(-1, -2)) * sm_scale - if causal: - q_pos = torch.arange(sq, device=q.device) - k_pos = torch.arange(sk, device=q.device) - allowed = k_pos[None, :] <= q_pos[:, None] - s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, float("-inf"))) - p = torch.exp(s - lse[..., None]) - p = torch.nan_to_num(p, nan=0.0, posinf=0.0, neginf=0.0) - - # dP = dO @ Vᵀ (reduction head_dim_v). - dp = torch.matmul(do_hd, v_hd.transpose(-1, -2)) - delta = (o_f * do_f).sum(dim=-1) - ds = p * (dp - delta[..., None]) # fp32 elementwise - - # dV = Pᵀ @ dO (reduction seqlen_q): P and dO 1D-block along sq. - p_sq = qdq(p, axis=-2, is_2d_block=False) - dv_full = torch.matmul(p_sq.transpose(-1, -2), do_sq) - - # dQ = sm_scale * dS @ K (reduction seqlen_k): dS along sk, K along sk. - ds_sk = qdq(ds, axis=-1, is_2d_block=False) - dq = sm_scale * torch.matmul(ds_sk, k_sk) - - # dK = sm_scale * dSᵀ @ Q (reduction seqlen_q): dS along sq, Q along sq. - ds_sq = qdq(ds, axis=-2, is_2d_block=False) - dk_full = sm_scale * torch.matmul(ds_sq.transpose(-1, -2), q_sq) - - if n_rep > 1: - dk = dk_full.view(b, hk, n_rep, sk, dqk).sum(dim=2) - dv = dv_full.view(b, hk, n_rep, sk, dvv).sum(dim=2) - else: - dk = dk_full - dv = dv_full - - return dq, dk, dv From 555aa7d9dd3cb9ec1c06150a422bca66d2c70d05 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Thu, 16 Jul 2026 03:42:01 -0500 Subject: [PATCH 08/14] test: add mxfp8 backward golden reference reusing forward e4m3 scales --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 20 +++- .../mxfp8/triton_flash_attention_mxfp8.py | 20 ++-- .../mxfp8/test_mxfp8_attention_reference.py | 110 ++++++++++++++++++ tests/unittest/mxfp8/utils.py | 101 ++++++++++++++++ 4 files changed, 239 insertions(+), 12 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 86af77ec..d7e536ca 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -441,6 +441,24 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - **A1 kernel `tl.dot_scaled` + 2D-scale 广播待 CDNA4 验**;`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 - padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 +### 10.1septies 截至 2026-07-16(A1 golden reference 已写并验、dispatch 已接) + +**A1 参照(stage-2 全量化):** 在 `tests/unittest/mxfp8/utils.py` 加 +`mxfp8_attention_backward_reference_stage2`——q/k/v 单份 2D-block dequant 全程复用(A1,**零 seqlen 重量化、无双重量化**),只有 dO/P/dS 现场 1D per-row 量化,按 §9.5 各 dot 的 reduction 轴。纯 fp32 matmul、无 `tl.dot_scaled`,MI250/CPU 可跑。 + +**测试:** +| 检查 | 结果 | +|---|---| +| `test_backward_stage2_reference_matches_sdpa`(A1 参照 vs bf16 SDPA autograd,MI250) | ✅ 6/6,dQ/dK≈23.2dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985(无双重量化,符合/略优于已删 option-A 的 ≈23dB 基线) | +| `test_backward_kernel_matches_reference`(A1 kernel vs A1 参照,CDNA4-only) | ⏳ 已写、待 CDNA4;MI250 上 `tl.dot_scaled` 编不过(预期,既有限制) | +| 全参照套件(纯 PyTorch) | ✅ 18 passed;14 failed 全是 fwd/bwd kernel 的 `tl.dot_scaled` gfx90a 编译失败(既有限制、非本次改动) | + +**dispatch:** `alto/kernels/dispatch/attention.py` 加 `precision == "mxfp8_e4m3"` 分支 → `triton_attention_mxfp8`(与 mxfp4 同 kwargs 调用路径,签名兼容)。 + +**仍待办:** +- **CDNA4 验收**:A1 kernel vs A1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差);`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 +- padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA 于 MI250 验 6/6(阶段2 全量化参照 vs SDPA 曾 6/6、cossim≈0.998/SNR≈23dB,已随 A 删除,A1 参照待补)。**A1 kernel 的 `tl.dot_scaled` + 2D-scale 广播在 MI250 编不了、当前无 A1 参照,正确性完全待 CDNA4 + A1 参照**(`_load_scale_sq` 为最大盲区)。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**A1 kernel 的 `tl.dot_scaled` + 2D-scale 广播在 MI250 编不了,kernel 正确性完全待 CDNA4(kernel vs A1 参照已写好待跑)**(`_load_scale_sq` 为最大盲区)。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index 23d88a42..60a25ae3 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -1810,20 +1810,18 @@ def attention_mxfp8_backward_triton_impl( max_seqlen_k: Optional[int], use_exp2: bool, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """MXFP8 flash-attention backward (stage-2, ``tl.dot_scaled``). + """MXFP8 flash-attention backward (stage-2 / A1, ``tl.dot_scaled``). - The saved e4m3 q/k/v are dequantized back to bf16 at entry (the exact inputs - the forward consumed; option A — no extra fwd memory), then **each of the 7 - backward dots quantizes its two operands in-kernel with 1D per-row MX scales - along that dot's reduction axis** and runs ``tl.dot_scaled`` (plan §9.5). - This uniform per-row layout is exactly what ``tl.dot_scaled`` consumes (no - broadcast), and matches ``mxfp8_attention_backward_reference_stage2``. + A1 operand source: the forward-saved e4m3 q/k/v and their compact 2D-block + scales are passed straight in and reused in *every* dot (no entry dequant, no + re-quantization of q/k/v). For head_dim-reduction dots the 2D scale already + groups along the reduction axis; for seqlen-reduction dots the same compact + 2D scale is re-indexed (``_load_scale_sq``) to broadcast per reduction row. + Only the backward-only tensors dO/P/dS are quantized in-kernel (1D per-row + along each dot's reduction axis), then fed to ``tl.dot_scaled`` (plan §9.5). + Matches ``mxfp8_attention_backward_reference_stage2``. Requires CDNA4 (native ``tl.dot_scaled``). - - A1 (AB decision): the saved e4m3 q/k/v and their compact 2D-block scales are - passed straight into the kernels, which reuse them in every dot (no entry - dequant, no re-quantization of q/k/v). Only dO/P/dS are quantized in-kernel. """ q = q.contiguous() k = k.contiguous() diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index 741e5694..35264391 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -28,6 +28,7 @@ calc_cossim, mxfp8_attention_forward_reference, mxfp8_attention_backward_reference, + mxfp8_attention_backward_reference_stage2, ) # Reuse the committed shape grid so Layer 2 stays in lock-step with Layer 3. from .test_mxfp8_attention import AttnConfig, test_cases @@ -186,3 +187,112 @@ def test_backward_reference_matches_sdpa(config, causal): snr = calc_snr(ref, got) assert sim > 0.97, f"{name} golden-vs-sdpa cosine-sim too low: {sim}" assert snr > 8, f"{name} golden-vs-sdpa SNR too low: {snr}" + + +# --------------------------------------------------------------------------- +# Backward Stage-2 (A1) Layer 1 — full-quant reference vs autograd fp32 SDPA +# Adds dO/P/dS quantization on top of Stage-1 (a numerical model of the +# kernel's tl.dot_scaled). q/k/v reuse the forward 2D-block dequant (A1, no +# double quant). The gap vs SDPA is the *full* mxfp8 backward quant error — +# the signal that all-e4m3 backward is viable. Pure PyTorch, runs on CPU/MI250. +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("config", reference_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_backward_stage2_reference_matches_sdpa(config, causal): + """Stage-2 (A1, all-operand-quantized) backward reference vs autograd fp32 SDPA.""" + device = "cuda" if torch.cuda.is_available() else "cpu" + dtype = torch.float32 if device == "cpu" else torch.bfloat16 + n_rep = config.num_head_q // config.num_head_kv + + q, k, v = _make_qkv_bhsd(2, config, device, dtype) + do = _make_do_bhsd(2, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + qg = q.clone().requires_grad_(True) + kg = k.clone().requires_grad_(True) + vg = v.clone().requires_grad_(True) + o_auto = torch.nn.functional.scaled_dot_product_attention( + qg, kg, vg, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) + o_auto.backward(do) + + o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + dq, dk, dv = mxfp8_attention_backward_reference_stage2(q, k, v, do, o_ref, lse_ref, sm_scale, causal) + + rows = [] + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + rows.append([f"{name} (stage2 vs sdpa)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + # A1 has no double-quant, so it should land at/above the retired option-A + # baseline (cossim~0.998 / SNR~23-26 dB). Thresholds keep margin. + for name, ref, got in [("dQ", qg.grad, dq), ("dK", kg.grad, dk), ("dV", vg.grad, dv)]: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > 0.995, f"{name} stage2-vs-sdpa cosine-sim too low: {sim}" + assert snr > 18, f"{name} stage2-vs-sdpa SNR too low: {snr}" + + +# --------------------------------------------------------------------------- +# Backward Layer 2 — kernel vs Stage-2 (A1) golden reference +# The A1 backward kernel runs tl.dot_scaled in-kernel, so this requires CDNA4 +# (like the forward Layer 2). Calls the backward op directly with the +# forward-reference o/lse (bypassing the forward kernel) and compares to the +# Stage-2 reference (same full quantization both sides -> gap = port bugs). +# --------------------------------------------------------------------------- + +@cuda_only +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True]) +def test_backward_kernel_matches_reference(config, causal): + """A1 backward kernel vs Stage-2 golden reference — isolates port bugs. CDNA4-only.""" + from alto.kernels.mxfp8.mxfp8_quantization import convert_to_mxfp8 + + device = "cuda" + dtype = torch.bfloat16 + q, k, v = _make_qkv_bhsd(4, config, device, dtype) + do = _make_do_bhsd(4, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) + + # Quantize operands the way the forward autograd does, then run the backward op. + q8, q_scale = convert_to_mxfp8(q, mxfp_format="e4m3", axis=-1, is_2d_block=True) + k8, k_scale = convert_to_mxfp8(k, mxfp_format="e4m3", axis=-1, is_2d_block=True) + v8, v_scale = convert_to_mxfp8(v, mxfp_format="e4m3", axis=-1, is_2d_block=True) + + dq_k, dk_k, dv_k = torch.ops.alto.attention_mxfp8_backward_triton_impl( + do.contiguous(), + q8.contiguous(), + k8.contiguous(), + v8.contiguous(), + o_ref.contiguous(), + lse_ref.contiguous(), + q_scale, + k_scale, + v_scale, + sm_scale=sm_scale, + causal=causal, + layout="bhsd", + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlen_q=config.seqlen_q, + max_seqlen_k=config.seqlen_kv, + use_exp2=True, + ) + dq_r, dk_r, dv_r = mxfp8_attention_backward_reference_stage2(q, k, v, do, o_ref, lse_ref, sm_scale, causal) + + rows = [] + pairs = [("dQ", dq_r, dq_k), ("dK", dk_r, dk_k), ("dV", dv_r, dv_k)] + for name, ref, got in pairs: + rows.append([f"{name} (kernel vs golden)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + # Threshold placeholder — calibrate on the first CDNA4 run. + for name, ref, got in pairs: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > 0.99, f"{name} kernel-vs-golden cosine-sim too low: {sim}" + assert snr > 25, f"{name} kernel-vs-golden SNR too low: {snr}" diff --git a/tests/unittest/mxfp8/utils.py b/tests/unittest/mxfp8/utils.py index 5990fe84..750d8761 100644 --- a/tests/unittest/mxfp8/utils.py +++ b/tests/unittest/mxfp8/utils.py @@ -377,3 +377,104 @@ def mxfp8_attention_backward_reference( dv = dv_full return dq, dk, dv + + +def mxfp8_attention_backward_reference_stage2( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + do: torch.Tensor, + o: torch.Tensor, + softmax_lse: torch.Tensor, + sm_scale: float, + causal: bool, + block_size: int = BLOCK_SIZE_DEFAULT, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Pure-PyTorch golden reference for the **Stage-2 (A1)** MXFP8 backward. + + Where Stage-1 (``mxfp8_attention_backward_reference``) runs every backward + matmul in fp32 on the dequantized Q/K/V, Stage-2 additionally quantizes the + backward-only tensors (dO / P / dS) along each dot's reduction axis — a + numerical model of what ``tl.dot_scaled`` does in the kernel. A kernel-vs-this + gap isolates Triton port bugs from mxfp8 error; a this-vs-bf16-SDPA gap + measures the full mxfp8 backward quantization error. No ``tl.dot_scaled`` + here (plain fp32 matmul on dequantized operands), so it runs on CPU / MI250. + + **A1 operand source (2026-07-16 decision):** the kernel reuses the forward's + saved e4m3 q/k/v + their 2D-block scale directly in *every* dot — including + the seqlen-reduction dots — with no re-quantization. So this reference uses a + single 2D-block dequant ``q_dq``/``k_dq``/``v_dq`` everywhere (no double + quantization along seqlen, unlike the retired option-A reference). Only + dO/P/dS are quantized fresh (1D per-row along their reduction axis), since the + forward never saw them. + + Quantization axis per dot (plan §9.5):: + + S = Q @ Kᵀ reduction head_dim : Q,K = 2D-block dequant (reuse) + dP = dO @ Vᵀ reduction head_dim_v : dO 1D axis=-1 ; V 2D reuse + dV = Pᵀ @ dO reduction seqlen_q : P,dO 1D along sq + dK = dSᵀ @ Q reduction seqlen_q : dS 1D along sq ; Q 2D reuse + dQ = dS @ K reduction seqlen_k : dS 1D along sk ; K 2D reuse + + Same top-left causal mask as the other references (matches SDPA on bhsd), so + causal tests must keep ``seqlen_q == seqlen_k``. Args mirror + ``mxfp8_attention_backward_reference``; ``o`` / ``softmax_lse`` must be the + forward-reference outputs. Returns ``(dq, dk, dv)``. + """ + b, hq, sq, dqk = q.shape + _, hk, sk, dvv = v.shape + assert hq % hk == 0, f"nheads_q ({hq}) must be a multiple of nheads_k ({hk})" + n_rep = hq // hk + + def qdq(x, axis): + return _mxfp8_qdq(x.to(torch.float32), axis=axis, is_2d_block=False, block_size=block_size) + + # q/k/v: single 2D-block dequant, reused in every dot (A1, no re-quant). + q_dq = _mxfp8_qdq(q, axis=-1, is_2d_block=True, block_size=block_size).to(torch.float32) + k_dq = _mxfp8_qdq(k, axis=-1, is_2d_block=True, block_size=block_size).to(torch.float32) + v_dq = _mxfp8_qdq(v, axis=-1, is_2d_block=True, block_size=block_size).to(torch.float32) + if n_rep > 1: + k_dq = k_dq.repeat_interleave(n_rep, dim=1) + v_dq = v_dq.repeat_interleave(n_rep, dim=1) + + do_f = do.to(torch.float32) + o_f = o.to(torch.float32) + lse = softmax_lse.to(torch.float32) + + # S -> P (recompute from saved lse; P feeding dS is elementwise, not quantized). + s = torch.matmul(q_dq, k_dq.transpose(-1, -2)) * sm_scale + if causal: + q_pos = torch.arange(sq, device=q.device) + k_pos = torch.arange(sk, device=q.device) + allowed = k_pos[None, :] <= q_pos[:, None] # top-left + s = torch.where(allowed[None, None, :, :], s, torch.full_like(s, float("-inf"))) + p = torch.exp(s - lse[..., None]) + p = torch.nan_to_num(p, nan=0.0, posinf=0.0, neginf=0.0) + + # dP = dO @ Vᵀ (reduction head_dim_v): dO 1D along head_dim_v, V 2D reuse. + do_hd = qdq(do_f, axis=-1) + dp = torch.matmul(do_hd, v_dq.transpose(-1, -2)) + delta = (o_f * do_f).sum(dim=-1) + ds = p * (dp - delta[..., None]) # fp32 elementwise + + # dV = Pᵀ @ dO (reduction seqlen_q): P and dO 1D along sq. + p_sq = qdq(p, axis=-2) + do_sq = qdq(do_f, axis=-2) + dv_full = torch.matmul(p_sq.transpose(-1, -2), do_sq) + + # dQ = sm_scale * dS @ K (reduction seqlen_k): dS 1D along sk, K 2D reuse. + ds_sk = qdq(ds, axis=-1) + dq = sm_scale * torch.matmul(ds_sk, k_dq) + + # dK = sm_scale * dSᵀ @ Q (reduction seqlen_q): dS 1D along sq, Q 2D reuse. + ds_sq = qdq(ds, axis=-2) + dk_full = sm_scale * torch.matmul(ds_sq.transpose(-1, -2), q_dq) + + if n_rep > 1: + dk = dk_full.view(b, hk, n_rep, sk, dqk).sum(dim=2) + dv = dv_full.view(b, hk, n_rep, sk, dvv).sum(dim=2) + else: + dk = dk_full + dv = dv_full + + return dq, dk, dv From 2260441b6c89861cc0ed0be417845d5287de3cf3 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Thu, 16 Jul 2026 03:43:29 -0500 Subject: [PATCH 09/14] feat: dispatch mxfp8_e4m3 attention to triton_attention_mxfp8 --- alto/kernels/dispatch/attention.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/alto/kernels/dispatch/attention.py b/alto/kernels/dispatch/attention.py index 33727623..c5d0ae3f 100644 --- a/alto/kernels/dispatch/attention.py +++ b/alto/kernels/dispatch/attention.py @@ -6,6 +6,7 @@ from torchtitan.models.common.attention import (ScaledDotProductAttentionWrapper) from alto.kernels.fp4.mxfp4.triton_flash_attention_mxfp4 import triton_attention_mxfp4 +from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 from .config import TrainingOpConfig __all__ = ["LPScaledDotProductAttentionWrapper"] @@ -20,6 +21,8 @@ def __init__(self, config: TrainingOpConfig): if isinstance(config, TrainingOpConfig) and config.precision == "mxfp4": self.attn_func = triton_attention_mxfp4 + elif isinstance(config, TrainingOpConfig) and config.precision == "mxfp8_e4m3": + self.attn_func = triton_attention_mxfp8 else: raise ValueError(f"Unsupported SDPA config: {config}") From f5bfedc9f8ae33307fbb22bfbccccdca78f2cea1 Mon Sep 17 00:00:00 2001 From: root Date: Tue, 21 Jul 2026 06:30:05 +0000 Subject: [PATCH 10/14] fix: validate mxfp8 attention public autograd path Remove a stale Triton JIT global dependency in backward quantization and add a user-facing autograd test so the MXFP8 attention wrapper is covered end to end. --- .../mxfp8/triton_flash_attention_mxfp8.py | 2 +- .../mxfp8/test_mxfp8_attention_reference.py | 78 +++++++++++++++++++ 2 files changed, 79 insertions(+), 1 deletion(-) diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index 60a25ae3..cee983c6 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -1057,7 +1057,7 @@ def _mx_quant( xq = _quantize_fp8( x, scales, - philox_seed, + 0, 0, BLOCK_M=BM, BLOCK_N=BN, diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index 35264391..9aa94d12 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -234,6 +234,84 @@ def test_backward_stage2_reference_matches_sdpa(config, causal): assert snr > 18, f"{name} stage2-vs-sdpa SNR too low: {snr}" +# --------------------------------------------------------------------------- +# Public autograd path — user-facing forward + backward wrapper +# --------------------------------------------------------------------------- + +public_autograd_cases = [ + AttnConfig(seqlen_q=128, seqlen_kv=128, num_head_q=4, num_head_kv=4, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=128, seqlen_kv=128, num_head_q=8, num_head_kv=2, head_dim_qk=128, head_dim_v=128), # GQA +] + + +@cuda_only +@pytest.mark.parametrize("config", public_autograd_cases) +@pytest.mark.parametrize("causal", [True]) +def test_public_autograd_matches_sdpa(config, causal): + """Public ``triton_attention_mxfp8`` forward/backward must wire gradients correctly. + + The lower-level backward op is tested across the large shape grid below. This + smaller test covers the user-facing autograd wrapper: saved tensors, ctx + metadata, backward op dispatch, and returned gradient positions. + """ + from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 + + device = "cuda" + dtype = torch.bfloat16 + n_rep = config.num_head_q // config.num_head_kv + + q, k, v = _make_qkv_bhsd(1, config, device, dtype) + do = _make_do_bhsd(1, config, device, dtype) + sm_scale = config.head_dim_qk**(-0.5) + + q_ref = q.clone().detach().requires_grad_(True) + k_ref = k.clone().detach().requires_grad_(True) + v_ref = v.clone().detach().requires_grad_(True) + o_ref = torch.nn.functional.scaled_dot_product_attention( + q_ref, k_ref, v_ref, is_causal=causal, scale=sm_scale, enable_gqa=n_rep > 1) + o_ref.backward(do) + + q_kernel = q.clone().detach().requires_grad_(True) + k_kernel = k.clone().detach().requires_grad_(True) + v_kernel = v.clone().detach().requires_grad_(True) + o_kernel = triton_attention_mxfp8( + q_kernel.contiguous(), + k_kernel.contiguous(), + v_kernel.contiguous(), + bias=None, + alibi_slopes=None, + sm_scale=sm_scale, + dropout_p=0.0, + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlens_q=config.seqlen_q, + max_seqlens_k=config.seqlen_kv, + causal=causal, + return_scores=False, + use_exp2=True, + layout="bhsd", + )[0] + o_kernel.backward(do) + + pairs = [ + ("O", o_ref, o_kernel, 0.99, 20), + ("dQ", q_ref.grad, q_kernel.grad, 0.995, 18), + ("dK", k_ref.grad, k_kernel.grad, 0.995, 18), + ("dV", v_ref.grad, v_kernel.grad, 0.995, 18), + ] + rows = [] + for name, ref, got, _, _ in pairs: + rows.append([f"{name} (public autograd vs sdpa)", calc_snr(ref, got), calc_cossim(ref, got)]) + print() + print(tabulate(rows, headers=["Tensor", "SNR", "Cosine Sim"], tablefmt="github")) + + for name, ref, got, sim_threshold, snr_threshold in pairs: + sim = calc_cossim(ref, got) + snr = calc_snr(ref, got) + assert sim > sim_threshold, f"{name} public-autograd-vs-sdpa cosine-sim too low: {sim}" + assert snr > snr_threshold, f"{name} public-autograd-vs-sdpa SNR too low: {snr}" + + # --------------------------------------------------------------------------- # Backward Layer 2 — kernel vs Stage-2 (A1) golden reference # The A1 backward kernel runs tl.dot_scaled in-kernel, so this requires CDNA4 From 6e9cb55bd4615bcf72d8f803eb5ad425ca730267 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 29 Jul 2026 03:17:44 -0500 Subject: [PATCH 11/14] fix: assert mxfp8 attention layout and causal shape constraints --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 68 ++++++++++++++++++- .../mxfp8/triton_flash_attention_mxfp8.py | 13 ++++ 2 files changed, 78 insertions(+), 3 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index d7e536ca..1d3e62de 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -97,7 +97,7 @@ mxfp4 attention 里与量化正交、原样保留的 FA v2 forward 骨架:caus - **V scale**:PV reduction 维是 seqlen_k,V scale 沿 seqlen_k:`[..., head_dim_v, seqlen_k/32]` - **P scale**:kernel 内动态生成,沿 `BLOCK_N`(seqlen_k)方向 -**约束**:head_dim 与 seqlen_k 必须被 `QUANT_BLOCK_SIZE=32` 整除。head_dim ∈ {128,192}、seqlen ∈ {1024,2048} 均满足。head_dim<64 因 `tl.dot_scaled` 限制不支持(mxfp4 亦然)。 +**约束**:head_dim 与 seqlen_k 必须被 `QUANT_BLOCK_SIZE=32` 整除。head_dim ∈ {128,192}、seqlen ∈ {1024,2048} 均满足。head_dim<64 因 `tl.dot_scaled` 限制不支持(mxfp4 亦然)。**layout 只支持 `bhsd`**:2D-block scale 按数据张量 `shape[-2] × shape[-1]` 分块,只有 bhsd 时 `shape[-2]` 才是 seqlen;已加入口断言,见 §10.1octies。 ### 用户 API(对齐 mxfp4,加一个格式开关,去掉 fallback 开关) @@ -117,7 +117,7 @@ def triton_attention_mxfp8( causal: bool, return_scores: bool, use_exp2: bool, - layout: str, # "bshd" / "bhsd" / "thd" + layout: str, # 仅支持 "bhsd"(入口断言,见 §2 约束 / §10.1octies) *, fwd_format: str = "e4m3", # Q/K/V/P 格式,V1 固定 e4m3;e5m2 预留 bwd_grad_format: str = "e4m3", # dO/dP/dS 格式,V1 固定 e4m3;e5m2 预留给未来 @@ -459,6 +459,68 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - **CDNA4 验收**:A1 kernel vs A1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差);`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 - padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 +### 10.1octies 截至 2026-07-29(CDNA4 真机验收通过;补 3 处护栏断言;causal 循环边界经分析判定**不修**) + +**环境:** gfx950(CDNA4)真机,`is_cdna4()=True`。Docker `exciting_kepler`(镜像 `wanghanthu/torchtitan:ubuntu22.04-pytorch2.12.0dev20260217-rocm7.2-patch`,`/home/yuesun/repos → /workspace`),共享集群上以 `HIP_VISIBLE_DEVICES=7` 单卡运行,其余 7 卡全程未占用。 + +**§10.1septies 的头号盲区已解除。** A1 backward kernel 的 `tl.dot_scaled` + 2D-scale 指针广播(含全仓无先例的 `_load_scale_sq` seqlen 轴 re-index)在真机跑通,且与 A1 参照吻合: + +| 检查 | 结果 | +|---|---| +| `test_backward_kernel_matches_reference`(A1 kernel vs A1 参照,全 `test_cases` 7 组) | ✅ **7 passed**,dQ/dK SNR 53~60 dB、dV 59~63 dB,cossim ≈ 1.0 | +| `test_public_autograd_matches_sdpa`(公开 autograd 端到端 vs bf16 SDPA) | ✅ **2 passed**,dQ/dK≈23.3 dB·cossim≈0.9976、dV≈26 dB·cossim≈0.9988 | +| `test_kernel_matches_reference`(forward Layer 2,§10.1sexies 在 gfx90a 上 7 failed) | ✅ **7 passed** | +| 全量 `test_mxfp8_attention_reference.py` + `test_mxfp8_attention.py` | ✅ **41 passed** | + +- kernel vs 参照 53~63 dB ⇒ **移植与 scale 广播无 bug**;端到端 23 dB 与 §10.1septies 参照预测的 23 dB 一致 ⇒ **误差全部来自 mxfp8 量化本身,不是 kernel 实现**。 +- **padded head_dim=192 已随 7 组网格覆盖**(config3/config4),§10.1septies「padded head 的 scale mask 仅基础兜底」在 backward 侧首次拿到真机信号;GQA(`num_head_kv` = 8/2 < q)同样覆盖。 +- 据此 **§7 验收标准第 1~5 项均已在 CDNA4 达成**;第 6 项(dispatch 路由 smoke)不在本次运行范围。 + +**backward causal 循环边界与自身掩码不一致:经分析判定不修(决策记录)** + +两个 backward kernel 的**掩码**按 bottom-right 对齐(`col_offset = N_CTX_Q - N_CTX_K`),但决定**循环范围**的两处把 `col_offset` 丢了,等价于硬编码了 `col_offset == 0`(即方阵): + +| 位置 | 现状(未改) | 非方阵下的后果 | +|---|---|---| +| `_bwd_kernel_dkdv` 的 `lo` | `(start_n*BLOCK_N - BLOCK_M + 1) // BLOCK_M * BLOCK_M` | `seqlen_k − seqlen_q > BLOCK_M` 时跳过必须计算的 query 块 → dK/dV 漏贡献 | +| `_bwd_kernel_dq` 的 `hi` | `BLOCK_M // BLOCK_N * (start_m+1) * BLOCK_N` | `seqlen_k > seqlen_q` 时只要差 1 就漏 key 列 → dQ 漏贡献 | + +一度改成了带 `col_offset` 的正确形式,**最终回退,只保留下面的断言**。理由按硬度排: + +1. **这条路不可达,所以它不是 bug。** 两个断言合起来把 `col_offset != 0` 完全堵死:`causal → seqlen_q == seqlen_k` 拦住定长;`layout == "bhsd"` 顺带杀掉 varlen(thd)——这点很关键,varlen 下每条序列实际长度不同,光断言 max_seqlen 是拦不住的。而 `col_offset == 0` 时新旧公式只差一块全掩块,输出必然逐位一致(已实测:撤回改动前后 21 个 SNR 数字一致到小数点后四位;因为 `p = exp(-inf) = 0`,量化成 e4m3 仍是 0,加进 fp32 累加器是精确的 0)。 +2. **它从来没有「独立可达」过。** 非方阵 + causal 这条路本来就因为 kernel(bottom-right)与参照/`F.sdpa`(top-left)的约定不一致而算不对,边界错只是被一个更大的破坏盖住。修边界并不能让这条路可用,必须先统一约定。 +3. **「负 `lo` 越界读内存」这个理由是假的,已实测推翻。** Triton 的整数 `//` 对负数**向零截断**(不是 Python 的下取整),实测 `BLOCK_M=BLOCK_N=64` 时原式在 `start_n=0..3` 给出 `[0, 0, 64, 128]`,floor 语义才会给 `[-64, 0, 64, 128]`。所以原式**不会**产生负 `lo`、不会越界;`tl.maximum(..., 0)` 是新公式自己引入的需求,不是旧代码的隐患。 +4. **唯一真实代价:方阵下每个 key block 多跑一个全掩 query block。** Triton 循环内无法提前退出,那块的 4 次 dot 实打实算完再被掩成 0。原式 `lo = max(0, (start_n−1)·64)`,正确式 `lo = start_n·64`,故 `start_n ≥ 1` 的每个 program 各多一次迭代:多出比例 `2(N−1) / (N(N+1))`,N = seqlen/64。seqlen 1024 → dkdv 内层迭代 **+11.0%**;2048 → **+5.9%**;4096 → **+3.0%**(迭代次数,非墙钟时间;dkdv 约占 backward 一半,实际再打对折)。序列越长越不值钱。 + +结论:这是一次**性能与可读性**改动,不是正确性修复;按 fix 提交会误导 reviewer 去找线上不存在的 bug。留待与下方 review 债一起独立评估(那批债里已有「`lo`/`hi` 隐含要求 `BLOCK_M == BLOCK_N`」一条,同源同处,应一并处理)。 + +**同源代码 `blockwise_fp8/triton_flash_attention_fp8_block.py` 有一字不差的两行**(1297、1665 行;mxfp8 backward 是从它 port 来的,**mxfp4 没有 backward,不是这段的来源**)。那份**没有** `seqlen_q == seqlen_k` 断言,autograd Function 的 `max_seqlens_q/k` 完全自由,但同样不可达:唯一生产入口 `blockwise_fa.py` 第 219 行 `assert query_states.shape[1] == key_states.shape[1]` 后把同一个 `seqlen` 传给 q 和 k;其 `test_attention.py` 的 10 组用例 `seqlen_q` 全等于 `seqlen_kv`。⇒ 同为潜伏、未暴露,本次不动。 + +**新增 3 处护栏断言(把「不报错但结果是垃圾」变成当场报错)** + +| 断言 | 位置 | 原因 | +|---|---|---| +| `layout == "bhsd"` | forward op + backward op 各一处 | `convert_to_mxfp8(is_2d_block=True)` 按数据张量 `shape[-2] × shape[-1]` 分 32×32 块,只有 bhsd 时 `shape[-2]` 才是 seqlen。`bshd`/`thd` 下它按 **nheads** 分组(nheads 恰为 32 倍数时连 `torch._check` 都拦不住),而 kernel 的 scale 指针数学假定按 seqlen 分组 → **静默读错 scale**。§2 契约里 `layout` 的三选一由此**收窄为只支持 bhsd**(生产 `dispatch/attention.py` 本就只传 bhsd) | +| `not causal or seqlen_q == seqlen_k` | backward op | §10.1ter 记录的「潜伏坑」落地:kernel 掩码 bottom-right、PyTorch 参照与 `F.sdpa` 掩码 top-left,两者只在方阵等价。V1 生产只跑 self-attention,故**不统一约定**(那是解决不存在的问题),改为断言拒绝。正向的 bottom-right 自身自洽、纯推理下 `sq != sk` 可用,**不拦 forward** | + +断言有效性单独验过:`causal=True, sq=64/sk=128` 走 backward → 按预期抛 AssertionError;`causal=False` 同形状 → 正常通过(非 causal 无对齐问题,不该拦);`layout="bshd"` 走 forward → 按预期抛 AssertionError。 + +**Code review 遗留债(已确认,本次未动,建议独立提交)** + +- `_MX_2D = tl.constexpr(False)`:名字叫 2D、值是 False,12 处调用全传它 ⇒ `_mx_quant` 的 `IS_2D_BLOCK` 永远走 1D 分支,而其 docstring 花两行解释了这个文件里不存在的 True 行为。参数只有一个取值就不该是参数。 +- `get_padded_headsize`(模块顶部)与 `get_padded_head_dim` 完全重复,前者无调用点。 +- scale 越界填充值不统一:forward `other=1`,backward `_load_scale_hd`/`_load_scale_sq` `other=127`。两者都不影响结果(对应数据已被掩成 0),但应统一为一个命名常量(E8M0 中性值是 127)。 +- `E4M3_TARGET_MAX_POW2` / `E4M3_MBITS` / `E4M3_FORMAT_ID` 手抄了 `mxfp8_quantization.FORMAT_TO_*` 的值,应直接 import。 +- backward docstring 大量引用外部文档编号(`plan §9.5`、`AB decision A1`、`dot a~g`),脱离本 plan 无法解读;建议改为按被算对象命名(qk / dp / dv / dk / dq)。 +- 两个 inner kernel 把运行时值标注成 `N_CTX_Q: tl.constexpr` / `N_CTX_K: tl.constexpr`(varlen 下由 `cu_seqlens` 运行时读出),类型标注不实。 +- `lo` / `hi` 隐含要求 `BLOCK_M == BLOCK_N`(若 `BLOCK_N=128`,`BLOCK_M//BLOCK_N` 变 0,原式直接给出 `hi=0` → dQ 全 0 且不报错);当前两者硬编码 64,无断言保护。 + +**仍待办:** + +- 上述 review 债清理(与本次断言分开提交)。`lo`/`hi` 的 `col_offset` 与 `BLOCK_M == BLOCK_N` 两条同源,建议合并评估。 +- 若未来要支持非方阵 causal(cross-attention / prefix),需先在 kernel 与参照之间统一 top-left 或 bottom-right 约定,**再把 `lo`/`hi` 补上 `col_offset`**(正确形式见本节决策记录,届时才是真修复),最后才放开该断言。三件事顺序不能颠倒。 +- §7 第 6 项 dispatch 路由 smoke 未跑。 + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**A1 kernel 的 `tl.dot_scaled` + 2D-scale 广播在 MI250 编不了,kernel 正确性完全待 CDNA4(kernel vs A1 参照已写好待跑)**(`_load_scale_sq` 为最大盲区)。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。 diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index cee983c6..1fdb6db8 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -899,6 +899,10 @@ def attention_mxfp8_forward_triton_impl( assert q_scale.is_contiguous() assert k_scale.is_contiguous() assert v_scale.is_contiguous() + assert layout == "bhsd", ( + f"MXFP8 attention requires layout='bhsd', got {layout!r}: the 2D-block scale groups the " + "data tensor's last two dims, which are (seqlen, head_dim) only for bhsd. Any other " + "layout groups across heads and the scale pointer math below reads the wrong scale.") # check if varlen is_varlen = layout == "thd" @@ -1833,8 +1837,17 @@ def attention_mxfp8_backward_triton_impl( if not do.is_contiguous(): do = do.contiguous() + assert layout == "bhsd", ( + f"MXFP8 attention requires layout='bhsd', got {layout!r}: the 2D-block scale groups the " + "data tensor's last two dims, which are (seqlen, head_dim) only for bhsd. Any other " + "layout groups across heads and the scale pointer math below reads the wrong scale.") + batch, nheads_q, nheads_k, head_size_qk, head_size_v, max_seqlen_q, max_seqlen_k = get_shape_from_layout( q, k, v, layout, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k) + assert not causal or max_seqlen_q == max_seqlen_k, ( + f"causal backward requires seqlen_q == seqlen_k, got {max_seqlen_q} vs {max_seqlen_k}: the " + "kernel masks bottom-right aligned while the PyTorch references and F.sdpa mask top-left, " + "and the two conventions only agree on square shapes.") stride_qz, stride_qh, stride_qm, stride_qk = get_strides_from_layout(q, layout) stride_kz, stride_kh, stride_kn, stride_kk = get_strides_from_layout(k, layout) stride_vz, stride_vh, stride_vn, stride_vk = get_strides_from_layout(v, layout) From acd2ed85cd4b446bb46c770a30513914723f5795 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 29 Jul 2026 03:56:59 -0500 Subject: [PATCH 12/14] perf: pin num_stages=1 for mxfp8 attention backward kernels --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 53 ++++++++++++++++++- .../mxfp8/triton_flash_attention_mxfp8.py | 4 ++ 2 files changed, 56 insertions(+), 1 deletion(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 1d3e62de..51dbe677 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -521,6 +521,57 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - 若未来要支持非方阵 causal(cross-attention / prefix),需先在 kernel 与参照之间统一 top-left 或 bottom-right 约定,**再把 `lo`/`hi` 补上 `col_offset`**(正确形式见本节决策记录,届时才是真修复),最后才放开该断言。三件事顺序不能颠倒。 - §7 第 6 项 dispatch 路由 smoke 未跑。 +### 10.1nonies 截至 2026-07-29(对照 blockwise_fp8 复审 backward;修掉 port 丢失的 `num_stages`,实测 1.15~1.51x) + +backward 是从 `blockwise_fp8/triton_flash_attention_fp8_block.py` port 来的(**不是 mxfp4——mxfp4 没有 backward**),故以它为基准做了第二轮逐行对照。结论:**无正确性 bug**,但发现一处实测性能损失。 + +**修掉:port 时丢失了 `num_stages=1`(唯一实质缺陷)** + +上游把这个启动参数藏在一个**只有一个空 config 的 `@triton.autotune`** 里(`get_autotune_bwd_configs`,1156~1166 行;挂在 `_bwd_kernel_dkdv` 1169 行、`_bwd_kernel_dq` 1524 行): + +```python +triton.Config({}, num_stages=1, num_warps=4) +``` + +它不调任何 BLOCK 尺寸,唯一作用就是钉住 `num_stages=1`。port 时这个装饰器被当作"可选的调优基础设施"删掉了,于是 mxfp8 的 backward 静默继承 Triton 默认值——实测(Triton 3.6.0)默认为 `num_warps=4, num_stages=2`,`num_warps` 恰好撞对,`num_stages` 翻倍。 + +`num_stages=2` 给内层循环加了两级软流水,多出的预取缓冲挤占寄存器;head_dim=192(padded 256)时 fp32 累加器与 scale tile 最多,退化最严重: + +| 形状(batch 4, causal, gfx950) | num_stages=2(继承默认) | num_stages=1(对齐上游) | 提速 | +|---|---|---|---| +| s1024, hq32/hkv8, d128 | 1.741 ms | 1.505 ms | **1.16x** | +| s2048, hq32/hkv8, d128 | 6.201 ms | 5.411 ms | **1.15x** | +| s2048, hq16/hkv16, d192 | 7.015 ms | 4.630 ms | **1.51x** | + +(`triton.testing.do_bench`,warmup 50 / rep 200,跑两遍复现,偏差 < 0.3%。同时试了 `num_warps=8`:2.084 / 7.538 / 9.091 ms,明显更差 ⇒ 上游的 4 是对的。) + +修法是给两个 launch 各加显式 `num_warps=4, num_stages=1`,**不照搬那个空 autotune**——参数只有一个取值就不该披着 autotune 的皮,那正是它被丢掉的原因。改后 41 passed。 + +**确认比上游更好的四处设计(不动)** + +- **`tl.dot_scaled` 消掉了一整套 scale 记账。** 上游是"普通 `tl.dot` + 事后乘标量 descale",于是 `p_scale` / `log_p_scale` / `acc_descale` 要在 forward inner、dkdv、dq 全程穿(softmax 后 `p *= p_scale`、epilogue `acc *= acc_descale`、backward 重算 p 补 `+ log_p_scale * RCP_LN2`、最后 `dq *= sm_scale / p_scale` 除回去)。换成硬件指令后这些**全部消失**——不是简化,是特殊情况不存在了。 +- **砍掉上游传了但没读的 kernel 参数**:上游给 dkdv 传 `Out`/`DO`/`DQ`、给 dq 传 `Out`/`DK`/`DV`。 +- **把上游的静默失效变成当场报错**:上游 forward 支持 dropout 但 backward 无 dropout 逻辑且不检查(开 dropout 训练即静默错梯度);alibi 同理(只进 DEBUG 打印,从不进 kernel)。mxfp8 在 autograd backward 三行 assert 挡住(2128~2130 行)。 +- **无 fallback**:上游留 `use_fp8` 开关让同一套 kernel 兼跑 bf16,代价是每个 kernel 里 `if USE_FP8` 分叉;mxfp8 只有 MX 一条路,符合 §1 定的"仅 CDNA4 无 fallback"。 + +**一处真 trade-off(未测,线索)**:dO 量化位置。上游在 preprocess 里量化一次、写出 `DO_FP8` + `do_scale`,两个 backward kernel 直接读。mxfp8 改为循环内现场量化,但 dkdv 内层**每次迭代量化 dO 两次**(1295 行按 head_dim 分组供 dp、1305 行转置后按 seqlen 分组供 dv),因为两个 dot 的归约轴不同、需要两套 1D scale 布局。若 dO 改用 2D-block scale 量化一次即可同时喂两个 dot ⇒ 直接落在 §10.1octies 债表里 `_MX_2D = tl.constexpr(False)` 那条上,值得一并测。 + +**本轮新增债(并入 §10.1octies 债表一起清)** + +- `USE_SR` 在两处调用点(forward 363 行、`_mx_quant` 内部 1072 行)硬编码 `False`:随机取整整条链路通着但永不启用,与 `_MX_2D=False` 同类。 +- `is_varlen` / thd 分支已被 `layout == "bhsd"` 断言证明**不可达**(forward 908/953-959/1015 行;backward driver 一路把 `cu_seqlens` 传进 kernel)。这是上一轮加断言的直接后果——按"不写兼容/回退代码"的标准该删,要留就得把 varlen 真正做完。 +- forward `register_fake`(2154 行起)仍保留 thd 的 LSE shape 分支,而真实 op 已 assert 拒绝非 bhsd。仅在 torch.compile trace 非 bhsd 时不一致(那种情况本也会 assert),化妆品级。 +- `AUTOTUNE` / `PERF` 模块常量无引用(与 `get_padded_headsize` 一样,都是随上游一起抄来的死物,上游也是死的)。 + +**四个已验证排除的虚警(记录在此,避免重复排查)** + +| 初查疑点 | 排除依据 | +|---|---| +| 上游 dkdv 的 exp2 分支对 `l_i` 双乘 `RCP_LN2` | 误读。1480 行是 `exp2(qk - l_i[:,None] + log_p_scale * RCP_LN2)`,`* RCP_LN2` 作用在 `log_p_scale` 上,`l_i` 只乘一次。两边都对 | +| 上游 autograd backward 返回 18 个梯度但 forward 只有 16 个输入 | 实测最小 `autograd.Function`:PyTorch **容忍**尾部多余的 `None`。mxfp8 这边是 15 对 15 精确匹配 | +| mxfp8 未 assert `o` / `softmax_lse` 连续(上游有) | `get_strides_from_layout` 与 `softmax_lse.stride()` 均直读张量真实 stride,非连续也读得对;`delta = empty_like(softmax_lse)` 继承同布局。无正确性风险 | +| mxfp8 缺 head_dim 整除检查(上游 assert `>=32` 且 `%2==0`) | mxfp8 需要的是更强的 `%32==0`,已由上游 `convert_to_mxfp8` 的 `torch._check(shape[-1] % block_size == 0)` 保证(2D 模式另检 `shape[-2]`)。上游保证的不在下游重复检查。这也解释了为何 layout 那条必须自己加:`convert_to_mxfp8` 检的是 `shape[-2] % 32`,bshd 下那是 nheads,32 头刚好整除,拦不住 | + ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。 +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。**2026-07-29 二轮对照 blockwise_fp8(backward 的真正上游)复审(§10.1nonies):无正确性 bug,修掉 port 时丢失的 `num_stages=1`(上游藏在单 config 空 autotune 里),实测 backward 提速 1.15~1.51x(d192 最显著)。** diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index 1fdb6db8..9a705a34 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -1965,6 +1965,8 @@ def attention_mxfp8_backward_triton_impl( IS_VARLEN=is_varlen, QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT, USE_ASM=is_cdna4(), + num_warps=4, + num_stages=1, ) wrap_triton(_bwd_kernel_dkdv)[(batch * nheads_k, num_block_n)]( @@ -2030,6 +2032,8 @@ def attention_mxfp8_backward_triton_impl( IS_VARLEN=is_varlen, QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT, USE_ASM=is_cdna4(), + num_warps=4, + num_stages=1, ) return dq, dk, dv From 943ab2644d1b8fb1d3cad7d615511933c48d02ed Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Wed, 29 Jul 2026 05:55:49 -0500 Subject: [PATCH 13/14] fix: correct mxfp8 attention GQA backward stride --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 63 ++++++++++-- .../mxfp8/triton_flash_attention_mxfp8.py | 6 +- tests/unittest/mxfp8/test_mxfp8_attention.py | 12 ++- .../mxfp8/test_mxfp8_attention_reference.py | 95 +++++++++++++++---- 4 files changed, 146 insertions(+), 30 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 51dbe677..05bde1d3 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -238,7 +238,9 @@ def triton_attention_mxfp8( - 门槛:cos-sim > 0.99 + SNR 硬断言(验收标准 3)。 3. **backward 校验**:完整 `autograd` 反向 vs bf16 SDPA autograd,比 dQ/dK/dV 的 cos-sim / SNR(验收标准 4)。 -参数网格:`batch=4` × mxfp4 的 `test_cases`(causal=True)× 首层额外跑 causal=False。 +参数网格:`batch=4` × mxfp4 的 `test_cases` × `causal ∈ {True, False}` 覆盖 forward、 +public autograd、kernel-vs-golden forward/backward;另设 `non_causal_cases` 覆盖 +`seqlen_q != seqlen_k`(只在 `causal=False` 下合法)。 --- @@ -501,7 +503,7 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` | 断言 | 位置 | 原因 | |---|---|---| | `layout == "bhsd"` | forward op + backward op 各一处 | `convert_to_mxfp8(is_2d_block=True)` 按数据张量 `shape[-2] × shape[-1]` 分 32×32 块,只有 bhsd 时 `shape[-2]` 才是 seqlen。`bshd`/`thd` 下它按 **nheads** 分组(nheads 恰为 32 倍数时连 `torch._check` 都拦不住),而 kernel 的 scale 指针数学假定按 seqlen 分组 → **静默读错 scale**。§2 契约里 `layout` 的三选一由此**收窄为只支持 bhsd**(生产 `dispatch/attention.py` 本就只传 bhsd) | -| `not causal or seqlen_q == seqlen_k` | backward op | §10.1ter 记录的「潜伏坑」落地:kernel 掩码 bottom-right、PyTorch 参照与 `F.sdpa` 掩码 top-left,两者只在方阵等价。V1 生产只跑 self-attention,故**不统一约定**(那是解决不存在的问题),改为断言拒绝。正向的 bottom-right 自身自洽、纯推理下 `sq != sk` 可用,**不拦 forward** | +| `not causal or seqlen_q == seqlen_k` | **forward op + backward op 各一处** | §10.1ter 记录的「潜伏坑」落地:kernel 掩码 bottom-right、PyTorch 参照与 `F.sdpa` 掩码 top-left,两者只在方阵等价。V1 生产只跑 self-attention,故**不统一约定**(那是解决不存在的问题),改为断言拒绝。
*(本节初稿曾写「正向自身自洽、纯推理下 `sq != sk` 可用,不拦 forward」,随后改为**两侧都拦**:一个 op 允许、另一个禁止同一种形状,调用者只能靠踩坑才知道边界在哪。forward 侧断言配 `test_causal_forward_rejects_non_square_shape` 回归。)* | 断言有效性单独验过:`causal=True, sq=64/sk=128` 走 backward → 按预期抛 AssertionError;`causal=False` 同形状 → 正常通过(非 causal 无对齐问题,不该拦);`layout="bshd"` 走 forward → 按预期抛 AssertionError。 @@ -519,11 +521,11 @@ forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias` - 上述 review 债清理(与本次断言分开提交)。`lo`/`hi` 的 `col_offset` 与 `BLOCK_M == BLOCK_N` 两条同源,建议合并评估。 - 若未来要支持非方阵 causal(cross-attention / prefix),需先在 kernel 与参照之间统一 top-left 或 bottom-right 约定,**再把 `lo`/`hi` 补上 `col_offset`**(正确形式见本节决策记录,届时才是真修复),最后才放开该断言。三件事顺序不能颠倒。 -- §7 第 6 项 dispatch 路由 smoke 未跑。 +- ~~§7 第 6 项 dispatch 路由 smoke 未跑~~ → 已于 §10.1nonies 补跑通过。 -### 10.1nonies 截至 2026-07-29(对照 blockwise_fp8 复审 backward;修掉 port 丢失的 `num_stages`,实测 1.15~1.51x) +### 10.1nonies 截至 2026-07-29(对照 blockwise_fp8 复审 backward;**修掉 dkdv 组循环的 dO head stride 真 bug**;修掉 port 丢失的 `num_stages`,实测 1.15~1.51x;补齐 §7 第 6 项) -backward 是从 `blockwise_fp8/triton_flash_attention_fp8_block.py` port 来的(**不是 mxfp4——mxfp4 没有 backward**),故以它为基准做了第二轮逐行对照。结论:**无正确性 bug**,但发现一处实测性能损失。 +backward 是从 `blockwise_fp8/triton_flash_attention_fp8_block.py` port 来的(**不是 mxfp4——mxfp4 没有 backward**),故以它为基准做了第二轮逐行对照。结论:**一处正确性 bug(GQA × 非对称 head_dim 下 dK/dV 为 nan 乃至进程崩溃,随上游一起继承)** + 一处实测性能损失。 **修掉:port 时丢失了 `num_stages=1`(唯一实质缺陷)** @@ -545,7 +547,7 @@ triton.Config({}, num_stages=1, num_warps=4) (`triton.testing.do_bench`,warmup 50 / rep 200,跑两遍复现,偏差 < 0.3%。同时试了 `num_warps=8`:2.084 / 7.538 / 9.091 ms,明显更差 ⇒ 上游的 4 是对的。) -修法是给两个 launch 各加显式 `num_warps=4, num_stages=1`,**不照搬那个空 autotune**——参数只有一个取值就不该披着 autotune 的皮,那正是它被丢掉的原因。改后 41 passed。 +修法是给两个 launch 各加显式 `num_warps=4, num_stages=1`,**不照搬那个空 autotune**——参数只有一个取值就不该披着 autotune 的皮,那正是它被丢掉的原因。改后全套通过(当时 41 个用例,测试矩阵扩到 75 见下)。 **确认比上游更好的四处设计(不动)** @@ -556,6 +558,53 @@ triton.Config({}, num_stages=1, num_warps=4) **一处真 trade-off(未测,线索)**:dO 量化位置。上游在 preprocess 里量化一次、写出 `DO_FP8` + `do_scale`,两个 backward kernel 直接读。mxfp8 改为循环内现场量化,但 dkdv 内层**每次迭代量化 dO 两次**(1295 行按 head_dim 分组供 dp、1305 行转置后按 seqlen 分组供 dv),因为两个 dot 的归约轴不同、需要两套 1D scale 布局。若 dO 改用 2D-block scale 量化一次即可同时喂两个 dot ⇒ 直接落在 §10.1octies 债表里 `_MX_2D = tl.constexpr(False)` 那条上,值得一并测。 +**补齐 §7 验收标准第 6 项:dispatch 路由 smoke(此前唯一未验项)** + +经真实 dispatch 路径 `LPScaledDotProductAttentionWrapper(TrainingOpConfig(precision="mxfp8_e4m3"))` 跑 forward + backward(b2 / hq32 / hkv8 / s1024 / d128 / bf16):路由确实落到 `triton_attention_mxfp8`;输出与 dQ/dK/dV 全部有限且非零;**dK/dV 形状为 kv-head 形状 `(2, 8, 1024, 128)` 而 dQ 为 `(2, 32, …)`,GQA 的 head 归约方向正确**;`is_causal=False` 同样跑通。⇒ §7 验收标准 **1~6 全部达成**。 + +顺带确认 dispatch 层与本次两处断言天然兼容:`attention.py` 第 68 行硬传 `layout="bhsd"`,第 63~64 行的 `max_seqlens_q/k` 取自真实形状,self-attention 下必然相等。 + +**修掉一个真 bug(本次 review 唯一的正确性缺陷):dkdv 组循环用错了 dO 的 head stride** + +`_bwd_kernel_dkdv` 的 GQA 组循环里,推进到下一个 query head 时把 dO 的指针按 **q 的** head stride 前进: + +```python +q_offset += stride_qh +do_offset += stride_qh # ← 应为 stride_doh +``` + +`stride_doh` 本来就传进了 kernel(1345 行)、初始化 `do_offset` 时也用对了(1419 行),只有这一步用错。q 的最后一维是 `head_dim_qk`、dO 的是 `head_dim_v`,所以 bhsd 连续布局下 `stride_qh = seqlen * head_dim_qk`、`stride_doh = seqlen * head_dim_v`,**两者仅在 `head_dim_qk == head_dim_v` 时相等**。 + +触发条件是两个条件**同时**成立:`GROUP_SIZE > 1`(GQA)**且** `head_dim_qk != head_dim_v`。原 `test_cases` 七组恰好把这两个轴**各自覆盖但从未交叉**(config2/3/6/7 是 GQA 但 128/128;config4/5 是 192/128 但 `num_head_q == num_head_kv`),所以 41 passed 完全盖不住它。 + +实测(`head_dim_qk=192, head_dim_v=128`,backward kernel vs A1 参照): + +| 形状 | 修前 dQ | 修前 dK / dV | 修后 dK / dV | +|---|---|---|---| +| hq8/hkv2(GROUP_SIZE=4) | 60.50 dB ✅ | **nan / nan** | 60.44 / 62.84 dB | +| hq8/hkv4(GROUP_SIZE=2) | 61.39 dB ✅ | **nan / nan** | 61.44 / 62.44 dB | +| hq16/hkv2(GROUP_SIZE=8) | 57.57 dB ✅ | **nan / nan** | 57.72 / 62.42 dB | + +dQ 全程正常——`_bwd_kernel_dq` 每个 query head 一个 program,没有这个组循环。GROUP_SIZE 与 seqlen 再大一些(seqlen 1024 / GROUP_SIZE 4)时越界读得更远,直接是 **HIP 内存访问错误、进程 abort**,不只是 nan。 + +**测试矩阵的两个单值轴一并补掉(同类隐患,不只是这一个 bug)** + +查覆盖时发现分层是**反过来的**:三个纯 PyTorch 参照测试都跑 `causal=[True, False]`,而**每一个真正启动 kernel 的测试都被钉在 `[True]`**——不可能有移植 bug 的那层测两个值,会有移植 bug 的那层只测一个。`blockwise_fp8` 的两个 kernel 测试都是 `[True, False]`;mxfp8 的测试骨架抄自 mxfp4(doc 原话 "mirrors the mxfp4 attention test"),把严格性一起抄丢了。(另查:mxfp4 测试里那句被注释掉的 `[True, False]` 属于整个 `test_attention_fp8_with_sparse_do` 函数被注掉,**不是**有人发现 causal=False 挂掉才关的。) + +四处 `[True]` → `[True, False]`(`test_kernel_matches_reference` / `test_backward_kernel_matches_reference` / `test_public_autograd_matches_sdpa` / `test_attention`)。非 causal 走的是不同代码:`lo = 0` 与 `hi = num_block_n * BLOCK_N` 替换掉两个 causal 边界公式,掩码跳过,forward 的分块切分与「整块被掩→写 0 提前退出」也绕开。实测全过。 + +**非方形形状补上 kernel 级覆盖**:`causal` 断言现在 forward/backward 两侧都拦,所以 `seqlen_q != seqlen_k` **只在非 causal 下合法**——而它此前只有纯 PyTorch 参照测试覆盖(`reference_cases` 里的 `sq128/sk256`),kernel 侧为零。新增 `non_causal_cases`(`sq256/sk512` 与 `sq512/sk256 + d192/128`)+ 两个测试;把两个 kernel 测试的主体抽成 `_check_forward_kernel_vs_reference` / `_check_backward_kernel_vs_reference` 复用,**没有在原测试里加 `if causal and not square: skip` 这类条件分支**(那是往正常路径塞特殊情况)。实测非方形 × GQA × 非对称 head_dim 三重交叉 56~61 dB、cossim≈1.0,即这条路本来就是好的,只是没人测过。 + +用例数 41 → **75**(`test_kernel_matches_reference` 16、`test_backward_kernel_matches_reference` 16、`test_attention` 16、三个参照测试各 6、`public_autograd` 4、两个非方形各 2、断言回归 1),全套 **11 秒**。 + +**测试矩阵补上 GQA × 非对称 head_dim 交叉点**:`test_cases` 新增 `num_head_q=32, num_head_kv=8, head_dim_qk=192, head_dim_v=128`。已验证这个配置在把修复撤回后确实会崩(不是静默通过)⇒ 是有效的回归防护。修复对原有七组配置**数值零影响**(`head_dim_qk == head_dim_v` 时两个 stride 本就相同,实测 SNR 逐位一致)。 + +**上游 `blockwise_fp8` 第 1378 行是一字不差的同一行**,其 `test_attention.py` 的 10 组用例也有完全相同的覆盖盲区(GQA 组全是 128/128,192/128 的两组全是 `num_head_q == num_head_kv`)⇒ 那份代码在 GQA + 非对称 head_dim 下同样会崩,且未被断言拦住。 + +**同类缺陷全仓排查(结论:无第二例)** + +既然 `num_stages` 是"必需参数藏在看起来可选的装饰器里"而丢失,扫了 `alto/kernels` 全部 20 个含 `@triton.jit` 的文件。循环密集型(attention / grouped GEMM)全部有显式 launch 配置:mxfp4 attention forward(392~393 行 `num_stages=1, num_warps=4`,无 backward)、blockwise attention fwd+bwd、三个 grouped_gemm 各有独立 `autotune.py`(6/20/28 configs)。其余无配置的文件都是单趟、访存受限的量化/elementwise kernel,`num_stages` 对它们无意义。**本缺陷仅此一处。** + **本轮新增债(并入 §10.1octies 债表一起清)** - `USE_SR` 在两处调用点(forward 363 行、`_mx_quant` 内部 1072 行)硬编码 `False`:随机取整整条链路通着但永不启用,与 `_MX_2D=False` 同类。 @@ -574,4 +623,4 @@ triton.Config({}, num_stages=1, num_warps=4) ### 10.2 一句话总结 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。**2026-07-29 二轮对照 blockwise_fp8(backward 的真正上游)复审(§10.1nonies):无正确性 bug,修掉 port 时丢失的 `num_stages=1`(上游藏在单 config 空 autotune 里),实测 backward 提速 1.15~1.51x(d192 最显著)。** +**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。**2026-07-29 二轮对照 blockwise_fp8(backward 的真正上游)复审(§10.1nonies):修掉 `_bwd_kernel_dkdv` 组循环误用 `stride_qh` 推进 dO 指针的真 bug(GQA × `head_dim_qk != head_dim_v` 下 dK/dV 为 nan 甚至进程 abort,测试矩阵两轴从未交叉故长期未暴露,已补交叉配置);修掉 port 时丢失的 `num_stages=1`(上游藏在单 config 空 autotune 里),实测 backward 提速 1.15~1.51x(d192 最显著);同类缺陷全仓排查无第二例;**补掉测试矩阵的两个单值轴(kernel 侧 causal 一直只测 True;非方形此前只有纯 PyTorch 参照覆盖),用例数 41 → 75、全套 11 秒**;补跑 dispatch 路由 smoke ⇒ **§7 验收标准 1~6 全部达成**。** diff --git a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py index 9a705a34..33828f1d 100644 --- a/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py +++ b/alto/kernels/mxfp8/triton_flash_attention_mxfp8.py @@ -913,6 +913,10 @@ def attention_mxfp8_forward_triton_impl( batch, nheads_q, nheads_k, head_size_qk, head_size_v, seqlen_q, seqlen_k = get_shape_from_layout( q, k, v, layout, cu_seqlens_q, cu_seqlens_k, max_seqlens_q, max_seqlens_k) + assert not causal or seqlen_q == seqlen_k, ( + f"causal forward requires seqlen_q == seqlen_k, got {seqlen_q} vs {seqlen_k}: the " + "kernel uses bottom-right causal masking while PyTorch SDPA/reference tests use top-left " + "masking for non-square shapes.") o_shape = (*q.shape[:-1], head_size_v) o = torch.empty( o_shape, @@ -1498,7 +1502,7 @@ def _bwd_kernel_dkdv( USE_ASM, ) q_offset += stride_qh - do_offset += stride_qh + do_offset += stride_doh q_scale_base += stride_qsh l_offset += stride_ldh d_offset += stride_ldh diff --git a/tests/unittest/mxfp8/test_mxfp8_attention.py b/tests/unittest/mxfp8/test_mxfp8_attention.py index 35abdcfd..0b7cf140 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention.py @@ -3,9 +3,9 @@ # SPDX-License-Identifier: MIT """Forward numerical tests for the MXFP8 (e4m3) flash attention kernel. -Mirrors ``tests/unittest/mxfp4/test_mxfp_attention.py``. The MXFP8 kernel is -forward-only (backward is not implemented), so only the forward output is -validated against a bf16 SDPA reference via SNR / cosine-similarity. +Mirrors ``tests/unittest/mxfp4/test_mxfp_attention.py`` for the forward +kernel-vs-bf16 SDPA check. Backward and kernel-vs-golden reference coverage live +in ``test_mxfp8_attention_reference.py``. """ import pytest @@ -61,12 +61,16 @@ def __init__(self, seqlen_q, seqlen_kv, num_head_q, num_head_kv, head_dim_qk, he AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=128, num_head_kv=128, head_dim_qk=192, head_dim_v=128), AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=48, num_head_kv=8, head_dim_qk=128, head_dim_v=128), AttnConfig(seqlen_q=2048, seqlen_kv=2048, num_head_q=64, num_head_kv=8, head_dim_qk=128, head_dim_v=128), + # GQA crossed with head_dim_qk != head_dim_v. The cases above cover those two + # axes only in isolation, which leaves the dO head-stride advance in the + # backward dkdv group loop untested (it produced NaN dK/dV). + AttnConfig(seqlen_q=1024, seqlen_kv=1024, num_head_q=32, num_head_kv=8, head_dim_qk=192, head_dim_v=128), ] @pytest.mark.parametrize("batch", [4]) @pytest.mark.parametrize("config", test_cases) -@pytest.mark.parametrize("causal", [True]) +@pytest.mark.parametrize("causal", [True, False]) def test_attention(batch, config, causal): device = "cuda" dtype = torch.bfloat16 diff --git a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py index 9aa94d12..f3fa9694 100644 --- a/tests/unittest/mxfp8/test_mxfp8_attention_reference.py +++ b/tests/unittest/mxfp8/test_mxfp8_attention_reference.py @@ -93,21 +93,21 @@ def test_reference_matches_bf16_sdpa(config, causal): # Layer 2 — kernel vs golden reference (requires CDNA4 native tl.dot_scaled) # --------------------------------------------------------------------------- -@cuda_only -@pytest.mark.parametrize("config", test_cases) -@pytest.mark.parametrize("causal", [True]) -def test_kernel_matches_reference(config, causal): - """Kernel vs golden reference — isolates Triton port bugs from quant error. +# Both ops assert seqlen_q == seqlen_k when causal, so non-square shapes are only +# reachable without causal masking. They stay out of the shared `test_cases` grid +# (which must remain valid for both causal values) and get their own coverage. +non_causal_cases = [ + AttnConfig(seqlen_q=256, seqlen_kv=512, num_head_q=8, num_head_kv=2, head_dim_qk=128, head_dim_v=128), + AttnConfig(seqlen_q=512, seqlen_kv=256, num_head_q=8, num_head_kv=2, head_dim_qk=192, head_dim_v=128), +] - Both sides apply the same mxfp8 quantization, so a large gap here points at a - Triton port bug (masking / online-softmax / LSE / strides) rather than - quantization error. - """ + +def _check_forward_kernel_vs_reference(config, causal, batch=4): from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 device = "cuda" dtype = torch.bfloat16 - q, k, v = _make_qkv_bhsd(4, config, device, dtype) + q, k, v = _make_qkv_bhsd(batch, config, device, dtype) sm_scale = config.head_dim_qk**(-0.5) o_kernel = triton_attention_mxfp8( @@ -139,6 +139,54 @@ def test_kernel_matches_reference(config, causal): assert snr > 30, f"kernel vs golden SNR too low: {snr}" +@cuda_only +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_kernel_matches_reference(config, causal): + """Kernel vs golden reference — isolates Triton port bugs from quant error. + + Both sides apply the same mxfp8 quantization, so a large gap here points at a + Triton port bug (masking / online-softmax / LSE / strides) rather than + quantization error. + """ + _check_forward_kernel_vs_reference(config, causal) + + +@cuda_only +@pytest.mark.parametrize("config", non_causal_cases) +def test_kernel_matches_reference_non_square(config): + """Forward kernel on seqlen_q != seqlen_k, the only shape family causal rejects.""" + _check_forward_kernel_vs_reference(config, causal=False, batch=2) + + +@cuda_only +def test_causal_forward_rejects_non_square_shape(): + """Non-square causal masking is intentionally blocked until its semantics are decided.""" + from alto.kernels.mxfp8.triton_flash_attention_mxfp8 import triton_attention_mxfp8 + + config = AttnConfig(seqlen_q=128, seqlen_kv=256, num_head_q=4, num_head_kv=4, head_dim_qk=128, head_dim_v=128) + q, k, v = _make_qkv_bhsd(1, config, "cuda", torch.bfloat16) + + with pytest.raises(AssertionError, match="causal forward requires seqlen_q == seqlen_k"): + triton_attention_mxfp8( + q.contiguous(), + k.contiguous(), + v.contiguous(), + bias=None, + alibi_slopes=None, + sm_scale=config.head_dim_qk**(-0.5), + dropout_p=0.0, + cu_seqlens_q=0, + cu_seqlens_k=0, + max_seqlens_q=config.seqlen_q, + max_seqlens_k=config.seqlen_kv, + causal=True, + return_scores=False, + use_exp2=True, + layout="bhsd", + ) + + # --------------------------------------------------------------------------- # Backward Layer 1 — golden reference vs autograd fp32 SDPA (pure PyTorch, CPU) # --------------------------------------------------------------------------- @@ -246,7 +294,7 @@ def test_backward_stage2_reference_matches_sdpa(config, causal): @cuda_only @pytest.mark.parametrize("config", public_autograd_cases) -@pytest.mark.parametrize("causal", [True]) +@pytest.mark.parametrize("causal", [True, False]) def test_public_autograd_matches_sdpa(config, causal): """Public ``triton_attention_mxfp8`` forward/backward must wire gradients correctly. @@ -320,17 +368,13 @@ def test_public_autograd_matches_sdpa(config, causal): # Stage-2 reference (same full quantization both sides -> gap = port bugs). # --------------------------------------------------------------------------- -@cuda_only -@pytest.mark.parametrize("config", test_cases) -@pytest.mark.parametrize("causal", [True]) -def test_backward_kernel_matches_reference(config, causal): - """A1 backward kernel vs Stage-2 golden reference — isolates port bugs. CDNA4-only.""" +def _check_backward_kernel_vs_reference(config, causal, batch=4): from alto.kernels.mxfp8.mxfp8_quantization import convert_to_mxfp8 device = "cuda" dtype = torch.bfloat16 - q, k, v = _make_qkv_bhsd(4, config, device, dtype) - do = _make_do_bhsd(4, config, device, dtype) + q, k, v = _make_qkv_bhsd(batch, config, device, dtype) + do = _make_do_bhsd(batch, config, device, dtype) sm_scale = config.head_dim_qk**(-0.5) o_ref, lse_ref = mxfp8_attention_forward_reference(q, k, v, sm_scale, causal) @@ -374,3 +418,18 @@ def test_backward_kernel_matches_reference(config, causal): snr = calc_snr(ref, got) assert sim > 0.99, f"{name} kernel-vs-golden cosine-sim too low: {sim}" assert snr > 25, f"{name} kernel-vs-golden SNR too low: {snr}" + + +@cuda_only +@pytest.mark.parametrize("config", test_cases) +@pytest.mark.parametrize("causal", [True, False]) +def test_backward_kernel_matches_reference(config, causal): + """A1 backward kernel vs Stage-2 golden reference — isolates port bugs. CDNA4-only.""" + _check_backward_kernel_vs_reference(config, causal) + + +@cuda_only +@pytest.mark.parametrize("config", non_causal_cases) +def test_backward_kernel_matches_reference_non_square(config): + """Backward kernel on seqlen_q != seqlen_k, the only shape family causal rejects.""" + _check_backward_kernel_vs_reference(config, causal=False, batch=2) From 8ded4157d98edcfcf063805851f992ed0ead5117 Mon Sep 17 00:00:00 2001 From: Yue Sun Date: Fri, 31 Jul 2026 02:41:48 -0500 Subject: [PATCH 14/14] docs: rewrite mxfp8 attention plan to describe the shipped state --- alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md | 637 +++++---------------- 1 file changed, 142 insertions(+), 495 deletions(-) diff --git a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md index 05bde1d3..a0a22195 100644 --- a/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md +++ b/alto/kernels/mxfp8/MXFP8_ATTENTION_PLAN.md @@ -1,105 +1,44 @@ -# MXFP8 E4M3 Flash-Attention(forward + backward)— Minimum Viable Plan +# MXFP8 E4M3 Flash-Attention(forward + backward) -目标:在 **AMD MI350 (CDNA4 / gfx950)** 上实现一个能支撑 mxfp8 训练的最小可用 mxfp8 flash-attention kernel,**包含前向与反向(dQ / dK / dV)**。 +在 **AMD MI350 (CDNA4 / gfx950)** 上支撑 mxfp8 训练的 flash-attention kernel,前向与反向(dQ / dK / dV)均已实现。 -## V1 三项决定(已与需求方确认) +**状态:已在 CDNA4 真机验收通过。** §7 的六项验收标准全部达成;测试 75 项全绿。本文档描述**已实现的现状**,末尾 §9 记录关键决策的理由。 -1. **forward + backward 都做**。⚠️ 关键前提:参照实现 `mxfp4 attention` 的 backward 是 17 个 `None` 占位、grad 断言全部注释——**backward 没有可机械改写的源**。V1 的 backward 需按标准 Flash-Attention v2 backward 数学**从零实现**,每个 dot 各自量化(见 §4.B / §9)。这是本 plan 相对 group gemm / mxfp4 attention 的**最大新增工作量与头号风险**。 - -2. **全 e4m3(已定)**。forward 的 Q/K/V/P 与 backward 的 dO/dP/dS **全部 e4m3**——整版 V1 单一 e4m3 格式,kernel 无 dtype 分发。e5m2 通道参数化预留但不启用(未来项)。 - -3. **仅 CDNA4,不加 fallback**。kernel 只走 `tl.dot_scaled`(e4m3×e4m3);**不写 CDNA3 dequant fallback**,不引入 `USE_DOT_SCALED` 开关、不引入 `use_dot_scaled` 参数。MI300 / CDNA3 明确不在 V1 范围。数值 ground-truth 由 **PyTorch 侧模拟 reference** 提供(见 §8),不再依赖 in-kernel fallback。 - -## 参照实现 +--- -- **kernel 移植源(仅 forward)**:`alto/kernels/fp4/mxfp4/triton_flash_attention_mxfp4.py`(FA v2 + mxfp4,forward-only)。 -- **backward 数学参照**:标准 Flash-Attention v2 backward(无 mxfp4 代码可抄,仅抄算法)。 -- **量化基础设施**:本目录 `mxfp8_quantization.py`(`convert_to_mxfp8` / `_calculate_scales` / `_quantize_fp8`)。**注意 V1 无 fallback,不使用 `_dequantize_fp8` 于 kernel 内**(`_dequantize_fp8` 只在 PyTorch 侧 reference 里用)。 -- **流程与验收范式**:本目录 `MXFP8_GROUPED_GEMM_PLAN.md`(「机械改写 + 分层校验 + 真机验证」方法论)。 +## 1. 范围与设计决定 -## 环境记录 +1. **forward + backward 都做。** forward 由 mxfp4 attention 机械改写,backward 由 `blockwise_fp8` attention backward 机械 port(三 kernel 结构同构)。 +2. **全 e4m3。** forward 的 Q/K/V/P 与 backward 的 dO/dP/dS 全部 e4m3,kernel 无 dtype 分发。 + ⚠️ **e5m2 没有参数化预留。** 格式以 `E4M3_TARGET_MAX_POW2 / E4M3_MBITS / E4M3_FORMAT_ID` 三个模块级 constexpr 硬编码,公开 API 也没有格式开关。若将来要切 e5m2,**必须改 kernel**,不是切参数。§3 的数值实测表明目前不需要切。 +3. **仅 CDNA4,无 fallback。** kernel 只走 `tl.dot_scaled`(e4m3×e4m3),不写 CDNA3 dequant fallback、无 `USE_DOT_SCALED` 开关。MI300 / CDNA3 不在范围内。数值 ground-truth 由 PyTorch 侧模拟 reference 提供(§6)。 -本 plan 撰写机为 gfx950(CDNA4/MI350),`is_cdna4()=True`。native `tl.dot_scaled` 路径可在本机验证;MI300/CDNA3 不在 V1 范围。 +## 2. 参照实现来源 -**2026-07-14 补充(MI250 离线开发机):** 当前开发/CI 环境为 gfx90a(MI250),`is_cdna4()=False`。Layer 2/3(kernel vs 黄金参照、kernel vs bf16 SDPA)在此硬件上跑不起来是预期行为,**不为 MI250 加 fallback 或 skip workaround**;待 CDNA4 硬件到位后一次性验证。Layer 1(纯 PyTorch 黄金参照 vs bf16 SDPA)可在任意设备离线跑。 +| 部分 | 来源 | +|---|---| +| forward kernel | `alto/kernels/fp4/mxfp4/triton_flash_attention_mxfp4.py`(FA v2 + mxfp4,forward-only) | +| backward 三 kernel | `alto/kernels/blockwise_fp8/triton_flash_attention_fp8_block.py`(`_bwd_preprocess` + `_bwd_kernel_dkdv` + `_bwd_kernel_dq`)。**mxfp4 没有 backward,不是这部分的来源** | +| 量化原语 | 本目录 `mxfp8_quantization.py`(`convert_to_mxfp8` / `_calculate_scales` / `_quantize_fp8`) | +| 流程与验收范式 | 本目录 `MXFP8_GROUPED_GEMM_PLAN.md`(「机械改写 + 分层校验 + 真机验证」) | ---- +**forward 相对 mxfp4 的改动**:删掉全部 head_dim packing(mxfp4 一 byte 装两个 e2m1,mxfp8 一 byte 一元素)——包括 `head_size_q/k/v *= 2`、`HALF_BLOCK_DMODEL_*`、`offs_d_*_pack`、`tl.dot_scaled` 的 `*_k_pack` 参数;`e2m1` → `e4m3`;`_pack_fp4` → `_quantize_fp8`;`convert_to_mxfp4` → `convert_to_mxfp8`。与量化正交的 FA v2 骨架(causal masking、GQA/MQA、LSE 写回、padded head、online-softmax、全 0 块 early-exit、`triton_op` + `autograd.Function` 三段式)原样保留。 -## 0. 格式选择:V1 全 e4m3,e5m2 留给后续 +## 3. 数值格式:全 e4m3 -**forward(Q/K/V/P)全 e4m3**——理由与 group gemm V1 一致:单一格式 → kernel 无 dtype 分发、最小版本最简单。attention forward 的四个 operand 都偏「动态范围小、单元素精度重要」,e4m3 天然合适: +attention 的四个 forward operand 都偏「动态范围小、单元素精度重要」,e4m3 天然合适: -| operand | 分布特征 | e4m3 是否够用 | +| operand | 分布特征 | e4m3 | |---|---|---| -| Q / K | LayerNorm 后激活,分布集中(±几十),动态范围需求低 | ✅ 单元素精度更重要 | +| Q / K | LayerNorm 后激活,分布集中(±几十) | ✅ 单元素精度更重要 | | V | 同上 | ✅ | -| P(softmax 概率) | ∈ [0,1],行和为 1;长尾但**有界**,且被 online-softmax 逐行重归一化 | ✅ ~6% 相对误差可接受 | - -**backward(dO / dP / dS)——V1 已定全 e4m3**(需求方确认,整版 V1 单一格式)。 - -风险记录(供未来升级参考,非 V1 待办):grad 是长尾 + 动态范围大的分布(`dP`/`dS` 尤甚),理论上是 e5m2 的场景(见 group gemm plan §0),全 e4m3 可能 underflow 小尾部。V1 接受此风险,理由是保持「单一格式、无 dtype 分发」的最小主题、且 backward 已是从零实现的大工程,不叠加混合格式复杂度。为让未来升级无痛,backward 的每个 `tl.dot_scaled` 同样把 dtype 参数化为 `LHS_FORMAT_ID` / `RHS_FORMAT_ID`(0=e4m3, 1=e5m2)constexpr,V1 默认全 0;若日后 backward 数值不过关,切 e5m2 即可,kernel 本体不动。 - -forward 的两个 `tl.dot_scaled` 也用同一套 `LHS_FORMAT_ID` / `RHS_FORMAT_ID` 参数化,V1 默认 e4m3,kernel 内 `if/else` 只走 e4m3 分支。 - ---- - -## 1. 复用 vs 改写 vs 新写 - -### 可直接复用(不动) - -`mxfp8_quantization.py` 提供 attention 需要的全部 device 端原语,无需新写量化代码: - -- `convert_to_mxfp8(x, axis=..., mxfp_format="e4m3", is_2d_block=...)`:wrapper 层量化 Q / K / V,以及 backward 里的 dO。 -- `_calculate_scales(x, ..., target_max_pow2, mbits, IS_2D_BLOCK=False)`:kernel 内算 P(及 backward 的 dP/dS)的 block scale。签名比 mxfp4 多 `target_max_pow2` / `mbits`,e4m3 传 `target_max_pow2=8, mbits=3`。 -- `_quantize_fp8(p, ps, ..., FP8_FORMAT=0, IS_2D_BLOCK=False)`:替代 mxfp4 的 `_pack_fp4`,kernel 内动态量化。 -- `is_cdna4()`:device 断言(V1 只支持 CDNA4)。 - -mxfp4 attention 里与量化正交、原样保留的 FA v2 forward 骨架:causal masking、GQA/MQA、varlen(thd)/bshd/bhsd、alibi/bias/dropout、LSE 写回、padded head、online-softmax 累积、全 0 块 early-exit、autotune wrapper、`triton_op` + `autograd.Function` + 用户入口三段式。 - -### 必须改写(forward,对应 group gemm plan §1) - -1. **删掉所有 head_dim packing**(mxfp4 沿 head_dim 一 byte 装两元素,mxfp8 一 byte 一元素): - - `get_shape_from_layout` 里的 `head_size_q/k/v *= 2` - - `HALF_BLOCK_DMODEL_QK/V`、`HALF_ACTUAL_BLOCK_DMODEL_QK/V`、`offs_d_qk_pack`、`offs_d_v_pack` - - `tl.dot_scaled(..., lhs_k_pack=, rhs_k_pack=)` 的 `*_k_pack` 参数 - - Q/K/V 指针里按 half-dim 的 stride,改回全 head_dim - - `SCALE_BLOCK_DMODEL_QK` 保留,但基于未 pack 的 head_dim 重算 -2. **两个 `tl.dot_scaled` 的 dtype**:mxfp4 写死 `"e2m1"` → mxfp8 参数化 `LHS_FORMAT_ID`/`RHS_FORMAT_ID`(V1 默认 e4m3),只走 e4m3×e4m3。 -3. **kernel 内 P 的动态量化**:`_calculate_scales`(带 `target_max_pow2=8, mbits=3`)→ `_pack_fp4` 换成 `_quantize_fp8(FP8_FORMAT=0)`。P 的 scale 沿 `BLOCK_N`(seqlen_k)方向,`IS_2D_BLOCK=False`。 -4. **wrapper 层量化**:`convert_to_mxfp4(axis=-1, is_2d_block=True)` → `convert_to_mxfp8(axis=-1, mxfp_format="e4m3", is_2d_block=...)`。 -5. **删除 fallback 相关**(本 plan 的净化项,非 mxfp4 有):不引入 `USE_DOT_SCALED`、`use_dot_scaled`、`_dequantize_fp8` 的 kernel 内调用。kernel 只有 `tl.dot_scaled` 一条路径。入口断言 `is_cdna4()`。 - -### 全新写(backward,无参照,头号工作量) - -按标准 FA v2 backward 数学从零实现三段(见 §9 完整推导): -- **preprocess kernel**:`delta = rowsum(dO ∘ O)` -- **dK/dV kernel**:遍历 K/V 块、重算 P、`dV += Pᵀ@dO`、`dP = dO@Vᵀ`、`dS = P∘(dP−delta)`、`dK += dSᵀ@Q` -- **dQ kernel**:`dQ += dS@K` +| P(softmax 概率) | ∈ [0,1],行和为 1;长尾但有界,被 online-softmax 逐行重归一化 | ✅ ~6% 相对误差可接受 | -每个 dot 用 `tl.dot_scaled`,operand 在对应 reduction 维上量化(forward 已存的 e4m3 版本复用,新产生的 dO/dP/dS 现场量化)。全部只走 CDNA4,无 fallback。 +**backward 全 e4m3 的顾虑已实测排除。** grad 是长尾 + 大动态范围分布(`dP`/`dS` 尤甚),理论上是 e5m2 的场景,曾担心小尾部 underflow。实测(stage-2 全量化参照 vs SDPA autograd):dQ/dK SNR ≈ 23 dB、cossim ≈ 0.9976;dV SNR ≈ 25~26 dB、cossim ≈ 0.9985。**未出现崩盘**,故不切 e5m2。 -### 全新写(其它) +## 4. 接口契约 -- `alto/kernels/mxfp8/triton_flash_attention_mxfp8.py`(forward 从 mxfp4 机械改写 + backward 从零写) -- `dispatch/attention.py` 加 `mxfp8_e4m3` 分支 -- `tests/unittest/mxfp8/test_mxfp8_attention.py` - -**预估**:forward kernel + wrapper ~450 行(删了 packing/fallback 比 mxfp4 少);backward 三个 kernel + wrapper ~500 行(全新);测试 ~250 行。 - ---- - -## 2. 接口契约 - -### scale 布局(沿用 mxfp4 注释) - -- **Q scale**:沿 head_dim(reduction 维)量化,`[..., seqlen_q, head_dim/32]` -- **K scale**:`"k scale is N×K even though k is K×N"`——scale 在非 reduction 维(seqlen_k)上 major:`[..., seqlen_k, head_dim/32]` -- **V scale**:PV reduction 维是 seqlen_k,V scale 沿 seqlen_k:`[..., head_dim_v, seqlen_k/32]` -- **P scale**:kernel 内动态生成,沿 `BLOCK_N`(seqlen_k)方向 - -**约束**:head_dim 与 seqlen_k 必须被 `QUANT_BLOCK_SIZE=32` 整除。head_dim ∈ {128,192}、seqlen ∈ {1024,2048} 均满足。head_dim<64 因 `tl.dot_scaled` 限制不支持(mxfp4 亦然)。**layout 只支持 `bhsd`**:2D-block scale 按数据张量 `shape[-2] × shape[-1]` 分块,只有 bhsd 时 `shape[-2]` 才是 seqlen;已加入口断言,见 §10.1octies。 - -### 用户 API(对齐 mxfp4,加一个格式开关,去掉 fallback 开关) +### 4.1 用户 API ```python def triton_attention_mxfp8( @@ -117,166 +56,50 @@ def triton_attention_mxfp8( causal: bool, return_scores: bool, use_exp2: bool, - layout: str, # 仅支持 "bhsd"(入口断言,见 §2 约束 / §10.1octies) - *, - fwd_format: str = "e4m3", # Q/K/V/P 格式,V1 固定 e4m3;e5m2 预留 - bwd_grad_format: str = "e4m3", # dO/dP/dS 格式,V1 固定 e4m3;e5m2 预留给未来 + layout: str, # 仅支持 "bhsd",见 §4.3 ) -> Tuple[Tensor, Tensor, Tensor]: # (o, softmax_lse, exp_scores) ``` -签名与 `triton_attention_mxfp4` 同形(方便 dispatch 直接替换),多 `fwd_format` / `bwd_grad_format` 两个 keyword-only 参数。**注意:不再有 `use_dot_scaled` 参数**(V1 只支持 CDNA4)。反向经 `autograd.Function` 自动触发,用户不直接调 backward kernel。 - ---- - -## 3. 各 dot 的 contraction & scale axis - -### Forward - -| dot | 计算 | reduction 维 | LHS quant axis | RHS quant axis | 一次 dot 跨几个 scale group | -|---|---|---|---|---|---| -| QK | `S = Q @ Kᵀ` | head_dim | Q: head_dim | K: head_dim | head_dim/32(=4 @ dim128)⚠️ | -| PV | `O += P @ V` | seqlen_k (BLOCK_N) | P: BLOCK_N | V: seqlen_k | BLOCK_N/32(=2 @ BLOCK_N=64) | - -### Backward(见 §9 推导) - -备注列为 **A1**(q/k/v 复用 forward 存的 e4m3 + 2D-block scale,dO/P/dS 现场量化): - -| dot | 计算 | reduction 维 | 备注 | -|---|---|---|---| -| dV | `dV += Pᵀ @ dO` | seqlen_q (BLOCK_M) | P/dO 现场 1D 量化沿 M | -| dP | `dP = dO @ Vᵀ` | head_dim_v | V 复用 fwd 2D scale(`_load_scale_hd`);dO 现场量化 | -| dK | `dK += dSᵀ @ Q` | seqlen_q (BLOCK_M) | dS 现场 1D 量化沿 M;**Q 复用 fwd 2D scale 转置 re-index**(`_load_scale_sq`) | -| dQ | `dQ += dS @ K` | seqlen_k (BLOCK_N) | dS 现场 1D 量化沿 N;**K 复用 fwd 2D scale 转置 re-index**(`_load_scale_sq`) | - -> ⚠️ **QK 一次 `dot_scaled` 跨 head_dim/32 个 32-wide scale group,是最大数值风险点**(group gemm plan §1.3 / §5 专门警告)。mxfp4 已接受此误差(未硬断言精度);mxfp8 更敏感。V1 无 fallback,改由 §8 的 **PyTorch 模拟 reference** 度量此误差;若发散,退路见 §5 风险 1。backward 的 dot 同样跨 group,同法度量。 - ---- - -## 4. 落地步骤 - -### 4.A Forward(机械改写 mxfp4) - -- **Step 1 — 骨架**:复制 `triton_flash_attention_mxfp4.py` → `triton_flash_attention_mxfp8.py`;改 import(`from .mxfp8_quantization import BLOCK_SIZE_DEFAULT, is_cdna4, _calculate_scales, _quantize_fp8`,删 `_pack_fp4`/`_unpack_fp4`);全局改名 `mxfp4→mxfp8`、`e2m1→e4m3`、op 名 `attention_mxfp8_forward_triton_impl`;kernel body 先占位跑通 import。 -- **Step 2 — 删 packing + QK dot**:按 §1 删所有 head_dim packing;QK 改 `tl.dot_scaled(q, qs, "e4m3", k, ks, "e4m3", out_dtype=fp32)`(无 `*_k_pack`)。验证:head_dim=128、单 head、无 causal,QK-only 输出 vs PyTorch 模拟 reference。 -- **Step 3 — PV dot + P 动态量化**:P 用 `_calculate_scales(..., target_max_pow2=8, mbits=3, IS_2D_BLOCK=False)` → `_quantize_fp8(..., FP8_FORMAT=0, IS_2D_BLOCK=False)`;PV 改 `tl.dot_scaled(p_fp8, ps, "e4m3", v, vs, "e4m3", out_dtype=fp32)`。验证:完整前向 vs PyTorch SDPA bf16(先不 causal)。 -- **Step 4 — forward wrapper + autograd.forward**:`attention_mxfp8_forward_triton_impl` 量化 Q/K/V(`axis=-1`),补输入契约检查(`head_dim % 32 == 0`、`seqlen_k % 32 == 0`、contiguous、`is_cdna4()`),launch。`_triton_attention_mxfp8.forward` 调 wrapper,`save_for_backward(q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias)`。`register_fake` 照搬 mxfp4。 - -### 4.B Backward(从零实现,见 §9) - -- **Step 5 — preprocess**:`delta = rowsum(dO ∘ O)`,`[batch, nheads, seqlen_q]`。dO 量化为 e4m3(`convert_to_mxfp8`)。 -- **Step 6 — dK/dV kernel**:遍历 K/V 块,重算 S→P(复用 fwd 的 q/k e4m3 + LSE),`dV += Pᵀ@dO`、`dP = dO@Vᵀ`、`dS = P∘(dP−delta)`、`dK += dSᵀ@Q`。每个 dot 用 `tl.dot_scaled`,operand 沿对应 reduction 维量化。验证:dK/dV vs bf16 SDPA autograd。 -- **Step 7 — dQ kernel**:`dQ += dS@K`。验证:dQ vs bf16 SDPA autograd。 -- **Step 8 — autograd.backward 接线**:`_triton_attention_mxfp8.backward` 从 ctx 取回张量、调 preprocess/dK-dV/dQ 三个 kernel,返回 `dq, dk, dv`(其余入参位 `None`)。 - -### 4.C Dispatch 接入 - -- **Step 9 — dispatch**:`dispatch/attention.py` 的 `LPScaledDotProductAttentionWrapper.__init__` 加分支: - ```python - elif config.precision in ("mxfp8_e4m3", "mxfp8"): - self.attn_func = triton_attention_mxfp8 - ``` - forward 已 layout-agnostic(传 `layout="bhsd"`),无需改。 +与 `triton_attention_mxfp4` **完全同形**,dispatch 可直接替换。无格式开关、无 `use_dot_scaled`。反向经 `autograd.Function` 自动触发,用户不直接调 backward kernel。 -### 4.D 测试 +`dispatch/attention.py` 的接入分支: -- **Step 10**:`test_mxfp8_attention.py`,见 §8。 - ---- - -## 5. 关键风险与对策 +```python +elif isinstance(config, TrainingOpConfig) and config.precision == "mxfp8_e4m3": + self.attn_func = triton_attention_mxfp8 +``` -| 风险 | 对策 | -|---|---| -| QK 一次 `dot_scaled` 跨 head_dim/32 个 scale group 发散(§3 ⚠️,头号数值风险) | 由 §8 PyTorch 模拟 reference 度量;若 native vs reference SNR 差距过大 → 退路是**沿 head_dim 分块累加 QK**(每 32-wide group 一次 `dot_scaled` + acc),代价是 QK 循环变长。V1 先测,超阈值再改 | -| **backward 无参照、从零写**(头号工作量风险) | 严格对标标准 FA v2 backward 数学(§9);每个 kernel 单独 vs bf16 SDPA autograd 校验(dV→dK→dQ 逐个隔离);先跑通非量化版(内部用高精度)再逐 dot 换 `tl.dot_scaled` | -| backward 全 e4m3 时 grad underflow / 数值不过关(§0 风险记录) | dtype 已参数化(`LHS/RHS_FORMAT_ID`);不过关时把 dO/dP/dS 切 e5m2(`bwd_grad_format="e5m2"`),kernel 本体不动 | -| P 在 kernel 内动态量化,`_quantize_fp8` 的 scale 广播语义与 mxfp4 `_pack_fp4` 不同 | Step 3 单独验证 P 量化路径;对照 `_quantize_fp8` 的 1D-block 分支(`IS_2D_BLOCK=False`) | -| head_dim / seqlen 非 32 倍数 | V1 断言拒绝(对齐 mxfp4 只测 ≥128 head_dim);padded-head 逻辑保留但要求 actual head_dim 仍 32 对齐 | +### 4.2 scale 布局 ---- +- **Q / K scale**:沿 head_dim(QK 的 reduction 维)量化。K 的 scale 在非 reduction 维(seqlen_k)上 major。 +- **V scale**:PV 的 reduction 维是 seqlen_k,V scale 沿 seqlen_k。 +- **P scale**:kernel 内动态生成,沿 `BLOCK_N`。 +- forward 把 Q/K/V 存成紧凑 2D-block scale `[.., seqlen/32, head_dim/32]`,backward 复用(§5.3)。 -## 6. 不做的事(V1 明确划线) +### 4.3 三类硬约束(入口断言) -- ❌ **CDNA3 / MI300 支持与 fallback 路径**——V1 仅 CDNA4,kernel 只有 `tl.dot_scaled` 一条路径 -- ❌ **混合格式**(forward P 或 backward grad 用 e5m2)——全 e4m3,e5m2 分支参数化预留不启用 -- ❌ head_dim < 64(`tl.dot_scaled` 限制) -- ❌ 2D-block P 量化(P 沿 `BLOCK_N` 一维即可) -- ❌ autotune 扩展(沿用 mxfp4 单 config:`BLOCK_M=BLOCK_N=64, PRE_LOAD_V=False`) -- ❌ TMA / async copy / pipelining 调优 -- ❌ 沿 head_dim 分块 QK(除非 §5 风险 1 触发) -- ❌ FSDP/TP 集成测试 - ---- - -## 7. 验收标准(V1 完成定义) - -| # | 标准 | 说明 | +| 断言 | 位置 | 理由 | |---|---|---| -| 1 | forward kernel 在 CDNA4 跑通 | 本机 gfx950 native `tl.dot_scaled` | -| 2 | backward(dQ/dK/dV)三个 kernel 在 CDNA4 跑通 | 本机 gfx950 | -| 3 | forward vs PyTorch SDPA bf16:cos-sim > 0.99、SNR > 阈值(硬断言) | 对齐 mxfp4 对比方式,但把 SNR/cossim 变成硬断言(mxfp4 只 print,此为主动加严) | -| 4 | backward dQ/dK/dV vs bf16 SDPA autograd:cos-sim / SNR 硬断言 | 阈值首跑标定,留裕度 | -| 5 | 覆盖 causal × GQA × head_dim{128,192} × seqlen{1024,2048} 网格 | 沿用 mxfp4 test_cases | -| 6 | dispatch 层 `mxfp8_e4m3` 分支可路由 | smoke:构造 config 走 `LPScaledDotProductAttentionWrapper` | - -> **相对 mxfp4 attention 的主动加严**:mxfp4 的 `test_mxfp_attention.py` 只 print、无 assert 且无 backward 测试;mxfp8 V1 把前向精度做成硬断言,并新增完整 backward 的梯度校验。 - ---- - -## 8. 测试组织(分层校验,归因清晰) - -文件:`tests/unittest/mxfp8/test_mxfp8_attention.py`,复用 `tests/unittest/mxfp8/utils.py` 的 `calc_snr` / `calc_cossim` / `prepare_data`(与 fp4 共用 `alto.kernels.fp4.testing_utils`)。 - -**无 fallback,ground-truth 改由 PyTorch 侧模拟 reference 提供。** 两层校验: - -1. **kernel 移植正确性**:native kernel vs **PyTorch 模拟 reference**——reference 用 `convert_from_mxfp8` dequant Q/K/V → 算 S=Q@Kᵀ → softmax → 用 PyTorch 版 mxfp8 量化 P → dequant → @V(即「在 PyTorch 里复刻 kernel 的量化流程」)。 - - 隔离:masking / online-softmax / LSE 的移植 bug,排除 mxfp8 量化误差本身。 - - 附带度量 §3 ⚠️ 的 QK 跨-group 误差(native kernel 的 `dot_scaled` 与 reference 的精确 dequant-matmul 之差即此误差)。 - - 门槛:SNR 高(> 40 dB 量级)。 -2. **端到端量化误差**:native kernel vs **PyTorch SDPA bf16**(无量化)。 - - 隔离:mxfp8 量化本身的总误差。 - - 门槛:cos-sim > 0.99 + SNR 硬断言(验收标准 3)。 -3. **backward 校验**:完整 `autograd` 反向 vs bf16 SDPA autograd,比 dQ/dK/dV 的 cos-sim / SNR(验收标准 4)。 - -参数网格:`batch=4` × mxfp4 的 `test_cases` × `causal ∈ {True, False}` 覆盖 forward、 -public autograd、kernel-vs-golden forward/backward;另设 `non_causal_cases` 覆盖 -`seqlen_q != seqlen_k`(只在 `causal=False` 下合法)。 - ---- - -## 9. Backward 设计(V1 实现,无参照,从零写) - -> backward 是 V1 交付物且无 mxfp4 代码可抄,本节锁定算法与量化决策。 - -### 9.1 标准 FA v2 backward 数学 - -给定前向已存的 `Q, K, V, O, softmax_lse`(LSE = log-sum-exp per row)与上游 `dO`: +| `layout == "bhsd"` | forward op + backward op | `convert_to_mxfp8(is_2d_block=True)` 按数据张量 `shape[-2] × shape[-1]` 分 32×32 块,只有 bhsd 时 `shape[-2]` 才是 seqlen。`bshd`/`thd` 下它按 **nheads** 分组(nheads 恰为 32 倍数时 `torch._check` 也拦不住),而 kernel 的 scale 指针数学假定按 seqlen 分组 → **静默读错 scale** | +| `not causal or seqlen_q == seqlen_k` | forward op + backward op | kernel 掩码用 bottom-right,PyTorch 参照与 `F.sdpa` 用 top-left,两者只在方阵下等价。V1 生产只跑 self-attention,故不统一约定(那是解决不存在的问题),改为断言拒绝。两侧都拦——一个 op 允许、另一个禁止同一种形状,调用者只能靠踩坑才知道边界 | +| backward 不支持 alibi / dropout | `autograd.Function.backward` | 上游 `blockwise_fp8` forward 支持 dropout 但 backward 无对应逻辑且不检查(开 dropout 训练即静默错梯度),alibi 同理。此处改为当场报错 | -- **preprocess**:`delta = rowsum(dO ∘ O)`,shape `[batch, nheads, seqlen_q]` -- **dK/dV kernel**(遍历 K/V 块,块内 recompute P): - - `S = Q @ Kᵀ * sm_scale`(+ mask/alibi/bias),`P = exp(S − softmax_lse)` - - `dV += Pᵀ @ dO` - - `dP = dO @ Vᵀ` - - `dS = P ∘ (dP − delta) * sm_scale` - - `dK += dSᵀ @ Q` -- **dQ kernel**:`dQ += dS @ K`(dS 同上重算) +维度约束:head_dim 与 seqlen_k 必须被 `QUANT_BLOCK_SIZE=32` 整除(由 `convert_to_mxfp8` 内的 `torch._check` 保证,不在下游重复检查);head_dim < 64 因 `tl.dot_scaled` 限制不支持(mxfp4 亦然)。 -### 9.2 V1 量化决策(全 e4m3) +断言的有效性单独验过:`causal=True, sq=64/sk=128` 走 backward 按预期抛 `AssertionError`;`causal=False` 同形状正常通过(非 causal 无对齐问题,不该拦);`layout="bshd"` 走 forward 按预期抛 `AssertionError`。forward 侧配 `test_causal_forward_rejects_non_square_shape` 做回归。 -- `Q / K / V` 复用 forward 已量化的 e4m3 版本(ctx 已存 `q, k, v, q_scale, k_scale, v_scale`)。 -- `dO` 在 backward 入口量化为 e4m3(`convert_to_mxfp8`),沿 dK/dV 需要的 reduction 维(seqlen_q)与 dP 需要的维(head_dim_v)——按 group gemm plan 的经验,可能需要**两套 dO 量化**(沿不同 axis),沿用其 autograd 已有逻辑。 -- `P / dS` 在 kernel 内现场量化(`_calculate_scales` + `_quantize_fp8`),沿各自 dot 的 reduction 维。 -- 每个 dot 的 dtype 走 `LHS/RHS_FORMAT_ID`,V1 全 0(e4m3)。 +## 5. 量化方案 -### 9.3 e5m2 升级路径(若 §0 的 e4m3 backward 数值不过关) +### 5.1 Forward 的两个 dot -把 `bwd_grad_format` 切 `"e5m2"` → dO/dP/dS 的量化与对应 `tl.dot_scaled` 的 `FORMAT_ID` 改 1,kernel 本体不动。这是 §5「backward 数值不过关」风险的兜底。 - -### 9.5 阶段2 具体方案(2026-07-15 细化,7-dot 量化轴对齐) +| dot | 计算 | reduction 维 | LHS quant axis | RHS quant axis | 跨几个 scale group | +|---|---|---|---|---|---| +| QK | `S = Q @ Kᵀ` | head_dim | Q: head_dim | K: head_dim | head_dim/32(=4 @ dim128) | +| PV | `O += P @ V` | seqlen_k (BLOCK_N) | P: BLOCK_N | V: seqlen_k | BLOCK_N/32(=2 @ BLOCK_N=64) | -阶段1(bf16 骨架)已于 MI250 验通移植正确性(§10.1ter)。阶段2 = 把 backward 两个 inner kernel 里的 7 个 `tl.dot` 逐个换成 `tl.dot_scaled`,每个 dot 的两个 operand **沿它的 reduction 轴** MX e4m3 量化。 +> 立项时把「QK 一次 `dot_scaled` 跨 head_dim/32 个 32-wide scale group」列为头号数值风险,退路是沿 head_dim 分块累加。**真机实测未触发**:kernel vs 参照 53~63 dB,误差全部来自量化本身而非跨 group 累加,退路未启用。 -**7 个 dot 与量化轴:** +### 5.2 Backward 的七个 dot | dot | kernel | 计算 | reduction 轴 | LHS 量化轴 | RHS 量化轴 | |---|---|---|---|---|---| @@ -288,283 +111,128 @@ public autograd、kernel-vs-golden forward/backward;另设 `non_causal_cases` | f | dq | `dp = do@vᵀ` | head_dim_v | 同 b | 同 b | | g | dq | `dq += ds@k` | seqlen_k (BLOCK_N) | ds 沿 seqlen_k | k 沿 seqlen_k | -**同一 operand 需多套量化(plan 早点名的核心工作量):** - -- **q**:沿 head_dim(a/e)+ 沿 seqlen_q(d)→ **2 套** -- **k**:沿 head_dim(a/e)+ 沿 seqlen_k(g)→ **2 套** -- **do**:沿 head_dim_v(b/f)+ 沿 seqlen_q(c)→ **2 套** -- **v**:仅 head_dim_v(b/f)→ 1 套 -- **p / ds**:kernel 内现算现量化,沿各自 reduction 轴(p 沿 seqlen_q;ds 沿 seqlen_q 供 d、沿 seqlen_k 供 g → ds 2 套) - -**量化 block 规则(2026-07-16 定案:A1——q/k/v 复用 forward 2D-block,dO/P/dS 现场 1D per-row):** +同一个 operand 被不同 dot 沿不同轴归约,因此需要多套 scale:q 沿 head_dim(a/e)+ 沿 seqlen_q(d);k 沿 head_dim(a/e)+ 沿 seqlen_k(g);do 沿 head_dim_v(b/f)+ 沿 seqlen_q(c);v 只需 head_dim_v。这是整个 backward 的核心工作量。 -- **q/k/v(forward 存过的)**:**直接复用 forward 的紧凑 2D-block scale `[.., seqlen/32, head_dim/32]`**,backward 不再重量化。每个 dot 用指针索引把这块 2D scale 广播成 `tl.dot_scaled` 要的 `[outer, reduction/32]`: - - **head_dim 收缩的 dot(a/b/e/f)**:`scale[outer//32, dgroup]`,与 forward QK 的 `qs_ptrs` 广播完全同款(`_load_scale_hd`)。 - - **seqlen 收缩的 dot(d/g)**:**转置对称复用**——同一个 32×32 块的 scale 换个轴索引成 `[head_dim, seqlen_block/32]`(`_load_scale_sq`),因为一个 32×32 块只有一个 scale,两轴共用。 -- **dO / P / dS(backward 新产生)**:forward 没存过,仍**现场 1D per-row 沿其 reduction 轴**量化(`_mx_quant`,`IS_2D_BLOCK=False`),scale `[outer, reduction/32]` 直接匹配 `dot_scaled`。 -- 早先(2026-07-15)为省事对所有 operand 统一 1D per-row(含 q/k/v 重量化),即 option A;2026-07-16 翻案改 A1(见 §10.1sexies / AB 决策稿 banner):直接复用 forward 那份 e4m3+2D scale,**零重量化、与 forward 逐位一致、无双重量化**,代价是 kernel 内两个 scale 广播辅助。 +标准 FA v2 backward 数学(`delta = rowsum(dO ∘ O)`;`dV += Pᵀ@dO`;`dP = dO@Vᵀ`;`dS = P∘(dP−delta)·sm_scale`;`dK += dSᵀ@Q`;`dQ += dS@K`)与上游一致,不赘述。 -**operand 来源(A1,2026-07-16 定,推翻 option A):** forward 存 e4m3 q/k/v + 2D-block scale(省显存,不变);backward driver **直接把 saved e4m3 + scale 传进 kernel**(无入口 dequant),每个 dot 复用之(见上)。dO 是 backward 的 bf16 输入,现场量化。**A(入口 dequant + 逐 dot 重量化)已删除,B(存 bf16)不做。** +### 5.3 scale 复用方案(A1) -**落地顺序(2026-07-16 改为「先 kernel 后参照」——需求方拍板,知情决策):** +**q/k/v 直接复用 forward 存的 e4m3 值 + 2D-block scale,backward 不重量化。** 每个 dot 用指针索引把紧凑的 2D scale 广播成 `tl.dot_scaled` 要的 `[outer, reduction/32]`: -1. ✅ **删 A**:删除 A 的 kernel 路径(入口 `convert_from_mxfp8` + 逐 dot 重量化 q/k/v)与 A 的黄金参照 `mxfp8_attention_backward_reference_stage2`(含 `operand_source` A/B 开关)+ 其两个测试。保留 stage-1 参照 `mxfp8_attention_backward_reference`(FA 数学基准,格式无关)。 -2. ✅ **写 A1 kernel**:新增 `_load_scale_hd` / `_load_scale_sq`(从紧凑 2D scale 重建各 dot 的 scale tile);两个 kernel + 两个 inner 改为复用 saved e4m3 q/k/v + 2D scale;driver 传 q/k/v 的 scale 指针 + 12 个 scale stride;删去不再用的 `convert_from_mxfp8` import。dO/P/dS 仍 `_mx_quant`。**MI250 import OK、语法/签名通;`tl.dot_scaled` 不编译,整体待 CDNA4。** -3. ⏳ **写 A1 golden reference**(下一步):纯 PyTorch 复刻「saved e4m3 值 + 2D tile scale 按轴索引」,MI250/CPU 可跑,对 bf16 SDPA autograd 度量数值。 -4. ⏳ **CDNA4 验收**:kernel vs A1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差)。 +- `_load_scale_hd`(head_dim 收缩,dot a/b/e/f):`scale[outer//32, dgroup]`,与 forward QK 的 `qs_ptrs` 广播同款。 +- `_load_scale_sq`(seqlen 收缩,dot d/g):**转置对称复用**——同一个 32×32 块只有一个 scale,两轴共用,换个轴索引成 `[head_dim, seqlen_block/32]`。 -> ⚠️ A1 kernel 的 `tl.dot_scaled` + 2D-scale 指针广播在 MI250 **一行都验不了**,且按新顺序**当前连 A1 参照都还没写**——A1 kernel 正确性完全待 CDNA4 + A1 参照。其中 seqlen 轴的 `_load_scale_sq` 广播全仓无先例,是最大盲区。这是需求方明确知情后的决策(宁可不留 A 冗余)。 +**dO / P / dS 是 backward 新产生的,forward 没存过**,仍现场 1D per-row 沿各自 reduction 轴量化(`_mx_quant`,`IS_2D_BLOCK=False`),scale 直接匹配 `dot_scaled`。 -### 9.4 save_for_backward 清单 +这套方案 correct-by-construction:与 forward 逐位一致、零重量化、无双重量化。演进过程见 §9.1。 -forward 存:`q, k, v, o, softmax_lse, q_scale, k_scale, v_scale, alibi, bias`(确保 backward 三个 kernel 所需张量齐全,避免中途改 forward 签名)。 +## 6. 测试组织 ---- +三层校验,归因清晰。文件:`tests/unittest/mxfp8/test_mxfp8_attention.py`(基线)、`test_mxfp8_attention_reference.py`(分层)、`utils.py`(黄金参照)。 -## 10. 实施进展记录 +| 层 | 对比 | 隔离出什么 | 硬件 | +|---|---|---|---| +| 1 | 纯 PyTorch 黄金参照 vs bf16 SDPA | 算法与量化放置是否正确(不含 Triton) | CPU / 任意设备 | +| 2 | kernel vs 黄金参照 | Triton 移植 bug(masking / online-softmax / LSE / strides / scale 广播),排除量化误差本身 | CDNA4 | +| 3 | kernel vs bf16 SDPA | mxfp8 量化的端到端总误差 | CDNA4 | -### 10.1 截至 2026-07-14 +黄金参照(`utils.py`)三个:`mxfp8_attention_forward_reference`、`mxfp8_attention_backward_reference`(stage-1,FA 数学基准)、`mxfp8_attention_backward_reference_stage2`(A1 全量化,复刻 §5.3)。全部纯 fp32 matmul、不含 `tl.dot_scaled`。 -**已落地(代码,未经硬件验证):** +**参数网格**:`test_cases` 八组(causal × GQA × head_dim{128,192} × seqlen{1024,2048},含 GQA 与非对称 head_dim 的交叉点)× `causal ∈ {True, False}`;`non_causal_cases` 两组覆盖 `seqlen_q != seqlen_k`(causal 断言使其只在非 causal 下合法);`reference_cases` 三组小形状供 CPU 层。共 **75 项**。 -- **forward kernel**:`alto/kernels/mxfp8/triton_flash_attention_mxfp8.py`——由 mxfp4 attention 机械改写(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`、Q/K/V 走 `convert_to_mxfp8`)。**forward-only**(backward 为 `None` 占位)、**仅 CDNA4 走 `tl.dot_scaled`,无 fallback**。`__init__.py` 导出 `triton_attention_mxfp8`。 -- **黄金参照**:`tests/unittest/mxfp8/utils.py` 的 `mxfp8_attention_forward_reference`——纯 PyTorch 复刻 forward(2D-block e4m3 Q/K/V、逐 key-block 在线 softmax + running-max 量化 P、全 fp32 matmul、无 `tl.dot_scaled`),可在 CPU / 任意设备跑。 -- **测试(三层,见 §8)**: - - 第 3 层 `test_mxfp8_attention.py::test_attention`(kernel vs bf16 SDPA,仿 mxfp4,硬断言 `cossim>0.99`/`SNR>20`)——**已 commit,作为基线不再改动**。 - - 第 1 层 `test_mxfp8_attention_reference.py::test_reference_matches_bf16_sdpa`(黄金参照 vs bf16,**CPU 现在可跑**)+ 第 2 层 `test_kernel_matches_reference`(kernel vs 黄金参照,CDNA4;`import` 复用第 3 层的 `test_cases` 网格)——新增独立文件,纯增量。 +> 测试分层曾经是反的:三个纯 PyTorch 参照测试跑 `causal=[True, False]`,而每一个真正启动 kernel 的测试都钉在 `[True]`——不可能有移植 bug 的那层测两个值,会有移植 bug 的那层只测一个。骨架抄自 mxfp4 时把严格性一起抄丢了。已全部改为两个值。 -**对 §4 步骤的进度映射:** +## 7. 验收标准与实测结果 -| 步骤 | 状态 | 备注 | +| # | 标准 | 结果 | |---|---|---| -| Step 1–4(forward kernel + wrapper + autograd.forward) | ✅ 代码完成 / ⏳ 未验证 | 无 CDNA4,尚未跑通任何 kernel 测试 | -| Step 5–8(backward:preprocess / dK-dV / dQ / autograd.backward) | ❌ 未开始 | 从零写,头号工作量,见 §9 | -| Step 9(dispatch 接入 `mxfp8_e4m3`) | ❌ 未开始 | — | -| Step 10(测试脚手架) | ✅ 三层脚手架就位 | 第 1 层已于 MI250 Docker 验证 6/6 通过;第 2/3 层待 CDNA4 | - -**已知偏差 / 修正记录:** - -- 写黄金参照时静态发现并修掉一处 masking bug:mask 值若用 `finfo.min`(有限)会让"整块被 mask 的行"算出 `exp(0)=1`(应为 0);改用真 `-inf`(masked 项 `exp(-inf)=0`,全 mask 行的 nan 由 `nan_to_num` 兜住)。此 bug 正是"无黄金参照、只对 bf16"抓不到的类型。 -- **2026-07-14 — 黄金参照 causal mask 与 SDPA 不一致(非 CDNA4,Layer 1)**:`mxfp8_attention_forward_reference` 原用 bottom-right 对齐(`key_j <= query_i + (seqlen_k - seqlen_q)`),但 PyTorch `F.sdpa(is_causal=True)` 在 **bhsd** 布局下实际为 **top-left**(`key_j <= query_i`)。`seqlen_q == seqlen_k` 时两者等价;`seqlen_kv > seqlen_q` 且 `causal=True` 时差异显著(`test_reference_matches_bf16_sdpa[True-config2]` cosine-sim 仅 ~0.39)。**修复**:`tests/unittest/mxfp8/utils.py` 改为 top-left mask。**验证**:MI250(gfx90a)Docker 上 Layer 1 六项全过。 -- **2026-07-14 — Triton kernel 编译期两处真 bug(非 CDNA4 专属,CDNA4 上同样会触发)**: - 1. **fp8 `tl.load` 的 `other` 类型**:`q = tl.load(q_ptrs, mask=..., other=0)` 中 `other=0`(int32)无法 cast 为 `fp8e4nv`,报 `cannot cast int32[...] to fp8e4nv`。**修复**:`triton_flash_attention_mxfp8.py` 中 Q load 及 `load_fn` 默认 `other` 改为 `0.0`。 - 2. **Triton constexpr 全局变量写法**:`E4M3_TARGET_MAX_POW2: tl.constexpr = FORMAT_TO_TARGET_MAX["e4m3"]` 等注解写法在 `@jit` 内不可见,报 `NameError: Cannot access global variable E4M3_TARGET_MAX_POW2`。**修复**:改为实例化写法 `E4M3_TARGET_MAX_POW2 = tl.constexpr(8)`、`E4M3_MBITS = tl.constexpr(3)`、`E4M3_FORMAT_ID = tl.constexpr(0)`。 - - 上述修复待在 CDNA4 真机上跑 Layer 2 时一并验收;MI250 上 Layer 2 仍因 `tl.dot_scaled` 不可用而无法运行,**属预期,未做 workaround**。 - -**Open items(阻塞真验证):** - -- **CDNA4 硬件**:M350/M355 卡预计约一周后到;在此之前所有 kernel 测试(第 2/3 层)无法运行,forward 代码处于"已写未验"状态(Triton 编译修复已合入,待真机确认 Layer 2)。 -- ~~**第 1 层 CPU 自证**~~:✅ 2026-07-14 于 MI250 Docker 跑通 `test_reference_matches_bf16_sdpa`(6/6 passed)。 -- **第 2 层 kernel vs 黄金参照**:待 CDNA4 硬件;届时验收 `test_kernel_matches_reference` 全网格。 -- **backward 全 e4m3 的数值风险**:见 §0 风险记录 / §5,真机验证后据实决定是否需切 e5m2。 - -### 10.1bis 截至 2026-07-15(backward 阶段1 + 参照) - -**关键决策修正:backward 不从零写。** 发现 `alto/kernels/blockwise_fp8/triton_flash_attention_fp8_block.py` 有一套完整 fp8 attention backward(`_bwd_preprocess` + `_bwd_kernel_dkdv` + `_bwd_kernel_dq`),结构与 §9 三 kernel 完全同构。因此 backward 改为**机械 port**(同 forward 从 mxfp4 改写的性质),§9"从零写、头号风险"前提作废。真正工作量集中在 **5 个 dot 的 MX scale 轴对齐**(②dp/③dv 的 dO、①qk/④dk 的 q 各需沿不同 reduction 轴的两套量化)。 - -**分两阶段落地,采用 port_now(骨架先行):** - -- **阶段1 ✅ 已落地(代码,未验)**:port 三 kernel(`_bwd_preprocess` / `_bwd_kernel_dq` + `_attn_bwd_dq_inner` / `_bwd_kernel_dkdv` + `_attn_bwd_dkdv_inner`)+ backward driver(注册 `alto::attention_mxfp8_backward_triton_impl` + `register_fake`)+ 接上 `autograd.backward`(返回 15 个梯度,顺手修正旧 stub 的 17 个 None bug)。 - - **dot 用高精度 bf16 `tl.dot` 占位**,每处留 `TODO(stage2)` 标注该 dot 的 reduction 轴。 - - **入口用 `convert_from_mxfp8` 把 saved e4m3 q/k/v 反量化回 bf16**——backward 吃的正是 forward 用过的那份量化输入,且**不含 `tl.dot_scaled`,MI250 可跑**。 - - causal 用 **bottom-right**(与 forward kernel 一致,非参照的 top-left;方阵下无差别)。 -- **阶段2 ✅ 已落地(代码,`tl.dot_scaled` 未验)**:逐 dot 换 `tl.dot_scaled` + 沿 reduction 轴的 MX e4m3 量化。详见 §10.1quinquies。需 CDNA4 才能验。 - -**backward 黄金参照 ✅ 已落地**:`tests/unittest/mxfp8/utils.py::mxfp8_attention_backward_reference`——纯 PyTorch 复刻阶段1 kernel(反量化 q/k/v、用 saved `lse` 在 fp32 重算 `P`、backward 不量化 P),标准 FA v2 backward(dV=Pᵀ@dO;dP=dO@Vᵀ;dS=P·(dP−delta);dQ=sm·dS@K;dK=sm·dSᵀ@Q)。喂 forward 参照的同一份 `o`/`lse` 时应与 kernel 高精度吻合。 - -**backward 测试(加在 `test_mxfp8_attention_reference.py`):** - -- **第 1 层 `test_backward_reference_matches_sdpa`(CPU)**:参照 vs autograd fp32 SDPA backward,验算法+量化误差(`cossim>0.97`/`SNR>8`)。**本机无 torch,未跑,待另一台机跑。** -- **第 2 层 `test_backward_kernel_matches_reference`(设备)**:直接调 backward op(喂 forward 参照的 `o`/`lse`,**绕过 forward 的 `tl.dot_scaled`**),故**阶段1 bf16 骨架 MI250 即可验**;kernel vs 参照(`cossim>0.99`/`SNR>25`,阈值待首跑校准)。 - -**Open items 更新:** -- **阶段1 backward 未经任何硬件验证**(本机无 torch);下一步在另一台机跑第 1 层(CPU)+ 第 2 层(MI250 即可,无需 CDNA4)拿反馈。 -- 阶段2(dot_scaled + MX 量化)仍待 CDNA4。 - -### 10.1ter 截至 2026-07-15(backward 阶段1 于 MI250 验通) - -**已在 MI250(gfx90a)Docker 跑通 backward 两层测试:** - -- **第 1 层 `test_backward_reference_matches_sdpa`:6/6 passed。** -- **第 2 层 `test_backward_kernel_matches_reference`:7/7 passed**(全 `test_cases` 网格),kernel vs 黄金参照 SNR≈55.5 dB、cos≈0.999999。证明阶段1 骨架的 FA v2 backward 数学 + strides + GQA + masking 移植正确。**注意:这验的是移植正确性,不是 mxfp8 backward 数值过关**——后者要阶段2 的 `tl.dot_scaled` + MX 量化,仍待 CDNA4。 - -**修掉一个真 bug(与 §10.1 forward 那次同类):** - -- **Triton constexpr 全局变量注解式写法**:backward 新增的模块级 `RCP_LN2: tl.constexpr = 1.4426950408889634` 被两个 backward inner kernel(`_attn_bwd_dq_inner` / `_attn_bwd_dkdv_inner`)访问时报 `NameError: Cannot access global variable RCP_LN2`。**修复**:改为实例化写法 `RCP_LN2 = tl.constexpr(1.4426950408889634)`。(forward 函数体内的局部 `RCP_LN2` 是局部变量,不受影响,未动。) - -**潜伏坑(非 bug,记录备查):** backward kernel 的 causal 用 **bottom-right**(`col_offset = N_CTX_Q - N_CTX_K`,与 forward kernel 一致),黄金参照用 **top-left**。当前 `test_cases` 全是方阵(seqlen_q==seqlen_kv),两者等价故测不出差异;若日后加**非方阵 causal** 用例,kernel 与参照会对不上,需统一对齐方式。 - -### 10.1quater 截至 2026-07-15(阶段2 数值先探路,全 e4m3 判定可行) - -**已落地(代码)+ 已于 MI250 验:** - -- **阶段2 PyTorch 参照** `mxfp8_attention_backward_reference_stage2`(`tests/unittest/mxfp8/utils.py`):按 §9.5 的 7-dot 量化轴表,把每个 backward matmul 的 operand 沿其 reduction 轴做 qdq(**统一 1D per-row**,见下文修订),fp32 matmul 模拟 `tl.dot_scaled`。不含 `tl.dot_scaled`,MI250/CPU 可跑。 -- **测试** `test_backward_stage2_reference_matches_sdpa`(stage-2 参照 vs autograd fp32 SDPA):MI250 **6/6 passed**。 - -**数值结论(关键决策依据):** 全 e4m3 backward 对 SDPA 的误差——**dQ/dK SNR≈23 dB、cossim≈0.9975;dV SNR≈25 dB、cossim≈0.9984**。相比阶段1(只量化 Q/K/V,SNR≈55dB),额外量化 dO/P/dS 把 SNR 拉到 ~23dB,但 cossim 稳在 0.998——**未出现 §0 担心的 dP/dS 长尾 underflow 崩盘**。⇒ **全 e4m3 backward 判定可行,不预防性切 e5m2**;e5m2 仍作为 §9.3 兜底保留。测试阈值据此标定为 `cossim>0.995`/`SNR>18`(留裕度)。 - -### 10.1quinquies 截至 2026-07-15(阶段2 kernel 已写,待 CDNA4 验) - -**已落地(代码,`tl.dot_scaled` 部分未验):** 按 §9.5 步骤2 把 backward 两个 inner kernel 的 7 个 `tl.dot` 全换 `tl.dot_scaled`。 - -- **新增 `_mx_quant` 内联辅助**(`triton_flash_attention_mxfp8.py`):包 `_calculate_scales`+`_quantize_fp8`,返回 `(e4m3 tile, uint8 scale)`。 -- **量化布局定案:统一 1D per-row**(`IS_2D_BLOCK=False`),scale `[outer, reduction/32]` 直接匹配 `dot_scaled`。**修正了盲写中发现的真 bug**:早先想让 head_dim dot 复用 forward 的 2D-block(`IS_2D_BLOCK=True`),但那返回紧凑 `[M/32, K/32]`,形状对不上 `dot_scaled`(forward 是靠指针把 2D scale 广播成 `[M, K/32]` 才喂进去的)。统一 1D 消除该特殊情况,参照实测数值几乎无差。 -- **operand 来源 = option A**:driver 入口仍 `convert_from_mxfp8` 反量化 saved e4m3 q/k/v→bf16(不额外存 bf16,省显存),kernel 内每个 dot 沿其 reduction 轴 1D 重量化。k/v/q/do 的 head_dim 套在 **outer kernel 量化一次复用**;p/ds 与 seqlen 套在 **inner 内**量化(转置使 reduction 轴落在最后一维,复用 last-axis 的量化辅助)。 -- 两个 launch 传 `QUANT_BLOCK_SIZE=BLOCK_SIZE_DEFAULT` / `USE_ASM=is_cdna4()`。 -- **参照同步**:`stage2` 参照全部 operand 改 1D per-row;q_sq/k_sk 按 option A 从 `q_hd`/`k_hd`(dequant 的 e4m3)再量化(双重量化),dO 单次。重跑 **6/6 passed,SNR≈23dB** 不变。 -- **`test_backward_kernel_matches_reference` 改对比 `stage2` 参照**(kernel 已是阶段2);该测试随之变为 **CDNA4-only**(`tl.dot_scaled` 在 MI250 不编译),与 forward Layer 2 同性质。 - -> ⚠️ 阶段2 kernel 的 `tl.dot_scaled` 与 scale 指针布局在 MI250 **无法编译/验证**,全靠对标 forward 已验证的 `dot_scaled` 用法 + stage2 参照「盲写」。待 CDNA4 跑步骤3 验收(kernel vs stage2 参照隔离移植 bug;kernel vs bf16 SDPA 端到端)。 - -**下一步:** CDNA4 到位后跑步骤3 验收;届时校准 `test_backward_kernel_matches_reference` 阈值。 - -### 10.1sexies 截至 2026-07-16(推翻 option A,改 A1;A 全删;A1 kernel 已写) - -**决策翻案(需求方拍板,知情决策):** 上会后定 **B 完全不要、A 也不留(视为冗余回退代码)**,直接上 AB 决策稿 §5 当初「不建议」的 **A1**。理由:留 A 当备用是冗余;A1 复用 forward 那份 e4m3+2D scale = correct-by-construction、与 forward 逐位一致、无双重量化。代价(kernel 内 scale 指针广播)已实现。详见 `MXFP8_BACKWARD_AB_DECISION.md` 顶部 banner。 - -**已落地(代码,未经硬件验证):** - -- **删 A**:删掉入口 `convert_from_mxfp8` 反量化 + 逐 dot 重量化 q/k/v 的 A 路径;删掉 A 的黄金参照 `mxfp8_attention_backward_reference_stage2`(含 `operand_source` A/B 开关)及其两个测试(`test_backward_stage2_reference_matches_sdpa` / `test_backward_kernel_matches_reference`);删掉 `triton_flash_attention_mxfp8.py` 里不再用的 `convert_from_mxfp8` import。保留 stage-1 参照 `mxfp8_attention_backward_reference`(FA 数学基准,非 A 专属)。 -- **A1 kernel**: - - 新增 `_load_scale_hd`(head_dim 收缩,照抄 forward `//32` 广播)/ `_load_scale_sq`(seqlen 收缩,转置对称 re-index),从紧凑 2D scale `[.., seqlen/32, head_dim/32]` 重建各 dot 的 scale tile,带越界/padded-head mask。 - - `_bwd_kernel_dkdv` / `_bwd_kernel_dq` + 两个 inner:q/k/v 改为**加载 saved e4m3 + 复用 2D scale**(不再 `_mx_quant`);dot a/b/e/f 用 `_load_scale_hd`,dot d/g 用 `_load_scale_sq`;dO/P/dS 仍 `_mx_quant` 现场 1D 量化。 - - driver:去掉入口 dequant,改传 saved e4m3 q/k/v + 三个 2D scale 及其 12 个 stride 进两个 launch。 - -**MI250(gfx90a) Docker 验证(2026-07-16):** +| 1 | forward kernel 在 CDNA4 跑通 | ✅ | +| 2 | backward(dQ/dK/dV)在 CDNA4 跑通 | ✅ | +| 3 | forward vs bf16 SDPA:cossim > 0.99、SNR 硬断言 | ✅ | +| 4 | backward dQ/dK/dV vs bf16 SDPA autograd:硬断言 | ✅ | +| 5 | 覆盖 causal × GQA × head_dim{128,192} × seqlen{1024,2048} | ✅ | +| 6 | dispatch 层 `mxfp8_e4m3` 分支可路由 | ✅ | -| 项 | 结果 | -|---|---| -| import `triton_flash_attention_mxfp8` | ✅ OK(A1 大改的签名/helper/装饰器解析注册通过) | -| 参照层 `test_mxfp8_attention_reference.py` | ✅ **12 passed**(含 stage-1 backward 参照 vs SDPA 6/6)→ 删 A 干净 | -| forward `test_kernel_matches_reference` | ❌ 7 failed(forward `attn_fwd` 的 `tl.dot_scaled` 在 gfx90a 编不过,**既有限制、非本次改动**;`git diff` 证实 forward kernel 本体无改动) | - -**仍待办:** -- **A1 golden reference 未写**(按新顺序「先 kernel 后参照」,下一步补)。当前 backward **无任何 A1 数值信号**。 -- **A1 kernel `tl.dot_scaled` + 2D-scale 广播待 CDNA4 验**;`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 -- padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 - -### 10.1septies 截至 2026-07-16(A1 golden reference 已写并验、dispatch 已接) +实测数值(gfx950): -**A1 参照(stage-2 全量化):** 在 `tests/unittest/mxfp8/utils.py` 加 -`mxfp8_attention_backward_reference_stage2`——q/k/v 单份 2D-block dequant 全程复用(A1,**零 seqlen 重量化、无双重量化**),只有 dO/P/dS 现场 1D per-row 量化,按 §9.5 各 dot 的 reduction 轴。纯 fp32 matmul、无 `tl.dot_scaled`,MI250/CPU 可跑。 - -**测试:** | 检查 | 结果 | |---|---| -| `test_backward_stage2_reference_matches_sdpa`(A1 参照 vs bf16 SDPA autograd,MI250) | ✅ 6/6,dQ/dK≈23.2dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985(无双重量化,符合/略优于已删 option-A 的 ≈23dB 基线) | -| `test_backward_kernel_matches_reference`(A1 kernel vs A1 参照,CDNA4-only) | ⏳ 已写、待 CDNA4;MI250 上 `tl.dot_scaled` 编不过(预期,既有限制) | -| 全参照套件(纯 PyTorch) | ✅ 18 passed;14 failed 全是 fwd/bwd kernel 的 `tl.dot_scaled` gfx90a 编译失败(既有限制、非本次改动) | - -**dispatch:** `alto/kernels/dispatch/attention.py` 加 `precision == "mxfp8_e4m3"` 分支 → `triton_attention_mxfp8`(与 mxfp4 同 kwargs 调用路径,签名兼容)。 +| backward kernel vs A1 参照 | dQ/dK **53~60 dB**、dV **59~63 dB**,cossim ≈ 1.0 | +| forward kernel vs 黄金参照 | ✅ 全网格通过 | +| 公开 autograd 端到端 vs bf16 SDPA | dQ/dK ≈ 23.3 dB·cossim ≈ 0.9976、dV ≈ 26 dB·cossim ≈ 0.9988 | +| 全量测试 | **75 passed**(约 11 秒) | -**仍待办:** -- **CDNA4 验收**:A1 kernel vs A1 参照(隔离移植 bug)+ vs bf16 SDPA autograd(端到端量化误差);`_load_scale_sq`(seqlen 轴广播)全仓无先例,最大盲区。 -- padded head_dim(如 192)的 scale mask 仅基础兜底,CDNA4 bring-up 时重点盯。 +kernel vs 参照 53~63 dB ⇒ 移植与 scale 广播无 bug;端到端 23 dB 与参照预测一致 ⇒ **误差全部来自 mxfp8 量化本身,不是 kernel 实现**。 -### 10.1octies 截至 2026-07-29(CDNA4 真机验收通过;补 3 处护栏断言;causal 循环边界经分析判定**不修**) +> **相对 mxfp4 attention 的主动加严**:mxfp4 的 `test_mxfp_attention.py` 只 print、无 assert 且无 backward 测试。mxfp8 把前向精度做成硬断言,并新增完整 backward 梯度校验。 -**环境:** gfx950(CDNA4)真机,`is_cdna4()=True`。Docker `exciting_kepler`(镜像 `wanghanthu/torchtitan:ubuntu22.04-pytorch2.12.0dev20260217-rocm7.2-patch`,`/home/yuesun/repos → /workspace`),共享集群上以 `HIP_VISIBLE_DEVICES=7` 单卡运行,其余 7 卡全程未占用。 +## 8. 不做的事(明确划线) -**§10.1septies 的头号盲区已解除。** A1 backward kernel 的 `tl.dot_scaled` + 2D-scale 指针广播(含全仓无先例的 `_load_scale_sq` seqlen 轴 re-index)在真机跑通,且与 A1 参照吻合: - -| 检查 | 结果 | -|---|---| -| `test_backward_kernel_matches_reference`(A1 kernel vs A1 参照,全 `test_cases` 7 组) | ✅ **7 passed**,dQ/dK SNR 53~60 dB、dV 59~63 dB,cossim ≈ 1.0 | -| `test_public_autograd_matches_sdpa`(公开 autograd 端到端 vs bf16 SDPA) | ✅ **2 passed**,dQ/dK≈23.3 dB·cossim≈0.9976、dV≈26 dB·cossim≈0.9988 | -| `test_kernel_matches_reference`(forward Layer 2,§10.1sexies 在 gfx90a 上 7 failed) | ✅ **7 passed** | -| 全量 `test_mxfp8_attention_reference.py` + `test_mxfp8_attention.py` | ✅ **41 passed** | - -- kernel vs 参照 53~63 dB ⇒ **移植与 scale 广播无 bug**;端到端 23 dB 与 §10.1septies 参照预测的 23 dB 一致 ⇒ **误差全部来自 mxfp8 量化本身,不是 kernel 实现**。 -- **padded head_dim=192 已随 7 组网格覆盖**(config3/config4),§10.1septies「padded head 的 scale mask 仅基础兜底」在 backward 侧首次拿到真机信号;GQA(`num_head_kv` = 8/2 < q)同样覆盖。 -- 据此 **§7 验收标准第 1~5 项均已在 CDNA4 达成**;第 6 项(dispatch 路由 smoke)不在本次运行范围。 +- ❌ CDNA3 / MI300 支持与 fallback 路径 +- ❌ 混合格式(e5m2)——见 §1 第 2 条,切换需改 kernel +- ❌ head_dim < 64(`tl.dot_scaled` 限制) +- ❌ 2D-block P 量化(P 沿 `BLOCK_N` 一维即可) +- ❌ autotune 扩展(`BLOCK_M = BLOCK_N = 64, PRE_LOAD_V=False` 单 config) +- ❌ TMA / async copy / pipelining 调优 +- ❌ 沿 head_dim 分块 QK(§5.1 风险未触发) +- ❌ 非方阵 causal(§4.3 断言拒绝,放开条件见 §9.2) +- ❌ FSDP/TP 集成测试 -**backward causal 循环边界与自身掩码不一致:经分析判定不修(决策记录)** +--- -两个 backward kernel 的**掩码**按 bottom-right 对齐(`col_offset = N_CTX_Q - N_CTX_K`),但决定**循环范围**的两处把 `col_offset` 丢了,等价于硬编码了 `col_offset == 0`(即方阵): +## 9. 关键决策记录 -| 位置 | 现状(未改) | 非方阵下的后果 | -|---|---|---| -| `_bwd_kernel_dkdv` 的 `lo` | `(start_n*BLOCK_N - BLOCK_M + 1) // BLOCK_M * BLOCK_M` | `seqlen_k − seqlen_q > BLOCK_M` 时跳过必须计算的 query 块 → dK/dV 漏贡献 | -| `_bwd_kernel_dq` 的 `hi` | `BLOCK_M // BLOCK_N * (start_m+1) * BLOCK_N` | `seqlen_k > seqlen_q` 时只要差 1 就漏 key 列 → dQ 漏贡献 | +### 9.1 backward 量化源:option A → A1 -一度改成了带 `col_offset` 的正确形式,**最终回退,只保留下面的断言**。理由按硬度排: +曾实现过 **option A**:backward 入口 `convert_from_mxfp8` 把 saved e4m3 q/k/v 反量化回 bf16,kernel 内每个 dot 沿其 reduction 轴 1D 重量化。它能跑,参照实测 SNR ≈ 23 dB。 -1. **这条路不可达,所以它不是 bug。** 两个断言合起来把 `col_offset != 0` 完全堵死:`causal → seqlen_q == seqlen_k` 拦住定长;`layout == "bhsd"` 顺带杀掉 varlen(thd)——这点很关键,varlen 下每条序列实际长度不同,光断言 max_seqlen 是拦不住的。而 `col_offset == 0` 时新旧公式只差一块全掩块,输出必然逐位一致(已实测:撤回改动前后 21 个 SNR 数字一致到小数点后四位;因为 `p = exp(-inf) = 0`,量化成 e4m3 仍是 0,加进 fp32 累加器是精确的 0)。 -2. **它从来没有「独立可达」过。** 非方阵 + causal 这条路本来就因为 kernel(bottom-right)与参照/`F.sdpa`(top-left)的约定不一致而算不对,边界错只是被一个更大的破坏盖住。修边界并不能让这条路可用,必须先统一约定。 -3. **「负 `lo` 越界读内存」这个理由是假的,已实测推翻。** Triton 的整数 `//` 对负数**向零截断**(不是 Python 的下取整),实测 `BLOCK_M=BLOCK_N=64` 时原式在 `start_n=0..3` 给出 `[0, 0, 64, 128]`,floor 语义才会给 `[-64, 0, 64, 128]`。所以原式**不会**产生负 `lo`、不会越界;`tl.maximum(..., 0)` 是新公式自己引入的需求,不是旧代码的隐患。 -4. **唯一真实代价:方阵下每个 key block 多跑一个全掩 query block。** Triton 循环内无法提前退出,那块的 4 次 dot 实打实算完再被掩成 0。原式 `lo = max(0, (start_n−1)·64)`,正确式 `lo = start_n·64`,故 `start_n ≥ 1` 的每个 program 各多一次迭代:多出比例 `2(N−1) / (N(N+1))`,N = seqlen/64。seqlen 1024 → dkdv 内层迭代 **+11.0%**;2048 → **+5.9%**;4096 → **+3.0%**(迭代次数,非墙钟时间;dkdv 约占 backward 一半,实际再打对折)。序列越长越不值钱。 +2026-07-16 翻案改 **A1**(复用 forward 的 e4m3 + 2D scale,§5.3),并**把 A 的代码与参照全部删除**,不留作备用。理由:A 有双重量化(dequant 后再量化),A1 与 forward 逐位一致、correct-by-construction;留 A 当回退是冗余代码。代价是 kernel 内多两个 scale 广播辅助(`_load_scale_hd` / `_load_scale_sq`)。 -结论:这是一次**性能与可读性**改动,不是正确性修复;按 fix 提交会误导 reviewer 去找线上不存在的 bug。留待与下方 review 债一起独立评估(那批债里已有「`lo`/`hi` 隐含要求 `BLOCK_M == BLOCK_N`」一条,同源同处,应一并处理)。 +A1 落地时采用「先 kernel 后参照」的顺序,意味着 kernel 写完后有一段时间没有任何数值信号,`_load_scale_sq`(seqlen 轴 re-index)全仓无先例、是当时最大盲区。CDNA4 验收 53~63 dB 后该盲区解除。 -**同源代码 `blockwise_fp8/triton_flash_attention_fp8_block.py` 有一字不差的两行**(1297、1665 行;mxfp8 backward 是从它 port 来的,**mxfp4 没有 backward,不是这段的来源**)。那份**没有** `seqlen_q == seqlen_k` 断言,autograd Function 的 `max_seqlens_q/k` 完全自由,但同样不可达:唯一生产入口 `blockwise_fa.py` 第 219 行 `assert query_states.shape[1] == key_states.shape[1]` 后把同一个 `seqlen` 传给 q 和 k;其 `test_attention.py` 的 10 组用例 `seqlen_q` 全等于 `seqlen_kv`。⇒ 同为潜伏、未暴露,本次不动。 +### 9.2 backward causal 循环边界:判定不修 -**新增 3 处护栏断言(把「不报错但结果是垃圾」变成当场报错)** +两个 backward kernel 的**掩码**按 bottom-right 对齐(`col_offset = N_CTX_Q - N_CTX_K`),但决定**循环范围**的两处把 `col_offset` 丢了,等价于硬编码 `col_offset == 0`: -| 断言 | 位置 | 原因 | +| 位置 | 现状 | 非方阵下的后果 | |---|---|---| -| `layout == "bhsd"` | forward op + backward op 各一处 | `convert_to_mxfp8(is_2d_block=True)` 按数据张量 `shape[-2] × shape[-1]` 分 32×32 块,只有 bhsd 时 `shape[-2]` 才是 seqlen。`bshd`/`thd` 下它按 **nheads** 分组(nheads 恰为 32 倍数时连 `torch._check` 都拦不住),而 kernel 的 scale 指针数学假定按 seqlen 分组 → **静默读错 scale**。§2 契约里 `layout` 的三选一由此**收窄为只支持 bhsd**(生产 `dispatch/attention.py` 本就只传 bhsd) | -| `not causal or seqlen_q == seqlen_k` | **forward op + backward op 各一处** | §10.1ter 记录的「潜伏坑」落地:kernel 掩码 bottom-right、PyTorch 参照与 `F.sdpa` 掩码 top-left,两者只在方阵等价。V1 生产只跑 self-attention,故**不统一约定**(那是解决不存在的问题),改为断言拒绝。
*(本节初稿曾写「正向自身自洽、纯推理下 `sq != sk` 可用,不拦 forward」,随后改为**两侧都拦**:一个 op 允许、另一个禁止同一种形状,调用者只能靠踩坑才知道边界在哪。forward 侧断言配 `test_causal_forward_rejects_non_square_shape` 回归。)* | - -断言有效性单独验过:`causal=True, sq=64/sk=128` 走 backward → 按预期抛 AssertionError;`causal=False` 同形状 → 正常通过(非 causal 无对齐问题,不该拦);`layout="bshd"` 走 forward → 按预期抛 AssertionError。 +| `_bwd_kernel_dkdv` 的 `lo` | `(start_n*BLOCK_N - BLOCK_M + 1) // BLOCK_M * BLOCK_M` | `seqlen_k − seqlen_q > BLOCK_M` 时跳过必须计算的 query 块 → dK/dV 漏贡献 | +| `_bwd_kernel_dq` 的 `hi` | `BLOCK_M // BLOCK_N * (start_m+1) * BLOCK_N` | `seqlen_k > seqlen_q` 时差 1 就漏 key 列 → dQ 漏贡献 | -**Code review 遗留债(已确认,本次未动,建议独立提交)** +一度改成带 `col_offset` 的正确形式,**最终回退**。理由按硬度排: -- `_MX_2D = tl.constexpr(False)`:名字叫 2D、值是 False,12 处调用全传它 ⇒ `_mx_quant` 的 `IS_2D_BLOCK` 永远走 1D 分支,而其 docstring 花两行解释了这个文件里不存在的 True 行为。参数只有一个取值就不该是参数。 -- `get_padded_headsize`(模块顶部)与 `get_padded_head_dim` 完全重复,前者无调用点。 -- scale 越界填充值不统一:forward `other=1`,backward `_load_scale_hd`/`_load_scale_sq` `other=127`。两者都不影响结果(对应数据已被掩成 0),但应统一为一个命名常量(E8M0 中性值是 127)。 -- `E4M3_TARGET_MAX_POW2` / `E4M3_MBITS` / `E4M3_FORMAT_ID` 手抄了 `mxfp8_quantization.FORMAT_TO_*` 的值,应直接 import。 -- backward docstring 大量引用外部文档编号(`plan §9.5`、`AB decision A1`、`dot a~g`),脱离本 plan 无法解读;建议改为按被算对象命名(qk / dp / dv / dk / dq)。 -- 两个 inner kernel 把运行时值标注成 `N_CTX_Q: tl.constexpr` / `N_CTX_K: tl.constexpr`(varlen 下由 `cu_seqlens` 运行时读出),类型标注不实。 -- `lo` / `hi` 隐含要求 `BLOCK_M == BLOCK_N`(若 `BLOCK_N=128`,`BLOCK_M//BLOCK_N` 变 0,原式直接给出 `hi=0` → dQ 全 0 且不报错);当前两者硬编码 64,无断言保护。 +1. **这条路不可达,所以它不是 bug。** §4.3 两条断言合起来把 `col_offset != 0` 完全堵死:`causal → seqlen_q == seqlen_k` 拦住定长,`layout == "bhsd"` 顺带杀掉 varlen(thd)——这点关键,varlen 下每条序列实际长度不同,光断言 max_seqlen 拦不住。而 `col_offset == 0` 时新旧公式只差一块全掩块,输出必然逐位一致(实测撤回改动前后 21 个 SNR 数字一致到小数点后四位:`p = exp(-inf) = 0`,量化成 e4m3 仍是 0,加进 fp32 累加器是精确的 0)。 +2. **它从来没有独立可达过。** 非方阵 + causal 本就因 kernel(bottom-right)与参照/`F.sdpa`(top-left)约定不一致而算不对,边界错只是被一个更大的破坏盖住。修边界不能让这条路可用,必须先统一约定。 +3. **「负 `lo` 越界读内存」这个理由是假的,已实测推翻。** Triton 的整数 `//` 对负数**向零截断**(不是 Python 的下取整),实测 `BLOCK_M=BLOCK_N=64` 时原式在 `start_n=0..3` 给出 `[0, 0, 64, 128]`,floor 语义才会给 `[-64, 0, 64, 128]`。原式不会产生负 `lo`、不会越界;`tl.maximum(..., 0)` 是新公式自己引入的需求。 +4. **唯一真实代价是性能。** 方阵下每个 key block 多跑一个全掩 query block(Triton 循环内无法提前退出,那块的 4 次 dot 实打实算完再被掩成 0)。多出比例 `2(N−1) / (N(N+1))`,N = seqlen/64:seqlen 1024 → dkdv 内层迭代 +11.0%;2048 → +5.9%;4096 → +3.0%。序列越长越不值钱。 -**仍待办:** +结论:这是**性能与可读性**改动,不是正确性修复;按 fix 提交会误导 reviewer 去找线上不存在的 bug。 -- 上述 review 债清理(与本次断言分开提交)。`lo`/`hi` 的 `col_offset` 与 `BLOCK_M == BLOCK_N` 两条同源,建议合并评估。 -- 若未来要支持非方阵 causal(cross-attention / prefix),需先在 kernel 与参照之间统一 top-left 或 bottom-right 约定,**再把 `lo`/`hi` 补上 `col_offset`**(正确形式见本节决策记录,届时才是真修复),最后才放开该断言。三件事顺序不能颠倒。 -- ~~§7 第 6 项 dispatch 路由 smoke 未跑~~ → 已于 §10.1nonies 补跑通过。 +同两行还隐含另一个约束:**`lo` / `hi` 要求 `BLOCK_M == BLOCK_N`**。若把 `BLOCK_N` 调成 128,`BLOCK_M // BLOCK_N` 整除成 0,`hi` 直接为 0 ⇒ **dQ 全 0 且不报错**。当前两者硬编码 64,无断言保护。调 BLOCK 尺寸前必须先处理这里。 -### 10.1nonies 截至 2026-07-29(对照 blockwise_fp8 复审 backward;**修掉 dkdv 组循环的 dO head stride 真 bug**;修掉 port 丢失的 `num_stages`,实测 1.15~1.51x;补齐 §7 第 6 项) +**若未来要支持非方阵 causal**(cross-attention / prefix),三件事顺序不能颠倒:① 先在 kernel 与参照之间统一 top-left 或 bottom-right 约定;② 再把 `lo`/`hi` 补上 `col_offset`(正确形式见本节,届时才是真修复);③ 最后才放开 §4.3 的断言。 -backward 是从 `blockwise_fp8/triton_flash_attention_fp8_block.py` port 来的(**不是 mxfp4——mxfp4 没有 backward**),故以它为基准做了第二轮逐行对照。结论:**一处正确性 bug(GQA × 非对称 head_dim 下 dK/dV 为 nan 乃至进程崩溃,随上游一起继承)** + 一处实测性能损失。 +上游 `blockwise_fp8` 有一字不差的两行(1297、1665 行),且**没有** `seqlen_q == seqlen_k` 断言,但同样不可达:唯一生产入口 `blockwise_fa.py` 第 219 行 assert 后把同一个 `seqlen` 传给 q 和 k。同为潜伏、未暴露。 -**修掉:port 时丢失了 `num_stages=1`(唯一实质缺陷)** +### 9.3 port 时丢失的 `num_stages=1` -上游把这个启动参数藏在一个**只有一个空 config 的 `@triton.autotune`** 里(`get_autotune_bwd_configs`,1156~1166 行;挂在 `_bwd_kernel_dkdv` 1169 行、`_bwd_kernel_dq` 1524 行): +上游把这个启动参数藏在一个**只有一个空 config 的 `@triton.autotune`** 里(`triton.Config({}, num_stages=1, num_warps=4)`),不调任何 BLOCK 尺寸,唯一作用就是钉住 `num_stages=1`。port 时它被当作「可选的调优基础设施」删掉,于是静默继承 Triton 默认值(3.6.0 实测 `num_warps=4, num_stages=2`——`num_warps` 恰好撞对,`num_stages` 翻倍)。 -```python -triton.Config({}, num_stages=1, num_warps=4) -``` +`num_stages=2` 给内层循环加两级软流水,多出的预取缓冲挤占寄存器;head_dim=192(padded 256)时 fp32 累加器与 scale tile 最多,退化最严重: -它不调任何 BLOCK 尺寸,唯一作用就是钉住 `num_stages=1`。port 时这个装饰器被当作"可选的调优基础设施"删掉了,于是 mxfp8 的 backward 静默继承 Triton 默认值——实测(Triton 3.6.0)默认为 `num_warps=4, num_stages=2`,`num_warps` 恰好撞对,`num_stages` 翻倍。 - -`num_stages=2` 给内层循环加了两级软流水,多出的预取缓冲挤占寄存器;head_dim=192(padded 256)时 fp32 累加器与 scale tile 最多,退化最严重: - -| 形状(batch 4, causal, gfx950) | num_stages=2(继承默认) | num_stages=1(对齐上游) | 提速 | +| 形状(batch 4, causal, gfx950) | num_stages=2 | num_stages=1 | 提速 | |---|---|---|---| | s1024, hq32/hkv8, d128 | 1.741 ms | 1.505 ms | **1.16x** | | s2048, hq32/hkv8, d128 | 6.201 ms | 5.411 ms | **1.15x** | | s2048, hq16/hkv16, d192 | 7.015 ms | 4.630 ms | **1.51x** | -(`triton.testing.do_bench`,warmup 50 / rep 200,跑两遍复现,偏差 < 0.3%。同时试了 `num_warps=8`:2.084 / 7.538 / 9.091 ms,明显更差 ⇒ 上游的 4 是对的。) - -修法是给两个 launch 各加显式 `num_warps=4, num_stages=1`,**不照搬那个空 autotune**——参数只有一个取值就不该披着 autotune 的皮,那正是它被丢掉的原因。改后全套通过(当时 41 个用例,测试矩阵扩到 75 见下)。 - -**确认比上游更好的四处设计(不动)** - -- **`tl.dot_scaled` 消掉了一整套 scale 记账。** 上游是"普通 `tl.dot` + 事后乘标量 descale",于是 `p_scale` / `log_p_scale` / `acc_descale` 要在 forward inner、dkdv、dq 全程穿(softmax 后 `p *= p_scale`、epilogue `acc *= acc_descale`、backward 重算 p 补 `+ log_p_scale * RCP_LN2`、最后 `dq *= sm_scale / p_scale` 除回去)。换成硬件指令后这些**全部消失**——不是简化,是特殊情况不存在了。 -- **砍掉上游传了但没读的 kernel 参数**:上游给 dkdv 传 `Out`/`DO`/`DQ`、给 dq 传 `Out`/`DK`/`DV`。 -- **把上游的静默失效变成当场报错**:上游 forward 支持 dropout 但 backward 无 dropout 逻辑且不检查(开 dropout 训练即静默错梯度);alibi 同理(只进 DEBUG 打印,从不进 kernel)。mxfp8 在 autograd backward 三行 assert 挡住(2128~2130 行)。 -- **无 fallback**:上游留 `use_fp8` 开关让同一套 kernel 兼跑 bf16,代价是每个 kernel 里 `if USE_FP8` 分叉;mxfp8 只有 MX 一条路,符合 §1 定的"仅 CDNA4 无 fallback"。 +(`triton.testing.do_bench`,warmup 50 / rep 200,跑两遍复现,偏差 < 0.3%。另试 `num_warps=8`:2.084 / 7.538 / 9.091 ms,明显更差 ⇒ 上游的 4 是对的。) -**一处真 trade-off(未测,线索)**:dO 量化位置。上游在 preprocess 里量化一次、写出 `DO_FP8` + `do_scale`,两个 backward kernel 直接读。mxfp8 改为循环内现场量化,但 dkdv 内层**每次迭代量化 dO 两次**(1295 行按 head_dim 分组供 dp、1305 行转置后按 seqlen 分组供 dv),因为两个 dot 的归约轴不同、需要两套 1D scale 布局。若 dO 改用 2D-block scale 量化一次即可同时喂两个 dot ⇒ 直接落在 §10.1octies 债表里 `_MX_2D = tl.constexpr(False)` 那条上,值得一并测。 +修法是给两个 launch 各加显式 `num_warps=4, num_stages=1`,**不照搬那个空 autotune**——参数只有一个取值就不该披着 autotune 的皮,那正是它被丢掉的原因。 -**补齐 §7 验收标准第 6 项:dispatch 路由 smoke(此前唯一未验项)** +**同类缺陷全仓排查无第二例**:扫了 `alto/kernels` 全部 20 个含 `@triton.jit` 的文件,循环密集型(attention / grouped GEMM)都有显式 launch 配置,其余是单趟访存受限的量化/elementwise kernel,`num_stages` 对它们无意义。 -经真实 dispatch 路径 `LPScaledDotProductAttentionWrapper(TrainingOpConfig(precision="mxfp8_e4m3"))` 跑 forward + backward(b2 / hq32 / hkv8 / s1024 / d128 / bf16):路由确实落到 `triton_attention_mxfp8`;输出与 dQ/dK/dV 全部有限且非零;**dK/dV 形状为 kv-head 形状 `(2, 8, 1024, 128)` 而 dQ 为 `(2, 32, …)`,GQA 的 head 归约方向正确**;`is_causal=False` 同样跑通。⇒ §7 验收标准 **1~6 全部达成**。 - -顺带确认 dispatch 层与本次两处断言天然兼容:`attention.py` 第 68 行硬传 `layout="bhsd"`,第 63~64 行的 `max_seqlens_q/k` 取自真实形状,self-attention 下必然相等。 - -**修掉一个真 bug(本次 review 唯一的正确性缺陷):dkdv 组循环用错了 dO 的 head stride** +### 9.4 修掉的真 bug:dkdv 组循环用错 dO 的 head stride `_bwd_kernel_dkdv` 的 GQA 组循环里,推进到下一个 query head 时把 dO 的指针按 **q 的** head stride 前进: @@ -573,54 +241,33 @@ q_offset += stride_qh do_offset += stride_qh # ← 应为 stride_doh ``` -`stride_doh` 本来就传进了 kernel(1345 行)、初始化 `do_offset` 时也用对了(1419 行),只有这一步用错。q 的最后一维是 `head_dim_qk`、dO 的是 `head_dim_v`,所以 bhsd 连续布局下 `stride_qh = seqlen * head_dim_qk`、`stride_doh = seqlen * head_dim_v`,**两者仅在 `head_dim_qk == head_dim_v` 时相等**。 - -触发条件是两个条件**同时**成立:`GROUP_SIZE > 1`(GQA)**且** `head_dim_qk != head_dim_v`。原 `test_cases` 七组恰好把这两个轴**各自覆盖但从未交叉**(config2/3/6/7 是 GQA 但 128/128;config4/5 是 192/128 但 `num_head_q == num_head_kv`),所以 41 passed 完全盖不住它。 +`stride_doh` 本就传进了 kernel、初始化 `do_offset` 时也用对了,只有这一步用错。q 的最后一维是 `head_dim_qk`、dO 的是 `head_dim_v`,bhsd 连续布局下两者**仅在 `head_dim_qk == head_dim_v` 时相等**。 -实测(`head_dim_qk=192, head_dim_v=128`,backward kernel vs A1 参照): - -| 形状 | 修前 dQ | 修前 dK / dV | 修后 dK / dV | -|---|---|---|---| -| hq8/hkv2(GROUP_SIZE=4) | 60.50 dB ✅ | **nan / nan** | 60.44 / 62.84 dB | -| hq8/hkv4(GROUP_SIZE=2) | 61.39 dB ✅ | **nan / nan** | 61.44 / 62.44 dB | -| hq16/hkv2(GROUP_SIZE=8) | 57.57 dB ✅ | **nan / nan** | 57.72 / 62.42 dB | +触发要两个条件**同时**成立:`GROUP_SIZE > 1`(GQA)**且** `head_dim_qk != head_dim_v`。原 `test_cases` 七组恰好把这两个轴**各自覆盖但从未交叉**,所以当时 41 passed 完全盖不住。实测(`head_dim_qk=192, head_dim_v=128`)修前 dK/dV 全 nan(dQ 正常——`_bwd_kernel_dq` 每个 query head 一个 program,没有这个组循环),GROUP_SIZE 与 seqlen 再大些直接 HIP 内存访问错误、进程 abort;修后 57~63 dB。 -dQ 全程正常——`_bwd_kernel_dq` 每个 query head 一个 program,没有这个组循环。GROUP_SIZE 与 seqlen 再大一些(seqlen 1024 / GROUP_SIZE 4)时越界读得更远,直接是 **HIP 内存访问错误、进程 abort**,不只是 nan。 +`test_cases` 已补上交叉配置(`num_head_q=32, num_head_kv=8, head_dim_qk=192, head_dim_v=128`),并验证过撤回修复后它确实会崩。修复对原有配置数值零影响。 -**测试矩阵的两个单值轴一并补掉(同类隐患,不只是这一个 bug)** +**上游 `blockwise_fp8` 第 1378 行是一字不差的同一行**,其测试的 10 组用例有完全相同的覆盖盲区 ⇒ 那份代码在 GQA + 非对称 head_dim 下同样会崩,且无断言拦截。 -查覆盖时发现分层是**反过来的**:三个纯 PyTorch 参照测试都跑 `causal=[True, False]`,而**每一个真正启动 kernel 的测试都被钉在 `[True]`**——不可能有移植 bug 的那层测两个值,会有移植 bug 的那层只测一个。`blockwise_fp8` 的两个 kernel 测试都是 `[True, False]`;mxfp8 的测试骨架抄自 mxfp4(doc 原话 "mirrors the mxfp4 attention test"),把严格性一起抄丢了。(另查:mxfp4 测试里那句被注释掉的 `[True, False]` 属于整个 `test_attention_fp8_with_sparse_do` 函数被注掉,**不是**有人发现 causal=False 挂掉才关的。) +### 9.5 相对上游 blockwise_fp8 更好的四处(有意保持) -四处 `[True]` → `[True, False]`(`test_kernel_matches_reference` / `test_backward_kernel_matches_reference` / `test_public_autograd_matches_sdpa` / `test_attention`)。非 causal 走的是不同代码:`lo = 0` 与 `hi = num_block_n * BLOCK_N` 替换掉两个 causal 边界公式,掩码跳过,forward 的分块切分与「整块被掩→写 0 提前退出」也绕开。实测全过。 +- **`tl.dot_scaled` 消掉了一整套 scale 记账。** 上游是「普通 `tl.dot` + 事后乘标量 descale」,于是 `p_scale` / `log_p_scale` / `acc_descale` 要在 forward inner、dkdv、dq 全程穿(softmax 后 `p *= p_scale`、epilogue `acc *= acc_descale`、backward 重算 p 补 `+ log_p_scale * RCP_LN2`、最后 `dq *= sm_scale / p_scale` 除回去)。换成硬件指令后这些全部消失——不是简化,是特殊情况不存在了。 +- **砍掉上游传了但没读的 kernel 参数**(给 dkdv 传 `Out`/`DO`/`DQ`、给 dq 传 `Out`/`DK`/`DV`)。 +- **把上游的静默失效变成当场报错**(§4.3 第三条)。 +- **无 fallback**:上游留 `use_fp8` 开关让同一套 kernel 兼跑 bf16,代价是每个 kernel 里 `if USE_FP8` 分叉。 -**非方形形状补上 kernel 级覆盖**:`causal` 断言现在 forward/backward 两侧都拦,所以 `seqlen_q != seqlen_k` **只在非 causal 下合法**——而它此前只有纯 PyTorch 参照测试覆盖(`reference_cases` 里的 `sq128/sk256`),kernel 侧为零。新增 `non_causal_cases`(`sq256/sk512` 与 `sq512/sk256 + d192/128`)+ 两个测试;把两个 kernel 测试的主体抽成 `_check_forward_kernel_vs_reference` / `_check_backward_kernel_vs_reference` 复用,**没有在原测试里加 `if causal and not square: skip` 这类条件分支**(那是往正常路径塞特殊情况)。实测非方形 × GQA × 非对称 head_dim 三重交叉 56~61 dB、cossim≈1.0,即这条路本来就是好的,只是没人测过。 - -用例数 41 → **75**(`test_kernel_matches_reference` 16、`test_backward_kernel_matches_reference` 16、`test_attention` 16、三个参照测试各 6、`public_autograd` 4、两个非方形各 2、断言回归 1),全套 **11 秒**。 - -**测试矩阵补上 GQA × 非对称 head_dim 交叉点**:`test_cases` 新增 `num_head_q=32, num_head_kv=8, head_dim_qk=192, head_dim_v=128`。已验证这个配置在把修复撤回后确实会崩(不是静默通过)⇒ 是有效的回归防护。修复对原有七组配置**数值零影响**(`head_dim_qk == head_dim_v` 时两个 stride 本就相同,实测 SNR 逐位一致)。 - -**上游 `blockwise_fp8` 第 1378 行是一字不差的同一行**,其 `test_attention.py` 的 10 组用例也有完全相同的覆盖盲区(GQA 组全是 128/128,192/128 的两组全是 `num_head_q == num_head_kv`)⇒ 那份代码在 GQA + 非对称 head_dim 下同样会崩,且未被断言拦住。 - -**同类缺陷全仓排查(结论:无第二例)** - -既然 `num_stages` 是"必需参数藏在看起来可选的装饰器里"而丢失,扫了 `alto/kernels` 全部 20 个含 `@triton.jit` 的文件。循环密集型(attention / grouped GEMM)全部有显式 launch 配置:mxfp4 attention forward(392~393 行 `num_stages=1, num_warps=4`,无 backward)、blockwise attention fwd+bwd、三个 grouped_gemm 各有独立 `autotune.py`(6/20/28 configs)。其余无配置的文件都是单趟、访存受限的量化/elementwise kernel,`num_stages` 对它们无意义。**本缺陷仅此一处。** - -**本轮新增债(并入 §10.1octies 债表一起清)** - -- `USE_SR` 在两处调用点(forward 363 行、`_mx_quant` 内部 1072 行)硬编码 `False`:随机取整整条链路通着但永不启用,与 `_MX_2D=False` 同类。 -- `is_varlen` / thd 分支已被 `layout == "bhsd"` 断言证明**不可达**(forward 908/953-959/1015 行;backward driver 一路把 `cu_seqlens` 传进 kernel)。这是上一轮加断言的直接后果——按"不写兼容/回退代码"的标准该删,要留就得把 varlen 真正做完。 -- forward `register_fake`(2154 行起)仍保留 thd 的 LSE shape 分支,而真实 op 已 assert 拒绝非 bhsd。仅在 torch.compile trace 非 bhsd 时不一致(那种情况本也会 assert),化妆品级。 -- `AUTOTUNE` / `PERF` 模块常量无引用(与 `get_padded_headsize` 一样,都是随上游一起抄来的死物,上游也是死的)。 - -**四个已验证排除的虚警(记录在此,避免重复排查)** +### 9.6 已排除的虚警(记录以免重复排查) | 初查疑点 | 排除依据 | |---|---| -| 上游 dkdv 的 exp2 分支对 `l_i` 双乘 `RCP_LN2` | 误读。1480 行是 `exp2(qk - l_i[:,None] + log_p_scale * RCP_LN2)`,`* RCP_LN2` 作用在 `log_p_scale` 上,`l_i` 只乘一次。两边都对 | -| 上游 autograd backward 返回 18 个梯度但 forward 只有 16 个输入 | 实测最小 `autograd.Function`:PyTorch **容忍**尾部多余的 `None`。mxfp8 这边是 15 对 15 精确匹配 | -| mxfp8 未 assert `o` / `softmax_lse` 连续(上游有) | `get_strides_from_layout` 与 `softmax_lse.stride()` 均直读张量真实 stride,非连续也读得对;`delta = empty_like(softmax_lse)` 继承同布局。无正确性风险 | -| mxfp8 缺 head_dim 整除检查(上游 assert `>=32` 且 `%2==0`) | mxfp8 需要的是更强的 `%32==0`,已由上游 `convert_to_mxfp8` 的 `torch._check(shape[-1] % block_size == 0)` 保证(2D 模式另检 `shape[-2]`)。上游保证的不在下游重复检查。这也解释了为何 layout 那条必须自己加:`convert_to_mxfp8` 检的是 `shape[-2] % 32`,bshd 下那是 nheads,32 头刚好整除,拦不住 | +| 上游 dkdv 的 exp2 分支对 `l_i` 双乘 `RCP_LN2` | 误读。`exp2(qk - l_i[:,None] + log_p_scale * RCP_LN2)` 里 `* RCP_LN2` 作用在 `log_p_scale` 上,`l_i` 只乘一次。两边都对 | +| 上游 autograd backward 返回 18 个梯度但 forward 只有 16 个输入 | 实测最小 `autograd.Function`:PyTorch 容忍尾部多余的 `None`。mxfp8 这边是 15 对 15 精确匹配 | +| mxfp8 未 assert `o` / `softmax_lse` 连续(上游有) | `get_strides_from_layout` 与 `softmax_lse.stride()` 均直读真实 stride,非连续也读得对;`delta = empty_like(softmax_lse)` 继承同布局 | +| mxfp8 缺 head_dim 整除检查(上游 assert `>=32` 且 `%2==0`) | mxfp8 需要更强的 `%32==0`,已由 `convert_to_mxfp8` 的 `torch._check` 保证。上游保证的不在下游重复检查。这也解释了为何 layout 那条必须自己加:`convert_to_mxfp8` 检的是 `shape[-2] % 32`,bshd 下那是 nheads,32 头刚好整除,拦不住 | +| mxfp4 测试里 `causal=[True, False]` 被注释掉,疑似 `causal=False` 有问题 | 那句注释属于整个 `test_attention_fp8_with_sparse_do` 函数被注掉,不是有人发现 `causal=False` 挂掉才关的。mxfp8 放开两个值后实测全过 | + +### 9.7 两处 Triton 语言坑(修复方式记录) -### 10.2 一句话总结 +- **fp8 `tl.load` 的 `other` 类型**:`other=0`(int32)无法 cast 为 `fp8e4nv`,报 `cannot cast int32[...] to fp8e4nv`。改为 `other=0.0`。 +- **模块级 constexpr 必须用实例化写法**:`X: tl.constexpr = 8` 这种注解式写法在 `@jit` 内不可见,报 `NameError: Cannot access global variable`。改为 `X = tl.constexpr(8)`。forward 的 `E4M3_*` 与 backward 的 `RCP_LN2` 都踩过。 -**forward 机械改写 mxfp4**(删 head_dim packing、`e2m1→e4m3`、`_pack_fp4→_quantize_fp8`);**backward 机械 port blockwise_fp8 backward**(三 kernel 结构白拿,非从零)——阶段1 骨架用高精度 bf16 `tl.dot` 接通并于 MI250 验通移植正确性;**阶段2 量化源经 option A → A1 翻案(2026-07-16):现为 backward 直接复用 forward 存的 e4m3 q/k/v + 2D-block scale(`_load_scale_hd`/`_load_scale_sq` 指针广播,零重量化),只有 dO/P/dS 现场 1D 量化;A 与其参照已删,B 不做**。全 e4m3(含 backward,已定)、**仅 CDNA4 无 fallback**。forward Layer 1 已于 MI250 验 6/6;backward stage-1 参照 vs SDPA、**A1 stage-2 全量化参照 vs SDPA** 均于 MI250 验 6/6(A1:dQ/dK≈23dB·cossim≈0.9976、dV≈25-26dB·cossim≈0.9985);dispatch 已接 `mxfp8_e4m3`。**2026-07-29 CDNA4 真机验收通过(§10.1octies):全套 41 passed,backward kernel vs A1 参照 53~63 dB·cossim≈1.0、端到端 vs SDPA≈23dB,`_load_scale_sq` 这个最大盲区解除,§7 验收标准 1~5 达成**;同时补 3 处护栏断言(`layout=="bhsd"` ×2;causal backward 要求方阵),把「不报错但结果是垃圾」变成当场报错。backward causal 循环边界与掩码不一致一事**经分析判定不修**(两断言已使其不可达,非正确性问题,详见 §10.1octies 决策记录)。**2026-07-29 二轮对照 blockwise_fp8(backward 的真正上游)复审(§10.1nonies):修掉 `_bwd_kernel_dkdv` 组循环误用 `stride_qh` 推进 dO 指针的真 bug(GQA × `head_dim_qk != head_dim_v` 下 dK/dV 为 nan 甚至进程 abort,测试矩阵两轴从未交叉故长期未暴露,已补交叉配置);修掉 port 时丢失的 `num_stages=1`(上游藏在单 config 空 autotune 里),实测 backward 提速 1.15~1.51x(d192 最显著);同类缺陷全仓排查无第二例;**补掉测试矩阵的两个单值轴(kernel 侧 causal 一直只测 True;非方形此前只有纯 PyTorch 参照覆盖),用例数 41 → 75、全套 11 秒**;补跑 dispatch 路由 smoke ⇒ **§7 验收标准 1~6 全部达成**。**