Avoid materializing GDN QKV tensors during target verification - #33778
Conversation
51814a4 to
c335e42
Compare
|
/tag-and-rerun-ci |
|
/rerun-failed-ci |
9 similar comments
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-skipped-ci |
43404d3 to
d3084ad
Compare
|
/rerun-skipped-ci |
|
/rerun-failed-ci |
4 similar comments
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
1 similar comment
|
/rerun-failed-ci |
| ) | ||
| if ( | ||
| (is_cuda() or is_hip() or is_xpu()) | ||
| and qkv_dim <= MAX_FUSED_QKV_SPLIT_DIM |
There was a problem hiding this comment.
when qkv_dim > MAX_FUSED_QKV_SPLIT_DIM, will the consumer side silently get wrong output?
There was a problem hiding this comment.
yes this can happen for larger models at low TP but it is pre-existing behavior; the qkv_dim > limit path already used strided split views before this change and this PR does not change that path.
| ) | ||
| query = query.view(1, actual_seq_len, layer.num_q_heads, layer.head_q_dim) | ||
| key = key.view(1, actual_seq_len, layer.num_k_heads, layer.head_k_dim) | ||
| value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim) |
There was a problem hiding this comment.
@Qiaolin-Yu should i add something like this here?
if not use_strided_target_verify_qkv:
query = query.contiguous()
key = key.contiguous()
value = value.contiguous()
|
/rerun-failed-ci |
2 similar comments
|
/rerun-failed-ci |
|
/rerun-failed-ci |
b8zhong
left a comment
There was a problem hiding this comment.
Hi, for performance benchmarking, we need fixed input/output benchmark, like bench_one_batch or similar, since we can't take the throughput number from GSM8K
|
hi @b8zhong, thanks for the review! Just added additional evidence for bench_one_batch with multiple batch sizes showing positive gains. please let me know if you need anything else! |
Avoid materializing contiguous Q, K, and V tensors when the selected target-verification kernel supports token-strided views. Preserve contiguous inputs for FlashInfer and CuTeDSL paths, and cover ReplaySSM routing and exact parity. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
438e02e to
d6b9744
Compare
|
/rerun-failed-ci |
2 similar comments
|
/rerun-failed-ci |
|
/rerun-failed-ci |
|
/rerun-failed-ci |
1 similar comment
|
/rerun-failed-ci |
…ernel strided q/k/v views Upstream 9fdb717 (main only, not in v0.5.20). On the 27B's D group every DFLASH verify step ran fused_qkv_split_gdn_prefill on every GDN layer, one copy kernel per layer per step, although the Triton verify kernel (fused_sigmoid_gating_delta_rule_update) reads q/k/v with their token strides (stride_q = q.stride()[1], runtime int, dense head axis). Target verify now passes torch.split views of the post-conv mixed_qkv when the chosen verify kernel opts in; prefill keeps the fused split (FLA chunk kernels want dense). Ported: LinearAttnKernelBase.supports_strided_target_verify_qkv (False), TritonGDNKernel opt-in, GDNKernelDispatcher._get_target_verify_kernel + target_verify_supports_strided_qkv, the routing in forward_extend. Adapted: the fork has no ReplaySSM spec-fold / circular verify routes (sgl-project#28695/sgl-project#35544 not on this line), so the routing is the dispatcher capability alone, decided inline per forward instead of in metadata init. The Triton opt-in is limited to the CUDA/HIP Triton kernel (NPU/CPU/XPU substitute other implementations in gdn_triton.py). Tree verify (topk>1) routes to the Triton tree kernel, which also opts in; FlashInfer/CuTe kernels inherit False. Tests (CPU, incg cgroup): test/registered/unit/layers/attention/test_gdn_verify_strided_qkv_33778.py - dispatcher capability (linear + tree Triton opt in, attribute-less kernel does not); - GDNAttnBackend.forward_extend TARGET_VERIFY: the verify kernel receives views sharing mixed_qkv's storage (token stride = qkv width), the fused split is never called; EXTEND still calls it; - the production verify kernel under TRITON_INTERPRET: strided views vs dense copies (ratio 3, 2 requests x 4 draft tokens, intermediate states cached, state update disabled) -> output and intermediate states bit-identical. 4 passed. On the base sources the capability and verify-routing tests fail.
…rides: fused split for head ratio 3 Upstream 8a1e6e4 (v0.5.20), the GDN part only, plus the row-stride hunk of upstream 52fecfd (sgl-project#37500) that the fork's port of sgl-project#37500 (e46a4db) left out. The 27B's GDN (16 k / 48 v heads, ratio 3) was outside the fused split list [1, 2, 4]: every decode / DFLASH verify step ran split + 2x .contiguous() + torch.cat + a z reshape copy on every GDN layer. The contiguous fused split kernel now walks a non-power-of-two head group one HEAD_V head at a time (V_POW2 constexpr prunes the dead branch; power-of-two groups keep the wide vector access), and qwen3_5 enables ratio 3 on CUDA (_GDN_FUSED_QKVZBA_RATIOS = (1, 2, 3, 4) on CUDA, (1, 2, 4) elsewhere). Prefill keeps the sgl-project#36267 strided views; verify then hands views on (sgl-project#33778). Row strides (sgl-project#37500): the kernel read rows at the dense widths. With finalize_fused_in_proj (called for the NF model, qwen4_exp; its GDN ratio is also 3) the split gets two column views of ONE GEMM output (row pitch = qkvz + ba width), so routing ratio 3 here without the strides would have read the wrong rows for every token after the first on NF. The 27B does not call finalize_fused_in_proj (dense rows); the wrapper now also asserts unit column strides. Not ported from sgl-project#34859 (not applicable to INT8 on sm86/sm120): sm120 FP8 GEMV, Hopper bf16 GEMV (SM90 only), SM90 FlashInfer GDN prefill default, NVFP4 SiLU+quant fusion, DSpark / qwen2_moe / deepseek / modelopt hunks. Tests (CPU, incg cgroup): test/registered/unit/layers/attention/test_gdn_fused_split_head_ratios_34859.py - production kernel under TRITON_INTERPRET vs plain slicing, ratios 1/2/3/4, uneven 27B rank-local heads 5/15 and 6/18, full 16/48, dense and fused-GEMM-strided rows: 14/14 bit-exact (worker run under incg; the shared 3 GiB cgroup OOM-killed the pytest wrapper of this class twice while other agents' suites ran). Base kernel: strided rows wrong for ratios 1/2/4, ratio 3 does not compile (arange not a power of two). - qwen3_5 dispatch: decode and target verify take the fused split for ratio 3, extend takes the views: 3 passed.
…ec ring 27B ReplaySSM package, slice S3 of REPLAYSSM_PLAN.md. With the flag off every path is unchanged (no rows planned -> recurrent route, intermediate commit). Ported: upstream GDNAttnBackend._replayssm_target_verify (same kernel call, launch_mode="verify", null block -1, request rows as replay indices) and the compact commit sequence of upstream's spec_utils (commit_gdn_replayssm_spec + commit_gdn_replayssm_circular + conv-window rollback). Adapted (this line's own wiring; upstream refuses GDN + DFLASH): - ForwardMetadata.replayssm_spec_rows: the verify metadata plans the request rows (req_pool_indices) in all three builders (eager, graph capture, graph replay -- there the static buffer with padded rows zeroed). The commit reads the rows the verify used from the metadata, like mamba_cache_indices, so no caller changes: DFLASH (dflash_worker_v2 untouched), the EAGLE/MTP commit and the lane all go through update_mamba_state_after_mtp_verify. - The commit sits in that shared hook (ring branch when intermediate_ssm is None). Linear chain -> accepted = last step + 1 (bonus included); -1 folds nothing, like the masked scatter. - FOLD EVERY COMMIT for any SSM dtype (upstream defers the fold for fp32): radix insert, HiCache backup, tail adopt and the flip all read `temporal` and assume it is the committed state after every verify, as the recurrent route makes it. The ring is then per-step scratch (write_pos 0 at every verify). - Heal: outside a capture the verify metadata re-states write_pos = is_flush = 0 for the step's rows (true by construction after every commit), so a verify never trusts cursor bytes from an earlier step (TMS restore, flip, reused row). cache_base is circular and harmless at write_pos 0. - forward_extend routes the verify to the ring when rows are planned; a draft tree on the ring, or a pool with neither intermediate state nor planned rows, is refused. The ring kernel honours token strides, so the verify split stays torch.split views (sgl-project#33778). - Advance cursor null block 0 (request row 0 is the ReqToTokenPool padding row); the widest verify window is kept on the pool as the advance's constant (one compiled variant across the adaptive ladder). - Ring length must be >= 16 (the compact commit's tl.dot) -- pool and static check; S2 test moved to L=16. - _replayssm_spec_for: a draft-KV-only producer (no verify workspace, sgl-project#1233 FIX 4) treats the flag as a no-op instead of hitting the pool's missing-window refusal. - getattr on the new metadata field in the two readers (sgl-project#624 stub drift: test_gdn_verify_strided_qkv_33778 doubles predate it). Tests (CPU, incg cgroup, one file at a time): test/registered/unit/layers/attention/test_replayssm_spec_route_s3.py: 10 passed. Hermetic: rows planned + healed (eager, replay; not in capture; padded row -> row 0), flag off / decode plan nothing, the three refusals, producer no-op. Interpreter (production route functions on real CPU pools, ring vs intermediate, 2 layers, head ratio 3, 2 requests + padded row, track crossing, second step after garbage cursors + heal): fp32 verify rel 3.7e-7 / 2.4e-7, commit (incl. track slot) 2.9e-7 / 3.1e-7 vs the recurrent route; fp16 (bf16 stand-in) verify 5.0e-4 vs recurrent 3.8e-4 of the fp32 truth, commit 3.6e-4 == recurrent 3.6e-4; conv rollback equal; untouched slots bit-exact; cursors 0 after every commit. On the S2 sources: 7 failed + 3 errors. Regression: test_replayssm_spec_ring_pool_s2 10, test_mamba_checkpoint_interval 58, test_dflash_mamba_track_post_verify_37818 5, test_conv_verify_private_ window_444 13, test_prefill_graph_stale_track_rows_34184 2, test_weg2_prefill_only_capture_1233 15, test_forward_metadata_plan_record 9, test_fused_replay_state_indices_32219 3, test_gdn_verify_strided_qkv_33778 4, test_verify_intermediate_row_ownership_450 12, test_mamba2_conv_verify_ private_window_450 10, test_dual_group_concurrency 159 -- all passed. ruff F: no new findings.
…ec ring 27B ReplaySSM package, slice S3 of REPLAYSSM_PLAN.md. With the flag off every path is unchanged (no rows planned -> recurrent route, intermediate commit). Ported: upstream GDNAttnBackend._replayssm_target_verify (same kernel call, launch_mode="verify", null block -1, request rows as replay indices) and the compact commit sequence of upstream's spec_utils (commit_gdn_replayssm_spec + commit_gdn_replayssm_circular + conv-window rollback). Adapted (this line's own wiring; upstream refuses GDN + DFLASH): - ForwardMetadata.replayssm_spec_rows: the verify metadata plans the request rows (req_pool_indices) in all three builders (eager, graph capture, graph replay -- there the static buffer with padded rows zeroed). The commit reads the rows the verify used from the metadata, like mamba_cache_indices, so no caller changes: DFLASH (dflash_worker_v2 untouched), the EAGLE/MTP commit and the lane all go through update_mamba_state_after_mtp_verify. - The commit sits in that shared hook (ring branch when intermediate_ssm is None). Linear chain -> accepted = last step + 1 (bonus included); -1 folds nothing, like the masked scatter. - FOLD EVERY COMMIT for any SSM dtype (upstream defers the fold for fp32): radix insert, HiCache backup, tail adopt and the flip all read `temporal` and assume it is the committed state after every verify, as the recurrent route makes it. The ring is then per-step scratch (write_pos 0 at every verify). - Heal: outside a capture the verify metadata re-states write_pos = is_flush = 0 for the step's rows (true by construction after every commit), so a verify never trusts cursor bytes from an earlier step (TMS restore, flip, reused row). cache_base is circular and harmless at write_pos 0. - forward_extend routes the verify to the ring when rows are planned; a draft tree on the ring, or a pool with neither intermediate state nor planned rows, is refused. The ring kernel honours token strides, so the verify split stays torch.split views (sgl-project#33778). - Advance cursor null block 0 (request row 0 is the ReqToTokenPool padding row); the widest verify window is kept on the pool as the advance's constant (one compiled variant across the adaptive ladder). - Ring length must be >= 16 (the compact commit's tl.dot) -- pool and static check; S2 test moved to L=16. - _replayssm_spec_for: a draft-KV-only producer (no verify workspace, sgl-project#1233 FIX 4) treats the flag as a no-op instead of hitting the pool's missing-window refusal. - getattr on the new metadata field in the two readers (sgl-project#624 stub drift: test_gdn_verify_strided_qkv_33778 doubles predate it). Tests (CPU, incg cgroup, one file at a time): test/registered/unit/layers/attention/test_replayssm_spec_route_s3.py: 10 passed. Hermetic: rows planned + healed (eager, replay; not in capture; padded row -> row 0), flag off / decode plan nothing, the three refusals, producer no-op. Interpreter (production route functions on real CPU pools, ring vs intermediate, 2 layers, head ratio 3, 2 requests + padded row, track crossing, second step after garbage cursors + heal): fp32 verify rel 3.7e-7 / 2.4e-7, commit (incl. track slot) 2.9e-7 / 3.1e-7 vs the recurrent route; fp16 (bf16 stand-in) verify 5.0e-4 vs recurrent 3.8e-4 of the fp32 truth, commit 3.6e-4 == recurrent 3.6e-4; conv rollback equal; untouched slots bit-exact; cursors 0 after every commit. On the S2 sources: 7 failed + 3 errors. Regression: test_replayssm_spec_ring_pool_s2 10, test_mamba_checkpoint_interval 58, test_dflash_mamba_track_post_verify_37818 5, test_conv_verify_private_ window_444 13, test_prefill_graph_stale_track_rows_34184 2, test_weg2_prefill_only_capture_1233 15, test_forward_metadata_plan_record 9, test_fused_replay_state_indices_32219 3, test_gdn_verify_strided_qkv_33778 4, test_verify_intermediate_row_ownership_450 12, test_mamba2_conv_verify_ private_window_450 10, test_dual_group_concurrency 159 -- all passed. ruff F: no new findings. 27B line (desk/27b-up-replayssm-line-0924): cherry-picked from 7d69bd8 (desk/27b-up-replayssm-0924, on the NF-based staging line fa757eb), clean. gdn_backend.py and mamba2_metadata.py are identical on both lines; dflash_worker_v2 differs here (sgl-project#31468, sgl-project#33459/sgl-project#30096 ports) but not in _update_target_mamba_state_after_verify, so the DFLASH commit reaches the shared hook exactly as on NF. One adaptation, text only: the commit docstring listed "tail adopt" among the readers of `temporal` -- that is the NF line's H21 install (weg2.tail_adopt), absent here; the 27B readers are the radix insert, the HiCache backup (the weg2 L2 arena write, xsn351) and the P/D flip. Tests on this line (CPU, incg cgroup, S3 tree, one file at a time): test_replayssm_spec_route_s3.py 10 passed -- interpreter, production route vs the recurrent route: fp32 verify rel 3.74e-7 / 2.40e-7 (step 1 / step 2 after garbage cursors + heal), commit incl. the track slot 2.94e-7 / 3.10e-7 (all < 4e-7); fp16 (bf16 stand-in) verify 5.03e-4 / 2.72e-4 vs recurrent 3.84e-4 / 2.55e-4 of the fp32 truth, commit 3.61e-4 / 3.27e-4 == recurrent; conv rollback equal, padded outputs zero, untouched slots bit-exact, cursors 0 after every commit; test_replayssm_spec_ring_pool_s2.py 10 passed (8 subtests). Regression (package tip = S6 tree, same cgroup, one file at a time): test_dflash_mamba_track_post_verify_37818 5, test_conv_verify_private_window_444 13, test_prefill_graph_stale_track_rows_34184 2, test_weg2_prefill_only_capture_1233 15, test_forward_metadata_plan_record 9, test_fused_replay_state_indices_32219 3, test_gdn_verify_strided_qkv_33778 4, test_verify_intermediate_row_ownership_450 12, test_mamba2_conv_verify_private_window_450 10, test_dual_group_concurrency 159, test_mamba_checkpoint_interval 58 -- all passed.
…ec ring 27B ReplaySSM package, slice S3 of REPLAYSSM_PLAN.md. With the flag off every path is unchanged (no rows planned -> recurrent route, intermediate commit). Ported: upstream GDNAttnBackend._replayssm_target_verify (same kernel call, launch_mode="verify", null block -1, request rows as replay indices) and the compact commit sequence of upstream's spec_utils (commit_gdn_replayssm_spec + commit_gdn_replayssm_circular + conv-window rollback). Adapted (this line's own wiring; upstream refuses GDN + DFLASH): - ForwardMetadata.replayssm_spec_rows: the verify metadata plans the request rows (req_pool_indices) in all three builders (eager, graph capture, graph replay -- there the static buffer with padded rows zeroed). The commit reads the rows the verify used from the metadata, like mamba_cache_indices, so no caller changes: DFLASH (dflash_worker_v2 untouched), the EAGLE/MTP commit and the lane all go through update_mamba_state_after_mtp_verify. - The commit sits in that shared hook (ring branch when intermediate_ssm is None). Linear chain -> accepted = last step + 1 (bonus included); -1 folds nothing, like the masked scatter. - FOLD EVERY COMMIT for any SSM dtype (upstream defers the fold for fp32): radix insert, HiCache backup, tail adopt and the flip all read `temporal` and assume it is the committed state after every verify, as the recurrent route makes it. The ring is then per-step scratch (write_pos 0 at every verify). - Heal: outside a capture the verify metadata re-states write_pos = is_flush = 0 for the step's rows (true by construction after every commit), so a verify never trusts cursor bytes from an earlier step (TMS restore, flip, reused row). cache_base is circular and harmless at write_pos 0. - forward_extend routes the verify to the ring when rows are planned; a draft tree on the ring, or a pool with neither intermediate state nor planned rows, is refused. The ring kernel honours token strides, so the verify split stays torch.split views (sgl-project#33778). - Advance cursor null block 0 (request row 0 is the ReqToTokenPool padding row); the widest verify window is kept on the pool as the advance's constant (one compiled variant across the adaptive ladder). - Ring length must be >= 16 (the compact commit's tl.dot) -- pool and static check; S2 test moved to L=16. - _replayssm_spec_for: a draft-KV-only producer (no verify workspace, sgl-project#1233 FIX 4) treats the flag as a no-op instead of hitting the pool's missing-window refusal. - getattr on the new metadata field in the two readers (sgl-project#624 stub drift: test_gdn_verify_strided_qkv_33778 doubles predate it). Tests (CPU, incg cgroup, one file at a time): test/registered/unit/layers/attention/test_replayssm_spec_route_s3.py: 10 passed. Hermetic: rows planned + healed (eager, replay; not in capture; padded row -> row 0), flag off / decode plan nothing, the three refusals, producer no-op. Interpreter (production route functions on real CPU pools, ring vs intermediate, 2 layers, head ratio 3, 2 requests + padded row, track crossing, second step after garbage cursors + heal): fp32 verify rel 3.7e-7 / 2.4e-7, commit (incl. track slot) 2.9e-7 / 3.1e-7 vs the recurrent route; fp16 (bf16 stand-in) verify 5.0e-4 vs recurrent 3.8e-4 of the fp32 truth, commit 3.6e-4 == recurrent 3.6e-4; conv rollback equal; untouched slots bit-exact; cursors 0 after every commit. On the S2 sources: 7 failed + 3 errors. Regression: test_replayssm_spec_ring_pool_s2 10, test_mamba_checkpoint_interval 58, test_dflash_mamba_track_post_verify_37818 5, test_conv_verify_private_ window_444 13, test_prefill_graph_stale_track_rows_34184 2, test_weg2_prefill_only_capture_1233 15, test_forward_metadata_plan_record 9, test_fused_replay_state_indices_32219 3, test_gdn_verify_strided_qkv_33778 4, test_verify_intermediate_row_ownership_450 12, test_mamba2_conv_verify_ private_window_450 10, test_dual_group_concurrency 159 -- all passed. ruff F: no new findings. 27B line (desk/27b-up-replayssm-line-0924): cherry-picked from 7d69bd8 (desk/27b-up-replayssm-0924, on the NF-based staging line fa757eb), clean. gdn_backend.py and mamba2_metadata.py are identical on both lines; dflash_worker_v2 differs here (sgl-project#31468, sgl-project#33459/sgl-project#30096 ports) but not in _update_target_mamba_state_after_verify, so the DFLASH commit reaches the shared hook exactly as on NF. One adaptation, text only: the commit docstring listed "tail adopt" among the readers of `temporal` -- that is the NF line's H21 install (weg2.tail_adopt), absent here; the 27B readers are the radix insert, the HiCache backup (the weg2 L2 arena write, xsn351) and the P/D flip. Tests on this line (CPU, incg cgroup, S3 tree, one file at a time): test_replayssm_spec_route_s3.py 10 passed -- interpreter, production route vs the recurrent route: fp32 verify rel 3.74e-7 / 2.40e-7 (step 1 / step 2 after garbage cursors + heal), commit incl. the track slot 2.94e-7 / 3.10e-7 (all < 4e-7); fp16 (bf16 stand-in) verify 5.03e-4 / 2.72e-4 vs recurrent 3.84e-4 / 2.55e-4 of the fp32 truth, commit 3.61e-4 / 3.27e-4 == recurrent; conv rollback equal, padded outputs zero, untouched slots bit-exact, cursors 0 after every commit; test_replayssm_spec_ring_pool_s2.py 10 passed (8 subtests). Regression (package tip = S6 tree, same cgroup, one file at a time): test_dflash_mamba_track_post_verify_37818 5, test_conv_verify_private_window_444 13, test_prefill_graph_stale_track_rows_34184 2, test_weg2_prefill_only_capture_1233 15, test_forward_metadata_plan_record 9, test_fused_replay_state_indices_32219 3, test_gdn_verify_strided_qkv_33778 4, test_verify_intermediate_row_ownership_450 12, test_mamba2_conv_verify_private_window_450 10, test_dual_group_concurrency 159, test_mamba_checkpoint_interval 58 -- all passed. NF line (H64, onto 88dec66): one conflict, resolved line by line in gdn_backend.forward_extend. This line carries no sgl-project#33778 port (no kernel_dispatcher.target_verify_supports_strided_qkv), so the 27B hunk that widens use_strided_verify_qkv for the ring has nothing to widen. Kept the NF condition `(is_cuda() or is_hip()) and qkv_dim <= MAX_FUSED_QKV_SPLIT_DIM`: the ring verify gets the same q/k/v preparation as the recurrent verify on this line (dense [1, T, H, D] from fused_qkv_split_gdn_prefill on CUDA, torch.split views elsewhere); _replayssm_target_verify flattens both layouts by numel. Text only: the commit docstring names the NF H21 tail adopt (weg2/tail_adopt.py reads and installs cache.temporal) among the readers of temporal again -- the 27B line had dropped it because it lacks H21. Tests (CPU, incg cgroup, one file at a time): test_replayssm_spec_route_s3.py 10 passed; test_replayssm_spec_ring_pool_s2.py 10 passed (8 subtests). (cherry picked from commit 974254d)
…cks, chronologisch) Grundlage: Präsenz-Scan aller 230 27B-Commits seit 76f8deb gegen diesen Baum (Stichprobe der hinzugefügten Zeilen je Commit); die 94 fehlenden minus die bewusst anders gewählten Formen (76e87ac/4ae11ababd -> S3 form.calibration_identity; 479f6ec/d7f588e017/d0fba8955f/34892e3017 -> S2 NF-Formen; 7f81f09/3c14481318 -> S4/S7a; 3dbb790 line_gate_27b (Werkzeug, Schritt 9); 8604d13 W100-by-name (Nutzer: bleibt aus); 6545e2c flashinfer-Pin in pyproject (Image-Frage, nicht Baum)). Liste: 92bbccb eb5d044 829ebd0 431fcbc ef4d11f 6816062 a233e50 2cc593c f0c8451 87cc4fb c529777 31f2dbe c50085a 3301a96 036b368 e1d1fe9 03c68af 6dddc2e 06932b5 87389c4 58a7490 f85ac55 fc64aa5 5aa24dd 97c0e9a 159333c d9f1532 f3c685b 8550655 e50fb59 db2c2ef f09dc0d c255e10 51b810e 28a55a2 34965fc ff3d9cc 340a018 bee5e10 67b6352 fdade85 ed6630d f1c9a43 b434831 517f26d 0b6b60b a40837f 644de86 aff2b7c 197b701 856024b 238512a 9738626 b857a22 1f8c24d d294b3e 810239d b429dfd e714c95 9efd974 3d63e0a d3cfcf3 fee6134 7985b56 49a14e9 fc45706 19c720e 5306bee 6f1235a c98eaa3 93bc802 328349e f9fb3a2 572af73 94fa8b4 d342caa 2fd7d7e 3ebbb96 871d55f 78c2f16 3babf51 196f6a8 e70af54 22eccfc Inhalt: Upstream-Ports (sgl-project#33758 sgl-project#37818 sgl-project#36738 sgl-project#33459/sgl-project#30096 sgl-project#34446 sgl-project#36267 sgl-project#33778 sgl-project#34859 sgl-project#36415 sgl-project#35255/sgl-project#36638 sgl-project#39858/sgl-project#40259 sgl-project#31417 sgl-project#34892 sgl-project#32225 sgl-project#30832/sgl-project#36626 sgl-project#39574 sgl-project#29579 sgl-project#31468 sgl-project#32575 sgl-project#31648); xsn409/410-Wake-Verdikte; Vision-Linie V1-V3b + xsn438 (SGLANG_WEG2_VISION_FLIP_URGENT); D-Planer L6 (159333c); DFLASH-Window-Pool sync-frei, PLAN_SYNC_FREE, D-Kollektive (vocab-argmax, a2a-Merge, deferred rebuild), #DGAP/D_DEFER_SEQ_LENS_CPU; Mamba-Anker Raster 4096 + Per-Path-Cap + Inner-Release; P-TRIM (--p-trim-end-anchor); FP8 uniform Marlin; ModelOpt/NVFP4 RadixArk; GGUF G1-G6 + F1/F2; native-mixed sgl-project#38 (sm_8x W4A8, sm_12x CUTLASS/W4A16); RC1-Capture-Set; sgl-project#49 Agent-Turns; dynchunk (--p-chunk-policy, --p-chunk-dynamic-min-tokens). Auflösungen (Gabel -> Form, Grund): - L6 d_operating_point_rows: 27B (d) "Token-Vektor auf jeder Position aus der Kapazität" nur bei TP-symmetrischem D (Profil d_layout paged_dcp); sonst NF-sgl-project#1293-Pin + NF-Anker- Klausel. mamba_ssm_dtype aus EARLY_READ_FACTS nur bei Profil early_read_flags. Overhead-Kalibrierung liest mit form.CalibrationIdentity statt LineIdentity. - RC1 Capture-Set: neuer RecordKey-Term d_capture_set (qwen27b), Leser CalibrationIdentity.d_max_running_requests; nextflash unverändert. - URC Carrier-Hold: 27B _weg2_carrier_hold entfällt (S2 NF-Rotation), Inner-Release und Per-Path-Cap bleiben (Env, Default aus; 27b.env setzt sie). - scheduler_pp_mixin/overlap_utils/batch_result_processor: NF H49/H58 und 27B #PGAP/#DGAP komponiert (beide Instrumente getrennt schaltbar). - schedule_policy: P-TRIM-Kurzschluss vor NF H63-Fold/QSA-Korn; sgl-project#36415 Hoist + NF computed_input_len. - gdn_backend sgl-project#33778: 27B-strided-Verify; NF-Ring flacht beide Layouts ab. - flashinfer_backend: RC9-Datei + NF-Form-A-Waiver (27B-intern mehrfach gegabelt). - checkpoint_census: GGUF-Leser + NF exclude_segments (PLE) in einer Aggregation. - xchg_manifest: FLAT_SEGMENTS in beiden (dst/src) NF-Breitenbedingungen ausgenommen. - vram_peak_window: NF-Kumulativ-Peak liest über den 27B-Fast-Read. - FP8: 8c86eb8-Rest nachgezogen (private Workspace-Registry entfernt, wie RC9). - argv_d: vision= an allen drei Aufrufstellen (inkl. NF --d-only). - census_checkpoint_decision (W161) jetzt für beide Profile aktiv. Gates: py_compile aller geänderten Dateien; ruff F821/F811 ohne neue Funde gegenüber dem Vorgänger (PendingSeqLensCpu ist String-Annotation wie in RC9); dup_defs_gate 0 neu. Tests: 73 portierte Testdateien, Lauf nach dem Ruhefenster (Boot aktiv).
Brings in 113 upstream commits (7ad55e4..d34f7b2). Upstream now contains three stack PRs as merged: sgl-project#40501, sgl-project#40175, sgl-project#33778; every line their squashes add is already present in the stack. Conflicts resolved: - arg_groups/fields/memory.py, managers/cache_controller.py: union of the stack's fast_file backend (sgl-project#39880) and upstream's tensorcast backend. - layers/attention/qsa/mqa.py: keep the stack's SM120 Triton imports and fp8 _scoring_dtype path; adopt upstream's ROCm 16-wide head alignment (sgl-project#38875) and its is_hip import. - models/qwen4_exp.py: all six hunks are stack-only additions (PLE host staging, CUDA-graph prewarm, replay prepare) against upstream's final sgl-project#40501; kept ours. The merged file equals the pre-merge stack version. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Summary
This PR removes redundant Q/K/V materialization from compatible GDN
speculative target-verification paths.
causal_conv1d_updatealready produces packed QKV. Previously, every GDN layerlaunched
fused_qkv_split_gdn_prefill_kernelto copy that output into threecontiguous tensors before Triton target verification. Triton accepts explicit
token strides, so it can consume zero-copy
torch.split/view tensors instead.The optimization is route-aware:
retain materialization.
and sampling are unchanged.
This eliminates one memory-copy kernel launch per GDN layer on every compatible
target-verification pass.
Profile evidence
Matched H200 CPU/GPU traces using Qwen3.5-4B, ReplaySSM, and NEXTN T=3:
The graph-span reduction's approximate 95% interval was
79.732-103.841 us.
Fixed input/output benchmark
The throughput results below use
sglang.benchmark.one_batch_server, notGSM8K. Unlike the in-process
one_batchpath, this runs against a realReplaySSM/NEXTN server and therefore exercises the speculative target-
verification route changed by this PR.
The baseline and patch were built from the same PR head; the baseline reverses
only this PR's six-file diff. Two alternating baseline/patch pairs ran on the
same H200 with a fresh server for every phase:
Each phase used one warmup followed by five measured single-batch requests per
shape. Inputs were exact random token IDs, outputs used exact lengths with EOS
ignored, and the server and benchmark seeds were fixed. The table aggregates
ten measurements per arm and shape.
At the representative 256/256 batch-64 shape, the paired-run 95% intervals
were +0.351% to +0.673% for output throughput and -0.642% to -0.330%
for total latency. Speculative acceptance was identical in every matched run.
Prior GSM8K end-to-end validation
These earlier GSM8K results are retained as secondary end-to-end and quality
context. Because GSM8K has variable output lengths, its throughput numbers are
not used as the primary performance evidence for this PR; the fixed
input/output results above are the primary benchmark.
Three alternating baseline/patch pairs ran on the same H200 with fresh
servers, seed 0, an empty prefix cache, and a 30-second cooldown. The benchmark
used its defaults: 200 questions, five shots, 512 maximum output tokens,
temperature 0, and parallelism 64.
Pair 1
Pair 2
Pair 3
Three-pair aggregate
GSM8K reproduction
Start the server:
CUDA_VISIBLE_DEVICES=0 \ sglang serve Qwen/Qwen3.5-4B \ --port 30000 \ --dtype bfloat16 \ --language-only \ --limit-mm-data-per-request '{"image":0,"video":0,"audio":0}' \ --context-length 32768 \ --mem-fraction-static 0.8 \ --max-running-requests 64 \ --linear-attn-decode-backend triton \ --random-seed 0 \ --speculative-algorithm NEXTN \ --speculative-draft-model-path Qwen/Qwen3.5-4B \ --speculative-num-steps 3 \ --speculative-eagle-topk 1 \ --speculative-num-draft-tokens 4 \ --enable-linear-replayssm-specRun the official SGLang GSM8K benchmark:
Dataset SHA-256:
Reproduction
The measurements above used PR head
f0b834be9b89f316f2daed9dc6f78e06df300e4band base referenced12b313b93e1547d9b02c3a84426aa88519fc494. The two benchmark trees can beconstructed without checking out an older, otherwise different SGLang tree:
The measured software versions were Python 3.12.6, PyTorch 2.13.0+cu130,
Triton 3.7.1, Transformers 5.12.1, and sglang-kernel 0.4.6.post1. A fresh
environment can be prepared from the PR tree with:
Start the server from either the baseline or patch source tree:
Run one warmup and five measured iterations for each fixed shape:
Restart the server between baseline and patch phases and run the order shown
above.
The following is the complete four-phase command used to automate that process.
It saves and terminates only the server process group it starts:
Aggregate the ten measurements per arm and verify acceptance parity:
Generate a fixed-shape CPU/GPU trace against either running phase with:
Run the focused H200 correctness suite from the patch tree:
Correctness validation
context; the fixed input/output benchmark is the primary performance
evidence.
CI States
Latest PR Test (Base): ✅ Run #35563025383
Latest PR Test (Extra): ❌ Run #35563025315
Latest PR Test (AMD ROCm 10): ⏳ Run #35563025321