Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 43 additions & 0 deletions tests/ut/worker/test_model_runner_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -557,6 +557,49 @@ def _super(self, kv_cache_config, kv_cache_allocation_context=None):
runner.init_routed_experts_capturer.assert_called_once_with()


@pytest.mark.parametrize("is_vllm_0_28_0", [True, False], ids=["v0.28.0", "newer"])
def test_initialize_kv_cache_forwards_allocation_context_by_vllm_version(is_vllm_0_28_0):
runner = _make_runner()
runner.vllm_config = SimpleNamespace()
runner.compilation_config = SimpleNamespace(static_forward_context={})
runner.pcp_manager = None
runner.model_state = SimpleNamespace(pcp_manager=None, kvpp_runtime=None)
runner.speculator = None
runner.model_config = SimpleNamespace(enable_return_routed_experts=False)
called = False
captured_kwargs: dict[str, object] = {}
allocation_context = object()
kv_cache_config = KVCacheConfig(
num_blocks=1,
kv_cache_tensors=[],
kv_cache_groups=[],
)

def _super(self, kv_cache_config, **kwargs):
nonlocal called
called = True
captured_kwargs.update(kwargs)
self.kv_cache_config = kv_cache_config
self.attn_groups = []

with (
patch("vllm_ascend.worker.v2.model_runner.vllm_version_is", return_value=is_vllm_0_28_0),
patch.object(GPUModelRunner, "initialize_kv_cache", _super),
patch("vllm_ascend.worker.v2.model_runner.ModelAclGraphManager", return_value="acl"),
patch(
"vllm_ascend.worker.v2.model_runner.KVPPRuntime.create_from_kv_cache",
return_value="kvpp",
),
):
Comment thread
yjyang62 marked this conversation as resolved.
runner.initialize_kv_cache(kv_cache_config, kv_cache_allocation_context=allocation_context)

assert called is True
if is_vllm_0_28_0:
assert "kv_cache_allocation_context" not in captured_kwargs
else:
assert captured_kwargs["kv_cache_allocation_context"] is allocation_context


@pytest.mark.parametrize("moe_type", [MoECommType.MC2, MoECommType.FUSED_MC2])
def test_profile_run_dummy_reserves_mc2(moe_type):
runner = _make_runner()
Expand Down
5 changes: 4 additions & 1 deletion vllm_ascend/worker/v2/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,9 +257,12 @@ def initialize_kv_cache(
if vllm_version_is("0.28.0"):
kv_cache_config = unwrap_mamba_kv_cache_groups(kv_cache_config)
with graph_manager_wrapper(self):
# vLLM 0.28 GPUModelRunner.initialize_kv_cache does not accept
# kv_cache_allocation_context. Gate it the same way as other
# 0.28 super() kwargs in this runner.
super().initialize_kv_cache(
kv_cache_config,
kv_cache_allocation_context=kv_cache_allocation_context,
**({} if vllm_version_is("0.28.0") else {"kv_cache_allocation_context": kv_cache_allocation_context}),
)
if self.pcp_manager is not None:
assert isinstance(self.pcp_manager, AscendPCPManager)
Expand Down
Loading