Skip to content
Draft
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
262 changes: 228 additions & 34 deletions tests/v1/kv_connector/unit/test_moriio_hma_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,16 +10,27 @@
"""

import importlib
import threading
from collections import defaultdict
from dataclasses import asdict
from types import SimpleNamespace

import pytest
import torch

from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import get_kv_cache_groups
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.core.single_type_kv_cache_manager import FullAttentionManager
from vllm.v1.kv_cache_interface import FullAttentionSpec, UniformTypeKVCacheSpecs
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KpoolTailSpec,
KVCacheConfig,
KVCacheSpec,
MLAAttentionSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.request import RequestStatus

moriio_connector = importlib.import_module(
"vllm.distributed.kv_transfer.kv_connector.v1.moriio.moriio_connector"
Expand Down Expand Up @@ -54,6 +65,7 @@ class _FakeScheduler(moriio_connector.MoRIIOConnectorScheduler): # type: ignore
def __init__(self, **attrs):
self._mamba_group_ids: list[int] = []
self._attn_group_ids: list[int] = [0]
self._attn_group_is_state: list[bool] = [False]
self._num_ssm_scratch_blocks = 0
self._is_hma_required = False
self.kv_cache_config = SimpleNamespace(
Expand Down Expand Up @@ -171,7 +183,7 @@ def test_split_block_groups_separates_attention_and_mamba():
sched = _FakeScheduler(_has_mamba=True, _attn_group_ids=[0], _mamba_group_ids=[1])
block_ids = ([1, 2], [99])
attn, mamba = sched.split_block_groups(block_ids)
assert attn == [1, 2]
assert attn == [[1, 2]]
assert mamba == [[99]]


Expand Down Expand Up @@ -208,7 +220,7 @@ def test_split_block_groups_ignores_transfer_disabled_group():
kv_cache_config=config,
)
assert sched.split_block_groups(([1, 2], [70, 71], [99])) == (
[1, 2],
[[1, 2]],
[[99]],
)

Expand All @@ -230,21 +242,220 @@ def test_exchange_blocks_ignore_transfer_disabled_group():
]


def test_scheduler_rejects_multiple_attention_groups_with_mamba():
"""A drafter owning its own transfer group is refused, with a hint."""
def test_exchange_clipped_blocks_match_select_for_single_full_attention():
# Back-compat invariant (K3/Qwen3.5 single full-attention group): when every
# group has blocks_per_sw == 0, get_exchange_clipped_blocks must reduce to
# select_transfer_block_ids, so the wire payload is byte-identical to the
# pre-multigroup path. Locks the single-group claim against future edits to
# get_exchange_clipped_blocks.
def select(block_ids):
# One full-attention attention group (0) + one mamba group (2); the
# middle group is not transfer-enabled.
return (block_ids[0], block_ids[2])

config = SimpleNamespace(
kv_cache_groups=[None, None, None],
select_transfer_block_ids=select,
)
raw = ([1, 2, 3], [70, 71], [90, 91, 92])

hma = _FakeScheduler(
kv_cache_config=config, _is_hma_required=True, blocks_per_sw=[0, 0]
)
plain = _FakeScheduler(kv_cache_config=config, _is_hma_required=False)

clipped = hma.get_exchange_clipped_blocks(raw)

# Equal to select, and identical to the non-HMA path (no clipping at all).
assert clipped == [list(group) for group in select(raw)]
assert clipped == plain.get_exchange_clipped_blocks(raw)


def test_scheduler_rejects_hybrid_without_attention_group():
config = SimpleNamespace(
kv_cache_groups=[
SimpleNamespace(enable_kv_transfer=True, kv_cache_spec=object()),
SimpleNamespace(enable_kv_transfer=True, kv_cache_spec=object()),
SimpleNamespace(enable_kv_transfer=True, kv_cache_spec=_mamba_spec()),
],
)
config.transfer_groups = tuple(config.kv_cache_groups)

with pytest.raises(moriio_common.MoRIIOError, match="separate attention group"):
with pytest.raises(moriio_common.MoRIIOError, match="attention group"):
moriio_connector.MoRIIOConnectorScheduler(_gate_vllm_config(), "engine", config)


def _glm5_next_mtp_kv_cache_config() -> KVCacheConfig:
"""GLM-5.3-Flash-shaped groups from the real grouping, with an MTP layer.

