Port Xe3 MOE kernel code - #494
Conversation
- 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>
There was a problem hiding this comment.
🟡 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 withAttributeErrorrather 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.cmakeemits 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.uint8packed weights even though the kernel treats them as raw bytes,fused_expertsaccepts 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_expertsbefore validating that the public argument is positive.n_experts == 0reaches undefined integer division instead of aTORCH_CHECK; validate it before computingavg_m.
int avg_m = total_m / n_experts;
src/sycl/kernels/moe/xe35/GroupGemmMxfp4W4A16.hpp:179
- The Xe35 MXFP4 entry point accepts only
int8packed weights, whilefused_expertsand the added MXFP4 fixture useuint8buffers with the same bits. Its following scale check also requires float32 direct multipliers, whereas the public MXFP4 contract provides UE8M0uint8/float8scales. 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_expertsinvokes 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 reducingmax_mon 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 inextension.hcan 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.
|
Hi @mkumargarg @airMeng @mingfeima Please review. |
Port Xe3 MOE kernel code from the private repo, including kernel code and test cases.
Changes other than porting
intel_gpu_criarchitecture-enum case behind#if __SYCL_COMPILER_VERSION >= 20260717, since older oneAPI toolchains (2026.0) don't define that enum member and failed to compile.Verification
tests/test_xe3_moe_smoke.pyKnown 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.