Skip to content

Enable and fuse the Inkling MoE gate epilogue on Intel XPU - #36385

Closed
jmunetong wants to merge 7 commits into
sgl-project:mainfrom
jmunetong:xpu/inkling-gate-epilogue
Closed

jmunetong wants to merge 7 commits into
sgl-project:mainfrom
jmunetong:xpu/inkling-gate-epilogue

Conversation

@jmunetong

@jmunetong jmunetong commented Aug 25, 2026 •

Copy link
Copy Markdown
Contributor

Enable and fuse the Inkling MoE gate epilogue on Intel XPU

The Inkling MoE gate epilogue — everything between the 258-column gate GEMM and the MoE
dispatch — does not run on XPU at all today, and once it does, the path it lands on is
22–64× slower than necessary. Three independent commits: unblock, correct, optimize. CUDA and
ROCm are unchanged.

1. Enablement: the dispatch predicate admits XPU

sigmoid_gate_topk_renorm gated its CUDA JIT path on torch.version.hip is None, which is also
true on XPU, so all 64 MoE layers took the CUDA path and died on assert logits.is_cuda
(inkling_gate_topk_renorm.py:77); fixed by gating on .is_cuda and torch.version.hip is None — a conjunction, not a swap, because either predicate alone still admits XPU or ROCm.

2. Correctness: the Triton renorm returns NaN on underflow

_sigmoid_gate_topk_renorm_kernel formed weights as sigmoid(x) / Σ sigmoid(x), and fp32
sigmoid flushes to zero below x ≈ −104, so once all K + S active logits are that negative the
denominator is exactly 0 and every routed weight comes back NaN; replaced with
exp(lp − logsumexp(lp)), lp = logsigmoid(x), which is algebraically identical and already
what the eager reference uses (inkling_common/moe.py:140).

3. Performance: give the unified router an Inkling epilogue

The fallback ranks the [T, 258] fp32 gate row with tl.topk, a bitonic sort over all 256
routed columns that is pathological on Xe, whereas _router_triton_kernel already ranks the
cheap way with masked-max passes — so this parameterizes that kernel's epilogue with four
constexprs (EPILOGUE, SHARED_SINK, HAS_GLOBAL_SCALE, RETURN_PACKED) rather than adding
a kernel.

Device time (profiler self-time), same process, interleaved, min of 5 reps, Arc Pro B60:

µs T=1 T=8 T=64 T=512 T=4096 T=8192
tl.topk 58.81 58.93 62.20 181.07 641.57 1237.55
router 2.43 2.46 2.77 3.40 15.64 28.02
speedup 24.2× 24.0× 22.5× 53.3× 41.0× 44.2×

Correctness

Verified against an fp64 oracle (rtol=atol=2e-3): 12 shapes × 6 token counts pass on all
12, max abs 3.6e-7, zero index flips
, including ties, a dominating shared sink, and a
non-multiple-of-BLOCK tail at T=37. 44 existing moe_fused_gate configs stay bit-identical
in weights and ids
with device time unchanged, since LOGSIGMOID_SINK lives entirely behind
the EPILOGUE / SHARED_SINK constexprs (both 0 for every existing caller) — so the result
carries to CUDA structurally.

Tests

Adds test/registered/xpu/test_inkling_gate_epilogue.py — 9 cases / 33 subtests, ~14 s on an
Arc Pro B60, covering both dispatch branches against the fp64 oracle plus the invariant, packed
mode, ties, a dominating sink, a BLOCK tail, and shared_sink == 0 non-regression; XPU CI
only per the register_xpu_ci BKM. Each case fails without its fix — reverting all three gives
21 failed / 4 passed — and the fp64 oracle is the gate because the sum invariant alone does
not catch a wrong carry.

Test plan

