[Refactor] Clarify DeepSeek V4 metadata names for V4.1 (#38947)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user