Skip to content

Add Inkling MoE gate top-k + logsigmoid renorm Triton kernel - #528

Open
jmunetong wants to merge 1 commit into
sgl-project:mainfrom
jmunetong:xpu/inkling-gate-topk-renorm
Open

jmunetong wants to merge 1 commit into
sgl-project:mainfrom
jmunetong:xpu/inkling-gate-topk-renorm

Conversation

@jmunetong

@jmunetong jmunetong commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

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's sigmoid_gate_topk_renorm on 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

logits is [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.

  • Ranking: sigmoid(logit) + bias over the N routed columns, top-k by k masked-max passes in registers, lowest expert id wins ties. The sort-based kernel it replaces uses tl.topk / tl.bitonic_merge, a bitonic sort over every routed column, and that sort is what's slow on Xe.
  • Renorm: over the winners' raw logits plus the S sink columns. The sinks take no bias and never enter the top-k but do join the normalizer. It is computed as exp(lp − logsumexp(lp)), lp = logsigmoid(x). This is algebraically sigmoid(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.
  • Scale: route_scale * global_scale. Optional packed output (id << 16) | bf16_bits(weight), bitwise identical to fused_pack_topk.

It's exported from sgl_kernel behind the same try/except ImportError as fp8_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_kernel from upstream main 968726f3ce, 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, so xpu:0 = physical device 4), torch 2.15.0.dev20260912+xpu, triton 3.8.0.

T sort-based (µs) inkling_gate_topk_renorm (µs) speedup
1 103.96 [92.65–113.00] 5.38 [4.41–6.06] 19.3× [17.3–21.0]
8 117.12 [117.12–117.13] 6.25 [6.25–6.27] 18.7× [18.7–18.7]
64 123.62 [123.57–123.63] 6.91 [6.90–6.95] 17.9× [17.8–17.9]
512 181.05 [180.95–213.99] 4.24 [4.22–4.28] 42.7× [42.3–50.7]
4096 638.72 [637.97–639.43] 20.40 [20.38–20.41] 31.3× [31.3–31.3]
8192 1231.68 [1231.51–1232.04] 37.10 [37.09–37.38] 33.2× [32.9–33.2]

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:

sort-based inkling_gate_topk_renorm ratio
decode (T=1), per call, n=384 64.69 µs median (69.30 mean) 5.26 µs median (5.07 mean) 12.3× median, 13.7× mean
prefill (T=4096), per call, n=48 654.58 µs median 22.92 µs median 28.6×
gate time per decode step −0.257 ms/step (−0.276 excluding one contaminated baseline session)
gate time per prefill −2.80 ms (−2.53 excluding it)

In 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 in run_suite.py, checked against an fp64 oracle:

  • production shape at T ∈ {1, 8, 37, 64, 512, 4096}, contiguous and strided
  • narrow rows (BLOCK_M > 1), bf16 logits, M == 0
  • exact ties (lowest id wins), a dominating sink, all-underflow logits (NaN-free)
  • packed output bitwise equal to the manual pack

The kernel ranks in fp32, where scores that differ in fp64 can tie (e.g. sigmoid rounding 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:

mutation result
sigmoid / Σ sigmoid renorm fails test_underflow_is_finite
highest-id-wins ties fails test_ties_lowest_id_wins
bias dropped from ranking 18 of 21 fail

Follow-up in sglang

sgl-project/sglang#42670 (superseding sgl-project/sglang#36385) calls this kernel from sigmoid_gate_topk_renorm under is_xpu(). It looks the op up with getattr(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 in pyproject_xpu.toml.

Pre-commit passes locally.

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).
@jmunetong
jmunetong marked this pull request as ready for review October 6, 2026 22:13
Copilot AI balanced review requested due to automatic review settings October 6, 2026 22:13

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.

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

No deployments
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.

2 participants