WT=<this branch>/python
# CI test added here
PYTHONPATH=$WT ZE_AFFINITY_MASK=4,5,6,7 python -m pytest test/registered/xpu/test_inkling_gate_epilogue.py
# fp64 oracle + device time, 12 shapes x 6 token counts
PYTHONPATH=$WT ZE_AFFINITY_MASK=4,5,6,7 python bench_gate_epilogue.py --out after.json
# the A/B that produces the speedup table (bench alone cannot: the JIT env var is inert on XPU)
PYTHONPATH=$WT ZE_AFFINITY_MASK=4,5,6,7 python verify/ablate_stages.py --repeats 5 --iters 400
# 44-config bit-exactness at shared_sink == 0
PYTHONPATH=$WT ZE_AFFINITY_MASK=4,5,6,7 python verify/nonreg_moe_fused_gate.py --check golden_moe_fused_gate.pt

Intel Arc Pro B60, torch 2.12.0+xpu, triton 3.7.1. ablate_stages.py reports the kernel that
actually ran (top=), so the dispatch branch is proven rather than inferred.


CI States

Latest PR Test (Base): ❌ Run #35787126629
Latest PR Test (Extra): ❌ Run #35787126272
Latest PR Test (AMD ROCm 10): ❌ Run #35787126473

@jmunetong
jmunetong force-pushed the xpu/inkling-gate-epilogue branch 2 times, most recently from df4237b to e37fa88 Compare August 28, 2026 21:35
…UDA JIT

`torch.version.hip is None` means "not ROCm", which is also true on XPU, so the
Inkling MoE gate entered the CUDA JIT path on XPU and died on an `is_cuda`
assert one frame later:

  - sigmoid_gate_topk_renorm -> inkling_gate_topk_renorm_v2 ->
    `assert logits.is_cuda` (inkling_gate_topk_renorm.py:77), at every token count
  - InklingGate.forward_fused -> inkling_gate_gemv[_fused] ->
    `assert x.is_cuda` (inkling_gate_topk_renorm.py:203), at <= 4 tokens

Every other condition guarding those JIT gates is satisfied by Inkling's real
inputs on XPU (k == 6, n_shared == 2, G == 258, stride % 8 == 0, ptr % 32 == 0,
and both SGLANG_OPT_* env vars default to true), so all 64 MoE layers raised at
default env, at any batch size. Measured on Arc Pro B60: AssertionError on all
12 benchmark shapes and every token count from 1 to 8192.

Gate on `.is_cuda and torch.version.hip is None` rather than swapping one for the
other. Neither conjunct is sufficient alone:

  - `torch.version.hip is None` admits XPU, which is the bug above.
  - `.is_cuda` admits ROCm. PyTorch gives HIP tensors device type `cuda`, so
    `.is_cuda` is true there; the repo already works around the same aliasing in
    `is_cuda()` (utils/common.py:151, which needs `torch.version.cuda is not
    None`) and in `is_arch_support_pdl()` (jit/utils/arch.py:179). Admitting ROCm
    would send it into inkling_gate_topk_renorm_v2_kernel, which uses
    `__reduce_max_sync` (inkling_gate_topk_renorm.cuh:301) with no HIP equivalent
    and no arch guard, and assumes a 32-lane warp (:22) against a 64-wide
    wavefront.

The conjunction also keeps `forward_fused` in agreement with
`InklingGate.__init__`, whose `ensure_gate_gemv_fused_scratch` call is guarded by
`torch.cuda.is_available() and torch.version.hip is None`. On ROCm a bare
`.is_cuda` would take the FUSED GEMV path with that scratch never allocated.

CUDA and ROCm behaviour is unchanged; only XPU moves, and it moves from raising
to the Triton kernel in the same module.
`_sigmoid_gate_topk_renorm_kernel` formed the weights as
`sigmoid(x) / sum(sigmoid(x))`. fp32 sigmoid goes subnormal below x ~ -87 and
flushes to zero below x ~ -104, so once all K + S active logits sit in that
region the denominator is exactly 0 and every routed weight and both shared
gammas come back NaN.

Normalize as `exp(lp - logsumexp(lp))` with `lp = logsigmoid(x)` instead. This
is algebraically identical -- exp(log s_i - log sum_j s_j) == s_i / sum_j s_j --
and is the same form the eager reference already uses
(`_logsigmoid_normalize`, srt/models/inkling_common/moe.py:140), so the fast
path was the one that disagreed with the reference, not the other way round.

