Skip to content

feat: W4A8 ARK XPU MoE kernel (int4 weight / int8 compute) with prefill + decode - #2143

Draft
a32543254 with Copilot wants to merge 77 commits into
mainfrom
copilot/copilotoptimize-int4-moe-performance
Draft

feat: W4A8 ARK XPU MoE kernel (int4 weight / int8 compute) with prefill + decode#2143
a32543254 with Copilot wants to merge 77 commits into
mainfrom
copilot/copilotoptimize-int4-moe-performance

Conversation

Copilot AI commented Aug 11, 2026

Copy link
Copy Markdown
Contributor

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's AUTO_S8 trick — 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/unpackq in xpu_wrapper.hpp:

sxt[e][n][j] = max_{g in block j} |s[e][n][g]| * 8 / 127
w8           = round(w4 * s / sxt)

Default block = K (group=-1) → blks == 1. Since _pack_int4_sym divides by 7.0, max|w8| = 7 * 127/8 ≈ 111 ≤ 127, so the conversion never clips. Epilogue is out = acc_s32 * scale_b[col] * scale_a[row].

Kernelwrapper/include/sycl_tla_moe_w4a8.hpp (new)

  • One-shot moe_w4a8_prepack: [E,N,K/2] int4 → [E,N,K] int8 + [E,N,blks] fp32 scales. Costs E*N*K bytes (2× the packed int4).
  • Per-token activation quant producing x_s8 and scale_a.
  • Prefill: grouped DPAS GEMM on XE_DPAS_TT<8, int32_t, int8_t, int8_t>, tile ladder mirroring the reference (m<16 → 8x128, m<128 → 64x128, m<=1024 → 128x128, else 256x128). num_tokens_per_expert is a device tensor, so it reuses the existing persistent work-stealing scheduler.
  • Decode: GEMV, SG_SIZE=16 / N_TILE=16, one output column per lane.
  • ARK_MOE_W4A8_AUTO_S8 overrides 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, 4 m.def registrations.
  • 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.
from auto_round_extension.ark import auto_round_kernel as ark

out = ark.moe_w4a8(x, weights, scales, num_tokens_per_expert, group_size=32, phase="prefill")

# or manage the prepack lifetime explicitly
w_s8, w_scale = ark.moe_w4a8_prepack(weights, scales, group_size=32)
out = ark.moe_gemm_w4a8(x, w_s8, w_scale, num_tokens_per_expert, phase="decode")

Benchmarktest/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.

pytest auto_round_extension/ark/test/test_moe_w4a8_perf.py -v
python auto_round_extension/ark/test/test_moe_w4a8_perf.py --warmup 10 --iters 50

Type of Change

New feature (Performance)

Checklist Before Submitting

  • My code has been tested locally.
  • Documentation has been updated as needed.
  • New or updated tests are included where applicable.
  • The CUDA CI has passed. You can trigger it by commenting /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-VALIDATION marker 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 of moe_w4a8_rescale_block_size across 11 cases, pybind argument order against the Python call sites, and offset overflow at E=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.

a32543254 and others added 30 commits July 31, 2026 11:16
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.
Revert commit f887763, restoring the fp8 per-group prefill dispatch threshold
to A_avg_M <= 512. The per-expert (per-tensor) fix from 93cde8c is retained.
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>
Copilot AI and others added 2 commits August 13, 2026 02:32
…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>
Copilot AI and others added 2 commits August 13, 2026 04:40
…-quant loads

Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Copilot AI and others added 2 commits August 13, 2026 05:26
…n (EN + CN)

Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… CN)

Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Copilot AI and others added 3 commits August 13, 2026 06:03
… 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>
Copilot AI and others added 2 commits August 13, 2026 07:18
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants