- fp8_e4m3_from_float32: manual FP8 E4M3 cast (bias=7, exp 0-15 valid, NaN guard for exp=15/mant=7, mantissa overflow handling) - fp8_e4m3_to_float32: dequantize FP8 E4M3 bit pattern back to Float32 - half_step_to_e2m1_idx: E2M1 step mapping (0-12 → 0-7) - quantize_e2m1_nibble: per-element E2M1 quantize + sign + pack - Verified 0/500 trial failures against Python reference - Key fixes discovered during validation: 1. FP8 E4M3 bias is 7, NOT 8 2. Exponent range is 0-15 (exp=15/mant=7 is NaN; others valid) 3. Subnormal formula: val = m * 2^(-9) = m/512 (NOT m/1024) 4. Round-to-nearest-even (not round-half-up) for half_step and mantissa 5. Mantissa overflow (round to 8) must increment exponent
69 lines
5.6 KiB
Markdown
69 lines
5.6 KiB
Markdown
OK. Here's the full status.
|
||
|
||
---
|
||
|
||
## STATUS UPDATE — DSV4 NVFP4 Kernel (as of 2026-05-28)
|
||
|
||
### What's DONE and PROVEN (cos 0.999+)
|
||
|
||
| Component | Status | Notes |
|
||
|---|---|---|
|
||
| **FMHA core** (hd=64/128/256) | ✅ | TMEM-P (hd≤64), SMEM-P (hd>64). cos 0.999997-0.999998 |
|
||
| **D1: multi-KV-tile via Python KV merge** | ✅ | Correct, cos 0.999998. 5–9 kernel launches per decode step |
|
||
| **D2: per-head multi-head launch** | ✅ | n_h=1–128, head-packed MQA/GQA |
|
||
| **D3: SWA length mask (in-kernel)** | ✅ | cos 0.999996 |
|
||
| **D4: causal mask on SWA (in-kernel)** | ✅ | cos 0.999996, combined with D3 |
|
||
| **D5c: sink merge** | ✅ | Single softmax over [S_comp, S_swa + attn_sink]. Single-tile and multi-tile both pass |
|
||
| **D5b: per-row LSE output** | ✅ | All 128 rows write correctly now |
|
||
| **Stage E: production wrapper** | ✅ | `production.py` — kernel cache, batch dim, MQA/GQA head-packed, `torch.library.custom_op` |
|
||
| **MoE fused SwiGLU (NVFP4)** | ✅ | Clamping in kernel, 7-warp specialization |
|
||
| **Router (hash + dense)** | ✅ | Both modes |
|
||
| **use_2cta_instrs conditional** | ✅ | 1.7–1.9× prefill speedup |
|
||
| **NVFP4 primitives** | ✅ | sf_dtype, TMA, MMA kind all verified |
|
||
| **GPU-only NVFP4 quantize** | ✅ | Byte-exact match with Python |
|
||
| **KV cache infra** | ✅ | allocator, paged_cache, state_cache, flush, schema, handle, manager |
|
||
| **Compressor (CSA/HCA)** | ✅ | flush_write kernels, FP8/FP4 quantize |
|
||
| **Indexer** | ✅ | gather_kv, score_topk (FP32 scalar). Compiles and runs on B200 |
|
||
| **Model assembly** | ✅ | config, layer_schedule, all 43 Flash + 61 Pro layers construct and validate |
|
||
|
||
### What's BLOCKED / UNDONE
|
||
|
||
| Item | Status | What's needed |
|
||
|---|---|---|
|
||
| **D1.5: in-kernel O rescale (TMEM round-trip)** | ❌ FUNDAMENTALLY BROKEN | TMEM load→store atoms have mismatched column mappings. NO-OP round-trip corrupts data. **Closed issue — do not re-attempt.** Production path is Python KV merge. |
|
||
| **Priority 1: Profile production decode** | ❌ Not done | Need to measure if 5–9 launch Python KV merge overhead actually matters. Gates Priority 8. |
|
||
| **Priority 2: One-way final-epilogue rewrite** | ❌ Not done | Replace `epilogue_tma_store` with MoE-style TMEM→REGS→SMEM→GMEM. Unlocks P4 (multi-CTA) and P6 (FP4 fuse). Attempted in `fmha_smem_acc.py` — many commits, unclear if it's working. |
|
||
| **Priority 3: NVFP4-1.1 FP4 quant in MoE epilogue** | ❌ Not done | Fuse amax+FP4 pack into L1→L2 path. Independent of FMHA. |
|
||
| **Priority 4: D2 multi-CTA grid** | ❌ Blocked on P2 | `epilogue_tma_store` can't accept `flat_divide` coordinates. |
|
||
| **Priority 5: Stage E cleanup** | ⚠️ Partial | production.py exists. E1✅ E2⚠️ E3⚠️ E4-E7 TODO. Many debug test files still in `tests/unit/` (66 test files, lots of diagnostic/debug artifacts). |
|
||
| **Priority 6: NVFP4-1.2 FP4 in FMHA output** | ❌ Blocked on P2 | Needs register slot in new final epilogue. |
|
||
| **Priority 7: NVFP4-2 FP4 KV pipeline** | ❌ Blocked on P2 | FP4 KV dequant in SMEM for deeper pipeline stages. |
|
||
| **Priority 8: Per-kt rescale fix** | ❌ Conditional on P1 | Three paths (A/B/C) if profiling shows Python merge overhead >5%. |
|
||
| **Priority 9: hd=512 single-kernel** | ❌ MLIR hang | CuTeDSL MLIR optimizer can't handle it. Decode works via head-packed M + hd≤256 chunks. |
|
||
| **Priority 10: Indexer FP4 tensor-core scoring** | ❌ Stage F | Scalar FP32 scoring today. Needs FP4 MMA + warp-level top-k. 2–3 weeks. |
|
||
|
||
### MAY_24_2026_PLAN_NEW.md (TMEM Round-Trip Investigation Plan)
|
||
|
||
This is a 5-phase investigation plan for Priority 8, Path A (CUTLASS atom replication). **Status: not executed.** The plan was written but MEMORY.md already contains the conclusion: TMEM round-trip is FUNDAMENTALLY BROKEN and will NEVER work. The plan's Phase 1 (read CUTLASS reference) and Phase 2 (NO-OP round-trip tests) were effectively done through the many `test_d15_*` and `test_tmem_roundtrip_minimal.py` test files. All variants failed. The plan is **historical** — its conclusion (escalate to Path C / Python KV merge) is what's already in production.
|
||
|
||
### ROADMAP.md
|
||
|
||
The roadmap is **current and accurate** as of 2026-05-26. The priority ordering and dependency chain are correct. No priorities have been completed since the last update — the latest work was Stage E production extraction (commit `b9f15c2`).
|
||
|
||
### Attention Kernel Folder
|
||
|
||
- `fmha.py` (592 lines) — **main production kernel**
|
||
- `fmha_smem_acc.py` (592 lines) — SMEM accumulator variant, attempted but unclear status (likely broken/in-progress for P2)
|
||
- `production.py` (427 lines) — **working production wrapper** with Python KV merge
|
||
- `fmha_backup_*.py` — historical backups
|
||
- 66 test files in `tests/unit/` — heavy diagnostic residue from the D1.5 investigation
|
||
|
||
### Summary: Where We're At
|
||
|
||
**The kernel works.** FMHA at hd=64/128/256 with all masks (SWA, causal, sink), multi-head, multi-KV-tile via Python merge — all proven on B200. The MoE, router, indexer, compressor, KV cache — all working.
|
||
|
||
**The next move is Priority 1 (profile production decode)** to determine if the Python KV merge overhead matters. If it's <5% of decode latency, you ship as-is and the D1.5 TMEM round-trip problem becomes irrelevant.
|
||
|
||
**If profiling says fix it, the path is Priority 2 first** (one-way final epilogue rewrite using MoE pattern), which unblocks the multi-CTA and FP4 fusion chain. The per-kt rescale (P8) is a separate, harder problem with three possible approaches.
|
||
|
||
**The `MAY_24_2026_PLAN_NEW.md` is dead** — its investigation was run, and the answer was "TMEM round-trip doesn't work." Should be archived or deleted. |