Measured on Arc Pro B60 against an fp64 oracle: the all-negative-logit cases at
base -90 and -200 go from NaN to 3.6e-7 / 4.8e-7 max abs error, base -40 was
already clean, and all 10 other benchmark shapes are unaffected. Cost is within
noise -- 58.74 -> 58.96 us at T=1, 181.31 -> 181.23 at T=512, 1237.28 ->
1239.06 at T=8192 (+0.37% / -0.04% / +0.14%).

The same algebra appears in the CUDA counterpart
(kernels/jit/csrc/moe/inkling_gate_topk_renorm.cuh:121-126 and :336-341). That
half is NOT touched here and has NOT been measured -- there is no NVIDIA GPU on
the machine this was developed on. It needs its own change from someone with
CUDA hardware.
`_sigmoid_gate_topk_renorm_kernel` builds its top-k on `tl.topk` /
`tl.bitonic_merge` over a packed (value << 16 | ~index) key. That codegen is
pathological on Intel Xe: 1237 us at T=8192 against 27 us for the unified
router's iterative masked-max on the same shape.

The unified router (`_router_triton_kernel`) is already structurally right for
this op -- iterative masked-max instead of `tl.topk`, BLOCK_M=1, num_warps=1 at
BLOCK_N=256. Only its epilogue is wrong for Inkling. So parameterize the
epilogue rather than adding a kernel:

  - `EPILOGUE` selects SUM_NORM (existing, default) or LOGSIGMOID_SINK.
  - `SHARED_SINK` counts trailing sink columns of the score row. They take no
    bias (bias stays [N], not [N + SHARED_SINK]) and never enter the top-k, but
    they do join the normalizer, in slots K_ROUTED .. K.
  - The value carried out of the top-k loop becomes the winner's RAW logit
    rather than its activated score, while ranking still happens on
    `sigmoid(logits) + bias`. Two quantities per column is why a plain
    "top-k and keep the winning score" kernel cannot serve this gate.
  - `HAS_GLOBAL_SCALE` applies the model's scalar global_scale, and
    `RETURN_PACKED` emits the (id << 16 | bf16 weight) form that
    InklingGate.emit_packed_topk consumes.

The wrapper gains `shared_sink`, `global_scale` and `return_packed` kwargs and
returns a 4-tuple (routed_weights, topk_indices, shared_weights, packed_topk)
when `shared_sink > 0`.

`sigmoid_gate_topk_renorm` dispatches here only when `not logits.is_cuda`.
CUDA is excluded on purpose: every number below is an Xe measurement, there is
no NVIDIA A/B for this epilogue, and the new tests are XPU-only, so a CUDA
caller that misses the JIT gate above (k != 6, G != 258, unaligned row, or
SGLANG_OPT_USE_GATE_TOPK_JIT=0) keeps the sort-based kernel it has today rather
than silently changing kernel and numerics. Lifting the exclusion is a separate
PR that needs NVIDIA hardware and a CUDA-side case in
test/registered/kernels/ops/moe/test_moe_fused_gate.py.

`SGLANG_OPT_USE_ROUTER_GATE_EPILOGUE` (EnvBool(True), registered in
srt/environ.py next to the other Inkling gate knobs) turns the epilogue off on
the devices that do reach it, leaving the sort-based kernel as the fallback.

Correctness, Arc Pro B60, fp64 oracle, rtol=atol=2e-3:
  - pass on all 12 shapes, max abs 3.6e-7, ZERO index flips at every token
    count, including the ties case (4 distinct selection values over 256
    columns), shared_outranks_routed, and a non-multiple-of-BLOCK tail.
  - free invariant holds: routed + shared weights sum to
    route_scale * global_scale (10.399998 .. 10.400002 vs an exact 10.400000).
  - packed mode: ids bit-match the plain path, weights bit-match
    `w.bfloat16()`, shared gammas bit-match.