KDA, KDA, MLA + kpool indexer + kpool tail, repeated twice, then the MTP
layer 6 (an MLA layer registered last). Grouping yields
[MLA + indexer, kpool tail, KDA, KDA].
"""
kpool, num_spec = 4, 1
mamba = _mamba_spec(num_speculative_blocks=num_spec)
specs: dict[str, KVCacheSpec] = {}
for i in range(7):
if i % 3 != 2 and i != 6:
specs[f"layers.{i}.linear_attn"] = mamba
continue
specs[f"layers.{i}.attn"] = MLAAttentionSpec(
block_size=1024, num_kv_heads=1, head_size=576, dtype=torch.bfloat16
)
specs[f"layers.{i}.indexer"] = MLAAttentionSpec(
block_size=1024,
num_kv_heads=1,
head_size=132,
dtype=torch.uint8,
tokens_per_state=kpool,
)
specs[f"layers.{i}.tail"] = KpoolTailSpec(
block_size=2 * kpool,
num_kv_heads=2,
head_size=128,
head_size_v=0,
dtype=torch.bfloat16,
sliding_window=2 * kpool,
)
vllm_config = SimpleNamespace(
scheduler_config=SimpleNamespace(disable_hybrid_kv_cache_manager=False),
parallel_config=SimpleNamespace(pipeline_parallel_size=1),
speculative_config=None,
)
groups = get_kv_cache_groups(vllm_config, specs)
return KVCacheConfig(num_blocks=64, kv_cache_tensors=[], kv_cache_groups=groups)


def _glm5_next_mtp_scheduler(kv_role: str = "kv_consumer"):
vllm_config = _gate_vllm_config(
speculative_config=_spec_config("mtp"),
block_size=1024,
num_lookahead_tokens=1,
num_speculative_tokens=1,
)
vllm_config.kv_transfer_config.kv_role = kv_role
config = _glm5_next_mtp_kv_cache_config()
return moriio_connector.MoRIIOConnectorScheduler(vllm_config, "engine", config)


@pytest.mark.usefixtures("restore_moriio_role")
def test_scheduler_accepts_glm5_next_groups_with_mtp():
"""Several attention groups (MLA + indexer, kpool tail) plus KDA, with MTP."""
scheduler = _glm5_next_mtp_scheduler()
groups = scheduler.kv_cache_config.transfer_groups

# The MTP layer shares the target's MLA and tail groups.
assert {"layers.6.attn", "layers.6.indexer"} <= set(groups[0].layer_names)
assert "layers.6.tail" in groups[1].layer_names
assert scheduler._attn_group_ids == [0, 1]
assert scheduler._mamba_group_ids == [2, 3]
assert scheduler._attn_group_is_state == [False, True]
assert scheduler.blocks_per_sw == [0, 2, 0, 0]
assert scheduler._num_ssm_scratch_blocks == 1
# 1 recomputed token + 1 padded spec row + 1 drafter lookahead slot.
assert scheduler._max_decode_tail_blocks == 1


@pytest.mark.usefixtures("restore_moriio_role")
@pytest.mark.parametrize("external_tokens", [0, 1500])
def test_glm5_next_mtp_read_pairs_every_group(external_tokens):
"""P payload -> D alignment -> zeroing filter across all transfer groups."""
producer = _glm5_next_mtp_scheduler("kv_producer")
p_request = SimpleNamespace(
request_id="req",
kv_transfer_params={"do_remote_decode": True, "transfer_id": "tx"},
status=RequestStatus.FINISHED_LENGTH_CAPPED,
)
# Mamba groups: [running state, MTP scratch slot].
p_blocks = ([10, 11], [50], [90, 91], [80, 81])
delay_free, p_params = _FakeConnector(producer).request_finished_all_groups(
p_request, p_blocks
)
assert delay_free is True
assert p_params["remote_block_ids"] == [[10, 11], [50], [90], [80]]

consumer = _glm5_next_mtp_scheduler()
d_request = _make_read_request(p_params["remote_block_ids"])
# The decode's MLA group holds one extra local block for the spec tail.
d_blocks = _FakeBlocks(([100, 101, 102], [300], [200, 201], [210, 211]))
consumer.update_state_after_alloc(
d_request, d_blocks, num_external_tokens=external_tokens
)

local = consumer._reqs_need_recv["req"][1]
if external_tokens:
assert local == [[100, 101], [300], [200], [210]]
assert consumer._req_kv_params["req"]["remote_block_ids"] == [
[10, 11],
[50],
[90],
[80],
]
else:
# A full local MLA hit still reads the kpool tail ring and KDA state.
assert local == [[], [300], [200], [210]]

output = SchedulerOutput.make_empty()
output.new_block_ids_to_zero = [100, 101, 102, 300]
consumer.build_connector_meta(output)
expected_zero = [102] if external_tokens else [100, 101, 102]
assert output.new_block_ids_to_zero == expected_zero


def test_read_blocks_routes_each_hybrid_layer_to_its_group():
offsets: dict[str, tuple[list[int], list[int]]] = {}
worker = _FakeWorker(
mode=MoRIIOMode.READ,
world_size=1,
_has_mamba=True,
_num_attn_transfer_groups=2,
_num_mamba_transfer_groups=1,
_hybrid_payload_index_by_layer={"mla.0": 0, "tail.0": 1, "kda.0": 2},
layer_name_to_local_kv_cache_metadata={"mla.0": 0, "tail.0": 0, "kda.0": 0},
moriio_wrapper=SimpleNamespace(lock=threading.Lock(), shutdown=lambda: None),
moriio_config=SimpleNamespace(transfer_timeout=1.0),
_recving_transfers=defaultdict(dict),
_recving_transfers_start={},
_recving_local_blocks={},
_reads_issued_this_step=[],
_mamba_reads_this_step=[],
_recving_transfers_callback_addr={},
)
worker._remote_tp_rank = lambda remote_tp_size: 0
worker.get_engine_name_with_dp = lambda engine_id, dp_rank: engine_id
worker._get_built_session = lambda engine_id: ([None, None], None)
worker._region_session_indices = lambda layer_name: [0, 1]
worker._is_mamba_layer = lambda layer_name: layer_name.startswith("kda")

def _offsets(layer_name, local, remote, meta, **kwargs):
offsets[layer_name] = (local, remote)
return [], [], []

def _mamba_reads(layer_name, sessions, regions, local, remote, *args):
offsets[layer_name] = (local, remote)
return []

worker._compute_block_transfer_offsets = _offsets
worker._post_read_with_backoff = lambda *args: None
worker._post_mamba_reads = _mamba_reads

worker._read_blocks(
[[100, 101], [300], [200]],
[[10, 11], [50], [90]],
dst_engine_id="prefill",
request_id="req",
transfer_id="tx",
remote_host="127.0.0.1",
remote_notify_port=6001,
remote_tp_size=1,
)

assert offsets == {
"mla.0": ([100, 101], [10, 11]),
"tail.0": ([300], [50]),
"kda.0": ([200], [90]),
}
assert worker._recving_local_blocks["req"] == [100, 101, 300]


def test_split_block_groups_preserves_multiple_mamba_groups():
sched = _FakeScheduler(
_has_mamba=True,
Expand All @@ -253,7 +464,7 @@ def test_split_block_groups_preserves_multiple_mamba_groups():
)

assert sched.split_block_groups(([1, 2], [40], [70])) == (
[1, 2],
[[1, 2]],
[[40], [70]],
)

Expand Down Expand Up @@ -287,23 +498,6 @@ def test_scheduler_rejects_mamba_group_without_two_states():
moriio_connector.MoRIIOConnectorScheduler(_gate_vllm_config(), "engine", config)


def test_scheduler_rejects_non_dspark_speculative_hybrid_read():
config = SimpleNamespace(
kv_cache_groups=[
SimpleNamespace(enable_kv_transfer=True, kv_cache_spec=object()),
SimpleNamespace(enable_kv_transfer=True, kv_cache_spec=_mamba_spec()),
],
)
config.transfer_groups = tuple(config.kv_cache_groups)

with pytest.raises(moriio_common.MoRIIOError, match="speculative decoding"):
moriio_connector.MoRIIOConnectorScheduler(
_gate_vllm_config(speculative_config=_spec_config("ngram")),
"engine",
config,
)


def _dspark_wrapped_attention_config(*, dcp_sharded: bool = True):
"""Mamba first, so a wrapped full-attention group at index 1 has to be found.

