diff --git a/docs/fused_fa_problem.md b/docs/fused_fa_problem.md new file mode 100644 index 0000000..7549824 --- /dev/null +++ b/docs/fused_fa_problem.md @@ -0,0 +1,435 @@ +# qxmx — Fused FlashAttention for prefill: problem statement + history + +> **STATUS (updated 2026-07-19): SOLVED.** Path A (`sycl::joint_matrix`) +> landed in commit `b375466` (M1.7c). Prefill 616 → **741 tok/s (+20%)** at +> the new default chunk=512, all gates green (`fa_kernel_test` ALL PASS, +> `attn_forward_batch_test` 5/5, greedy ' Paris', logits mean|diff|=0.1082). +> See memory `qxmx_m1_7c_fused_fa_landed.md` for the shipped design + +> `qxmx_joint_matrix_b70.md` for the B70 tile-support probe. The doc below +> is the original problem statement as written BEFORE the attempt; kept as +> the record of what was known/tried and as the brief that informed the +> successful approach. +> +> **Key finding the doc anticipated but didn't know:** B70's +> `matrix_combinations` query OVER-PROMISES. Only fp16/bf16 16×16×16 +> SG16 tiles actually JIT; tf32 (m=0 n=16 k=8) and fp16 32×32×64 are +> listed but JIT-fail ("undefined builtin" / "not supported"). The shipped +> kernel uses fp16 16×16×16, NOT the tf32 the doc's Path A speculated +> about. A second finding: K/V dequant was 76% of v1 kernel time (48× WG +> redundancy) — hoisting it to a per-chunk global fp16 shadow was the +> bigger win than the DPAS fusion itself. + +--- + +**Status:** open problem, hard, blocked once. This document is a +self-contained brief for an agent asked to attempt it fresh. +**Target:** Intel Arc Pro B70 (Xe2/Battlemage, 256 EU, 128 KB SLM/WG, +sub_group sizes {16, 32}), oneAPI 2026.x icpx `-fsycl`. +**Goal:** replace the current 3-kernel FA split with a fused tiled kernel +that keeps S/P in registers and reaches parity with llama.cpp's FA on the +same hardware. Measured prize: prefill 616 → ~900+ tok/s (llama-bench +pp4096 ub=2048 hits 980 tok/s on this card with the same model). + +--- + +## 1. The gap, precisely + +Prefill on the 5289-token `bench/session_4mb.txt` prompt, naive DeltaNet +(the shipped SSM path), default chunk=256: + +- **ours:** 616 tok/s (8.6 s), attn 24.7 ms/chunk-1, growing to 172 ms by + chunk 20 (attn is 9% of chunk-1, ~42% by chunk-20). +- **llama-bench** (Ternary-Bonsai-27B Q2_0, same B70): ub=256 → 630 tok/s, + ub=2048 → 980 tok/s (plateau). `-fa auto` (their FA path is on). + +At matched ub=256 we are within ~2% of llama.cpp. The 37% gap to 980 is +llama.cpp running ub=2048, where their FA scales and ours does not. + +### Why we don't scale with chunk (per-phase profile, 2026-07-19) + +Per-token attention cost at chunk-1 (no prior context — pure in-chunk +causal work) grows with chunk size: + +| chunk | attn/tok (chunk 1) | total tok/s | +|------:|-------------------:|------------:| +| 256 | 0.096 ms | 616 | +| 512 | 0.121 ms | 623 | +| 1024 | 0.183 ms | 607 | +| 2048 | 0.305 ms | 578 | +| 4096 | 0.550 ms | 522 | + +Bigger chunk amortizes per-tile launch overhead (good — FFN drops 8% +per-token from 256→512) but the in-chunk O(N²) causal attention work +dominates and gets *worse* per-token. Attention is the wall that blocks +chunk-size scaling. + +### The structural reason: our FA is a per-chunk re-FA, not a tiled pass + +`gpu_flash_attn` (src/qxmx_fa.cpp:299) runs, per chunk, a K-tile loop: + +```cpp +int64_t total = pos_start + n_tok; // walks 0..pos_start+n_tok-1 +int n_tiles = (total + TILE_T - 1) / TILE_T; +for (int ti = 0; ti < n_tiles; ti++) { // rescans prior context EVERY chunk + fa_dequant_kt(q, kv_k, t0, tile_len, ws); + fa_qk_gemm(q, n_tok, ws); // small GEMM: M=n_tok, K=HD, N=TILE_T + fa_softmax_merge(q, pos_start, t0, tile_len, n_tok, scale, ws); + fa_dequant_vt(q, kv_v, t0, tile_len, ws); + fa_pv_gemm(q, n_tok, ws); // small GEMM: M=n_tok, K=TILE_T, N=HD +} +fa_finalize(q, out, n_tok, ws); +``` + +Per K-tile that's **6 kernel launches**, ~16 tiles/chunk at chunk=256 → +~96 launches/chunk. Each QK/PV GEMM is small (`M=256, K=256, N=128`, +`M=256, K=128, N=256`) and runs through `gemm_tf32_batched` at 4.2 TFLOPs +(see qxmx_gemm_tf32.cpp) — **3% of the 133 T-MAC/s Xe2 ceiling**. S and P +are materialized to global memory between the QK and PV GEMMs (~3 MB per +tile at CHUNK=256). + +A **fused** FA-2 kernel (one WG per (q-head, Q-tile)) keeps Q in registers, +streams K/V through SLM, holds S/P in registers, and writes only the final +O. No per-tile launch overhead, no S/P materialization, the QK/PV compute +is the same DPAS work but issued by one WG with full occupancy. That is +what llama.cpp's FA does, and it is why they scale to ub=2048. + +### Important: incremental FA (carrying (m,l,O) across chunks) is NOT the answer + +A no-rescan probe (`QXMX_FA_NO_RESCAN`, skip prior-context tiles, OUTPUT +WRONG) showed chunk-20 attn 172→24 ms and prefill 616→754 tok/s. It is +tempting to conclude "ship incremental FA carrying (m,l,O) state for +22%." + +That conclusion is **wrong**. The prior-context QK/PV GEMMs are necessary +per-query work (each new query must attend to prior keys — causal +attention). The no-rescan probe skipped necessary work, not just +redundant work. The 754 figure is an unachievable ceiling. + +Reconciliation: for 2048 tokens, llama.cpp one-pass ub=2048 does ~2.1M QK +dots/query-head; our 8× chunk=256 over growing context does ~2.36M. **We +do less total work than llama.cpp** (we only scan the growing prefix, not +the full context every time), yet they are faster. The gap is per-op +efficiency (launch overhead + small-tile GEMM + S/P materialization), not +algorithm. **Do not propose incremental FA as the fix.** The fix is +fusion. + +--- + +## 2. What we tried (the history — read before re-attempting) + +The full git log: `git log --oneline | grep -E "M1.7b"`. Key commits: + +### M1.7b.1 (`bcef29a`) — fused FlashAttention-2 kernel, plain SYCL, scalar QK + +Aila-style: one WG per (q-head, query-token), LWS=128. K walked in +TILE_T=128 blocks. **Scalar** `for d+=8` unrolled QK dot (no DPAS). +Online-softmax via SLM merge_state; `reduce_over_group` for tile_max/sum; +`sycl::native::exp`. T13 rotated-domain V (accumulate P·V_rotated_scaled, +butterfly once per output row at the end). + +**Correct.** Greedy gate passed. But the scalar HD=256 dot capped prefill +at ~230 tok/s. No DPAS = no XMX throughput. (M1.7b.6 shipped this as the +fallback; 230 tok/s — slower than the 3-kernel split that replaced it.) + +### M1.7b.4 (`7303d84`) — ESIMD-vectorized FA (no DPAS) + +Rewrote the M1.7b.1 scalar K-dot with ESIMD vector ops. **2.35× faster +attn** but still no DPAS. Hit the ESIMD constraints hard: +- **64-lane cap per WG** (ESIMD LWS ≤ 64). +- **No `sycl::exp` / `reduce_over_group` / sub_group shuffles** — ESIMD + and plain-SYCL subgroup ops are mutually exclusive in the same kernel. +- **Register pressure:** 8×256 float O accumulator = 8 KB = a full GRF. + +### M1.7b.5 — inline-asm DPAS in plain SYCL (THE BLOCKED ATTEMPT) + +The idea: crib just the DPAS inline-asm invocation from sycl-tla +(`include/cute/arch/mma_xe.hpp`, `XE_DPAS_TT`), with NO cute dependency, +so a plain `sycl::parallel_for` kernel can issue DPAS *without* +`[[intel::sycl_explicit_simd]]`. That would let the same kernel also use +`sycl::exp` / `reduce_over_group` / sub_group shuffles (which ESIMD +forbids) — exactly what a fused FA needs (DPAS for QK/PV, subgroup ops +for softmax). + +Code: `src/qxmx_dpas.h` (deleted in `ee09d1c` cleanup; recover with +`git show ee09d1c^:src/qxmx_dpas.h`). Wrapper: + +```cpp +struct dpas_tf32_m1 { + static constexpr int M = 1, N = 16, K = 8; + using AV = dpas_vec; /* GRF, M*K=8 tf32 padded */ + using BV = dpas_vec; /* K=8 tf32 */ + using DV = dpas_vec; /* M*N=16 fp32 */ + static void fma(DV& d, const AV& a, const BV& b) { + DV dl = d; +#ifdef __SYCL_DEVICE_ONLY__ + asm("{\n" + ".decl DST v_type=G type=f num_elts=16 alias=<%0,0>\n" + ".decl SRC1_UD v_type=G type=UD num_elts=8 alias=<%1,0>\n" + ".decl SRC2_UD v_type=G type=UD num_elts=128 alias=<%2,0>\n" + "dpas.tf.tf.8.1 (M1, 16) DST.0 DST.0 SRC1_UD.0 SRC2_UD(0,0)\n" + "}\n" + : "+rw"(dl) : "rw"(b), "rw"(a)); +#endif + d = dl; + } +}; +``` + +`dpas_vec` is `T __attribute__((ext_vector_type(N)))` — OpenCL-style +vector that lowers to a GRF region for inline asm (this is what sycl-tla's +`vector_t` expands to in device code; `sycl::marray` hits "impossible +constraint: cannot store value into a register" for the 128-element A/D +operands). + +Probe: `tests/dpas_asm_probe.cpp` (also deleted; `git show ee09d1c^`). +Validates `qxmx_dpas.h::dpas_tf32_m1::fma` against a CPU fp32 reference in a +`sycl::single_task`. + +**THE FAILURE (reproduced 2026-07-19, exact output):** + +``` +Build program log for 'Intel(R) Graphics [0xe223]': +error: parsing vISA inline assembly failed: +near line 775: syntax error, unexpected IDENT +error: backend compiler failed build. +``` + +The vISA JIT rejects the asm dialect. Host compilation (`icpx -fsycl`) +**succeeds** — the `#ifdef __SYCL_DEVICE_ONLY__` block is hidden on host, +so the failure is only at device-JIT time. This is the wall. + +### M1.7b.6 (`5b1a490`) — plain-SYCL scalar FA (fallback, shipped briefly) + +Aila-style, one WG per (q-head, query-token), LWS=128, scalar QK dot (no +DPAS). 230 tok/s — **slower than the 3-kernel split** (439 at the time, 616 +now). Fusion without DPAS is a loss: the scalar HD=256 dot is too slow. + +### M1.7b.7 (`d9fd3df`) — 3-kernel split (shipped, current path) + +Decomposed FA into per-K-tile pipeline: dequant K → QK GEMM +(`gemm_tf32_batched`, DPAS) → softmax merge → dequant V → PV GEMM +(`gemm_tf32_batched`) → finalize. **439 tok/s then, 616 now** (after the +naive DeltaNet + GEMM tuning wins). This is the current code in +`src/qxmx_fa.cpp`. It works, it's correct, it uses DPAS — but at the cost +of per-tile launch overhead and S/P materialization that a fused kernel +would avoid. + +--- + +## 3. The working DPAS primitive (use this) + +The DPAS path that DOES work is `xmx::dpas` from ESIMD. See +`src/qxmx_gemm_tf32.cpp:84`: + +```cpp +q.parallel_for(sycl::nd_range<1>(...), + [=](sycl::nd_item<1> it) [[intel::sycl_explicit_simd]] { + esimd::simd acc(0.f); + for (int k = 0; k < K8; k++) { + esimd::simd b; /* load */ + esimd::simd a; /* load */ + acc = xmx::dpas(acc, b, a); + } + /* epilogue: write acc */ + }); +``` + +- `xmx::dpas(acc, b, a)` — ESIMD DPAS, tf32 operands. +- Requires `[[intel::sycl_explicit_simd]]` on the kernel. +- This runs at **4.2 TFLOPs** in `gemm_tf32_batched` (small GEMM, M=256); + the M1.2 spike `gemm_gpu_v4` (ternary u2 DPAS) reaches 52-71 TFLOPs. +- **The constraint that blocks fusion:** an ESIMD kernel cannot use + `sycl::exp`, `reduce_over_group`, or sub_group shuffles. So a fused FA + that wants DPAS for QK/PV *and* subgroup softmax in the same kernel is + impossible via this path. That is exactly why M1.7b.5 tried inline-asm + DPAS in plain SYCL — to escape the ESIMD restriction. + +`gemm_gpu_v4` (src/qxmx_gemm.cpp) is the **ternary** u2 DPAS path for the +FFN/attn projections (weight GEMMs). It is irrelevant to FA's QK/PV which +are fp32×fp32 (tf32) activations, not ternary weights. + +--- + +## 4. What to try (the unexplored paths, ranked) + +### Path A — `sycl::joint_matrix` (plain SYCL DPAS, the untried middle) + +`sycl::joint_matrix` +is the SYCL 2020 standard matrix-multiply API. It issues DPAS **without** +`[[intel::sycl_explicit_simd]]` and **without** inline asm — plain SYCL, +works in the same kernel as `sycl::exp`, `reduce_over_group`, sub_group +shuffles. This is exactly the property M1.7b.5 was reaching for via inline +asm. + +**Reference implementations in the tree:** +- `src/aila/src/ops/AttentionOps.cpp` — `attention_decode_joint_matrix_tiled` + (lines 323-455) uses `joint_matrix` + + `joint_matrix_mad`. Read this first; it is a working DPAS attention + path in plain SYCL on this exact hardware family. +- `src/sycl-tla/applications/flash_attention_v2/` — Intel's BMG-tuned FA-2. + See `kernel/legacy/xe_flash_attn_prefill.hpp` and + `collective/legacy/xe_flash_attn_prefill_mma.hpp`. This is the reference + FA-2 prefill kernel for Battlemage. The `examples/06_bmg_flash_attention/` + directory has runnable harnesses (not currently built; `06_xe_fmha_fwd.cpp` + is the prefill entry). + +**The plan:** write a fused FA kernel using `joint_matrix` for QK and PV, +with softmax as plain-SYCL subgroup ops in the same kernel. One WG per +(q-head, Q-tile). Tile sizes to crib from sycl-tla's BMG FA config (see +`build-bmg/examples/06_bmg_flash_attention/` ctest names for the hdim +variants). Validate against the M1.7b.7 3-kernel split as the correctness +oracle (it is already in tree, well-tested: `fa_kernel_test` 15/15, +`attn_forward_batch_test` 5/5, greedy ' Paris' unchanged). + +**Why this is the best first try:** it is the path M1.7b.5 was trying to +synth via inline asm, now available as a supported API. aila + sycl-tla +give concrete crib-able reference code. The only reason it wasn't tried +before is that M1.7b.5 went down the inline-asm dead end first. + +**Risk:** `joint_matrix` tile support is device-dependent. aila's code +falls back to a baseline kernel when "no compatible joint_matrix tile on +this device" (see AttentionOps.cpp:567-598). Verify the B70 (device 0xE223) +exposes a compatible tf32 or bf16 tile for the shapes we need (TM×TK×TN). +The aila `attention_decode_joint_matrix` path runs with TM=32, TK=32, +TN=16, SG=16 — check those work on B70 before committing to a tile config. + +### Path B — ESIMD DPAS with an ESIMD-native softmax + +Stay in ESIMD (the working `xmx::dpas` path) but replace the forbidden +`sycl::exp`/`reduce_over_group` with ESIMD equivalents: +- `esimd::math` has `exp`/`log2` (ESIMD-native, no `sycl::exp` needed). +- Subgroup reductions in ESIMD: `esimd::reduce` + `esimd::merge` + lane + shuffles via `esimd::shuffle` (cf. sycl-tla `include/cute/atom/...`). +- Online softmax is per-row (one row of S per lane), so the "reduce" is + within a subgroup, not across the WG — ESIMD `reduce_over_sub_group` if + it exists, or a manual 5-step shuffle tree. + +Reference: sycl-tla's non-legacy FA kernels +(`applications/flash_attention_v2/collective/`) are ESIMD + ESIMD-native +softmax. This is the harder path (more ESIMD idiom to learn) but avoids +the `joint_matrix` device-tile-support question. + +### Path C — fix the inline-asm dialect + +The M1.7b.5 asm failed with "parsing vISA inline assembly failed: syntax +error, unexpected IDENT" at line 775. The likely cause: the asm dialect +in `qxmx_dpas.h` is from an older IGC/vISA; oneAPI 2026.x's vISA may have +changed the syntax (or the `.decl`/`dpas` token forms). To revive this +path: find a working inline-asm DPAS example that compiles+JITs on +oneAPI 2026.x — search `intel/llvm` sycl_ext_intel_esimd.md DPAS API +section and the IGC DPAS.md for the current vISA dialect. Compare +verbatim against `qxmx_dpas.h`'s asm. This is the riskiest path (vISA +asm is brittle, undocumented surface, easy to get wrong); only attempt +if A and B both fail. + +--- + +## 5. Constraints + gates (do not violate) + +### Hard constraints (architecture, see qxmx_sycl_gotchas.md) +- **In-order `sycl::queue`.** Never `.wait()` between device kernels + except the final prefill sync and the `phase()` profiling waits. +- **USM shared for GPU-RMW buffers; malloc_device for GPU-only.** The + existing `fa_split_workspace` (src/qxmx_gpu.h:42) is all `malloc_shared`. +- **No per-call malloc.** Workspace is an engine member, init-allocated + for CHUNK + TILE_T (see `fa_split_alloc` in qxmx_gpu.h:55). Any new + workspace goes there. +- **M must be a multiple of 16** (gemm_gpu_v4 N-tile constraint; engine + chunks are multiples of 16, with a <16 tail falling back to decode). + +### Correctness gate (unchanged from M1.7) +- `fa_kernel_test` 15/15 (tests/qxmx_fa... — actually the current test is + `attn_forward_batch_test` + `fa_kernel_test`; check meson.build). +- Greedy output unchanged: `qxmx_diff "The capital of France is"` + → ' Paris' (mean|diff| = 0.108 baseline, unchanged). +- Reassociation drift OK in the GEMM epilogue; **greedy tokens must match.** +- Causal: query at pos `qpos` attends only to keys `<= qpos`. +- T13 rotated-domain V composition (validated in `fa_t13_test`): accumulate + P·V_rotated_scaled, butterfly once per output row, ×NRM ×vnorm_corr at + the end. Do NOT butterfly per V-block. See memory + `qxmx_m1_7b_flash_attention.md` for the T13 math. + +### Perf gate +- **Prefill ≥ 700 tok/s** on `bench/session_4mb.txt` (5289 tokens). Current + is 616; the 3-kernel split's per-tile overhead is the ~100 ms/chunk we're + chasing. A fused kernel that matches llama.cpp's per-tile efficiency + should clear 700 and ideally reach toward 900. +- **Attn must NOT grow with context** at fixed chunk (the rescan + symptom). At chunk=256, chunk-20 attn should be within ~2× of chunk-1 + attn (currently 7× — 172 vs 25 ms). This is the structural test. +- Cross-check microbench against in-engine `QXMX_PROFILE=1 QXMX_FFN_DEBUG=1` + numbers. **Never benchmark with trivial-zero data** (see memory + `qxmx_gemm_value_dependence.md` — a zero A+B+1.0-scales GEMM runs 65% + faster than real values on Xe2 due to FMA dependency-chain stalls). + +### Do not waste time on +- **Incremental FA carrying (m,l,O) state across chunks.** The prior-context + QK/PV GEMMs are necessary per-query work; carrying state doesn't avoid + them. The no-rescan probe's 754 tok/s is an unachievable ceiling. See + §1.3 above + memory `qxmx_m1_7b_flash_attention.md`. +- **Chunk size >256 as a primary lever.** Attention's in-chunk O(N²) grows + faster than FFN's amortization win. Chunk size only helps AFTER fusion + lands (a fused FA scales because the per-tile launch overhead is gone). +- **Split-K / A-load prefetch on the FFN GEMMs.** Both measured NO-GO + (2026-07-19, `gemm_sk_test`). FFN is at its kernel-design ceiling; not + the lever. See memory `qxmx_gemm_kernel_facts.md`. +- **Re-doing the M1.7b.5 inline-asm approach verbatim.** It failed at JIT + with a specific vISA syntax error. If attempting Path C, first identify + what changed in the vISA dialect — do not copy `qxmx_dpas.h` and + expect it to work. + +--- + +## 6. Reference file map (all in /home/clee/src/qxmx unless noted) + +| what | path | +|------|------| +| current FA (3-kernel split, shipped) | `src/qxmx_fa.cpp`, `src/qxmx_fa.h` | +| FA workspace (init-allocated) | `src/qxmx_gpu.h:42-73` | +| FA call site (attn_forward_b) | `src/qxmx_gpu.cpp:771-815` (grep `gpu_flash_attn`) | +| engine prefill chunk loop | `src/qxmx_gpu.cpp:1240-1280` (`engine::prefill`, `run_chunk_b`) | +| working tf32 DPAS primitive | `src/qxmx_gemm_tf32.cpp` (`gemm_tf32_batched`, `xmx::dpas`) | +| working ternary u2 DPAS (FFN) | `src/qxmx_gemm.cpp` (`gemm_gpu_v4`, `gemm_gpu_v4_dual`) | +| M1.7b.6 plain-SYCL scalar FA (fallback) | `git show 5b1a490:src/qxmx_fa.cpp` | +| M1.7b.5 inline-asm DPAS (BLOCKED) | `git show ee09d1c^:src/qxmx_dpas.h`, `tests/dpas_asm_probe.cpp` | +| M1.7b.1 fused scalar FA (original) | `git show bcef29a:qxmx_fa.cpp` (pre-`src/` reorg) | +| **aila joint_matrix attention reference** | `/home/clee/src/aila/src/ops/AttentionOps.cpp:323-455` | +| **sycl-tla BMG FA-2 reference** | `/home/clee/src/sycl-tla/applications/flash_attention_v2/` | +| sycl-tla FA runnable harness | `/home/clee/src/sycl-tla/examples/06_bmg_flash_attention/06_xe_fmha_fwd.cpp` | +| KV cache layout (K q8_0/fp8/fp16, V TurboQuant 4-bit WHT) | `src/qxmx_gpu.h:75+` (grep `KROWBYTES`, `VROW_U32`, `VBLK`) | +| model dims | `src/qxmx.h` (D_MODEL 5120, N_HEAD 24, N_HEAD_KV 4, HEAD_DIM 256, ROT_DIM 64) | +| FA test harnesses | `tests/attn_forward_batch_test.cpp`, `tests/fa_kernel_test.cpp` (meson targets) | + +### Memory files (in ~/.config/maki/, project-scoped) +- `qxmx_m1_7b_flash_attention.md` — this problem's analysis + the + no-rescan-probe-is-misleading finding. +- `qxmx_gemm_kernel_facts.md` — the FFN GEMM levers + the dead ones. +- `qxmx_sycl_gotchas.md` — in-order queue, USM kinds, the malloc_host RMW bug. +- `qxmx_gemm_value_dependence.md` — the trivial-zero-data microbench trap. +- `qxmx_b70_hardware_facts.md` — 256 EU, 128 KB SLM/WG, subgroups {16,32}. + +### Commands +```bash +source /opt/intel/oneapi/setvars.sh >/dev/null 2>&1 +meson setup --reconfigure build >/dev/null 2>&1 && meson compile -C build +M=~/models/bonsai/Ternary-Bonsai-27B-Q2_g64.gguf + +./build/qxmx_diff "$M" "The capital of France is" # decode gate (0.108, ' Paris') +QXMX_PROFILE=1 ./build/qxmx_run "$M" -f bench/session_4mb.txt -n 1 # prefill tok/s + phase breakdown +QXMX_PROFILE=1 QXMX_FFN_DEBUG=1 ./build/qxmx_run "$M" -f bench/session_4mb.txt -n 1 # + per-layer ffn +./build/attn_forward_batch_test # FA bit-exactness vs decode oracle (5/5) +``` +**icpx gotcha:** for any ad-hoc icpx compile outside meson, pass +`--gcc-install-dir=/usr/lib/gcc/x86_64-linux-gnu/14` (see +`shell_gotchas.md`). Inside meson it's handled by `meson/icpx.ini`. + +--- + +## 7. The one-sentence summary + +The gap to llama.cpp is not algorithmic — it is that our FA is a 3-kernel +split with per-tile launch overhead and small-tile GEMM inefficiency, +while theirs is a fused kernel that keeps S/P in registers; the fix is a +fused tiled FA using `sycl::joint_matrix` (Path A, untried, with aila + +sycl-tla as reference) or ESIMD-native softmax (Path B), and the inline-asm +DPAS path (M1.7b.5, Path C) is blocked by a vISA JIT syntax error that +needs a fresh dialect investigation before it's worth re-attempting. \ No newline at end of file diff --git a/meson.build b/meson.build index 1d368fd..9018090 100644 --- a/meson.build +++ b/meson.build @@ -70,6 +70,7 @@ libmodel = static_library('qxmx_model', src/'qxmx_model.c', c_args : c_warn) # engine_sources_naive = common + deltanet_common + deltanet_naive. engine_sources_common = files( 'src/qxmx_gpu.cpp', 'src/qxmx_gemm.cpp', 'src/qxmx_gemm_tf32.cpp', 'src/qxmx_fa.cpp', + 'src/qxmx_fa_fused.cpp', 'src/qxmx_ref.cpp', 'src/qxmx_tokenizer.cpp', 'third_party/unicode.cpp', 'third_party/unicode-data.cpp', ) @@ -245,4 +246,17 @@ executable('dpas_tf32_probe', sources : tests/'dpas_tf32_probe.cpp', cpp_args : sycl_args, link_args : sycl_args, + include_directories : inc) + +# M1.7c fused FA microbench + phase isolation (QXMX_FA_PHASES). +executable('fa_fused_bench', + sources : [tests/'fa_fused_bench.cpp'] + engine_sources_naive, + kwargs : test_engine_args) + +# joint_matrix probe (SYCL): enumerate ext_oneapi_matrix tile combos on B70 + +# validate joint_matrix_mad bf16/tf32 (Path A enabler, docs/fused_fa_problem.md). +executable('jm_probe', + sources : tests/'jm_probe.cpp', + cpp_args : sycl_args, + link_args : sycl_args, include_directories : inc) \ No newline at end of file diff --git a/plan.md b/plan.md index 5321a04..4f5fbc0 100644 --- a/plan.md +++ b/plan.md @@ -411,20 +411,27 @@ are flat. FFN is 48% of chunk-1 — the largest fixed share, and it's GEMM-bound sweep harness — it uses REAL A+B (see the trap below). **Next levers (ranked by expected payoff ÷ effort):** -1. **FFN GEMM kernel rewrite** — the big remaining lever (121 ms/chunk, 48% - of chunk-1). Tile-param tuning is tapped out (see "Tried + abandoned"). - The win requires a different kernel: - - (a) SLM-staged B with double-buffering — B is 1.31–4.46 MB, fits in SLM - per WG; current kernel reloads B from L2/HBM each K-group pass. - - (b) Split-K for the down-proj (K=17408 is large; split-K across WGs + - reduction). - - (c) Fused gate+up+swiglu — avoid materializing the 2× 17.8 MB g/u to - HBM; keep in registers/SLM, write only the swiglu output t. - Each is a real kernel rewrite, moderate-high effort. Profile with - `QXMX_PROFILE=1 QXMX_FFN_DEBUG=1` to confirm the gain is in the GEMM. -2. **Attention** (9% at chunk-1, grows to ~22% by chunk 20). The M1.7b.7 - 3-kernel split's per-tile launch overhead. A fused tiled kernel would help - but hits the M1.7b.5 inline-asm/ESIMD wall. Moderate effort, context-scaling. +1. ~~**FFN GEMM kernel rewrite**~~ — ALL THREE sub-levers measured DEAD + (2026-07-19; full record in memory `qxmx_gemm_kernel_facts.md`): + - (a) ~~SLM-staged B~~ — mislabeled; B is already SLM-staged + (qxmx_gemm.cpp:67-78). The unstaged operand is A; tested as (a'). + - (a') ~~A-load prefetch / double-buffer~~ — NO-GO. Bit-identical, perf + neutral (0.831 vs 0.820 ms). ESIMD `copy_from` with L2-cached hints is + already HW/compiler-pipelined; XMX throughput is the ceiling. + - (b) ~~Split-K for the down-proj~~ — NO-GO. NS=2/4/8 all neutral-to-worse + (1.175/0.823/0.827 vs 0.820 ms). DPAS-throughput-bound, not occupancy- + bound, despite 80 WGs / 256 EUs. + - (c) ~~Fused gate+up+swiglu~~ — half already done (`gemm_gpu_v4_dual` + fuses gate+up); remaining swiglu fusion saves ~0.17 ms/layer, small. + Both FFN GEMMs are at the kernel design's ceiling (gate+up 102 TFLOPs + combined, down 95 TFLOPs = 36-38% of the 133 T-MAC bar). Further gains + need a fundamentally different algorithm — outside MVP scope. +2. **Attention** — RESOLVED by M1.7c (commit a4eb2a9): fused FA via + `sycl::joint_matrix` (the Path A from docs/fused_fa_problem.md, NOT the + M1.7b.5 inline-asm wall that blocked earlier). Prefill 616→741 tok/s. + See memory `qxmx_m1_7c_fused_fa_landed.md`. Remaining sub-levers listed + there (DPAS throughput 10.4→20+ TFLOPs, cross-chunk shadow re-dequant, + GQA batching). 3. **SSM naive further** (1.1 ms/layer, ~36% of chunk-1). Warp-per-column restructure (16 lanes × 8 rows, `reduce_over_group` within the warp) is the remaining SSM lever. Moderate effort, uncertain payoff (the 128-row @@ -433,12 +440,21 @@ are flat. FFN is 48% of chunk-1 — the largest fixed share, and it's GEMM-bound **Tried + abandoned (do not revisit without new evidence):** - Chunk size >256 (sweep: 256→598, 512→597, 1024→577 — neutral-to-negative; naive recurrence is O(chunk×HD²) per head, sequential in-chunk, no amortization). + NOTE: measured pre-M1.7c; post-fusion FFN's 8% amortization at 512 survives + and chunk=512 is the new default (741 tok/s). - L2-norm hoist to a pre-kernel (net neutral; the 3× GQA redundancy isn't wasteful enough to pay for a kernel launch). - Per-shape GEMM dispatch (real-value sweep: per-shape best is only 0.3 ms/chunk better than one global 32/4/8 + 32/4/6 — not worth the code). - GEMM tile-param tuning beyond 32/4/8 + 32/4/6 (all 11 dispatch-table configs swept across all 11 engine shapes under real load; 32/4/8+32/4/6 is optimal). +- Split-K down-proj (NO-GO: NS=2/4/8 neutral-to-worse vs oracle; DPAS- + throughput-bound not occupancy-bound despite 0.31 WGs/EU). +- A-load prefetch / double-buffer (NO-GO: bit-identical, perf neutral — ESIMD + copy_from already HW/compiler-pipelined, A-load not the limiter). +- Inline-asm DPAS in plain SYCL (M1.7b.5): vISA JIT rejects the asm dialect + ("parsing vISA inline assembly failed: syntax error, unexpected IDENT"). + Obsoleted by M1.7c's `sycl::joint_matrix` path — do NOT revive inline asm. ### `gemm_gpu_v4` sharp edges (actionable for anyone touching the kernel) 1. **A-operand loads are unguarded.** The kernel loads A_codes/A_sc_t for a diff --git a/src/qxmx_fa.cpp b/src/qxmx_fa.cpp index 795c008..bf84350 100644 --- a/src/qxmx_fa.cpp +++ b/src/qxmx_fa.cpp @@ -31,6 +31,7 @@ #include "qxmx_fa.h" #include +#include #include @@ -296,12 +297,23 @@ static void fa_finalize(sycl::queue& q, float* out, int n_tok, }); } -/* ---- Entry: orchestrate the per-tile loop ------------------------------- */ +/* ---- Entry: fused kernel by default; QXMX_FA_SPLIT=1 keeps the M1.7b.7 + * 3-kernel split as an A/B oracle. The fused path fills Oacc + l_state and + * reuses this file's fa_finalize for /l + T13 inverse WHT. */ template void gpu_flash_attn(sycl::queue& q, const float* Q, const uint8_t* kv_k, const uint32_t* kv_v, float* out, int64_t pos_start, int n_tok, fa_split_workspace* ws) { + static const bool use_split = [] { + const char* e = std::getenv("QXMX_FA_SPLIT"); + return e && std::atoi(e) != 0; + }(); + if (!use_split) { + gpu_flash_attn_fused(q, Q, kv_k, kv_v, pos_start, n_tok, ws); + fa_finalize(q, out, n_tok, ws); + return; + } const float scale = 1.0f / sycl::sqrt((float)HD); /* Prep Q + init state. */ fa_prep_q(q, Q, n_tok, ws); diff --git a/src/qxmx_fa.h b/src/qxmx_fa.h index e5381b8..7b2281d 100644 --- a/src/qxmx_fa.h +++ b/src/qxmx_fa.h @@ -62,5 +62,15 @@ void gpu_flash_attn(sycl::queue& q, const float* Q, float* out, int64_t pos_start, int n_tok, fa_split_workspace* ws); +/* M1.7c fused FA-2 kernel (qxmx_fa_fused.cpp): one launch per chunk, fp16 + * joint_matrix DPAS, S/P/O in SLM. Fills ws->Oacc (unnormalized, rotated + * domain) and ws->l_state; the caller runs fa_finalize for /l + T13 WHT. + * Same cache/causal contract as gpu_flash_attn; n_tok need only be >= 1. */ +template +void gpu_flash_attn_fused(sycl::queue& q, const float* Q, + const uint8_t* kv_k, const uint32_t* kv_v, + int64_t pos_start, int n_tok, + fa_split_workspace* ws); + } // namespace qx #endif \ No newline at end of file diff --git a/src/qxmx_fa_fused.cpp b/src/qxmx_fa_fused.cpp new file mode 100644 index 0000000..c01247c --- /dev/null +++ b/src/qxmx_fa_fused.cpp @@ -0,0 +1,370 @@ +/* qxmx_fa_fused.cpp: M1.7c fused FlashAttention-2 prefill kernel. + * + * One WG per (q-head, Q-tile of TQ=32 rows), 8 sub_groups of 16 (LWS=128). + * The whole FA pass for a chunk is 3 launches (shadow dequant, fused FA, + * fa_finalize in qxmx_fa.cpp) -- no per-tile GEMM launches, no S/P global + * materialization (the M1.7b.7 split's overheads). + * + * DPAS via sycl::joint_matrix fp16 16x16x16 SG16 -- the ONLY tile family + * that JITs on B70 (0xE223); see memory qxmx_joint_matrix_b70.md. Plain + * SYCL kernel (NOT ESIMD), so sycl::native::exp + reduce_over_group + + * sub_group ops coexist with the matrix ops in the same kernel -- the + * property the blocked M1.7b.5 inline-asm path was reaching for. + * + * v2 (2026-07-19): the in-kernel K/V dequant was 76% of kernel time + * (fa_fused_bench QXMX_FA_PHASES isolation: 20.3 ms of 26.6 ms at + * pos_start=4864) -- byte-granular cache reads x 48-fold redundant across + * the (head, q-tile) WGs. Now K/V are dequantized ONCE per chunk into a + * global fp16 shadow (Kh/Vh [total][KV_DIM]) by fa_dequant_kv_shadow, and + * the fused kernel's joint_matrix B-loads read the shadow directly from + * global (L2-hot; each K/V tile is re-read by 48 WGs). No Ks/Vs SLM: + * SLM use drops to 60 KB/WG (Qs 16K, Sf 8K, Pf 4K, Of 32K). + * + * fp16 operands: Q/K/V are O(1) post-norm values (K cache is fp16-native + * for KCT::fp16; q8_0/fp8 dequant -> fp16 loses nothing meaningful: fp16 + * has the same 10-bit mantissa the tf32 GEMM path truncates to). + * + * Per-WG dataflow per K-tile (TK=64 keys), 3 barriers: + * 1. QK: SG s -> S-tile (16q x 16k), 16 mads over HD. A from SLM Qs, + * B from GLOBAL Kh (col_major, stride KV_DIM). Store Sf (fp32). + * 2. softmax+rescale: SG s -> 4 rows; causal mask, online (m,l) merge, + * Pf (fp16); rescale the SG's own Of rows by alpha inline (each SG + * owns its rows -- no cross-SG alpha exchange, no separate phase). + * 3. PV: SG s -> 4 O-tiles (16x16), accumulator load from Of, A from + * SLM Pf, B from GLOBAL Vh (row_major), mad, store back. + * Epilogue: Of -> ws->Oacc, l -> ws->l_state (valid rows only; the caller + * runs fa_finalize for /l + T13 inverse WHT). + * + * Correctness contract identical to gpu_flash_attn (qxmx_fa.cpp): Q is + * [n_tok][NQ*HD] fp32, kv_k/kv_v the caches with history at [0..pos_start), + * causal cap per query, V in TurboQuant rotated domain. n_tok need only be + * >= 1 -- the final q-tile is padded (clamped Q rows, masked writes), so + * the engine's %16 mid-chunks work unchanged. + */ +#include "qxmx_fa.h" + +#include +#include + +#include + +#include "qxmx_fp8.h" + +namespace qx { + +namespace jm = sycl::ext::oneapi::experimental::matrix; + +static constexpr int HD = QX_HEAD_DIM; /* 256 */ +static constexpr int NQ = QX_N_HEAD; /* 24 */ +static constexpr int NKV = QX_N_HEAD_KV; /* 4 */ +static constexpr int GQA = NQ / NKV; /* 6 */ +static constexpr int KV_DIM = QX_KV_DIM; /* 1024 */ +static constexpr int BLK_PER_HEAD = HD / VBLK; /* 8 */ +static constexpr int TQ = 32; /* Q rows per WG tile */ +static constexpr int TK = 64; /* K keys per streamed tile */ +static constexpr int NSG = 8; /* sub_groups per WG */ +static constexpr int SG = 16; /* sub_group size */ +static constexpr int LWS = NSG * SG; /* 128 */ +static constexpr int JT = 16; /* joint_matrix tile edge */ + +using fp16 = sycl::half; +using jm_a = jm::joint_matrix; +using jm_b_row = jm::joint_matrix; +using jm_b_col = jm::joint_matrix; +using jm_c = jm::joint_matrix; + +/* Dequant one K element (kglob, d) of kv-head hk. Same math as + * fa_dequant_kt (qxmx_fa.cpp), untransposed. */ +template +static inline float fa_fused_dequant_k(const uint8_t* kv_k, int64_t kglob, + int hk, int d) { + if constexpr (std::is_same_v) { + const fp16* kp = reinterpret_cast(kv_k); + return float(kp[(size_t)kglob * KV_DIM + (size_t)hk * HD + d]); + } else { + const uint8_t* kb = kv_k + (size_t)kglob * Kct::row_bytes + + (size_t)(hk * BLK_PER_HEAD + (d >> 5)) * KBLK_BYTES; + uint16_t sbits = (uint16_t)kb[0] | ((uint16_t)kb[1] << 8); + float dd = float(sycl::bit_cast(sbits)); + uint8_t code = kb[2 + (d & 31)]; + if constexpr (std::is_same_v) + return (float)(int8_t)code * dd; + else + return e4m3_to_fp32(code) * dd; + } +} + +/* Dequant one V element (kglob, d), T13 rotated domain (same math as + * fa_dequant_vt): LUT[code] * vscale * vnorm_corr. */ +static inline float fa_fused_dequant_v(const uint32_t* kv_v, int64_t kglob, + int hk, int d) { + int blk = hk * BLK_PER_HEAD + (d >> 5); + int vlane = d & 31; + const uint32_t* db = kv_v + (size_t)kglob * VROW_U32 + blk * 5; + uint16_t sc_u = (uint16_t)(db[0] & 0xFFFFu); + uint16_t nc_u = (uint16_t)((db[0] >> 16) & 0xFFFFu); + float vscale = float(sycl::bit_cast(sc_u)); + float vnorm_corr = float(sycl::bit_cast(nc_u)); + uint32_t c0 = db[1], c1 = db[2], c2 = db[3], c3 = db[4]; + int w = vlane >> 3; + int off = (vlane & 7) * 4; + uint32_t word = (w == 0) ? c0 : (w == 1) ? c1 : (w == 2) ? c2 : c3; + uint32_t code = (word >> off) & 0xFu; + return V_LM_LUT[code] * vscale * vnorm_corr; +} + +/* Grow the fp16 shadow to hold `total` tokens. The one q.wait() before the + * free is the only pipeline drain; growth is O(log) per prompt. */ +static void fa_shadow_ensure(sycl::queue& q, fa_split_workspace* ws, + int64_t total) { + if (total <= ws->shadow_tok) return; + int64_t cap = ws->shadow_tok > 0 ? ws->shadow_tok : 8192; + while (cap < total) cap *= 2; + q.wait(); + if (ws->Kh) sycl::free(ws->Kh, q); + if (ws->Vh) sycl::free(ws->Vh, q); + ws->Kh = sycl::malloc_shared((size_t)cap * KV_DIM, q); + ws->Vh = sycl::malloc_shared((size_t)cap * KV_DIM, q); + ws->shadow_tok = cap; +} + +/* Dequant the visible K/V range [0, roundup(total,TK)) into the fp16 shadow. + * One launch per chunk; ~0.3 ms at total=5120 (bandwidth-bound). Rows + * [total, roundup) are zero-filled so the fused kernel's PV never sees stale + * bits (0 x stale = NaN if the stale pattern is inf/NaN). */ +template +static void fa_dequant_kv_shadow(sycl::queue& q, const uint8_t* kv_k, + const uint32_t* kv_v, int64_t total, + fa_split_workspace* ws) { + fp16* Kh = ws->Kh; + fp16* Vh = ws->Vh; + const int64_t padded = (total + TK - 1) / TK * TK; + const size_t n = (size_t)padded * KV_DIM; + q.parallel_for(sycl::nd_range<1>((n + 255) / 256 * 256, 256), + [=](sycl::nd_item<1> it) { + size_t idx = it.get_global_id(0); + if (idx >= n) return; + int64_t t = (int64_t)(idx / KV_DIM); + int d = (int)(idx % KV_DIM); + int hk = d / HD, dd = d % HD; + float kval = 0.0f, vval = 0.0f; + if (t < total) { + kval = fa_fused_dequant_k(kv_k, t, hk, dd); + vval = fa_fused_dequant_v(kv_v, t, hk, dd); + } + Kh[idx] = fp16(kval); + Vh[idx] = fp16(vval); + }); +} + +template +void gpu_flash_attn_fused(sycl::queue& q, const float* Q, + const uint8_t* kv_k, const uint32_t* kv_v, + int64_t pos_start, int n_tok, + fa_split_workspace* ws) { + const float scale = 1.0f / sycl::sqrt((float)HD); + float* Oacc = ws->Oacc; + float* l_state = ws->l_state; + const int n_qtiles = (n_tok + TQ - 1) / TQ; + const int64_t total = pos_start + n_tok; /* keys 0..total-1 visible */ + /* QXMX_FA_PHASES: profiling bitmask (bit1 QK, 2 softmax, 4 PV; bit0 is + * the shadow dequant launch). Default 31 = all. Skipped phases leave + * garbage -- timing probes only, results wrong. */ + const char* phe = std::getenv("QXMX_FA_PHASES"); + const unsigned phases = phe ? (unsigned)std::atoi(phe) : 31u; + + fa_shadow_ensure(q, ws, total); + if (phases & 1u) fa_dequant_kv_shadow(q, kv_k, kv_v, total, ws); + const fp16* Kh = ws->Kh; + const fp16* Vh = ws->Vh; + + q.submit([&](sycl::handler& h) { + sycl::local_accessor Qs(sycl::range<1>(TQ * HD), h); + sycl::local_accessor Sf(sycl::range<1>(TQ * TK), h); + sycl::local_accessor Pf(sycl::range<1>(TQ * TK), h); + sycl::local_accessor Of(sycl::range<1>(TQ * HD), h); + sycl::local_accessor mSt(sycl::range<1>(TQ), h); + sycl::local_accessor lSt(sycl::range<1>(TQ), h); + + h.parallel_for( + sycl::nd_range<1>((size_t)NQ * n_qtiles * LWS, LWS), + [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(SG)]] { + const int gid = it.get_group(0); + const int hh = gid / n_qtiles; + const int q0 = (gid % n_qtiles) * TQ; + const int hk = hh / GQA; + const int lid = it.get_local_id(0); + const int sg_id = lid / SG; + const int lane = lid % SG; + auto sg = it.get_sub_group(); + auto grp = it.get_group(); + + auto q_p = Qs.get_multi_ptr(); + auto s_p = Sf.get_multi_ptr(); + auto p_p = Pf.get_multi_ptr(); + auto o_p = Of.get_multi_ptr(); + + /* Stage Q (fp32 global -> fp16 SLM), clamp padding rows; + * init (m, l, O). */ + for (int idx = lid; idx < TQ * HD; idx += LWS) { + int r = idx / HD, d = idx % HD; + int qg = q0 + r; + if (qg >= n_tok) qg = n_tok - 1; + Qs[idx] = fp16(Q[(size_t)qg * (NQ * HD) + + (size_t)hh * HD + d]); + Of[idx] = 0.0f; + } + if (lid < TQ) { mSt[lid] = -1e30f; lSt[lid] = 0.0f; } + sycl::group_barrier(grp); + + jm_a sub_a; + jm_b_col sub_b_col; + jm_b_row sub_b_row; + jm_c sub_c; + + for (int64_t t0 = 0; t0 < total; t0 += TK) { + const int tile_len = (int)sycl::min((int64_t)TK, total - t0); + + /* 1. QK: SG s -> S-tile (rb=s/4, kband=s%4), 16 mads. + * B col_major from global Kh: element (k,n) = + * Kh[(t0+kband+n)*KV_DIM + hk*HD + ks + k]. */ + if (phases & 2u) + { + const int rb = (sg_id >> 2) * JT; + const int64_t krow = t0 + (sg_id & 3) * JT; + jm::joint_matrix_fill(sg, sub_c, 0.0f); + for (int ks = 0; ks < HD; ks += JT) { + jm::joint_matrix_load(sg, sub_a, + q_p + rb * HD + ks, HD); + auto b_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>( + Kh + krow * KV_DIM + hk * HD + ks); + jm::joint_matrix_load(sg, sub_b_col, b_ptr, + KV_DIM); + jm::joint_matrix_mad(sg, sub_c, sub_a, + sub_b_col, sub_c); + } + jm::joint_matrix_store(sg, sub_c, + s_p + rb * TK + (sg_id & 3) * JT, + TK, jm::layout::row_major); + } + sycl::group_barrier(grp); + + /* 2. softmax + Of rescale: SG s -> rows [s*4, s*4+4). + * 4 lanes per row, 16 consecutive cols each; quad + * xor-shuffle reductions (2+2 steps) instead of per-row + * reduce_over_group (~10 steps x 4 rows). */ + if (phases & 4u) + { + constexpr int RPS = TQ / NSG; /* 4 rows/SG */ + const int r = sg_id * RPS + (lane >> 2); + const int cq = (lane & 3) * JT; /* col base */ + const int64_t qpos = pos_start + q0 + r; + float sv[JT]; + float part_max = -1e30f; + for (int j = 0; j < JT; j++) { + int c = cq + j; + bool valid = (c < tile_len) + && (t0 + c <= qpos); + sv[j] = valid ? Sf[r * TK + c] * scale + : -1e30f; + part_max = sycl::fmax(part_max, sv[j]); + } + part_max = sycl::fmax(part_max, + sycl::permute_group_by_xor(sg, part_max, 1)); + part_max = sycl::fmax(part_max, + sycl::permute_group_by_xor(sg, part_max, 2)); + float prev_m = mSt[r]; + float new_m = sycl::fmax(prev_m, part_max); + float alpha_r = sycl::native::exp(prev_m - new_m); + float part_sum = 0.0f; + for (int j = 0; j < JT; j++) { + float e = (sv[j] > -1e29f) + ? sycl::native::exp(sv[j] - new_m) + : 0.0f; + Pf[r * TK + cq + j] = fp16(e); + part_sum += e; + } + part_sum += sycl::permute_group_by_xor(sg, part_sum, 1); + part_sum += sycl::permute_group_by_xor(sg, part_sum, 2); + if ((lane & 3) == 0) { + mSt[r] = new_m; + lSt[r] = alpha_r * lSt[r] + part_sum; + } + /* Rescale Of for this SG's 4 rows. All lanes need + * each row's alpha: broadcast from the quad leaders + * (lanes 0,4,8,12 hold rows 0..3 of this SG). */ + float a_own = alpha_r; + for (int rr = 0; rr < RPS; rr++) { + float a = sycl::group_broadcast(sg, a_own, rr * 4); + int orow = (sg_id * RPS + rr) * HD; + for (int d = lane; d < HD; d += SG) + Of[orow + d] *= a; + } + } + sycl::group_barrier(grp); + + /* 3. PV: SG s -> row-band (s&1), col-tiles (s>>1)*4+j. + * B row_major from global Vh: element (k,n) = + * Vh[(t0+ks+k)*KV_DIM + hk*HD + ct + n]. */ + if (phases & 16u) + { + const int rb = (sg_id & 1) * JT; + for (int j = 0; j < 4; j++) { + const int ct = ((sg_id >> 1) * 4 + j) * JT; + jm::joint_matrix_load(sg, sub_c, + o_p + rb * HD + ct, HD, + jm::layout::row_major); + for (int ks = 0; ks < TK; ks += JT) { + jm::joint_matrix_load(sg, sub_a, + p_p + rb * TK + ks, TK); + auto b_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>( + Vh + (t0 + ks) * KV_DIM + hk * HD + ct); + jm::joint_matrix_load(sg, sub_b_row, b_ptr, + KV_DIM); + jm::joint_matrix_mad(sg, sub_c, sub_a, + sub_b_row, sub_c); + } + jm::joint_matrix_store(sg, sub_c, + o_p + rb * HD + ct, HD, + jm::layout::row_major); + } + } + sycl::group_barrier(grp); + } + + /* Epilogue: Of -> Oacc, l -> l_state (valid rows only). */ + for (int idx = lid; idx < TQ * HD; idx += LWS) { + int r = idx / HD; + int qi = q0 + r; + if (qi < n_tok) + Oacc[(size_t)hh * n_tok * HD + (size_t)qi * HD + + (idx % HD)] = Of[idx]; + } + if (lid < TQ) { + int qi = q0 + lid; + if (qi < n_tok) + l_state[(size_t)hh * n_tok + qi] = lSt[lid]; + } + }); + }); +} + +template void gpu_flash_attn_fused(sycl::queue&, const float*, + const uint8_t*, const uint32_t*, + int64_t, int, fa_split_workspace*); +template void gpu_flash_attn_fused(sycl::queue&, const float*, + const uint8_t*, const uint32_t*, + int64_t, int, fa_split_workspace*); +template void gpu_flash_attn_fused(sycl::queue&, const float*, + const uint8_t*, const uint32_t*, + int64_t, int, fa_split_workspace*); + +} // namespace qx diff --git a/src/qxmx_gpu.cpp b/src/qxmx_gpu.cpp index 5c38faf..b7de25b 100644 --- a/src/qxmx_gpu.cpp +++ b/src/qxmx_gpu.cpp @@ -1086,6 +1086,11 @@ bool engine::init(qx_model* m, int mc) { fa_ws.m_state = sycl::malloc_shared((size_t)QX_N_HEAD * T, q); fa_ws.l_state = sycl::malloc_shared((size_t)QX_N_HEAD * T, q); fa_ws.Qsep = sycl::malloc_shared((size_t)QX_N_HEAD * T * QX_HEAD_DIM, q); + /* M1.7c fused-FA fp16 K/V shadow (grown on demand by + * fa_shadow_ensure). 8192 tokens covers the bench prompt. */ + fa_ws.shadow_tok = 8192; + fa_ws.Kh = sycl::malloc_shared((size_t)fa_ws.shadow_tok * QX_KV_DIM, q); + fa_ws.Vh = sycl::malloc_shared((size_t)fa_ws.shadow_tok * QX_KV_DIM, q); } g_bvnni_d = bvnni_d; g_suma_d = suma_d; g_act_scale_d = act_scale_d; n_tokens = 0; diff --git a/src/qxmx_gpu.h b/src/qxmx_gpu.h index 3742877..00e3263 100644 --- a/src/qxmx_gpu.h +++ b/src/qxmx_gpu.h @@ -45,6 +45,13 @@ struct fa_split_workspace { float *O, *Oacc; /* PV accumulator (rotated domain) + scratch */ float *m_state, *l_state; /* online-softmax max/sum per (q-head, query tok) */ float *Qsep; /* Q transposed to [NQ][n_tok][HD] for batched GEMM */ + /* M1.7c fused-FA fp16 shadow of the dequantized K/V caches, + * [shadow_tok][KV_DIM] each (K fp16-native, V T13 rotated domain). + * Written once per chunk by fa_dequant_kv_shadow; the fused kernel's + * joint_matrix B-loads read it directly from global. Grown by + * fa_shadow_ensure (qxmx_fa_fused.cpp). */ + sycl::half *Kh, *Vh; + int64_t shadow_tok; }; /* Allocate a fa_split_workspace sized for n_tok_max query tokens and the @@ -64,12 +71,16 @@ inline fa_split_workspace fa_split_alloc(sycl::queue& q, int n_tok_max) { w.m_state = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max, q); w.l_state = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max, q); w.Qsep = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, q); + w.shadow_tok = 8192; + w.Kh = sycl::malloc_shared((size_t)w.shadow_tok * QX_KV_DIM, q); + w.Vh = sycl::malloc_shared((size_t)w.shadow_tok * QX_KV_DIM, q); return w; } inline void fa_split_free(sycl::queue& q, fa_split_workspace& w) { sycl::free(w.Kt, q); sycl::free(w.Vt, q); sycl::free(w.S, q); sycl::free(w.P, q); sycl::free(w.O, q); sycl::free(w.Oacc, q); sycl::free(w.m_state, q); sycl::free(w.l_state, q); sycl::free(w.Qsep, q); + sycl::free(w.Kh, q); sycl::free(w.Vh, q); } /* ---- KV cache block layout (shared by decode qxmx_gpu.cpp and prefill @@ -134,7 +145,7 @@ struct dev_ffn_blk { }; struct engine : engine_i { - int chunk = 256; /* runtime prefill chunk size (QXMX_CHUNK; must be %16) */ + int chunk = 512; /* runtime prefill chunk size (QXMX_CHUNK; must be %16). 512: FFN amortizes 8% better than 256 and the M1.7c fused FA keeps attn flat there (741 vs 716 tok/s, 2026-07-19) */ sycl::queue q; int max_ctx = 0; diff --git a/tests/fa_fused_bench.cpp b/tests/fa_fused_bench.cpp new file mode 100644 index 0000000..80fd445 --- /dev/null +++ b/tests/fa_fused_bench.cpp @@ -0,0 +1,88 @@ +/* fa_fused_bench: microbenchmark for the M1.7c fused FA kernel. + * One config: pos_start=4864, n_tok=256 (chunk-20-like), q8_0 K cache. + * Times gpu_flash_attn over ITERS launches (in-order queue, single + * wait at the end). Use with QXMX_FA_PHASES bitmask for phase isolation + * (bit0 dequant, 1 QK, 2 softmax, 3 rescale, 4 PV; 31=all) and + * QXMX_FA_SPLIT=1 for the M1.7b.7 baseline. + * + * Build: meson compile -C build fa_fused_bench && ./build/fa_fused_bench + */ +#include +#include +#include +#include +#include +#include +#include +#include + +#include "qxmx.h" +#include "qxmx_fa.h" +#include "qxmx_gpu.h" + +using namespace qx; + +int main() { + constexpr int64_t pos_start = 4864; + constexpr int n_tok = 256; + constexpr int KV = QX_KV_DIM, NQ = QX_N_HEAD, HD = QX_HEAD_DIM; + constexpr size_t QOUT = (size_t)NQ * HD; + const int max_ctx = (int)(pos_start + n_tok); + constexpr int ITERS = 5; + + sycl::queue q(sycl::gpu_selector_v, sycl::property::queue::in_order{}); + printf("fa_fused_bench: pos_start=%lld n_tok=%d K=q8_0 iters=%d\n", + (long long)pos_start, n_tok, ITERS); + printf(" phases mask (QXMX_FA_PHASES): bit0 shadow dequant, 1 QK, 2 softmax," + " 4 PV\n"); + + /* Non-trivial data (varied scales -- see qxmx_gemm_value_dependence). */ + std::mt19937_64 rng(0xFA17ULL); + std::normal_distribution g(0.0f, 1.0f); + std::vector kf((size_t)max_ctx * KV), vf((size_t)max_ctx * KV), + qf((size_t)n_tok * QOUT); + for (auto& x : kf) x = g(rng); + for (auto& x : vf) x = g(rng); + for (auto& x : qf) x = g(rng); + + float* d_kf = sycl::malloc_shared(kf.size(), q); + float* d_vf = sycl::malloc_shared(vf.size(), q); + float* d_q = sycl::malloc_shared(qf.size(), q); + memcpy(d_kf, kf.data(), kf.size() * 4); + memcpy(d_vf, vf.data(), vf.size() * 4); + memcpy(d_q, qf.data(), qf.size() * 4); + uint8_t* kvk = sycl::malloc_shared((size_t)max_ctx * KROW_BYTES_Q8, q); + uint32_t* kvv = sycl::malloc_shared((size_t)max_ctx * VROW_U32, q); + float* out = sycl::malloc_shared((size_t)n_tok * QOUT, q); + for (int t = 0; t < max_ctx; t++) { + gpu_quant_k_row_q8(q, d_kf + (size_t)t * KV, + kvk + (size_t)t * KROW_BYTES_Q8); + gpu_quant_v_row(q, d_vf + (size_t)t * KV, + kvv + (size_t)t * VROW_U32); + } + q.wait(); + + fa_split_workspace ws = fa_split_alloc(q, n_tok); + /* warmup + JIT */ + gpu_flash_attn(q, d_q, kvk, kvv, out, pos_start, n_tok, &ws); + q.wait(); + auto t0 = std::chrono::steady_clock::now(); + for (int i = 0; i < ITERS; i++) + gpu_flash_attn(q, d_q, kvk, kvv, out, pos_start, n_tok, &ws); + q.wait(); + auto t1 = std::chrono::steady_clock::now(); + double ms = std::chrono::duration(t1 - t0).count() + / ITERS; + int n_tiles = (int)((pos_start + n_tok + 63) / 64); + printf(" fused+finalize: %.3f ms/call (%.3f ms per K-tile x %d tiles)\n", + ms, ms / n_tiles, n_tiles); + /* rough DPAS throughput of QK+PV: 2 * 2 * NQ * n_tok * total * HD */ + double flops = 4.0 * NQ * n_tok * (double)(pos_start + n_tok) * HD; + printf(" effective QK+PV throughput: %.2f TFLOPs\n", + flops / (ms * 1e-3) / 1e12); + + fa_split_free(q, ws); + sycl::free(d_kf, q); sycl::free(d_vf, q); sycl::free(d_q, q); + sycl::free(kvk, q); sycl::free(kvv, q); sycl::free(out, q); + return 0; +} diff --git a/tests/jm_probe.cpp b/tests/jm_probe.cpp new file mode 100644 index 0000000..2c20c35 --- /dev/null +++ b/tests/jm_probe.cpp @@ -0,0 +1,244 @@ +/* jm_probe: enumerate sycl::joint_matrix (ext_oneapi_matrix) support on the + * B70 and validate joint_matrix_mad against a CPU fp32 reference. + * Path A enabler for the fused FA kernel (docs/fused_fa_problem.md §4): + * joint_matrix issues DPAS from plain SYCL (no ESIMD attr, no inline asm), + * so it can share a kernel with sycl::exp / reduce_over_group softmax. + * + * Validates the two load patterns a fused FA needs: + * QK: C[M][N] = A[M][K] @ B^T (use::b col_major over row-major K rows) + * PV: C[M][N] = A[M][K] @ B (use::b row_major over row-major V rows) + * + * Build: meson compile -C build jm_probe && ./build/jm_probe + */ +#include +#include +#include +#include +#include + +#include + +namespace jm = sycl::ext::oneapi::experimental::matrix; +using bf16 = sycl::ext::oneapi::bfloat16; + +static const char* mt_name(jm::matrix_type t) { + using mt = jm::matrix_type; + switch (t) { + case mt::bf16: return "bf16"; + case mt::fp16: return "fp16"; + case mt::tf32: return "tf32"; + case mt::sint8: return "s8"; + case mt::uint8: return "u8"; + case mt::fp32: return "fp32"; + case mt::fp64: return "fp64"; + default: return "other"; + } +} + +/* One WG = one sub_group of SG lanes. C[TM][TN] = A[TM][TK] (op) B[TK][TN]. + * b_is_kT=true: B stored row-major as [TN][TK] (K-cache rows), loaded + * col_major -> computes A @ B^T (the QK pattern). + * b_is_kT=false: B stored row-major as [TK][TN] (V-tile rows), loaded + * row_major -> computes A @ B (the PV pattern). */ +/* Storage type for a joint_matrix element type: precision::tf32 is a tag; + * its in-memory storage is fp32. bf16/fp16 store as themselves. */ +template struct jm_storage { using type = T; }; +template <> struct jm_storage { using type = float; }; + +static bool combo_supported( + const std::vector& combos, + jm::matrix_type t, int m, int n, int k) { + for (const auto& c : combos) + if ((c.msize == 0 || (int)c.msize == m) && (int)c.nsize == n && + (int)c.ksize == k && c.atype == t && c.btype == t && + c.ctype == jm::matrix_type::fp32 && c.dtype == jm::matrix_type::fp32) + return true; + return false; +} + +template +static int validate(sycl::queue& q, const char* name, bool b_is_kT) { + using S = typename jm_storage::type; + std::vector Ah(TM * TK), Bh(TN * TK), Ch(TM * TN, 0.f), + ref(TM * TN, 0.f); + /* Values exactly representable in bf16/tf32 (multiples of 1/16, |v|<1): + * products + fp32 accumulation are exact, so ref == device bit-for-bit. */ + for (int i = 0; i < TM * TK; i++) Ah[i] = ((i % 7) - 3) * 0.25f; + for (int i = 0; i < TN * TK; i++) Bh[i] = (((i * 3) % 5) - 2) * 0.25f; + for (int m = 0; m < TM; m++) + for (int n = 0; n < TN; n++) { + float acc = 0.f; + for (int k = 0; k < TK; k++) + acc += Ah[m * TK + k] * + (b_is_kT ? Bh[n * TK + k] : Bh[k * TN + n]); + ref[m * TN + n] = acc; + } + + S* dA = sycl::malloc_shared(TM * TK, q); + S* dB = sycl::malloc_shared(TN * TK, q); + float* dC = sycl::malloc_shared(TM * TN, q); + for (int i = 0; i < TM * TK; i++) dA[i] = S(Ah[i]); + for (int i = 0; i < TN * TK; i++) dB[i] = S(Bh[i]); + memset(dC, 0, TM * TN * sizeof(float)); + + q.submit([&](sycl::handler& h) { + h.parallel_for(sycl::nd_range<1>(SG, SG), + [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(SG)]] { + auto sg = it.get_sub_group(); + jm::joint_matrix sub_a; + jm::joint_matrix sub_b_row; + jm::joint_matrix sub_b_col; + jm::joint_matrix sub_c; + jm::joint_matrix_fill(sg, sub_c, 0.0f); + auto a_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dA); + auto b_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dB); + jm::joint_matrix_load(sg, sub_a, a_ptr, TK); + if (b_is_kT) + jm::joint_matrix_load(sg, sub_b_col, b_ptr, TK); + else + jm::joint_matrix_load(sg, sub_b_row, b_ptr, TN); + if (b_is_kT) + jm::joint_matrix_mad(sg, sub_c, sub_a, sub_b_col, sub_c); + else + jm::joint_matrix_mad(sg, sub_c, sub_a, sub_b_row, sub_c); + auto c_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dC); + jm::joint_matrix_store(sg, sub_c, c_ptr, TN, + jm::layout::row_major); + }); + }).wait(); + + double max_abs = 0.0; + for (int i = 0; i < TM * TN; i++) + max_abs = std::fmax(max_abs, std::fabs((double)dC[i] - ref[i])); + printf(" %-28s SG=%-2d TM=%-2d TK=%-2d TN=%-2d : max|d|=%.3e %s\n", + name, SG, TM, TK, TN, max_abs, max_abs < 1e-3 ? "PASS" : "FAIL"); + sycl::free(dA, q); sycl::free(dB, q); sycl::free(dC, q); + return max_abs < 1e-3 ? 0 : 1; +} + +/* Accumulator-load test: C initialized to 1.0 in memory, joint_matrix_load + * the accumulator, one mad, store. Validates the PV epilogue pattern + * (O staged in SLM, loaded as C, accumulated, stored back). */ +template +static int validate_acc(sycl::queue& q, const char* name) { + using S = typename jm_storage::type; + std::vector Ah(TM * TK), Bh(TK * TN), ref(TM * TN, 0.f); + for (int i = 0; i < TM * TK; i++) Ah[i] = ((i % 7) - 3) * 0.25f; + for (int i = 0; i < TK * TN; i++) Bh[i] = (((i * 3) % 5) - 2) * 0.25f; + for (int m = 0; m < TM; m++) + for (int n = 0; n < TN; n++) { + float acc = 1.0f; + for (int k = 0; k < TK; k++) + acc += Ah[m * TK + k] * Bh[k * TN + n]; + ref[m * TN + n] = acc; + } + S* dA = sycl::malloc_shared(TM * TK, q); + S* dB = sycl::malloc_shared(TK * TN, q); + float* dC = sycl::malloc_shared(TM * TN, q); + for (int i = 0; i < TM * TK; i++) dA[i] = S(Ah[i]); + for (int i = 0; i < TK * TN; i++) dB[i] = S(Bh[i]); + for (int i = 0; i < TM * TN; i++) dC[i] = 1.0f; + q.submit([&](sycl::handler& h) { + h.parallel_for(sycl::nd_range<1>(SG, SG), + [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(SG)]] { + auto sg = it.get_sub_group(); + jm::joint_matrix sub_a; + jm::joint_matrix sub_b; + jm::joint_matrix sub_c; + auto a_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dA); + auto b_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dB); + auto c_ptr = sycl::address_space_cast< + sycl::access::address_space::global_space, + sycl::access::decorated::no>(dC); + jm::joint_matrix_load(sg, sub_a, a_ptr, TK); + jm::joint_matrix_load(sg, sub_b, b_ptr, TN); + jm::joint_matrix_load(sg, sub_c, c_ptr, TN, + jm::layout::row_major); + jm::joint_matrix_mad(sg, sub_c, sub_a, sub_b, sub_c); + jm::joint_matrix_store(sg, sub_c, c_ptr, TN, + jm::layout::row_major); + }); + }).wait(); + double max_abs = 0.0; + for (int i = 0; i < TM * TN; i++) + max_abs = std::fmax(max_abs, std::fabs((double)dC[i] - ref[i])); + printf(" %-28s SG=%-2d TM=%-2d TK=%-2d TN=%-2d : max|d|=%.3e %s\n", + name, SG, TM, TK, TN, max_abs, max_abs < 1e-3 ? "PASS" : "FAIL"); + sycl::free(dA, q); sycl::free(dB, q); sycl::free(dC, q); + return max_abs < 1e-3 ? 0 : 1; +} + +int main() { + sycl::queue q(sycl::gpu_selector_v, sycl::property::queue::in_order{}); + sycl::device dev = q.get_device(); + printf("jm_probe: %s\n", dev.get_info().c_str()); + + if (!dev.has(sycl::aspect::ext_intel_matrix)) { + printf(" NO ext_intel_matrix aspect -- joint_matrix unsupported\n"); + return 1; + } + printf(" ext_intel_matrix aspect: yes\n supported combinations:\n"); + auto combos = + dev.get_info(); + for (const auto& c : combos) { + printf(" m=%-3llu n=%-3llu k=%-3llu %s x %s -> %s (acc %s)\n", + (unsigned long long)c.msize, (unsigned long long)c.nsize, + (unsigned long long)c.ksize, mt_name(c.atype), mt_name(c.btype), + mt_name(c.dtype), mt_name(c.ctype)); + } + fflush(stdout); + + /* Each validation is exception-isolated: an unsupported shape throws at + * JIT/submit time; catch it and report instead of losing the run. + * (tf32 deliberately NOT instantiated: its combo is listed but the driver + * JIT fails on the PackedB col-major builtin and poisons the program.) */ + int fails = 0; + auto run = [&](const char* name, auto&& fn) { + try { + fails += fn(); + } catch (const sycl::exception& e) { + printf(" %-28s : THREW %.60s -> unsupported\n", name, e.what()); + fails += 1; + } + fflush(stdout); + }; + using fp16 = sycl::half; + /* v1 fused-FA shapes (TQ=32, TK=64, 8xSG16): + * QK: A(16,16) row_major Q, B(16,16) col_major K-rows, k-steps over HD. + * PV: A(16,16) row_major P, B(16,16) row_major V, C(16,16) accum. */ + run("fp16 16x16x16 QK colB", [&] { + return validate(q, "fp16 16x16x16 QK colB", true); }); + run("fp16 16x16x16 PV rowB", [&] { + return validate(q, "fp16 16x16x16 PV rowB", false); }); + run("fp16 16x16 accum-load", [&] { + return validate_acc(q, "fp16 16x16 accum-load"); }); + /* bigger tiles (future TQ=64 / wide-N variants) */ + run("fp16 32x32x64 QK colB", [&] { + return validate(q, "fp16 32x32x64 QK colB", true); }); + run("fp16 32x32x64 PV rowB", [&] { + return validate(q, "fp16 32x32x64 PV rowB", false); }); + run("fp16 32x64 accum-load", [&] { + return validate_acc(q, "fp16 32x64 accum-load"); }); + run("bf16 16x16x16 QK colB", [&] { + return validate(q, "bf16 16x16x16 QK colB", true); }); + printf(fails ? "FAILED\n" : "ALL PASSED\n"); + return fails ? 1 : 0; +}