[GLM-5.3 Flash] Restore and enable KPool metadata fusion (#38845)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2@liverpool.ac.uk>
Co-authored-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com>
Co-authored-by: zanes-ops <zanes@nvidia.com>
This commit is contained in:
Baizhou Zhang
2026-09-12 16:01:11 -07:00
committed by GitHub
co-authored by Xinyuan Tong zRzRzRzRzRzRzR Shijin Zhang zanes-ops
parent 288627e400
commit a66451c058
12 changed files with 1580 additions and 145 deletions
@@ -0,0 +1,98 @@
"""KPool fused metadata must retain live tails and refresh captured buffers."""
import unittest
import torch
from sglang.kernels.ops.attention.dsa_kpool_metadata.verify import (
fused_dsa_target_verify_metadata,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")
class TestKPoolMetadataFusion(CustomTestCase):
def test_verify_replay_boundaries_and_request_remapping(self):
device = "cuda"
bs, next_n, width, topk, pool_size = 4, 6, 16384, 2048, 4
seq = torch.tensor([1, 61, 2047, 8191], device=device, dtype=torch.int64)
req = torch.tensor([3, 1, 6, 0], device=device, dtype=torch.int64)
table = torch.arange(8 * width, device=device, dtype=torch.int32).view(8, width)
def empty(*shape):
return torch.full(shape, -1, device=device, dtype=torch.int32)
buffers = dict(
cache_seqlens=empty(bs),
cu_seqlens_k=empty(bs + 1),
page_table_1=empty(bs * next_n, width),
seqlens_expanded=empty(bs * next_n),
dsa_cache_seqlens=empty(bs * next_n),
dsa_cu_seqlens_k=empty(bs * next_n + 1),
real_page_table=empty(bs * next_n, width // 64),
paged_mqa_ctx_lens_2d=empty(bs, next_n),
)
addresses = {key: value.data_ptr() for key, value in buffers.items()}
def refresh():
fused_dsa_target_verify_metadata(
seq_lens=seq,
req_pool_indices=req,
req_to_token=table,
bs=bs,
max_seqlen_k=width,
dsa_index_topk=topk,
real_page_size=64,
next_n=next_n,
index_kpool=pool_size,
**buffers,
)
refresh()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
refresh()
for lengths, requests in [
([2, 63, 2048, 8193], [0, 6, 1, 3]),
([64, 128, 2051, 8190], [5, 2, 7, 4]),
([0, 3, 2045, 9000], [7, 0, 3, 2]),
]:
seq.copy_(torch.tensor(lengths, device=device))
req.copy_(torch.tensor(requests, device=device))
graph.replay()
expanded = (
seq[:, None] + torch.arange(1, next_n + 1, device=device)
).flatten()
expected = torch.minimum(expanded, topk + expanded % pool_size).int()
torch.testing.assert_close(buffers["seqlens_expanded"], expanded.int())
torch.testing.assert_close(buffers["dsa_cache_seqlens"], expected)
torch.testing.assert_close(
buffers["dsa_cu_seqlens_k"][1:], expected.cumsum(0).int()
)
torch.testing.assert_close(buffers["cache_seqlens"], (seq + next_n).int())
torch.testing.assert_close(
buffers["paged_mqa_ctx_lens_2d"],
(seq + next_n).int()[:, None].expand(bs, next_n),
)
expected_pages = table[req].repeat_interleave(next_n, dim=0)
row_lens = (seq + next_n).repeat_interleave(next_n)
live = torch.arange(width, device=device)[None, :] < row_lens[:, None]
torch.testing.assert_close(
buffers["page_table_1"][live], expected_pages[live]
)
real_live = (
torch.arange(0, width, 64, device=device)[None, :] < row_lens[:, None]
)
torch.testing.assert_close(
buffers["real_page_table"][real_live],
(expected_pages[:, ::64] // 64)[real_live],
)
self.assertEqual(
addresses, {key: value.data_ptr() for key, value in buffers.items()}
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,117 @@
"""Fused KPool replay and MTP sibling copies preserve captured buffer identity."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.dsa_metadata_kit import (
BS,
NEXT_N,
ROUNDS,
addresses,
apply_metadata,
assert_metadata_equal,
inputs,
make_backend,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
class TestDSAMetadataReplay(CustomTestCase):
def test_fusion_matches_ordinary_metadata(self):
for mode in (
ForwardMode.DECODE,
ForwardMode.TARGET_VERIFY,
ForwardMode.DRAFT_EXTEND_V2,
):
with self.subTest(mode=mode):
seq, req = inputs(*ROUNDS[0])
fused = make_backend(mode, seq, req)
ordinary = make_backend(mode, seq, req, fusion=False)
pointers = addresses(fused.forward_metadata)
for lengths, requests in ROUNDS:
seq.copy_(torch.tensor(lengths, device="cuda"))
req.copy_(torch.tensor(requests, device="cuda"))
spec = None
if mode.is_draft_extend_v2():
spec = SimpleNamespace(
num_accept_tokens=torch.tensor(
[1, 2, 5, NEXT_N], device="cuda", dtype=torch.int32
)
)
apply_metadata(fused, mode, seq, req, spec)
apply_metadata(ordinary, mode, seq, req, spec)
assert_metadata_equal(
self, fused.forward_metadata, ordinary.forward_metadata
)
self.assertEqual(pointers, addresses(fused.forward_metadata))
def test_precomputed_verify_retains_live_tail(self):
mode = ForwardMode.TARGET_VERIFY
seq, req = inputs(*ROUNDS[0])
fused = make_backend(mode, seq, req)
ordinary = make_backend(mode, seq, req, fusion=False)
pointers = addresses(fused.forward_metadata)
for lengths, requests in ROUNDS[1:]:
seq.copy_(torch.tensor(lengths, device="cuda"))
req.copy_(torch.tensor(requests, device="cuda"))
precomputed = fused._precompute_replay_metadata(
BS, req, seq, seq.cpu(), mode
)
fused.init_forward_metadata_replay_cuda_graph_from_precomputed(
BS, precomputed, mode
)
apply_metadata(ordinary, mode, seq, req)
assert_metadata_equal(
self, fused.forward_metadata, ordinary.forward_metadata
)
self.assertEqual(pointers, addresses(fused.forward_metadata))
def test_precomputed_and_sibling_copy_refresh_derived_metadata(self):
mode = ForwardMode.DECODE
seq, req = inputs(*ROUNDS[0])
source = make_backend(mode, seq, req)
sibling = make_backend(mode, seq, req)
ordinary = make_backend(mode, seq, req, fusion=False)
pointers = addresses(sibling.forward_metadata)
for lengths, requests in ROUNDS[1:]:
seq.copy_(torch.tensor(lengths, device="cuda"))
req.copy_(torch.tensor(requests, device="cuda"))
precomputed = source._precompute_replay_metadata(
BS, req, seq, seq.cpu(), mode
)
source.init_forward_metadata_replay_cuda_graph_from_precomputed(
BS, precomputed, mode
)
# An eligible sibling must reuse the derived results, not silently
# fall through to the full recomputation path.
with patch.object(
sibling,
"init_forward_metadata_replay_cuda_graph_from_precomputed",
side_effect=AssertionError("unexpected sibling fallback"),
):
sibling._copy_replay_metadata_from_sibling(
source, BS, precomputed, mode
)
apply_metadata(ordinary, mode, seq, req)
assert_metadata_equal(
self, source.forward_metadata, ordinary.forward_metadata
)
assert_metadata_equal(
self, sibling.forward_metadata, ordinary.forward_metadata
)
self.assertEqual(pointers, addresses(sibling.forward_metadata))
self.assertIsNot(
sibling.forward_metadata.kpool_write_plan,
source.forward_metadata.kpool_write_plan,
)
if __name__ == "__main__":
unittest.main()