Draft
feat: W4A8 ARK XPU MoE kernel (int4 weight / int8 compute) with prefill + decode#2143
Conversation
Signed-off-by: Dong, Bo1 <bo1.dong@intel.com>
Signed-off-by: Dong, Bo1 <bo1.dong@intel.com>
Merge ec61621 accidentally widened the dpas_w8a16_policy_m_32 bucket in the fp8 per-tensor (per-expert) prefill dispatch from A_avg_M <= 32 to <= 512, routing large-M prefill through the small 32x64 tile instead of the large-M 128x128 default tile and regressing performance. Restore the <= 32 threshold.
Merge ec61621 also widened the dpas_w8a16_policy_m_32 bucket in the fp8 per-group prefill dispatch from A_avg_M <= 32 to <= 512, regressing large-M group-size prefill for the same reason as the per-tensor path. Restore the <= 32 threshold so large-M prefill uses the 128x128 default tile.
Migrate the auto-dispatch logic from branch copilot/update-phase-auto-dispatch-logic (commit 9605fe4): phase="auto" now dispatches to decode when activations.shape[0] <= threshold (total tokens) instead of inspecting num_tokens_per_expert.max(), avoiding a host-device sync. Adds ARK_MOE_AUTO_DECODE_MAX_TOKENS env override (default 256) and updates test_moe_unified.py accordingly.
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
The shared S4 DPAS grouped-GEMM (prefill path) already beats the scalar GEMV decode kernel by ~2x at 256 tokens (bs32) and only loses at the single-stream bs1 (8-token) extreme. Routing 256-token batches to decode was leaving ~2x on the table, so lower the auto-dispatch default and update coupled unified-dispatch tests and perf-test notes. Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…reshold tuning Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… GEMV Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…mangled names Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…atch regression Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…at blks==1 Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…tity; docs(ark): document both prefill changes Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… epilogue Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…-quant loads Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…n (EN + CN) Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… CN) Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… for W4A8 MoE prefill Replaces the W4A8 prefill epilogue's 64 scalar 32-byte stores per sub-group fragment with the hardware 2D block store (`make_block_2d_copy_D` + `copy(copy_d, tCrD, tCgC)`, the sequence already compiled in `sycl_tla_dense_gemm.hpp` for the same accumulator/output widths), and removes the activation quantizer's second read of every row by holding it in registers between the absmax and the quantize pass. Also corrects the harness roofline, which counted only weight bytes and so understated the bandwidth a shape needs to hit 100 TFLOPS by up to 2.2x. Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
…ted roofline Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… sweep Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Adds a W4A8 MoE kernel to the ARK XPU backend covering both prefill and decode: int4 symmetric weights, int8 compute dtype, activations dynamically quantized per token to int8. Follows the W4A8 weight-only GEMM in
zhenzhong/woqgemm_s8_update, and uses ARK'sAUTO_S8trick — re-scaling int4 group=32 weights into int8 group=-1 so the K loop needs a single full-width int32 accumulation instead of per-group folding.Built on top of
copilot/optimize-int4-moe-performance.Numerics
AUTO_S8 re-scale, ported from
packscale/unpackqinxpu_wrapper.hpp:Default block = K (
group=-1) →blks == 1. Since_pack_int4_symdivides by 7.0,max|w8| = 7 * 127/8 ≈ 111 ≤ 127, so the conversion never clips. Epilogue isout = acc_s32 * scale_b[col] * scale_a[row].Kernel —
wrapper/include/sycl_tla_moe_w4a8.hpp(new)moe_w4a8_prepack:[E,N,K/2]int4 →[E,N,K]int8 +[E,N,blks]fp32 scales. CostsE*N*Kbytes (2× the packed int4).x_s8andscale_a.XE_DPAS_TT<8, int32_t, int8_t, int8_t>, tile ladder mirroring the reference (m<16 → 8x128,m<128 → 64x128,m<=1024 → 128x128, else256x128).num_tokens_per_expertis a device tensor, so it reuses the existing persistent work-stealing scheduler.SG_SIZE=16/N_TILE=16, one output column per lane.ARK_MOE_W4A8_AUTO_S8overrides the rescale block size; invalid values fall back to K. Shape gate:N % 16 == 0,K % 64 == 0,group_size % 8 == 0,K % group_size == 0.Plumbing
sycl_tla_common.hpp: 4 public declarations.ark.cpp: include, 2 wrappers, 4m.defregistrations.auto_round_kernel/__init__.py: Python API plus a prepack cache. The cache key includes device type/index and the entry pins the source tensors, so a freed-and-reallocated weight buffer can't alias another layer's int8 weights.Benchmark —
test/test_moe_w4a8_perf.py(new)Standalone perf + accuracy harness, runnable under pytest or directly. Qwen3-MoE shapes (
E=128, hidden=2048, inter=768, top_k=8, group_size=32). Reports SNR/cosine/max-rel-err against a torch bf16 baseline and latency/TFLOPS/speedup against the W4A16 kernel, for prefill across batch×seq and decode across batch.Type of Change
New feature (Performance)
Checklist Before Submitting
/azp run Unit-Test-CUDA-AutoRound.Important
Needs hardware validation. No XPU or SYCL compiler was available, so the kernel is unbuilt and has not run on device; the header carries the same
STATUS: NEEDS-HARDWARE-VALIDATIONmarker as the sibling MoE DPAS headers. What was checked offline: a numpy replay of the full AUTO_S8 + int8-activation path (no clipping, ~38.5 dB end-to-end output SNR vs. the script's 20 dB gate), Python↔C++ parity ofmoe_w4a8_rescale_block_sizeacross 11 cases, pybind argument order against the Python call sites, and offset overflow atE=128/N=2048/K=2048.Worth a close look on hardware: CuTe tile/policy tuning, and the 128-token cutoff for decode auto-dispatch.
Docs:
test/README_MOE_W4A8.md+README_MOE_W4A8_CN.md.