[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:
co-authored by
Xinyuan Tong
zRzRzRzRzRzRzR
Shijin Zhang
zanes-ops
parent
288627e400
commit
a66451c058
@@ -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()
|
||||
Reference in New Issue
Block a user