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 comentários · No GitHub

AMD / HIPNVIDIA / CUDAModels & quants

Descrição

## 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)

No site

Links install, modelos, releases.