Non-regression at the default `shared_sink == 0`: 44 configs (11 shapes x 4
token counts) captured before and after are bit-identical in both weights and
ids -- DeepSeek-V3 grouped routing, grouped + fused shared experts, ungrouped
256/384, Kimi-K2-class 896 experts (the BLOCK_N=1024 / num_warps=4 branch), all
three scoring functions, tanh softcapping, renorm on/off, scale on/off. Device
time for those callers is unchanged (2.13 -> 2.34 us at T=1, 28.64 -> 26.76 at
T=8192, i.e. noise in both directions). That evidence is from XPU; the argument
that it carries to CUDA rests on LOGSIGMOID_SINK living entirely behind the
`EPILOGUE` / `SHARED_SINK` constexprs, which are 0 for every existing caller.

Device us (profiler self-time) through `sigmoid_gate_topk_renorm`, same
process, interleaved, min of 3 reps:

               T=1     T=8    T=64   T=512  T=4096  T=8192
  tl.topk    58.81   58.93   62.22  181.17  641.59 1237.80
  router      2.43    2.45    2.75    3.43   15.63   28.09
  speedup     24.2x   24.0x   22.6x   52.8x   41.1x   44.1x

In packed mode -- what InklingGate actually requests -- the old kernel costs
37-41% more while the router costs 0.4% more, so the speedup is 34.1x / 32.6x /
31.5x / 72.6x / 41.2x / 44.3x on the same token counts.

Two caveats. This is reachable on XPU only together with the `.is_cuda`
dispatch fix earlier in this branch; without it the CUDA JIT branch is selected
first and still asserts. And the op is submission-bound: 49.6 us of host submit
against 2.45 us of device time at T=1, so this kernel win is NOT visible
end-to-end until the outputs are preallocated per layer or the decode step is
captured into a graph. Do not read these ratios as end-to-end numbers.
`sigmoid_gate_topk_renorm` had no test references anywhere in the tree, and
`moe_fused_gate` was covered only by test_moe_fused_gate.py, which hardcodes
DEVICE = "cuda" and registers register_cuda_ci only. So all three changes
earlier in this branch shipped with no automated coverage on XPU.

Adds test/registered/xpu/test_inkling_gate_epilogue.py: 9 cases / 33 subtests,
~20 s on an Arc Pro B60. Both dispatch branches (the unified router's
LOGSIGMOID_SINK epilogue and the tl.topk fallback) are checked against an fp64
recompute of the gate contract, plus the free normalization invariant, packed
mode, ties, a dominating shared sink, a non-multiple-of-BLOCK tail, and the
requirement that moe_fused_gate is unchanged at the default shared_sink == 0.
The branch is forced per case with
`envs.SGLANG_OPT_USE_ROUTER_GATE_EPILOGUE.override(...)`.

Per the register_xpu_ci BKM this is a dedicated file under
test/registered/xpu/ registering XPU CI only, so no per-class device skip is
needed. Deliberately no register_cuda_ci: CI collects by file, so adding it
would make CUDA runners pick up these XPU-only cases. A CUDA-side counterpart
for the shared _router_triton_kernel epilogue belongs in
test/registered/kernels/ops/moe/test_moe_fused_gate.py instead.

Each case was verified to fail without the fix it guards:

  - Reverting all three changes: 21 failed / 4 passed.
  - With only the dispatch fix applied, so the NaN is reachable:
    all_negative_90 and all_negative_200 fail on "NaN in routed weights" while
    all_negative_40 passes -- the intended discrimination.
  - Mutating the router epilogue to carry the activated score instead of the
    winner's raw logit (the mistake that "silently changes routing"): every
    router subtest fails against the oracle, tl.topk subtests still pass.
    Note the sum invariant alone does NOT catch this -- a wrong carry still
    yields a normalized set -- which is why the oracle check is the gate.

Not covered here, and not coverable from an XPU runner: the ROCm half of the
`.is_cuda and torch.version.hip is None` dispatch predicate, and the CUDA
exclusion on the router epilogue. Both are argued in their own commit messages.
@jmunetong
jmunetong force-pushed the xpu/inkling-gate-epilogue branch from e37fa88 to 26ba18d Compare September 3, 2026 22:29
@jmunetong
jmunetong marked this pull request as ready for review September 3, 2026 22:29
Conflict: python/sglang/kernels/ops/moe/moe_fused_gate.py. Upstream rewrote the
unified router (token-dependent bias, padding mask, packed output, V4.1
sqrtsoftplus, duplicate-max fix in the top-k loop); this branch added the
LOGSIGMOID_SINK epilogue to the same function.

