Skip to content

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

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

feat: W4A8 ARK XPU MoE kernel (int4 weight / int8 compute) with prefill + decode#2143
a32543254 with Copilot wants to merge 86 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>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Copilot AI and others added 2 commits August 14, 2026 04:48
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
… fusion claim

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

Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
Co-authored-by: a32543254 <53296245+a32543254@users.noreply.github.com>
@a32543254
a32543254 marked this pull request as ready for review August 31, 2026 06:42
@a32543254
a32543254 requested a review from Zhenzhong1 August 31, 2026 06:42
@a32543254

Copy link
Copy Markdown
Contributor

@copilot resolve the merge conflicts in this pull request

Copilot AI and others added 2 commits August 31, 2026 06:58
Resolve conflicts with main's SYCL-TLA translation-unit refactor, which
moved kernel entry points out of the headers into CMake-generated TUs
(sycl_tla_moe.cpp.in + MOE_SOURCE_MODE).

Textual conflicts:
- ark.cpp: keep main's include list (kernel headers are no longer pulled
  into the pybind TU).
- sycl_tla_common.hpp: keep both sides -- main's dense-GEMM / igemm
  declarations plus this branch's W4A8 declarations.
- sycl_tla_moe.hpp / sycl_tla_moe_decode.hpp: drop the header-resident
  ark::moe_gemm and ark::moe_gemm_decode definitions (main now emits them
  non-inline from modes 1 and 3; keeping the header copies would be an ODR
  violation). The moe_detail / moe_decode_detail internals are unchanged.

Semantic ports (not flagged by git, since the files differ):
- mode 2 called moe_gemm_launcher<'R','R',Element,...>, but this branch
  changed that template to take WGTile/SGLayout and added
  moe_gemm_dispatch<Element> for tile selection -- switch mode 2 over.
- MoeDecodeParams gains num_tokens_per_expert / num_experts; the int4-sym
  and FP8 DPAS fast paths are driven by per-expert row counts, not the
  per-token expert map.
- New shared predicates sycl_tla_moe_decode_{int4,fp8}_dpas_fastpath so
  mode 3 can skip the fill_expert_id launch on exactly the paths that
  never read it, without the skip and the routing decision diverging.
- mode 5 (fp8): add the FP8 DPAS fast path and route the fallback through
  launch_fp8_by_mode.
- new mode 18: int4 decode gets its own TU (mode 6 is shared with
  int2/int8) carrying the S4 DPAS fast path and the coalesced GEMV.
- new mode 19: the four W4A8 entry points. The header versions move into
  namespace moe_w4a8_detail and mode 19 emits the external ark:: symbols,
  since pybind takes their address and inline definitions would leave the
  symbols unemitted.

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

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

Copilot AI commented Aug 31, 2026

Copy link
Copy Markdown
Contributor Author

@copilot resolve the merge conflicts in this pull request

Merged in 81aa4e67 (follow-up whitespace fix in 1649b92f). Branch is now 0 commits behind main.

This wasn't purely textual. main refactored SYCL-TLA so kernel entry points no longer live in headers — they're emitted from CMake-generated TUs (sycl_tla_moe.cpp.in, selected by MOE_SOURCE_MODE) — so the header-resident entry points and decode fast paths on this branch had to be ported into that architecture.

Textual conflicts (5 regions, 4 files)

  • ark.cpp — took main's include list; kernel headers are no longer pulled into the pybind TU.
  • sycl_tla_common.hpp — kept both sides (main's dense-GEMM/igemm decls + this branch's W4A8 decls).
  • sycl_tla_moe.hpp / sycl_tla_moe_decode.hpp — dropped the header-resident ark::moe_gemm / ark::moe_gemm_decode; main now emits both non-inline from modes 1 and 3, so keeping the header copies would be an ODR violation. The moe_detail / moe_decode_detail internals were kept.

Semantic breakages git did not flag (the files differ, so no conflict was reported)

  • Mode 2 called moe_gemm_launcher<'R','R',Element,...> (6 template params), but this branch had changed that template to take WGTile/SGLayout (8) and added moe_gemm_dispatch<Element> for tile selection. Switched mode 2 over — this would otherwise have been a hard compile error.
  • MoeDecodeParams gains num_tokens_per_expert / num_experts; the int4-sym and FP8 DPAS fast paths key off per-expert row counts, not the per-token expert map.
  • Added shared predicates sycl_tla_moe_decode_{int4,fp8}_dpas_fastpath so mode 3 can skip the fill_expert_id launch on exactly the paths that never read it, without the skip and the routing decision diverging across TUs.
  • Mode 5 (fp8): restored the FP8 DPAS fast path; fallback routes through launch_fp8_by_mode.
  • New mode 18: int4 decode gets its own TU (mode 6 is shared with int2/int8, so int4-only logic couldn't live there), carrying the S4 DPAS fast path and the coalesced GEMV.
  • New mode 19: the four W4A8 entry points. Header impls moved to namespace moe_w4a8_detail; mode 19 emits the external ark:: symbols, since pybind takes their address and inline definitions would leave them unemitted at link time.

Verification — no SYCL compiler in this environment, so checks were static: every namespaced symbol in the template resolves against the merged headers; exactly one definition site per entry point; all call arities and decl/defn signatures match; namespace and per-mode brace balance clean; and the auto-merged __init__.py kept all 11 functions this branch added plus all 38 from main, with main's one removal honored. CodeQL (cpp) is clean.

The kernel still needs the hardware validation called out in the description — that's unchanged by this merge.

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

Copy link
Copy Markdown
Collaborator

/azp run Unit-Test-CUDA-AutoRound

@azure-pipelines

Copy link
Copy Markdown
Azure Pipelines could not run because the pipeline triggers exclude this branch/path.

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.

3 participants