Repository navigation
Conversation
inkling_gate_topk_renorm fuses sigmoid + bias ranking, top-k over the routed experts, and a logsigmoid renorm over the winners' raw logits plus the shared-expert sink columns. It is a drop-in for sglang's sigmoid_gate_topk_renorm on XPU (same arguments, return tuple and dtypes). The top-k is k masked-max passes in registers instead of tl.topk / tl.bitonic_merge, whose bitonic sort over every routed column is slow on Xe: 18-53x lower device time at the production shape (256 routed + 2 sink, k=6). The renorm is exp(lp - logsumexp(lp)) with lp = logsigmoid(x), so it stays finite where sigmoid(x) / sum(sigmoid(x)) divides 0/0 (all logits < ~-104).
Contributor
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
Low-level XPU-specific numerical and performance behavior warrants final human review and device CI validation.
Review effort: Balanced
Findings: None
What changed in this PR
Adds an XPU-optimized Triton kernel for Inkling MoE gate top-k selection and stable logsigmoid renormalization.
Changes:
- Implements fused ranking, renormalization, scaling, and packed output.
- Exports the kernel through
sgl_kernel. - Adds 21 accuracy cases and per-commit suite registration.
| File | Description |
|---|---|
python/sgl_kernel/inkling_gate.py |
Implements the Triton kernel and wrapper. |
python/sgl_kernel/__init__.py |
Exports the optional Triton operation. |
tests/test_inkling_gate.py |
Tests shapes, numerics, ties, underflow, and packing. |
tests/run_suite.py |
Registers the new tests. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Adds
inkling_gate_topk_renorm, a Triton kernel for the Inkling MoE gate epilogue, meaning everything between the 258-column gate GEMM and the MoE dispatch. It is a drop-in for sglang'ssigmoid_gate_topk_renormon XPU: same arguments, same(routed_weights, topk_indices, shared_weights, packed_topk)return, same output dtypes. This replaces the approach in sgl-project/sglang#36385. That PR put the optimization inside sglang's shared_router_triton_kernel. Here it lives in its own module, so sglang's CUDA/ROCm kernels are untouched and sglang needs only a few lines of XPU dispatch.What the kernel does
logitsis[M, N + S]: N routed experts followed by S shared-expert sink columns. It needs only column stride 1, so Inkling's[T, 258]slice of a padded[T, 264]GEMM output takes no copy.sigmoid(logit) + biasover the N routed columns, top-k by k masked-max passes in registers, lowest expert id wins ties. The sort-based kernel it replaces usestl.topk/tl.bitonic_merge, a bitonic sort over every routed column, and that sort is what's slow on Xe.exp(lp − logsumexp(lp)),lp = logsigmoid(x). This is algebraicallysigmoid(x) / Σ sigmoid(x)but stays finite once every active logit is below ≈ −104, where fp32 sigmoid flushes to 0 and the explicit form divides 0/0 into NaN.route_scale * global_scale. Optional packed output(id << 16) | bf16_bits(weight), bitwise identical tofused_pack_topk.It's exported from
sgl_kernelbehind the sametry/except ImportErrorasfp8_paged_mqa_logits_triton.Speed Tests
Isolated kernel. Device time (profiler self-time) at the production shape: 256 routed + 2 sink, k=6, fp32, a
[T, 258]slice of[T, 264], the same layout as the model's gate GEMM output. Baseline is sglang's_sigmoid_gate_topk_renorm_kernelfrom upstreammain968726f3ce, the Triton path XPU falls back to today. Five independent processes, each the minimum of 5 interleaved reps × 200 iters; median [range] across the five. Intel Arc Pro B60 (ZE_AFFINITY_MASK=4,5,6,7, soxpu:0= physical device 4), torch 2.15.0.dev20260912+xpu, triton 3.8.0.inkling_gate_topk_renorm(µs)At T=512 the sort-based kernel alternates between ≈181 µs and ≈214 µs across processes, so that row is a range, not a point. A sixth process ran with both kernels slower at T ≥ 512 (33–88×); an idle-then-rerun did not reproduce it, so it is excluded.
In the model. 6-layer reduced Inkling (
--load-format dummy), tp=4 on the same B60s, one sglang build with this op injected or not (sgl-project/sglang#42670), 3 interleaved sessions per arm. Profiler device time per call, pooled over all ranks:inkling_gate_topk_renormIn the model the sort-based kernel runs faster at T=1 than in isolation (64.69 vs 103.96 µs; cause not identified), so the in-model decode ratio is lower than the isolated one. No end-to-end throughput or latency change is resolved: the decode reduction is ≈1% of the 24.4 ms median ITL, below what this run can detect.
Indices are bit-identical to the sort-based kernel at every T, and weights match within 2e-3.
Accuracy Tests
tests/test_inkling_gate.py, 21 cases, registered inrun_suite.py, checked against an fp64 oracle:BLOCK_M > 1), bf16 logits,M == 0The kernel ranks in fp32, where scores that differ in fp64 can tie (e.g.
sigmoidrounding to exactly 1.0). So the ranking check requires a valid top-k to 1e-6 rather than the fp64 order bit for bit, and weights are checked at the kernel's own picks. The exact-tie test still requires exact indices. 30 consecutive runs pass. Each check was confirmed to bite by mutating the kernel:sigmoid / Σ sigmoidrenormtest_underflow_is_finitetest_ties_lowest_id_winsFollow-up in sglang
sgl-project/sglang#42670 (superseding sgl-project/sglang#36385) calls this kernel from
sigmoid_gate_topk_renormunderis_xpu(). It looks the op up withgetattr(sgl_kernel, "inkling_gate_topk_renorm", None), so against a wheel without this op, XPU keeps using the existing Triton kernel. It becomes the XPU fast path once a wheel with this PR is pinned inpyproject_xpu.toml.Pre-commit passes locally.