Expand Down Expand Up @@ -418,7 +612,7 @@ def test_split_block_groups_keeps_only_running_state():
_attn_group_ids=[0],
_mamba_group_ids=[1],
)
assert sched.split_block_groups(([1], [40, 41])) == ([1], [[41]])
assert sched.split_block_groups(([1], [40, 41])) == ([[1]], [[41]])


@pytest.mark.parametrize(
Expand All @@ -439,7 +633,7 @@ def test_split_block_groups_strips_dspark_scratch_slots(blocks, expected):
_num_ssm_scratch_blocks=2,
)

assert sched.split_block_groups(([1], blocks)) == ([1], [expected])
assert sched.split_block_groups(([1], blocks)) == ([[1]], [expected])


# --------------------------------------------------------------------------
Expand All @@ -448,12 +642,12 @@ def test_split_block_groups_strips_dspark_scratch_slots(blocks, expected):
def test_request_finished_all_groups_carries_attn_and_mamba_in_one_field():
seen = {}

def _fake_request_finished(request, attn_block_ids, mamba_block_groups):
seen["attn"] = list(attn_block_ids)
def _fake_request_finished(request, attn_block_groups, mamba_block_groups):
seen["attn"] = [list(group) for group in attn_block_groups]
seen["mamba"] = [list(group) for group in mamba_block_groups]
return True, {
"do_remote_prefill": True,
"remote_block_ids": [attn_block_ids, *mamba_block_groups],
"remote_block_ids": [*attn_block_groups, *mamba_block_groups],
}

