[Fix] Use real ReqKvInfo in unit-test req mocks (#37339)
This commit is contained in:
@@ -14,6 +14,7 @@ from sglang.srt.beam_search.fork import (
|
||||
neutral_member_sampling_params,
|
||||
remap_kv_mapping,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -119,7 +120,7 @@ class _FakeAllocator:
|
||||
class TestFreeMemberRows(CustomTestCase):
|
||||
def _make_group(self, req_to_token, allocated_len):
|
||||
leader = SimpleNamespace(
|
||||
kv=SimpleNamespace(
|
||||
kv=ReqKvInfo(
|
||||
kv_allocated_len=allocated_len, kv_committed_len=allocated_len
|
||||
),
|
||||
)
|
||||
@@ -208,7 +209,7 @@ class TestRetireReclaimsStagedOrphans(CustomTestCase):
|
||||
allocator = _FakeAllocator()
|
||||
group = SimpleNamespace(
|
||||
leader=SimpleNamespace(
|
||||
kv=SimpleNamespace(kv_allocated_len=8, kv_committed_len=8),
|
||||
kv=ReqKvInfo(kv_allocated_len=8, kv_committed_len=8),
|
||||
),
|
||||
prompt_len=5,
|
||||
member_rows=torch.tensor([1, 2], dtype=torch.int64),
|
||||
|
||||
@@ -38,6 +38,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsa.utils import should_use_dsa_fused_topk
|
||||
from sglang.srt.managers.overlap_utils import FutureMap, RelayPayload
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.srt.speculative.eagle_disaggregation import (
|
||||
@@ -106,7 +107,7 @@ class TestDisaggregationWire(unittest.TestCase):
|
||||
|
||||
def test_prebuilt_skips_unused_prompt_tensor(self):
|
||||
req = SimpleNamespace(
|
||||
kv=SimpleNamespace(req_pool_idx=0),
|
||||
kv=ReqKvInfo(req_pool_idx=0),
|
||||
prefix_indices=[0, 1],
|
||||
extend_range=SimpleNamespace(length=3),
|
||||
origin_input_ids=[0, 1, 2, 3, 4],
|
||||
|
||||
@@ -14,6 +14,7 @@ from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||
from sglang.srt.utils.common import Range
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
@@ -49,7 +50,7 @@ def _make_req(rid="test-req-0", origin_input_ids=None, output_ids=None):
|
||||
output_ids=output_ids,
|
||||
fill_ids=origin_input_ids + output_ids,
|
||||
seqlen=len(origin_input_ids) + len(output_ids),
|
||||
kv=SimpleNamespace(req_pool_idx=None, kv_allocated_len=0, kv_committed_len=0),
|
||||
kv=ReqKvInfo(),
|
||||
finished_reason=None,
|
||||
hisparse_staging=False,
|
||||
staging=False,
|
||||
|
||||
@@ -13,7 +13,11 @@ from sglang.srt.disaggregation.decode import ( # noqa: E402
|
||||
SchedulerDisaggregationDecodeMixin,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode # noqa: E402
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req # noqa: E402
|
||||
from sglang.srt.managers.schedule_batch import ( # noqa: E402
|
||||
FINISH_ABORT,
|
||||
Req,
|
||||
ReqKvInfo,
|
||||
)
|
||||
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
||||
from sglang.srt.runtime_context import get_context # noqa: E402
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -96,7 +100,7 @@ class TestDecodePreallocQueuePriority(unittest.TestCase):
|
||||
finished_reason=FINISH_ABORT("failed") if failed else None,
|
||||
return_logprob=False,
|
||||
sampling_params=SimpleNamespace(max_new_tokens=8),
|
||||
kv=SimpleNamespace(req_pool_idx=int(priority) % 8, cache_protected_len=0),
|
||||
kv=ReqKvInfo(req_pool_idx=int(priority) % 8, cache_protected_len=0),
|
||||
time_stats=MagicMock(),
|
||||
)
|
||||
return SimpleNamespace(
|
||||
|
||||
@@ -11,6 +11,7 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.common import _release_overallocated_kv_indices
|
||||
@@ -131,7 +132,7 @@ class TestFreeSegment(unittest.TestCase):
|
||||
token_to_kv_pool_allocator=alloc,
|
||||
req_to_token_pool=SimpleNamespace(req_to_token=row.unsqueeze(0)),
|
||||
)
|
||||
req = SimpleNamespace(kv=SimpleNamespace(req_pool_idx=0))
|
||||
req = SimpleNamespace(kv=ReqKvInfo(req_pool_idx=0))
|
||||
|
||||
before = len(alloc.free_pages)
|
||||
alloc.free_group_begin()
|
||||
|
||||
@@ -27,7 +27,6 @@ import random
|
||||
import unittest
|
||||
import unittest.mock
|
||||
from array import array
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
@@ -39,6 +38,7 @@ from sglang.srt.disaggregation.kv_events import (
|
||||
BlockStoredWithMetadata,
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
EvictParams,
|
||||
@@ -518,7 +518,7 @@ class TestRadixCache(unittest.TestCase):
|
||||
)
|
||||
cache.req_to_token_pool = ReqToTokenPool(request_indices.clone())
|
||||
req = unittest.mock.Mock(
|
||||
kv=SimpleNamespace(req_pool_idx=0, cache_protected_len=0),
|
||||
kv=ReqKvInfo(req_pool_idx=0, cache_protected_len=0),
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
priority=0,
|
||||
|
||||
@@ -19,6 +19,7 @@ from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.test.ci.ci_register import register_xpu_ci
|
||||
|
||||
# Must be set before lmcache imports. Save prior values so tearDownModule can
|
||||
@@ -71,7 +72,7 @@ def _make_req(rid, req_pool_idx, token_ids, tree):
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
last_node=tree.root_node,
|
||||
kv=SimpleNamespace(
|
||||
kv=ReqKvInfo(
|
||||
req_pool_idx=req_pool_idx,
|
||||
cache_protected_len=0,
|
||||
kv_committed_len=len(token_ids),
|
||||
@@ -167,7 +168,7 @@ class TestLMCRadixCacheXPU(unittest.TestCase):
|
||||
req_pool_idx = req_to_token_pool.alloc(
|
||||
[
|
||||
SimpleNamespace(
|
||||
inflight_middle_chunks=0, kv=SimpleNamespace(req_pool_idx=None)
|
||||
inflight_middle_chunks=0, kv=ReqKvInfo(req_pool_idx=None)
|
||||
)
|
||||
]
|
||||
)[0]
|
||||
|
||||
Reference in New Issue
Block a user