Pull requests / #888

#888 hip: chunk-parallel GDN recurrence (WY form), two parity-gated arms - opt-in STRATA_GDN_WY (measured on gfx1100)

closed · @xyzzing · 0 评论 · 在 GitHub 查看

BenchmarksAMD / HIPModels & quants

描述

## What this is

The gated-delta-net recurrence (~7.4% of a 128K prefill, latency-bound: one serial token chain per head) rewritten in the chunk-parallel **WY form** (fla's `recompute_w_u_fwd` structure, reimplemented against this engine's exact contract): per 64-token sub-chunk one C×C unit-lower-triangular solve plus GEMM-shaped folds, one block per (value head, state column slice), state register-resident.

Two arms behind one opt-in env, `STRATA_GDN_WY=1` (env absent = today's serial `cols_pipe`, dispatch byte-identical):

- **fma arm** — pure-fp32 folds.
- **wmma arm** (gfx1100, runtime-arch-guarded, falls back to fma with a printed arm line) — rocWMMA 16×16×16 fp16×fp16→f32 folds for the A-build (`K̃ᵀK̃`), the rhs/pq folds, `P = QK̃ᵀ`, O_intra, and the rank-BT state fold; the triangular solve stays fp32 scalar; ~61.4 KiB LDS via barrier-separated overlays.

## The parity discipline (the substance)

Two HIP-visible CTest harnesses ship with the kernels:

- `gdn_wy_check` — all-heads fp64 oracle + a **derived** reassociation bound `(T/WY+1)·2⁻¹⁴·(1+scale)`, NaN-aware metrics, negative controls, 9 fp64 host fixtures.
- `gdn_wy_wmma_check` — the **fp16-class bound** `⌈T/BT⌉·(2⁻⁶+2⁻¹⁴)·(1+scale)` was derived and recorded **before the wmma kernel was written** (the precision-contract rule), with an explicit +2.5·B negative control. Green: worst oracle-bound utilisation **3.5%** across T ∈ {1,37,64,65,200,4099} × decay modes.

## Measured verdict on gfx1100 (7900 XTX, nightly ROCm): the WY form loses on this geometry

32k-token engine spot pairs, arms differing only in the env:

| arm | gdn recurrence (32k) | result |
| --- | --- | --- |
| serial `cols_pipe` | 1,318 ms | baseline |
| WY fma folds | 189,878 ms | 144× slower — REJECT |
| WY wmma folds | ~3,050 ms eq. | ~6× slower per layer-chunk than cols_pipe (~127 ms vs ~13.7 ms); engine prefill −39..−40% reproducibly — REJECT |

rocWMMA 16×16×16 at 64×64/64×32 shapes with 4-warp blocks does not amortize fragment traffic at S=128/HV=48. The kernels ship **opt-in default-off and parity-protected** so the next attempt (C=32 LDS-side restructure, GEMM fusion) starts from evidence instead of folklore.

## gfx1100 / rocWMMA 7.1.1 notes baked into the code

- rocwmma headers must be included at **global scope** (their `std::` names resolve against an enclosing namespace otherwise).
- `store_matrix_sync` static-asserts fragment/pointer type equality — no converting stores; fold results stay fp32 in LDS overlays and P16 is hand-converted through an f32 scratch region.
- An mma operand tile must never overlap the same GEMM's store target (the T=1 oracle-FAIL class this campaign hit and fixed).

## Testing

`gdn_wy_check` and `gdn_wy_wmma_check` (unconditional CTest, HIP-visible scope); both green on this branch. Engine runs: `STRATA_GDN_WY=1` prints `gdn wy arm = wmma (gfx1100, opt-in)` (or the fma line where rocwmma is absent); unset behaves byte-identically to main.

Related context from the same campaign: the model's rope scaling (`--rope-scaling yarn`) was validated to 512K prompt positions end-to-end on this hardware (needle recall at 256K/524K depths, 1.5-1.9K tok/s prefill via `--kv-resident` streaming) — decode at that depth is gated by MTP draft acceptance (0.80 → 0.06), a separate problem from this PR.

站内延伸阅读

链到安装、模型与版本说明,便于 SEO/GEO,非官方 issue 正文。