Skip to content

Port Xe3 MOE kernel code - #494

Merged
Xia-Weiwen merged 8 commits into
mainfrom
port-xe3-moe
Sep 21, 2026
Merged

Xia-Weiwen merged 8 commits into
mainfrom
port-xe3-moe

Conversation

@Xia-Weiwen

@Xia-Weiwen Xia-Weiwen commented Sep 17, 2026 •

Copy link
Copy Markdown
Collaborator

Port Xe3 MOE kernel code from the private repo, including kernel code and test cases.

Changes other than porting

  1. src/sycl/Device.cpp -- Guarded the intel_gpu_cri architecture-enum case behind #if __SYCL_COMPILER_VERSION >= 20260717, since older oneAPI toolchains (2026.0) don't define that enum member and failed to compile.
  2. tests/test_xe3_moe_smoke.py  (new) -- minimal-shape Xe3 smoke tests for the two Xe3 grouped-GEMM ops

Verification

  • Build on a CRI simulator with oneAPI 2026.2 ✅
  • Run smoke tests for MOE on a CRI simulator ✅
    • tests/test_xe3_moe_smoke.py

Known issue
MXFP8/MXFP4 blockwise-scaled grouped GEMM kernels needed a CMake exclusion because the public repo's pinned cutlass-sycl commit lacks a kernel specialization that only exists in Intel's internal cutlass fork (used by the private repo). Files remain in the tree, ready to re-enable once the public dependency updates. The core bf16 and mxfp4-w4a16 grouped GEMMs (actually used by fused_experts) work fine.

Xia-Weiwen and others added 2 commits September 16, 2026 17:13
- Add xe35 grouped GEMM (bf16) and mxfp4 W4A16 grouped GEMM kernels
- Add MXFP8/MXFP4 blockwise-scaled grouped GEMM kernels for Xe3
- Wire up CMake AOT instantiation for the new xe35 kernels
- Register new xe35 ops in torch_extension_sycl.cc
- Add Python dispatch (is_xe3_arch, moe.py xe35 branch, mxfp4 wrapper)
- Port blockwise MOE tests
The MXFP8/MXFP4 blockwise-scaled grouped GEMM kernels
(src/sycl/xe35/blockwise_moe_mxfp{4,8}.cpp) require a
cutlass::gemm::kernel::GemmUniversal<GroupProblemShape, ...,
IntelXeGenericGroup epilogue, ..., GroupScheduler> kernel-level
specialization that is not present in the public intel/sycl-tla
commit currently pinned for DPCPP_SYCL_TARGET=cri.

Exclude the two .cpp files from compilation (keeping them in the
source tree) and remove the corresponding torch op registrations
so no undefined-symbol link errors occur. Re-enable once the
public sycl-tla dependency provides the missing specialization.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Critical dispatch, operator availability, test, and runtime validation issues remain unresolved.

Get a fresh assessment by requesting another Copilot review.

Pull request overview

Ports Xe3/CRI MoE grouped-GEMM kernels, Python dispatch, build integration, and tests. Blockwise MXFP4/MXFP8 sources remain excluded pending public cutlass-sycl support.

Changes:

  • Adds Xe35 BF16 and MXFP4 W4A16 kernels and launchers.
  • Adds blockwise MXFP4/MXFP8 infrastructure and tests.
  • Adds Xe3 detection, operator registration, and Python wrappers.
File summaries
File Reviewed changes and final notes
tests/test_mxfp4_blockwise_moe.py Adds blockwise MXFP4 tests. Critical (3 votes): does not skip when the operator is unavailable.
tests/test_cutlass_moe.py Adds FP8/MXFP4 tests. Critical (3 votes each): imports undefined, unexported wrappers.
src/torch_extension_sycl.cc Registers Xe35 operators. Moderate (3 votes): mutated output schema omits the Tensor! alias.
src/sycl/xe35/moe_group_gemm_helper.hpp Adds device-pointer and scale-layout helpers.
src/sycl/xe35/GroupGemmMxfp4W4A16.cpp Adds the Xe35 MXFP4-W4A16 dispatcher shim.
src/sycl/xe35/GroupGemm.cpp Adds the Xe35 BF16 dispatcher shim.
src/sycl/xe35/blockwise_moe_runner.hpp Adds the shared blockwise runner. Nit (1 vote): uses torch/extension.h instead of the established SABI include.
src/sycl/xe35/blockwise_moe_mxfp8.cpp Adds the MXFP8 blockwise implementation.
src/sycl/xe35/blockwise_moe_mxfp4.cpp Adds the MXFP4 blockwise implementation. Critical (2 votes): unaligned scale-storage pitch for ragged counts. Moderate (1 vote): synchronizes through a CPU copy on every call.
src/sycl/kernels/moe/xe35/mxfp4_w4a16/moe_mainloop.hpp Adds the MXFP4 dequantization mainloop.
src/sycl/kernels/moe/xe35/mxfp4_w4a16/moe_kernel.hpp Adds the MXFP4 kernel scheduler.
src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp Adds MXFP4 dispatch and validation. Moderate (2 votes): indexes the shape before checking rank. Moderate (1 vote each): rejects public packed-weight/scale representations and divides by n_experts before validation.
src/sycl/kernels/moe/xe35/GroupGemm.hpp Adds BF16 dispatch and tile selection. Moderate (2 votes): computes avg_m before validating n_experts.
src/sycl/kernels/moe/xe35/common/activation.hpp Adds shared activation helpers.
src/sycl/kernels/moe/xe35/bf16/moe_mainloop.hpp Adds the Xe35 BF16 mainloop.
src/sycl/kernels/moe/xe35/bf16/moe_kernel.hpp Adds the Xe35 BF16 kernel scheduler.
src/sycl/GroupGemmXe35LauncherInstance.cpp.in Adds the BF16 launcher template.
src/sycl/GroupGemmMxfp4W4A16Xe35LauncherInstance.cpp.in Adds the MXFP4 launcher template.
src/GroupGemmXe35.cmake Generates Xe35 BF16 instantiations. Moderate (1 vote): separate libraries recreate module-count and memory pressure.
src/GroupGemmMxfp4W4A16Xe35.cmake Generates Xe35 MXFP4 instantiations.
src/CMakeLists.txt Integrates Xe35 builds and exclusions. Moderate (1 vote): blockwise implementations are excluded while wrappers and tests remain exposed.
python/sgl_kernel/utils.py Adds Xe3 architecture detection. Critical (2 votes): CRI architecture mapping is missing, causing the query to raise.
python/sgl_kernel/moe.py Adds Xe3 MoE dispatch. Critical (3 votes): Xe3 W4A16 flags are rejected while later branches use the Xe2 op. Moderate (2 votes): blockwise dispatch calls an unregistered operator.
python/sgl_kernel/__init__.py Adds public Python exports.
include/sgl_kernel_ops.h Adds public Xe35 operator declarations.
Review details