sched = _FakeScheduler(_has_mamba=True, _attn_group_ids=[0], _mamba_group_ids=[1])
Expand All @@ -468,7 +662,7 @@ def _fake_request_finished(request, attn_block_ids, mamba_block_groups):
# Attention ids drive request_finished; the mamba slot rides the SAME
# remote_block_ids channel as [attn, *mamba_groups], so the proxy/router
# needs no KDA-specific field.
assert seen["attn"] == [10, 11]
assert seen["attn"] == [[10, 11]]
assert seen["mamba"] == [[42]]
assert params["remote_block_ids"] == [[10, 11], [42]]
assert "remote_mamba_block_ids" not in params
Expand Down Expand Up @@ -966,15 +1160,15 @@ def test_mamba_tp_ratio_rejects_heterogeneous_tp(remote_tp_size):


def test_worker_selects_each_mamba_layers_own_block_group():
worker = _FakeWorker(_mamba_payload_index_by_layer={"kda.0": 1, "kda.1": 2})
worker = _FakeWorker(_hybrid_payload_index_by_layer={"kda.0": 1, "kda.1": 2})
groups = ([10], [40], [70])

assert worker._mamba_blocks_for_layer("kda.0", groups) == [40]
assert worker._mamba_blocks_for_layer("kda.1", groups) == [70]


def test_worker_rejects_missing_mamba_block_group():
worker = _FakeWorker(_mamba_payload_index_by_layer={"kda.0": 2})
worker = _FakeWorker(_hybrid_payload_index_by_layer={"kda.0": 2})

with pytest.raises(moriio_common.MoRIIOError, match="no block group"):
worker._mamba_blocks_for_layer("kda.0", ([10], [40]))
Loading
Loading