feat(deepgemm): migrate to sgl-deep-gemm 0.1.2 wheel API

Port the DeepGEMM "deprecate-from-sgl-kernel, separate sgl-deep-gemm wheel"
migration (upstream PR #24268 and follow-ons) so our code runs against the
torch-2.11 dev-cu13 image, which ships DeepGEMM as the separate sgl-deep-gemm
0.1.2 wheel (import name still `deep_gemm`). Verified against upstream/main HEAD,
NOT the introducing PR — the wheel API drifted between 0.0.1 and 0.1.2.

Ports (all verified against HEAD = wheel 0.1.2):
- nsa_indexer/nsa_backend: paged-MQA context_lens must be (N_total, 1). The
  0.0.1 form (batch_size, next_n) DEADLOCKS fp8_paged_mqa_logits on next_n>=2
  (our EAGLE deploy uses next_n=4 on SM90/H200, which does not take the SM100-
  only DG-native broadcast path). Matches HEAD's _to_2d_context_lens.
- compile_utils warmup: hasattr-guard the dropped get/set_compile_mode API;
  pass m_indices positionally to m_grouped_fp8_gemm_nt_contiguous.
- fp8_utils.transform_scale_ue8m0: restore TMA-aligned stride when the DLPack
  round-trip collapses a size-1 trailing dim.
- moe_runner/deep_gemm: guard the SBO masked-gemm return unpack when overlap is
  inactive (#26839) — reachable via our --enable-single-batch-overlap.
- entrypoint: enable DeepGEMM PDL by default, hasattr-guarded (#23979).
- engine: require sglang-kernel >= 0.4.3 (first sgl-deep-gemm-era kernel); the
  migrated (N_total,1) paths would misbehave on the old bundled DeepGEMM.

Verified-unchanged symbols left as-is: fp8_mqa_logits, fp8_gemm_nt, bf16_gemm_*,
get_mk_alignment_for_contiguous_layout, transform_sf_into_required_layout,
get/set_num_sms, the masked-gemm signature. Runtime validation pending the
torch-2.11 image (tai-kernel rebuild + harness).

Refs: WI-2026-06-07-001

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
(cherry picked from commit 32818784d332b5c48faf8027247cd4cd9a0f48cd)
This commit is contained in:
2026-06-08 23:46:59 +08:00
committed by laoyao0822
parent 430ed85083
commit db5ab48b19
9 changed files with 121 additions and 11 deletions
@@ -154,6 +154,11 @@ date = "2026-06-07"
scope = "cp_shared_kv"
content = "Post-rebase test validation on f75ffff8d. Root-caused + FIXED a test-isolation bug in f75ffff8d's test_cp_shared_kv_runtime.py: it installs CPU-CI sgl_kernel stubs at import (sys.modules.setdefault + torch.library FRAGMENT). On GPU, if collected before real sgl_kernel loads, the empty stub shadows it process-wide -> later real-kernel test (fast_topk_transform_ragged_fused) calls a None lambda -> TypeError; sys.modules cleanup makes it WORSE (segfault from FRAGMENT double-registration). Proven via probe (sgl_kernel.__file__=None, fast_topk=<lambda>). Fix: import real sgl_kernel first (setdefault keeps real, FRAGMENT hits already-registered, setattr fills only missing); CPU-CI unaffected. Verified GPU: cp_shared_kv_runtime alone 120 pass; combined cp+topk 125 pass. Commit b1cbacffa. SEPARATE pre-existing: test_nsa_pool_host_unit.py 3 fail IN ISOLATION (env-specific, NOT mine): one is f75ffff8d's own fail-fast rejecting CUDA src_indices; two are cudaErrorHostMemoryAlreadyRegistered (host-pinned-mem conflict from running tests on g0034 while the live prefill server pins hicache memory). per-layer suite still 24/24. Branch rebased on f75ffff8d, docs_internal stripped+gitignored; user handles push."
[[content.journal]]
date = "2026-06-07"
scope = "deepgemm-port"
content = "DeepGEMM/env-rebase investigation. Upstream renamed deep_gemm: PR#24268 (ecb786c8d7, 2026-05-06) deprecated DeepGemm bundled in sgl-kernel and moved it to a SEPARATE 'sgl-deep-gemm' pip wheel (import name STILL 'deep_gemm'; wheel is py3-none = torch-ABI-agnostic, cu-version-tagged). Wheel API changed (release-0426): get_paged_mqa_logits_metadata now needs 2D context_lens [bs,next_n]; get_compile_mode/set_compile_mode optional (hasattr-guard); transform_scale_ue8m0 DLPack stride fix when shape[-1]==1; configurer non-cuda guard; warmup m_indices kwarg->positional. Our runtime hot path entrypoint.py:81 already positional; only compile_utils warmup needs it. preload_kernels is commented out (non-issue). Our code uses OLD API at all these sites; fork-base=2d288ba8c9 (#15852, 2026-03-23); #24268 NOT in our history. KEY ENV FINDING: torch 2.11 bump #21247 (2026-05-02) PRECEDES the deepgemm split, so any upstream image with sgl-deep-gemm is ALSO torch 2.11. Current dev-cu13-2 = torch2.9.1+cu130/sgl-kernel0.4.0 (predates both). Upstream/main now: torch2.11.0, cuda-python>=13, sgl-kernel0.4.3, sgl-deep-gemm0.1.2, flashinfer0.6.12[cu13], transformers5.8.1, xgrammar0.2.1, torchao0.17.0, mooncake0.3.11.post1(cu13 prebuilt wheel), new deps tilelang/tokenspeed_mla/kernels. torch2.9->2.11 = ABI break forcing rebuild of tai-kernel + native ext. sgl-deep-gemm wheel is py3-none so CAN be installed on torch2.9 (decouples deepgemm port from torch jump). BLOCKED on: exact target image versions (ssh to inspect denied by classifier) - need user to authorize introspection or name the target image tag."
[[content.acceptance_criteria]]
text = "govctl check passes"
status = "pending"
@@ -0,0 +1,36 @@
#:schema ../schema/work.schema.json
[govctl]
id = "WI-2026-06-07-001"
title = "Rebase runtime env to torch-2.11 dev-cu13 image (DeepGEMM sgl-deep-gemm port + compat)"
status = "active"
created = "2026-06-07"
started = "2026-06-07"
[content]
description = "Rebase the runtime environment to the torch-2.11 dev-cu13 image (built from upstream/main's current Dockerfile: torch 2.11.0, sgl-kernel 0.4.3, sgl-deep-gemm 0.1.2, flashinfer 0.6.12, transformers 5.8.1, mooncake 0.3.11.post1) WITHOUT a full code rebase. Phase 1 (done): port the DeepGEMM sgl-deep-gemm wheel migration (API compat + coupled correctness fixes) verified against upstream HEAD = wheel 0.1.2. Phase 2/3 (in image): rebuild tai-kernel against torch 2.11, import-smoke + fix any transformers-5.8/xgrammar-0.2/torch-2.11 breakages empirically, then harness coldchunk byte-equality + perf validate."
[[content.journal]]
date = "2026-06-07"
scope = "deepgemm-port"
content = "History review (user: don't rush, dig commit history of files to change) caught a CRITICAL bug in the initial port. Verified all 4 edited functions against upstream HEAD (= sgl-deep-gemm 0.1.2), NOT just PR#24268 (which targeted wheel 0.0.1). Findings: (1) compile_utils warmup hasattr-compile-mode-guard + positional m_indices = byte-identical to HEAD. OK. (2) fp8_utils transform_scale_ue8m0 DLPack stride fix = byte-identical to HEAD. OK. (3) entrypoint masked/contig deep_gemm signatures unchanged at HEAD (fp8_m_grouped_gemm_nt_masked enable_overlap/max_block_n; m_grouped_fp8_gemm_nt_contiguous positional) - our entrypoint already matches, no change. (4) CRITICAL: _to_2d_context_lens layout CHANGED between 0.0.1 and 0.1.2. PR#24268 used (batch_size, next_n); HEAD uses ALWAYS (N_total,1) with comment 'avoid deadlock at deep_gemm.fp8_paged_mqa_logits'. Passing (bs,next_n>=2) DEADLOCKS the kernel. Our EAGLE deploy = next_n=4 on SM90/H200 (DG-native broadcast path is SM100-only), so it takes the per-token (N_total,1) path. FIXED nsa_backend _to_2d_context_lens to HEAD's (N_total,1) form. nsa_indexer unsqueeze(-1) already yields (N_total,1) - consistent. Lesson: verify against HEAD/target-wheel, not the intro PR; wheel APIs drift between minor versions."
[[content.journal]]
date = "2026-06-07"
scope = "deepgemm-port"
content = "Complete-migration sweep (user: 要迁移就完整迁移). Verified EVERY deep_gemm symbol our code calls against HEAD/0.1.2. UNCHANGED+confirmed: fp8_mqa_logits (ks/ke cu_seqlens, not paged ctx - no 2D issue), fp8_gemm_nt, bf16_gemm_nt/nn, get_mk_alignment_for_contiguous_layout (no-arg), transform_sf_into_required_layout (same kwargs), get/set_num_sms, fp8_m_grouped_gemm_nt_masked (HEAD only ADDS optional recipe_a/b NVFP4 + _ensure_cuda, not required for FP8). PORTED beyond the 4 compat edits: (5) #26839 SBO masked-return unpack guard in moe_runner/deep_gemm.py (our line 388 had the unguarded unpack; our launch enables --enable-single-batch-overlap). (6) #23979 PDL-on-by-default in entrypoint.py (hasattr-guarded set_pdl(True), get_bool_env_var SGLANG_DEEPGEMM_PDL default true - perf default matching new wheel). (7) engine.py sglang-kernel assert 0.4.0->0.4.3 (migrated code requires sgl-deep-gemm-era wheel; (N_total,1) ctx would misbehave on old bundled DeepGEMM). SKIPPED with rationale (not on GLM-5.1-FP8/NSA/deepgemm/H200-SM90 path): #26025 fallback-unsupported-shapes (N/A - our tree has no _varlen_deep_gemm_silu_mul_quant/SGLANG_OPT_USE_JIT_EP_ACTIVATION), #25286 Gemma4 triton_scaled_mm scale layout (triton path not deepgemm), #22300 Minimax fp16/flashinfer-trtllm fallback, #26473/#17392 BF16 features, #26496/#24692 SM120/NVFP4, #26238/#25884 dsv4, AMD/ROCm. All 7 edited files py_compile clean. Runtime validation pending the torch-2.11 image (Phase 2/3)."
[[content.acceptance_criteria]]
text = "Port deep_gemm sgl-deep-gemm 0.1.2 API compat (2D->N_total,1 ctx_lens, compile-mode hasattr guard, positional m_indices, DLPack stride fix, SBO masked unpack guard, PDL default, kernel version assert)"
status = "pending"
category = "added"
[[content.acceptance_criteria]]
text = "Rebuild tai-kernel against torch 2.11 and import-smoke + harness-validate in the new dev-cu13 image"
status = "pending"
category = "added"
[[content.acceptance_criteria]]
text = "govctl check passes"
status = "pending"
category = "chore"