Suppressed comments (7)

src/CMakeLists.txt:131

  • This filter removes both blockwise implementations from the CRI build, while their Python wrappers remain exported and the existing/new CRI test classes are still collected. On CRI those callers reach an unregistered torch.ops.sgl_kernel.* operator and fail with AttributeError rather than skipping or reporting a deliberate unsupported configuration. Gate the wrappers/tests with the same feature availability, or add an explicit unsupported path.
list(FILTER device_cpp EXCLUDE REGEX "/xe35/blockwise_moe_mxfp[48]\\.cpp$")

src/GroupGemmXe35.cmake:21

  • These loops generate the same 84 BF16 launcher translation units that the Xe20 build explicitly bundles to avoid Level Zero module pressure, but BuildOnLinux.cmake emits each Xe35 source as its own shared/device library. This recreates the module-count and memory risk that motivated the Xe20 bundling. Bundle the Xe35 launcher family or otherwise keep the generated modules under the driver limit.
foreach(act_type 0 1 2 3)
    if(act_type EQUAL 3)
        set(with_bias_list false)
    else()
        set(with_bias_list true false)

src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp:177

  • Unlike the existing W4A16 entry point, this new Xe3 op rejects torch.uint8 packed weights even though the kernel treats them as raw bytes, fused_experts accepts both int8/uint8, and the repository's MXFP4 quantizer returns uint8. Valid packed weights therefore fail validation before launch. Accept both Char and Byte.
  TORCH_CHECK(packed_weights.scalar_type() == at::ScalarType::Char, "packed_weights must be int8");

src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp:213

  • As in the BF16 dispatcher, this divides by n_experts before validating that the public argument is positive. n_experts == 0 reaches undefined integer division instead of a TORCH_CHECK; validate it before computing avg_m.
  int avg_m = total_m / n_experts;

src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp:179

  • The Xe35 MXFP4 entry point accepts only int8 packed weights, while fused_experts and the added MXFP4 fixture use uint8 buffers with the same bits. Its following scale check also requires float32 direct multipliers, whereas the public MXFP4 contract provides UE8M0 uint8/float8 scales. Routing the public path to this op would therefore reject the format under test; accept the public representations or normalize them before dispatch.
  TORCH_CHECK(packed_weights.scalar_type() == at::ScalarType::Char, "packed_weights must be int8");

  auto sc_shape = scales.sizes().vec();

src/sycl/xe35/blockwise_moe_mxfp4.cpp:212

  • Unlike the MXFP8 implementation, this per-call problem_sizes.to(torch::kCPU) forces a device-to-host synchronization before launching the GEMM. fused_experts invokes MoE GEMMs frequently, so this serializes every call; use a device-safe upper bound for the scratch allocation (as the MXFP8 path does) instead of reducing max_m on the host.
    auto problem_sizes_cpu = problem_sizes.to(torch::kCPU);
    const int32_t* psz = problem_sizes_cpu.data_ptr<int32_t>();
    int max_m = 0;
    for (int e = 0; e < E; ++e) {
      if (psz[e * 3] > max_m) max_m = psz[e * 3];

src/sycl/xe35/blockwise_moe_runner.hpp:21

  • This is the only source-side use of <torch/extension.h> in the repository; the other new SABI-facing Xe3 dispatchers use <torch/all.h>. When the blockwise sources are re-enabled, pulling in extension.h can conflict with the SABI build and deviates from the established binding include.
#include <torch/extension.h>
  • Files reviewed: 26/26 changed files
  • Comments generated: 10
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread python/sgl_kernel/moe.py Outdated
Comment thread python/sgl_kernel/utils.py
Comment thread src/sycl/xe35/blockwise_moe_mxfp4.cpp
Comment thread tests/test_cutlass_moe.py
Comment thread tests/test_cutlass_moe.py
Comment thread tests/test_mxfp4_blockwise_moe.py
Comment thread python/sgl_kernel/moe.py
Comment thread src/sycl/kernels/moe/xe35/GroupGemm.hpp
Comment thread src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp
Comment thread src/torch_extension_sycl.cc Outdated
@mkumargarg
mkumargarg requested a review from sspintel September 17, 2026 11:25
@Xia-Weiwen
Xia-Weiwen marked this pull request as ready for review September 18, 2026 07:06
@Xia-Weiwen

Copy link
Copy Markdown
Collaborator Author

Hi @mkumargarg @airMeng @mingfeima Please review.

@Xia-Weiwen
Xia-Weiwen merged commit 739d22e into main Sep 21, 2026
6 of 7 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants