Pull requests / #1296

#1296 hip: chunk-parallel GDN recurrence (WY form), two parity-gated arms (STRATA_GDN_WY)

open · @xyzzing · 0 Kommentare · Auf GitHub

BenchmarksAMD / HIPNVIDIA / CUDAModels & quants

Beschreibung

Re-file of #888 — that PR was auto-closed on 2026-10-06 when this repository's main history was cleaned up (per the maintainer's note there, not a rejection). Same two commits, re-based onto the new main. Original discussion: #888.

What the re-base changed, explicitly:

- The old branch's first commit had swept in two local agent-onboarding files (`.agentic/`, `CLAUDE.md`); they are dropped here — they were never part of the PR.
- `src/prefill/kernels.cu` had grown the Aurora gfx11 HEAD/quad kernels in the same region the WY arm lands. The resolution keeps both: this branch adds the WY kernels + the `STRATA_GDN_WY` branch of the dispatch *on top of* the current Aurora code, with the composition untouched (the WY kernels write the same pre-norm output the cols kernels write; the following `gdn_out_norm_kernel` call is the existing one). Default path (no env, `STRATA_GDN_HEAD` default on this card) is byte-identical to main.
- Gates re-run on this base: full HIP build green; `gdn_wy_check` all checks PASS (0 mismatches, fp64 oracle 9/9 fixtures); `gdn_wy_wmma_check` PASS (worst bound utilisation 3.5%, negative control ok).

The measured verdict in the original description stands: on gfx1100 both arms lose to the serial cols_pipe kernel (fma ~144x slower, tensor-core folds ~6x slower per layer-chunk, engine prefill -39..-40%) — this ships opt-in default-off + parity-protected so the next attempt starts from evidence, not folklore.

---

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


*Prepared with an AI engineering agent (GLM-5.3 / Z.ai) under human direction; every number is from our own recorded runs.*

Mehr auf der Site

Links zu Install, Modellen, Releases.