Resolved by keeping upstream's SUM_NORM epilogue verbatim and re-applying the
sink epilogue as the EPILOGUE == 1 branch:
- the top-k loop keeps upstream's `remaining` mask, which now also fixes
  duplicate-max selection on the sink path
- the sink path reuses upstream's out_packed_ptr and stride_pm/pk instead of
  adding a second packed pointer, and adopts its `& 0xFFFF` mask
- unused pointer args take upstream's _dummy_i32 rather than aliasing an output,
  per its Dynamo write-back note
- shared_sink / global_scale / return_packed are now keyword-only, alongside
  upstream's other Triton-only extras; all call sites already pass them by keyword
@jmunetong

Copy link
Copy Markdown
Contributor Author

Verdict: measured on Intel XPU at tp=4, this PR is worth +8.29% output tok/s and −7.85% TPOT with no resolved change in TTFT (p = 0.37) — and it is essentially the whole of the measured end-to-end effect.

End-to-end — n = 12 per arm (6 interleaved sessions × 2 unprofiled reps; p computed on session means, n = 6)

metric gate off (_streaming_topk_kernel) gate on (_router_triton_kernel) Δ % p
Output token throughput (tok/s) 35.20 ± 0.84 38.12 ± 0.93 +2.92 +8.29% 0.0002
TPOT (ms, mean) 27.76 ± 0.68 25.58 ± 0.63 −2.18 −7.85% 0.0001
Median ITL (ms) 27.64 ± 0.72 25.41 ± 0.57 −2.23 −8.06% 0.0002
TTFT (ms, mean) 682.07 ± 18.81 687.41 ± 28.17 +5.34 +0.78% 0.37 — not resolved

Per-op device time, decode per step — n = 6 sessions/arm; 🔴 = buckets this PR computes differently, decided by kernel identity and the calls column (router kernel replaced 1:1; elementwise/memcpy/other launches fused away)

bucket gate off gate on Δ % noise calls
🔴 moe_router_topk 3.663 0.041 −3.622 −98.9% ±0.289 8 → 8
🔴 elementwise 0.205 0.082 −0.122 −59.5% ±0.014 93.8 → 25.8
🔴 memcpy 0.101 0.025 −0.076 −75.2% ±0.005 18.8 → 6.8
🔴 other 0.043 0.035 −0.008 −18.6% ±0.003 20.1 → 16.1
NON-COLLECTIVE TOTAL 8.241 4.512 −3.729 −45.2% ±0.458 301.9 → 217.9
comm_allreduce 21.240 23.349 +2.108 +9.9% ±7.302 13 → 13 — not resolved, not a regression

Prefill: moe_router_topk goes 3.497 → 0.075 ms, −3.422 (−97.9%) against a ±0.017 bar — but that is only −0.53% of the prefill device total, which is why TTFT shows nothing.

Attribution: the A/B is both-optimizations-vs-neither from one binary, env-gated, no rebuild between arms. The companion q-norm PR #37323 is −0.011% of prefill device time and +0.0012 ms/step (p = 0.26, not resolved) in decode, so the end-to-end effect above is attributable to this PR.

Environment: Intel Arc Pro B60, tp=4, oneAPI 2026.1, 6-layer reduced Inkling checkpoint, --attention-backend triton, --page-size 64, decode CUDA graphs on.
Exact repro (env vars, aggregation, gotchas) lives in the Reproduce section of our internal report AB_gate_plus_norm_tp4.md plus its sweep_run2.sh; happy to paste either here on request.

@jmunetong

Copy link
Copy Markdown
Contributor Author

@airMeng could you look into this PR when you get a chance?

@airMeng airMeng left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I think these are huge changes to CUDA/ROCm as well. Can we contribute these optimized triton functions to sgl-kernel-xpu then we can dispatch to the optimized path only for XPU here?

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants