Pull requests / #999
#999 hip: the fused-int8 MoE kernels on gfx1100 rocWMMA - opt-in STRATA_PF_FUSED arm + parity fixture (engagement: parity, REJECT-at-parity recorded)
closed · @xyzzing · 0 评论 · 在 GitHub 查看
AMD / HIPNVIDIA / CUDAModels & quants
描述
## What this is The fused-int8 MoE expert kernels (`moe_fused_iq.cu`, CUDA sm_80-only upstream) translated to **rocWMMA 16×16×16 INT8→INT32 for gfx1100/RDNA3**: `native_kernel_wmma<T, GU, WW>` behind the existing opt-in env (`STRATA_PF_FUSED=1`, plus `STRATA_PF_FUSED_NATIVE` semantics preserved) — same batch machinery (grouping tables, work-item walk, per-row IQ decode via the project's `convert()`/`signed8` semantics), plain LDS int8 tiles (the ldmatrix swizzles dropped — rocWMMA reads unswizzled tiles), the `mma_tile_16` rocWMMA primitive chained per 16-k slice, and the original per-half/per-sub-block fp32 scale folds applied per-lane from the i32 readbacks (the m16n8 lane mapping — so the SwiGLU/requant and down-scatter epilogues run verbatim). **Verdict up front: engagement REJECT at parity** — measured, not promoted, per the project's arm pattern. On gfx1100 (7900 XTX, nightly ROCm, 128K × 3 paired engine cells): the production MMQ path and this arm are at parity (-0.8% end-to-end; gate/up +1.7%, noise). The pinned ggml's MMQ tiles are **already RDNA3-aware** (`mmq-config-rdna3.cuh`), so the generic-tile headroom this translation targeted does not exist. The arm ships opt-in dormant and parity-protected. ## The fixture (the substance) `prefill_fused_iq_wmma_test` (HIP CTest): the IQ3_S gate/up chain — decode (verbatim `convert()`/`signed8` semantics) → plain LDS int8 tiles → rocWMMA mma → per-half/per-sub-block fp32 scale folds — against a double reference on the same quantized activations: - IQ3_S gate/up: **rel RMS 0.000e+00** (the int8 dots and INT32 accumulation are integer-exact, as the arithmetic gives), worst row 1.27e-05 (fp32 accumulation order only) - Q2_0 down decode: bit-exact vs ggml `to_float` - deterministic: a double-launch probe diffs 0/327680 output elements Bounds pre-declared before the kernel existed (1e-4 RMS / 1e-3 worst row); no tolerance moved. ## Traps the fixture caught (all in the translation, all fixed) 1. `signed8` mis-port: the sign spread is **one bit per value** (`__vcmpne4(s & 0x08040201, 0)` masks), not bit-7 replication — half the values decoded wrong with the naive port 2. `__half2float((unsigned short) x)` on HIP converts the **integer**, not the bits — needs `__ushort_as_half` (the CUDA arm never runs this path on HIP, so it never showed) 3. A decode-coverage pass missing (rows 64–127 never decoded) 4. Per-sub-block scale slots (`ws[2*uj]`, `ws[2*uj+1]`) vs per-16-group reads — the fold must index the sub-block's own scale 5. rocWMMA headers at global scope; `__half2float`-class traps aside, the fixture also pinned `matrix_b col_major` = `ptr[n·ld + k]` and 16×16 accumulators (ld = N) ## Testing - `prefill_fused_iq_wmma_test`: PASS (bounds 1e-4/1e-3 RMS/worst-row), deterministic ×3 - `gdn_wy_check`, `gdn_wy_wmma_test` suites: unchanged, green on the same tree - Engagement: 3 paired 128K cells, `STRATA_PF_FUSED=1` vs absent, same binary — parity within noise; decode untouched (prefill-only arm, `--max-new 8` cells)
站内延伸阅读
链到安装、模型与版本说明,便于 SEO/GEO,非官方 issue 正文。