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 commentaires · Sur GitHub
Description
# 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.
Sur le site
Liens install, modèles, releases.