[Refactor] Clarify DeepSeek V4 metadata names for V4.1 (#38947)

This commit is contained in:
Liangsheng Yin
2026-09-10 15:42:52 -07:00
committed by GitHub
parent d076eec427
commit dc5f59c3a2
16 changed files with 123 additions and 220 deletions
@@ -365,7 +365,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
metadata.c128_topk_lengths_clamp1 = torch.tensor(
[base + 39, base + 40], dtype=torch.int32
)
metadata.c1_flashmla_metadata = object()
metadata.c0_flashmla_metadata = object()
metadata.c4_flashmla_metadata = object()
metadata.c128_flashmla_metadata = object()
return metadata
@@ -377,10 +377,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
)
from sglang.srt.server_args import ServerArgs
# cg-refactor folded the legacy enable_breakable_cuda_graph flag
# into cuda_graph_config. Verify the per-phase backend selectors
# default to None (i.e. nothing opted into BREAKABLE without an
# explicit CLI flag).
# Breakable graphs require explicit opt-in for each phase.
sa = ServerArgs(model_path="dummy")
self.assertNotEqual(sa.cuda_graph_backend_decode, "breakable")
self.assertNotEqual(sa.cuda_graph_backend_prefill, "breakable")
@@ -519,7 +516,7 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
"swa_topk_lengths",
"c128_page_indices",
"c128_topk_lengths_clamp1",
"c1_flashmla_metadata",
"c0_flashmla_metadata",
"c4_flashmla_metadata",
"c128_flashmla_metadata",
]
@@ -688,13 +685,8 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
class TestDSV4SwaOutCacheLocResolution(CustomTestCase):
"""`get_swa_out_cache_loc`: cached fast path vs store-time fallback.
The KV-store consumers run in paths that never invoke
`init_forward_metadata_in_graph` (eager idle, runners that only run the
out-graph prep) or whose batch is re-padded after init (DP attention).
The resolver must use the per-forward cached value only when it is
provably current and fall back to translating `out_cache_loc` otherwise.
"""SWA writes must translate live locations for idle or missing/mismatched caches.
A matching cache on an active forward must be reused.
"""
def _make_backend(self, mapping: torch.Tensor):
@@ -654,7 +654,7 @@ def _make_dsv4_target(*, unified, mapping=None):
pool.get_unified_swa_ring_buf_infos = lambda: (
_buf_infos(12) if unified else ([], [], [])
)
pool.get_c128_state_buf_infos = lambda: ([], [], [])
pool.get_request_state_buf_infos = lambda: ([], [], [])
return pool
@@ -46,7 +46,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
force_deep_gemm_metadata=True,
)
@@ -68,7 +68,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
)
@@ -84,7 +84,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
use_topk_v2=False,
)
@@ -253,7 +253,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
extend_start_loc=torch.tensor([0], dtype=torch.int32),
extend_num_tokens=query_rows,
)
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)
metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64)
page_table = torch.tensor([[3, 1]], dtype=torch.int32).repeat(query_rows, 1)
c4_seq_lens = torch.tensor([62, 63, 64, 65], dtype=torch.int32)
@@ -301,7 +301,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
extend_start_loc=torch.tensor([0], dtype=torch.int32),
extend_num_tokens=query_rows,
)
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)
metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64)
page_table = torch.zeros((query_rows, 1), dtype=torch.int32)
c4_seq_lens = torch.tensor(
[124_997, 124_998, 124_999, 125_000], dtype=torch.int32
@@ -339,7 +339,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
backend = SimpleNamespace(_can_use_nonpaged_indexer=can_use_nonpaged_indexer)
backend.dsa_topk_backend = SimpleNamespace(is_sgl_kernel=lambda: True)
c4_indexer = SimpleNamespace(use_fp4_indexer=False, index_topk=512)
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)
metadata = SimpleNamespace(nonpaged_plan=None, compressed_page_size=64)
def build_plan(query_rows):
batch = SimpleNamespace(