Pull requests / #1215

#1215 prefill MoE (SYCL): batch the per-expert GEMMs via oneMKL strided gemm_batch (+43% prefill on Arc Pro B70)

closed · @verycaptain · 0 comentários · No GitHub

Models & quants

Descrição

# Batch the per-expert MoE GEMMs on the SYCL prefill walk (+43% prefill on Arc Pro B70)

## What

`STRATA_MOE_BATCH=1` (default off, SYCL builds only): in the prefill MoE fallback walk, group 16 experts per layer, sort them by row count descending, pad each expert's rows to the group's max, and run gate/up and down as **one strided `gemm_batch` call per group** instead of 16 sequential `Gemm::f16` calls. New `Gemm::f16_batch` (oneMKL strided batch, falls back to sequential calls on any MKL exception).

Supporting changes:
- dequant ring widens 2 -> 16 slots when the flag is on (feeds the group without re-dequanting)
- the four row buffers (Xs/GU/Hh/Dm) get 15% slack for the padded layout when the flag is on
- a fits-guard falls back to the sequential walk when a layer's padding exceeds capacity
- sorting respects the walk: the routed-only staging walk takes a **global** cnt-desc sort (the stager's jobs and the compute loop both follow the `order` vector, so it is safe); the stream-all walk takes **windowed** cnt-desc sorts with W <= ring/2 so the slot-release walk keeps its id-order invariant

## Why

oneMKL's fp16 batch efficiency scales with the per-matrix size: on Arc Pro B70, 39.9 TF/s at M=96 vs 123 TF/s at M=1024 (single-GEMM M-sweep). The engine's per-expert GEMMs at live chunk sizes sit far below that. Batching 16 experts with sorted order (padding ~1.03x; unsorted is 1.63x and a net loss) recovers most of the gap.

## Measured (1x Intel Arc Pro B70, qwen3.8-flash-next IQ3_S/q4km, 124k cold prompt, alternating flag A/B on the same binary; the box has a second B70 running a separate production instance, not used for these numbers)

End-to-end numbers below were measured on our fork tree (0.1.39-merged, same engine code paths and the same port); the parity gates were then re-run against current upstream/main + this commit with the flag on.

- prefill: **1,591 -> 2,291 t/s (+43%)**, reproducible (1,597/1,585 off vs 2,291/2,291 on)
- decode: flat (~46 t/s either way; decode is spec-drafter-bound)
- GPU phase table: gemm gate/up 15.1% -> 6.5%, gemm down 15.1% -> 3.3%, dequant 8.8% -> 1.0%; the MoE chain drops from ~34% to ~11% of prefill
- gates with the flag ON: qsa_parity 0 failures, iq_parity 0 failures, qsa_prompt_attn_parity PASS

## Things I measured that shaped the design (worth knowing before reviewing)

- **group-form `gemm_batch` (pointer arrays, per-expert n) is mediocre** on this driver: 1.21x gate/up, 0.67x down. The strided form with uniform n per call is what MKL's batch kernel exploits.
- **padding decides win vs loss**: unsorted batch16 at the live chunk is a net loss on the down GEMM (0.80x). cnt-desc sorting is the unlock, not the batching itself.
- **host descriptors passed to async `gemm_batch` must outlive the call**: stack vectors destroyed after the call returns make MKL read freed memory and fake a driver DEVICE_LOST at large group counts. The port keeps them persistent; benches in `sycl/src/kernels/moe_bench/` (available if you want them in-tree) reproduce both the crash and the fix.

## Not included

XMX alternatives were measured and regress on bmg-g31 (quantized GEMM 0.20-0.25x) — not used here.

## Build note (pre-existing, not from this PR)

The SYCL tree does not build as-is against a current oneAPI nightly (icpx 7.1 / clang 22): `-fp-model=precise` is not a valid icpx argument, `-qmkl` was dropped from the compiler, `cpuid.h`'s static `__cpuidex` collides with the new builtin, `__spirv_GroupNonUniformShuffle*` calls are ambiguous for non-long `T`, and several shared headers have drifted past the SYCL copies (`Gemm::bf16`/`native` gained `ldx`, `load_experts_gguf` gained `ready`, `fused_gr_read_multi`'s return type differs). I patched all of these locally to build and gate this change; they are kept out of this PR to keep it scoped to the MoE batching. The parity gates (qsa_parity, iq_parity, qsa_prompt_attn_parity) build and pass with the flag on.

No site

Links install, modelos, releases.