Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai> Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
1408 lines
50 KiB
Python
1408 lines
50 KiB
Python
"""Unit tests for MLX attention discovery and generic cache handling."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
|
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
|
|
|
_HAS_MLX = importlib.util.find_spec("mlx") is not None
|
|
_SKIP_REASON = "requires mlx"
|
|
|
|
if _HAS_MLX:
|
|
import mlx.core as mx
|
|
import mlx.nn as nn
|
|
import torch
|
|
from mlx_lm.models.cache import ArraysCache
|
|
|
|
import sglang.srt.hardware_backend.mlx.aot as mlx_aot
|
|
from sglang.srt.hardware_backend.mlx.aot import (
|
|
MlxAOTKernelSet,
|
|
MlxAOTRoPEKernel,
|
|
)
|
|
from sglang.srt.hardware_backend.mlx.kv_cache import (
|
|
BatchedDecodeContext,
|
|
ContiguousAttentionKVCache,
|
|
MlxAttentionKVPool,
|
|
MLXAttentionWrapper,
|
|
MlxAuxiliaryStateComponent,
|
|
MlxAuxiliaryStatePool,
|
|
MlxAuxiliaryStateReqToTokenPool,
|
|
MlxModelCacheLayout,
|
|
find_attention_layers,
|
|
is_attention_module,
|
|
patch_model_attention,
|
|
)
|
|
from sglang.srt.hardware_backend.mlx.model_runner import (
|
|
MlxModelRunner,
|
|
MlxPendingDecode,
|
|
)
|
|
from sglang.srt.hardware_backend.mlx.scheduler_mixin import (
|
|
MlxPendingJob,
|
|
SchedulerMlxOverlapMixin,
|
|
)
|
|
from sglang.srt.managers.scheduler_components import (
|
|
batch_result_processor as batch_result_processor_module,
|
|
)
|
|
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
|
SchedulerBatchResultProcessor,
|
|
)
|
|
from sglang.srt.managers.utils import GenerationBatchResult
|
|
from sglang.srt.mem_cache.base_prefix_cache import InsertParams, InsertResult
|
|
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
|
|
|
|
|
def _set_runner_cache_layout(
|
|
runner,
|
|
*,
|
|
num_layers: int,
|
|
attention_layer_indices: list[int],
|
|
attention_modules: dict[int, object] | None = None,
|
|
) -> None:
|
|
attention_modules = attention_modules or {}
|
|
attention_set = set(attention_layer_indices)
|
|
layers = []
|
|
attrs = []
|
|
for layer_idx in range(num_layers):
|
|
if layer_idx in attention_set:
|
|
attrs.append("self_attn")
|
|
layers.append(
|
|
SimpleNamespace(self_attn=attention_modules.get(layer_idx, object()))
|
|
)
|
|
else:
|
|
attrs.append(None)
|
|
layers.append(SimpleNamespace())
|
|
runner._cache_layout = MlxModelCacheLayout.from_attention_discovery(layers, attrs)
|
|
|
|
|
|
def _set_runner_decode_context_defaults(runner) -> None:
|
|
runner._aot_kernels = MlxAOTKernelSet()
|
|
runner._attention_kv_pool = None
|
|
runner._req_pool_idx = {}
|
|
runner._req_to_token_pool = None
|
|
|
|
|
|
def _set_dummy_server_args_for_auxiliary_state_tests() -> None:
|
|
server_args = ServerArgs(model_path="dummy", page_size=1)
|
|
server_args._mamba_cache_chunk_size = 64
|
|
set_global_server_args_for_scheduler(server_args)
|
|
|
|
|
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
|
class TestMlxAttentionPatching(unittest.TestCase):
|
|
def test_standard_attention_is_patched_once(self):
|
|
model = FakeModel(
|
|
[
|
|
FakeLayer("self_attn", FakeAttention()),
|
|
FakeLayer("self_attn", FakeAttention()),
|
|
]
|
|
)
|
|
|
|
layers, attrs = find_attention_layers(model)
|
|
|
|
self.assertEqual(len(layers), 2)
|
|
self.assertEqual(attrs, ["self_attn", "self_attn"])
|
|
self.assertEqual(patch_model_attention(model), 2)
|
|
self.assertIsInstance(model.layers[0].self_attn, MLXAttentionWrapper)
|
|
self.assertIsInstance(model.layers[1].self_attn, MLXAttentionWrapper)
|
|
self.assertEqual(patch_model_attention(model), 0)
|
|
|
|
def test_alias_head_names_are_supported(self):
|
|
model = FakeModel([FakeLayer("attention", FakeAttention(use_aliases=True))])
|
|
|
|
_, attrs = find_attention_layers(model)
|
|
|
|
self.assertEqual(attrs, ["attention"])
|
|
self.assertEqual(patch_model_attention(model), 1)
|
|
self.assertIsInstance(model.layers[0].attention, MLXAttentionWrapper)
|
|
|
|
def test_aot_rope_kernel_build_uses_head_aliases(self):
|
|
attn = FakeAttention(use_aliases=True)
|
|
attn.rope = SimpleNamespace(dims=2, traditional=False, base=10000.0)
|
|
original_loader = mlx_aot._load_metal_rope_pool_fused
|
|
mlx_aot._load_metal_rope_pool_fused = lambda: object()
|
|
try:
|
|
kernel = mlx_aot._build_rope_kernel(
|
|
mlx_aot.MlxAOTKernelBuildInputs(
|
|
sample_attn=attn,
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
)
|
|
)
|
|
finally:
|
|
mlx_aot._load_metal_rope_pool_fused = original_loader
|
|
|
|
self.assertTrue(kernel.enabled)
|
|
self.assertEqual(kernel.config["num_qo_heads"], 2)
|
|
|
|
def test_auxiliary_state_model_returns_per_layer_attention_attrs(self):
|
|
model = FakeModel(
|
|
[
|
|
FakeLayer("linear_attn", ProjectionOnlyMixer()),
|
|
FakeLayer("self_attn", FakeAttention()),
|
|
FakeLayer("linear_attn", ProjectionOnlyMixer()),
|
|
]
|
|
)
|
|
|
|
_, attrs = find_attention_layers(model)
|
|
|
|
self.assertEqual(attrs, [None, "self_attn", None])
|
|
self.assertEqual(patch_model_attention(model), 1)
|
|
self.assertFalse(isinstance(model.layers[0].linear_attn, MLXAttentionWrapper))
|
|
self.assertIsInstance(model.layers[1].self_attn, MLXAttentionWrapper)
|
|
|
|
def test_projection_only_mixer_is_not_attention(self):
|
|
self.assertFalse(is_attention_module(ProjectionOnlyMixer()))
|
|
|
|
def test_cache_layout_separates_attention_and_auxiliary_layers(self):
|
|
layout = MlxModelCacheLayout.from_attention_discovery(
|
|
[object(), object(), object(), object()],
|
|
[None, "self_attn", None, "self_attn"],
|
|
)
|
|
|
|
self.assertEqual(layout.num_layers, 4)
|
|
self.assertEqual(layout.attention_layer_indices, (1, 3))
|
|
self.assertEqual(layout.auxiliary_layer_indices, (0, 2))
|
|
self.assertEqual(layout.attention_pool_index(1), 0)
|
|
self.assertEqual(layout.attention_pool_index(3), 1)
|
|
self.assertTrue(layout.has_auxiliary_state)
|
|
|
|
def test_gated_query_projection_keeps_attention_width(self):
|
|
inner = FakeGatedAttention()
|
|
wrapper = MLXAttentionWrapper(inner, layer_idx=0)
|
|
cache = ContiguousAttentionKVCache(
|
|
n_kv_heads=1, head_dim=2, max_seq_len=4, dtype=mx.float32
|
|
)
|
|
ctx = BatchedDecodeContext(
|
|
batch_size=1,
|
|
seq_lens=[0],
|
|
attention_layer_caches=[[cache]],
|
|
)
|
|
|
|
out = wrapper._batched_decode(mx.zeros((1, 1, 4), dtype=mx.float32), ctx)
|
|
mx.eval(out)
|
|
|
|
self.assertEqual(out.shape, (1, 1, 4))
|
|
self.assertEqual(inner.o_proj.last_input_shape, (1, 1, 4))
|
|
|
|
def test_attn_config_uses_float_dtype_for_quantized_projection(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
attn = FakeAttention()
|
|
attn.k_proj.weight = mx.zeros((2, 4), dtype=mx.uint32)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=1,
|
|
attention_layer_indices=[0],
|
|
attention_modules={0: attn},
|
|
)
|
|
|
|
n_kv_heads, head_dim, dtype = MlxModelRunner._get_attn_config(runner)
|
|
|
|
self.assertEqual(n_kv_heads, 1)
|
|
self.assertEqual(head_dim, 2)
|
|
self.assertEqual(dtype, mx.float32)
|
|
|
|
def test_attn_config_rejects_heterogeneous_kv_shapes(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
first = FakeAttention()
|
|
second = FakeAttention()
|
|
second.n_kv_heads = 2
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[0, 1],
|
|
attention_modules={0: first, 1: second},
|
|
)
|
|
|
|
with self.assertRaisesRegex(
|
|
NotImplementedError,
|
|
"uniform softmax-attention KV shape",
|
|
):
|
|
MlxModelRunner._get_attn_config(runner)
|
|
|
|
def test_attn_config_rejects_sliding_window_attention(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=1,
|
|
attention_layer_indices=[0],
|
|
attention_modules={0: FakeAttention()},
|
|
)
|
|
runner._cache_layout.layers[0].use_sliding = True
|
|
|
|
with self.assertRaisesRegex(
|
|
NotImplementedError,
|
|
"sliding-window attention",
|
|
):
|
|
MlxModelRunner._get_attn_config(runner)
|
|
|
|
|
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
|
class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
|
|
def test_dense_prefill_keeps_pool_backed_radix_path(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
runner.model = FakeDenseModel(num_layers=2)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[0, 1],
|
|
)
|
|
runner._max_seq_len = 8
|
|
runner._cache_pool = []
|
|
runner.disable_radix_cache = False
|
|
runner._attention_kv_pool = MlxAttentionKVPool(
|
|
pool_size=8,
|
|
num_layers=2,
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
dtype=mx.float32,
|
|
)
|
|
runner._req_to_token_pool = None
|
|
runner._req_caches = {}
|
|
runner._req_token_ids = {}
|
|
runner._req_pool_idx = {}
|
|
runner._req_synced_offset = {}
|
|
prefix_slots = mx.array([2, 3], dtype=mx.int32)
|
|
k_prefix = mx.stack(
|
|
[
|
|
mx.ones((2, 1, 2), dtype=mx.float32) * 10,
|
|
mx.ones((2, 1, 2), dtype=mx.float32) * 20,
|
|
]
|
|
)
|
|
runner._attention_kv_pool.set_kv_all_layers(
|
|
prefix_slots, k_prefix, k_prefix * 2
|
|
)
|
|
mx.eval(*runner._attention_kv_pool.all_buffers())
|
|
|
|
pending = runner.prefill_start(
|
|
req_id="r0",
|
|
new_token_ids=[13],
|
|
full_token_ids=[11, 12, 13],
|
|
prefix_slot_ids=[2, 3],
|
|
new_slot_ids=[4],
|
|
req_pool_idx=0,
|
|
)
|
|
MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache)
|
|
mx.eval(*runner._attention_kv_pool.all_buffers())
|
|
runner.prefill_finalize(pending)
|
|
|
|
self.assertEqual(runner.model.seen_inputs, [[[13]]])
|
|
self.assertEqual(runner.model.seen_offsets, [[2, 2]])
|
|
self.assertEqual(pending.synced_offset, 3)
|
|
self.assertTrue(
|
|
all(isinstance(c, ContiguousAttentionKVCache) for c in pending.cache)
|
|
)
|
|
layer0_k, layer0_v = runner._attention_kv_pool.get_kv(
|
|
0, mx.array([4], dtype=mx.int32)
|
|
)
|
|
layer1_k, layer1_v = runner._attention_kv_pool.get_kv(
|
|
1, mx.array([4], dtype=mx.int32)
|
|
)
|
|
mx.eval(layer0_k, layer0_v, layer1_k, layer1_v)
|
|
self.assertEqual(layer0_k.tolist(), [[[1.0, 1.0]]])
|
|
self.assertEqual(layer0_v.tolist(), [[[2.0, 2.0]]])
|
|
self.assertEqual(layer1_k.tolist(), [[[2.0, 2.0]]])
|
|
self.assertEqual(layer1_v.tolist(), [[[4.0, 4.0]]])
|
|
|
|
def test_dense_decode_uses_batched_attention_for_single_and_multi_request(self):
|
|
for req_ids in (["r0"], ["r0", "r1"]):
|
|
with self.subTest(req_ids=req_ids):
|
|
runner = object.__new__(MlxModelRunner)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=1,
|
|
attention_layer_indices=[0],
|
|
)
|
|
runner._req_caches = {rid: [object()] for rid in req_ids}
|
|
runner._req_token_ids = {
|
|
rid: [idx + 10] for idx, rid in enumerate(req_ids)
|
|
}
|
|
calls = []
|
|
|
|
def fake_batched(caches, batched_input, helper_req_ids):
|
|
calls.append(
|
|
(len(caches), batched_input.tolist(), list(helper_req_ids))
|
|
)
|
|
return mx.array(list(range(len(caches))), dtype=mx.int32)
|
|
|
|
def fail_native(*args, **kwargs):
|
|
raise AssertionError("dense decode should use batched attention")
|
|
|
|
runner._decode_with_batched_attention = fake_batched
|
|
runner._decode_with_native_cache = fail_native
|
|
|
|
pending = runner.decode_batch_start(req_ids)
|
|
|
|
self.assertEqual(
|
|
calls,
|
|
[
|
|
(
|
|
len(req_ids),
|
|
[[idx + 10] for idx in range(len(req_ids))],
|
|
req_ids,
|
|
)
|
|
],
|
|
)
|
|
self.assertEqual(
|
|
pending.lazy_tokens.tolist(),
|
|
list(range(len(req_ids))),
|
|
)
|
|
|
|
def test_dense_chained_decode_uses_batched_attention_for_single_request(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=1,
|
|
attention_layer_indices=[0],
|
|
)
|
|
calls = []
|
|
|
|
def fake_batched(caches, batched_input, helper_req_ids):
|
|
calls.append((len(caches), batched_input.tolist(), list(helper_req_ids)))
|
|
return mx.array([8], dtype=mx.int32)
|
|
|
|
def fail_native(*args, **kwargs):
|
|
raise AssertionError("dense chained decode should use batched attention")
|
|
|
|
runner._decode_with_batched_attention = fake_batched
|
|
runner._decode_with_native_cache = fail_native
|
|
prev = MlxPendingDecode(
|
|
lazy_tokens=mx.array([7], dtype=mx.int32),
|
|
req_ids=["r0"],
|
|
caches=[[object()]],
|
|
)
|
|
|
|
pending = runner.decode_batch_start_chained(prev)
|
|
|
|
self.assertEqual(calls, [(1, [[7]], ["r0"])])
|
|
self.assertEqual(pending.lazy_tokens.tolist(), [8])
|
|
|
|
def test_decode_finalize_does_not_snapshot_auxiliary_state(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
runner._req_token_ids = {"r0": [8]}
|
|
runner._decode_step_ct = 0
|
|
calls = []
|
|
runner._store_auxiliary_state = lambda req_pool_idx, cache: calls.append(
|
|
(req_pool_idx, cache)
|
|
)
|
|
pending = MlxPendingDecode(
|
|
lazy_tokens=mx.array([9], dtype=mx.int32),
|
|
req_ids=["r0"],
|
|
caches=[[object()]],
|
|
)
|
|
|
|
next_tokens = runner.decode_batch_finalize(pending)
|
|
|
|
self.assertEqual(next_tokens, [9])
|
|
self.assertEqual(runner._req_token_ids["r0"], [8, 9])
|
|
self.assertEqual(calls, [])
|
|
|
|
def test_store_auxiliary_state_for_request_snapshots_on_demand(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
cache = [object()]
|
|
runner._req_pool_idx = {"r0": 3}
|
|
runner._req_caches = {"r0": cache}
|
|
calls = []
|
|
runner._store_auxiliary_state = lambda req_pool_idx, cache_arg: calls.append(
|
|
(req_pool_idx, cache_arg)
|
|
)
|
|
|
|
runner.store_auxiliary_state_for_request("r0")
|
|
runner.store_auxiliary_state_for_request("missing")
|
|
|
|
self.assertEqual(calls, [(3, cache)])
|
|
|
|
def test_dense_batched_attention_helper_supports_single_request(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
model = FakeWrappedAttentionModel()
|
|
runner.model = model
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=1,
|
|
attention_layer_indices=[0],
|
|
)
|
|
_set_runner_decode_context_defaults(runner)
|
|
cache = [
|
|
[
|
|
ContiguousAttentionKVCache(
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
max_seq_len=4,
|
|
dtype=mx.float32,
|
|
)
|
|
]
|
|
]
|
|
|
|
lazy_tokens = runner._decode_with_batched_attention(
|
|
cache,
|
|
mx.array([[7]], dtype=mx.int32),
|
|
["r0"],
|
|
)
|
|
mx.eval(lazy_tokens, *MlxModelRunner._cache_state_arrays(cache))
|
|
|
|
self.assertEqual(lazy_tokens.tolist(), [0])
|
|
self.assertEqual(cache[0][0].offset, 1)
|
|
self.assertEqual(model.seen_inputs, [[[7]]])
|
|
self.assertEqual(model.seen_cache_types, [["AttentionOffsetCache"]])
|
|
|
|
def test_batched_decode_context_resolves_aot_rope_slots_from_request_ids(self):
|
|
cache0 = ContiguousAttentionKVCache(
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
max_seq_len=4,
|
|
dtype=mx.float32,
|
|
)
|
|
cache1 = ContiguousAttentionKVCache(
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
max_seq_len=4,
|
|
dtype=mx.float32,
|
|
)
|
|
cache0.offset = 1
|
|
cache1.offset = 2
|
|
kernel_set = MlxAOTKernelSet(
|
|
rope=MlxAOTRoPEKernel(
|
|
base=10000.0,
|
|
config={
|
|
"head_dim": 2,
|
|
"num_qo_heads": 1,
|
|
"num_kv_heads": 1,
|
|
},
|
|
rope_pool_fused=object(),
|
|
)
|
|
)
|
|
req_to_token_pool = SimpleNamespace(
|
|
req_to_token=torch.tensor(
|
|
[
|
|
[0, 41, 42],
|
|
[0, 51, 52],
|
|
],
|
|
dtype=torch.int64,
|
|
)
|
|
)
|
|
|
|
ctx = BatchedDecodeContext.from_decode(
|
|
caches=[[cache0], [cache1]],
|
|
req_ids=["r0", "r1"],
|
|
aot_kernels=kernel_set,
|
|
kv_pool=object(),
|
|
req_pool_idx={"r0": 0, "r1": 1},
|
|
req_to_token_pool=req_to_token_pool,
|
|
attention_layer_indices=[0],
|
|
)
|
|
|
|
self.assertEqual(ctx.seq_lens, [1, 2])
|
|
self.assertIsNotNone(ctx.aot.rope)
|
|
self.assertEqual(ctx.aot.rope.new_token_slots.tolist(), [41, 52])
|
|
|
|
def test_auxiliary_decode_uses_hybrid_batching_for_multi_request(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[1],
|
|
)
|
|
req_ids = ["r0", "r1"]
|
|
runner._req_caches = {rid: [object(), object()] for rid in req_ids}
|
|
runner._req_token_ids = {rid: [idx + 20] for idx, rid in enumerate(req_ids)}
|
|
calls = []
|
|
|
|
def fake_hybrid(caches, batched_input, helper_req_ids):
|
|
calls.append((len(caches), batched_input.tolist(), list(helper_req_ids)))
|
|
return mx.array([4, 5], dtype=mx.int32)
|
|
|
|
def fail_batched(*args, **kwargs):
|
|
raise AssertionError(
|
|
"auxiliary decode should use hybrid batching, not full batched"
|
|
)
|
|
|
|
runner._decode_with_hybrid_batching = fake_hybrid
|
|
runner._decode_with_batched_attention = fail_batched
|
|
|
|
pending = runner.decode_batch_start(req_ids)
|
|
|
|
self.assertEqual(calls, [(2, [[20], [21]], req_ids)])
|
|
self.assertEqual(pending.lazy_tokens.tolist(), [4, 5])
|
|
|
|
def test_auxiliary_layer_batches_mergeable_native_cache(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
layer = FakeBatchableAuxiliaryLayer()
|
|
cache0 = ArraysCache(size=1)
|
|
cache1 = ArraysCache(size=1)
|
|
|
|
out = runner._decode_auxiliary_layer(
|
|
layer,
|
|
mx.zeros((2, 1, 4), dtype=mx.float32),
|
|
[cache0, cache1],
|
|
)
|
|
mx.eval(out, cache0[0], cache1[0])
|
|
|
|
self.assertEqual(layer.input_layernorm.seen_shapes, [(2, 1, 4)])
|
|
self.assertEqual(layer.linear_attn.seen_shapes, [(2, 1, 4)])
|
|
self.assertEqual(layer.post_attention_layernorm.seen_shapes, [(2, 1, 4)])
|
|
self.assertEqual(layer.mlp.seen_shapes, [(2, 1, 4)])
|
|
self.assertEqual(layer.linear_attn.cache_type, "ArraysCache")
|
|
self.assertEqual(
|
|
out.tolist(),
|
|
[[[2.0, 2.0, 2.0, 2.0]], [[2.0, 2.0, 2.0, 2.0]]],
|
|
)
|
|
self.assertEqual(cache0[0].tolist(), [[0.0]])
|
|
self.assertEqual(cache1[0].tolist(), [[1.0]])
|
|
|
|
def test_arrays_cache_auxiliary_batching_uses_fast_merge(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
layer = FakeBatchableAuxiliaryLayer()
|
|
cache0 = ArraysCache(size=1)
|
|
cache1 = ArraysCache(size=1)
|
|
original_merge = ArraysCache.merge
|
|
|
|
def fail_merge(cls, caches):
|
|
raise AssertionError("ArraysCache fast path should not call merge()")
|
|
|
|
ArraysCache.merge = classmethod(fail_merge)
|
|
try:
|
|
out = runner._decode_auxiliary_layer(
|
|
layer,
|
|
mx.zeros((2, 1, 4), dtype=mx.float32),
|
|
[cache0, cache1],
|
|
)
|
|
mx.eval(out, cache0[0], cache1[0])
|
|
finally:
|
|
ArraysCache.merge = original_merge
|
|
|
|
self.assertEqual(layer.linear_attn.cache_type, "ArraysCache")
|
|
self.assertEqual(
|
|
out.tolist(),
|
|
[[[2.0, 2.0, 2.0, 2.0]], [[2.0, 2.0, 2.0, 2.0]]],
|
|
)
|
|
self.assertEqual(cache0[0].tolist(), [[0.0]])
|
|
self.assertEqual(cache1[0].tolist(), [[1.0]])
|
|
|
|
def test_auxiliary_layer_split_back_copies_cache_metadata(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
layer = FakeBatchableAuxiliaryLayer()
|
|
cache0 = FakeMergeableAuxiliaryCache(tag="old0")
|
|
cache1 = FakeMergeableAuxiliaryCache(tag="old1")
|
|
|
|
out = runner._decode_auxiliary_layer(
|
|
layer,
|
|
mx.zeros((2, 1, 4), dtype=mx.float32),
|
|
[cache0, cache1],
|
|
)
|
|
mx.eval(out, cache0[0], cache1[0])
|
|
|
|
self.assertEqual(layer.linear_attn.cache_type, "FakeMergeableAuxiliaryCache")
|
|
self.assertEqual(cache0.tag, "split-0")
|
|
self.assertEqual(cache1.tag, "split-1")
|
|
self.assertEqual(cache0.extra_metadata, {"idx": 0})
|
|
self.assertEqual(cache1.extra_metadata, {"idx": 1})
|
|
self.assertEqual(cache0[0].tolist(), [[0.0]])
|
|
self.assertEqual(cache1[0].tolist(), [[1.0]])
|
|
|
|
def test_auxiliary_state_prefill_restores_prefix_state(self):
|
|
runner = object.__new__(MlxModelRunner)
|
|
runner.model = FakeAuxiliaryStateModel()
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[1],
|
|
)
|
|
runner._max_seq_len = 8
|
|
runner._cache_pool = []
|
|
runner.disable_radix_cache = False
|
|
runner._attention_kv_pool = MlxAttentionKVPool(
|
|
pool_size=8,
|
|
num_layers=1,
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
dtype=mx.float32,
|
|
)
|
|
runner._req_to_token_pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
runner._req_caches = {}
|
|
runner._req_token_ids = {}
|
|
runner._req_pool_idx = {}
|
|
runner._req_synced_offset = {}
|
|
req = FakeRequest()
|
|
runner._req_to_token_pool.alloc([req])
|
|
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
|
req.mamba_pool_idx,
|
|
[FakeNativeCache(mx.array([42.0], dtype=mx.float32)), None],
|
|
[0],
|
|
)
|
|
|
|
pending = runner.prefill_start(
|
|
req_id="r0",
|
|
new_token_ids=[13],
|
|
full_token_ids=[11, 12, 13],
|
|
prefix_slot_ids=[2, 3],
|
|
new_slot_ids=[4],
|
|
req_pool_idx=req.req_pool_idx,
|
|
)
|
|
MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache)
|
|
runner.prefill_finalize(pending)
|
|
|
|
self.assertEqual(runner.model.seen_inputs, [[[13]]])
|
|
self.assertEqual(runner.model.seen_auxiliary_states, [[42.0]])
|
|
self.assertEqual(pending.synced_offset, 3)
|
|
self.assertIsInstance(pending.cache[0], FakeNativeCache)
|
|
self.assertIsInstance(pending.cache[1], ContiguousAttentionKVCache)
|
|
restored = [FakeNativeCache(), None]
|
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
|
req.mamba_pool_idx, restored, [0]
|
|
)
|
|
self.assertEqual(restored[0].state[0].tolist(), [1.0])
|
|
|
|
def test_auxiliary_state_prefill_tracks_chunk_aligned_auxiliary_state(self):
|
|
_set_dummy_server_args_for_auxiliary_state_tests()
|
|
runner = object.__new__(MlxModelRunner)
|
|
runner.model = FakeAuxiliaryStateModel()
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[1],
|
|
)
|
|
runner._max_seq_len = 128
|
|
runner._cache_pool = []
|
|
runner.disable_radix_cache = False
|
|
runner._attention_kv_pool = MlxAttentionKVPool(
|
|
pool_size=96,
|
|
num_layers=1,
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
dtype=mx.float32,
|
|
)
|
|
runner._req_to_token_pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=128,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
runner._req_caches = {}
|
|
runner._req_token_ids = {}
|
|
runner._req_pool_idx = {}
|
|
runner._req_synced_offset = {}
|
|
req = FakeRequest()
|
|
runner._req_to_token_pool.alloc([req])
|
|
token_ids = list(range(70))
|
|
|
|
pending = runner.prefill_start(
|
|
req_id="r0",
|
|
new_token_ids=token_ids,
|
|
full_token_ids=token_ids,
|
|
prefix_slot_ids=[],
|
|
new_slot_ids=list(range(1, 71)),
|
|
req_pool_idx=req.req_pool_idx,
|
|
req=req,
|
|
)
|
|
MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache)
|
|
runner.prefill_finalize(pending)
|
|
tracked = [FakeNativeCache(), None]
|
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
|
req.mamba_ping_pong_track_buffer[0], tracked, [0]
|
|
)
|
|
|
|
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [64, 6])
|
|
self.assertEqual(req.mamba_last_track_seqlen, 64)
|
|
self.assertEqual(tracked[0].state[0].tolist(), [64.0])
|
|
self.assertEqual(pending.synced_offset, 70)
|
|
|
|
def test_auxiliary_state_prefill_advances_tracked_boundary_after_cached_prefix(
|
|
self,
|
|
):
|
|
_set_dummy_server_args_for_auxiliary_state_tests()
|
|
runner = object.__new__(MlxModelRunner)
|
|
runner.model = FakeAuxiliaryStateModel()
|
|
_set_runner_cache_layout(
|
|
runner,
|
|
num_layers=2,
|
|
attention_layer_indices=[1],
|
|
)
|
|
runner._max_seq_len = 512
|
|
runner._cache_pool = []
|
|
runner.disable_radix_cache = False
|
|
runner._attention_kv_pool = MlxAttentionKVPool(
|
|
pool_size=320,
|
|
num_layers=1,
|
|
n_kv_heads=1,
|
|
head_dim=2,
|
|
dtype=mx.float32,
|
|
)
|
|
runner._req_to_token_pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=512,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
runner._req_caches = {}
|
|
runner._req_token_ids = {}
|
|
runner._req_pool_idx = {}
|
|
runner._req_synced_offset = {}
|
|
req = FakeRequest()
|
|
runner._req_to_token_pool.alloc([req])
|
|
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
|
|
req.mamba_pool_idx,
|
|
[FakeNativeCache(mx.array([64.0], dtype=mx.float32)), None],
|
|
[0],
|
|
)
|
|
token_ids = list(range(257))
|
|
|
|
pending = runner.prefill_start(
|
|
req_id="r0",
|
|
new_token_ids=token_ids[64:],
|
|
full_token_ids=token_ids,
|
|
prefix_slot_ids=list(range(1, 65)),
|
|
new_slot_ids=list(range(65, 258)),
|
|
req_pool_idx=req.req_pool_idx,
|
|
req=req,
|
|
)
|
|
MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache)
|
|
runner.prefill_finalize(pending)
|
|
tracked = [FakeNativeCache(), None]
|
|
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
|
|
req.mamba_ping_pong_track_buffer[0], tracked, [0]
|
|
)
|
|
|
|
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [192, 1])
|
|
self.assertEqual(runner.model.seen_auxiliary_states, [[64.0], [192.0]])
|
|
self.assertEqual(req.mamba_last_track_seqlen, 256)
|
|
self.assertEqual(tracked[0].state[0].tolist(), [192.0])
|
|
self.assertEqual(pending.synced_offset, 257)
|
|
|
|
def test_cache_arrays_flattens_native_array_cache_state(self):
|
|
cache = FakeNestedStateCache()
|
|
|
|
arrays = MlxModelRunner._cache_arrays(cache)
|
|
|
|
self.assertEqual(len(arrays), 2)
|
|
self.assertTrue(all(isinstance(arr, mx.array) for arr in arrays))
|
|
|
|
def test_auxiliary_state_pool_tracks_scheduler_slots_and_snapshots(self):
|
|
pool = MlxAuxiliaryStatePool(size=4, device="cpu")
|
|
|
|
first = pool.alloc(2)
|
|
cache = [FakeNativeCache(mx.array([1.0], dtype=mx.float32))]
|
|
pool.store_cache(first[0], cache, [0])
|
|
cache[0].state[0][0] = 9.0
|
|
forked = pool.fork_from(first[0].unsqueeze(0))
|
|
restored = [FakeNativeCache()]
|
|
pool.restore_cache(forked[0], restored, [0])
|
|
pool.free(first)
|
|
|
|
self.assertEqual(first.tolist(), [1, 2])
|
|
self.assertEqual(forked.tolist(), [3])
|
|
self.assertEqual(restored[0].state[0].tolist(), [1.0])
|
|
self.assertEqual(pool.available_size(), 3)
|
|
|
|
def test_auxiliary_state_pool_restores_instance_meta_state(self):
|
|
pool = MlxAuxiliaryStatePool(size=2, device="cpu")
|
|
slot = pool.alloc(1)
|
|
cache = [
|
|
FakeNativeCache(
|
|
mx.array([1.0], dtype=mx.float32),
|
|
meta_state={"seen": mx.array([3.0], dtype=mx.float32)},
|
|
)
|
|
]
|
|
pool.store_cache(slot[0], cache, [0])
|
|
cache[0].meta_state["seen"][0] = 9.0
|
|
|
|
restored = [FakeNativeCache(meta_state={})]
|
|
pool.restore_cache(slot[0], restored, [0])
|
|
|
|
self.assertEqual(restored[0].meta_state["seen"].tolist(), [3.0])
|
|
|
|
def test_auxiliary_state_req_pool_maps_request_indices(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
|
|
req_indices = pool.alloc([req])
|
|
auxiliary_state_idx = pool.get_auxiliary_state_indices(req.req_pool_idx)
|
|
pool.free(req)
|
|
|
|
self.assertEqual(req_indices, [1])
|
|
self.assertIsNotNone(auxiliary_state_idx)
|
|
self.assertIsNone(req.req_pool_idx)
|
|
self.assertIsNotNone(req.mamba_pool_idx)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
|
pool.free_auxiliary_state_cache(req)
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertEqual(pool.available_size(), 2)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
|
|
|
def test_auxiliary_state_req_pool_can_keep_tracked_auxiliary_slot(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
|
req.mamba_next_track_idx = 0
|
|
|
|
pool.free_auxiliary_state_cache(req, track_buffer_to_keep=0)
|
|
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
|
|
|
def test_auxiliary_state_component_inserts_tracked_slot_and_frees_live_slot(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
|
req.mamba_next_track_idx = 0
|
|
req.mamba_last_track_seqlen = 64
|
|
component = MlxAuxiliaryStateComponent(
|
|
SimpleNamespace(req_to_token_pool=pool),
|
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
|
)
|
|
insert_params = InsertParams()
|
|
|
|
cache_len = component.prepare_for_caching_req(
|
|
req=req,
|
|
insert_params=insert_params,
|
|
token_ids_len=70,
|
|
is_finished=True,
|
|
)
|
|
component.cleanup_after_caching_req(
|
|
req=req,
|
|
is_finished=True,
|
|
insert_result=InsertResult(prefix_len=0, mamba_exist=False),
|
|
insert_params=insert_params,
|
|
)
|
|
|
|
self.assertEqual(cache_len, 64)
|
|
self.assertTrue(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
|
self.assertEqual(insert_params.mamba_value.tolist(), [2])
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
|
self.assertIsNone(req.mamba_last_track_seqlen)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
|
|
|
def test_auxiliary_state_component_unfinished_frees_tracked_source_slot(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
|
req.mamba_next_track_idx = 0
|
|
req.mamba_last_track_seqlen = 64
|
|
component = MlxAuxiliaryStateComponent(
|
|
SimpleNamespace(req_to_token_pool=pool),
|
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
|
)
|
|
insert_params = InsertParams()
|
|
|
|
cache_len = component.prepare_for_caching_req(
|
|
req=req,
|
|
insert_params=insert_params,
|
|
token_ids_len=70,
|
|
is_finished=False,
|
|
)
|
|
component.cleanup_after_caching_req(
|
|
req=req,
|
|
is_finished=False,
|
|
insert_result=InsertResult(prefix_len=0, mamba_exist=False),
|
|
insert_params=insert_params,
|
|
)
|
|
|
|
self.assertEqual(cache_len, 64)
|
|
self.assertEqual(insert_params.mamba_value.tolist(), [3])
|
|
self.assertIsNotNone(req.mamba_pool_idx)
|
|
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
|
self.assertIsNone(req.mamba_last_track_seqlen)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 2)
|
|
|
|
def test_auxiliary_state_component_keeps_new_live_slot_owned_by_radix(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
component = MlxAuxiliaryStateComponent(
|
|
SimpleNamespace(req_to_token_pool=pool),
|
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
|
)
|
|
insert_params = InsertParams()
|
|
|
|
cache_len = component.prepare_for_caching_req(
|
|
req=req,
|
|
insert_params=insert_params,
|
|
token_ids_len=7,
|
|
is_finished=True,
|
|
)
|
|
component.cleanup_after_caching_req(
|
|
req=req,
|
|
is_finished=True,
|
|
insert_result=InsertResult(prefix_len=0, mamba_exist=False),
|
|
insert_params=insert_params,
|
|
)
|
|
|
|
self.assertEqual(cache_len, 7)
|
|
self.assertFalse(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
|
self.assertEqual(insert_params.mamba_value.tolist(), [1])
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
|
|
|
def test_auxiliary_state_component_frees_stale_track_slot_when_live_slot_inserted(
|
|
self,
|
|
):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
|
|
req.mamba_next_track_idx = 0
|
|
component = MlxAuxiliaryStateComponent(
|
|
SimpleNamespace(req_to_token_pool=pool),
|
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
|
)
|
|
insert_params = InsertParams()
|
|
|
|
cache_len = component.prepare_for_caching_req(
|
|
req=req,
|
|
insert_params=insert_params,
|
|
token_ids_len=7,
|
|
is_finished=True,
|
|
)
|
|
component.cleanup_after_caching_req(
|
|
req=req,
|
|
is_finished=True,
|
|
insert_result=InsertResult(prefix_len=0, mamba_exist=False),
|
|
insert_params=insert_params,
|
|
)
|
|
|
|
self.assertEqual(cache_len, 7)
|
|
self.assertFalse(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
|
|
self.assertEqual(insert_params.mamba_value.tolist(), [1])
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertIsNone(req.mamba_ping_pong_track_buffer)
|
|
self.assertIsNone(req.mamba_next_track_idx)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
|
|
|
|
def test_auxiliary_state_component_frees_duplicate_live_slot(self):
|
|
pool = MlxAuxiliaryStateReqToTokenPool(
|
|
size=2,
|
|
max_context_len=8,
|
|
device="cpu",
|
|
enable_memory_saver=False,
|
|
auxiliary_state_size=4,
|
|
)
|
|
req = FakeRequest()
|
|
pool.alloc([req])
|
|
component = MlxAuxiliaryStateComponent(
|
|
SimpleNamespace(req_to_token_pool=pool),
|
|
SimpleNamespace(enable_mamba_extra_buffer=False),
|
|
)
|
|
insert_params = InsertParams()
|
|
|
|
component.prepare_for_caching_req(
|
|
req=req,
|
|
insert_params=insert_params,
|
|
token_ids_len=7,
|
|
is_finished=True,
|
|
)
|
|
component.cleanup_after_caching_req(
|
|
req=req,
|
|
is_finished=True,
|
|
insert_result=InsertResult(prefix_len=7, mamba_exist=True),
|
|
insert_params=insert_params,
|
|
)
|
|
|
|
self.assertIsNone(req.mamba_pool_idx)
|
|
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
|
|
|
|
|
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
|
class TestMlxOverlapScheduler(unittest.TestCase):
|
|
def test_finalize_pending_job_updates_scheduler_last_batch(self):
|
|
token_ids = torch.tensor([7], dtype=torch.long)
|
|
scheduler = FakeOverlapScheduler(token_ids)
|
|
stale_batch = SimpleNamespace(input_ids=None)
|
|
batch_copy = SimpleNamespace(input_ids=None)
|
|
schedule_batch = SimpleNamespace(input_ids=None)
|
|
scheduler.last_batch = stale_batch
|
|
|
|
pending = MlxPendingJob(
|
|
lazy_tokens=None,
|
|
prefills=["prefill"],
|
|
extends=[],
|
|
decode=None,
|
|
mode="extend",
|
|
batch_copy=batch_copy,
|
|
schedule_batch=schedule_batch,
|
|
reqs=[SimpleNamespace(rid="r0")],
|
|
)
|
|
|
|
scheduler._finalize_mlx_pending_job(pending)
|
|
|
|
self.assertIs(scheduler.last_batch, schedule_batch)
|
|
self.assertTrue(torch.equal(batch_copy.input_ids, token_ids))
|
|
self.assertTrue(torch.equal(schedule_batch.input_ids, token_ids))
|
|
self.assertIs(scheduler.processed_batch, batch_copy)
|
|
self.assertIs(scheduler.processed_result, scheduler.tp_worker.result)
|
|
|
|
def test_finished_request_snapshots_before_release(self):
|
|
events = []
|
|
tree_cache = object()
|
|
processor = SchedulerBatchResultProcessor(
|
|
is_generation=True,
|
|
disaggregation_mode=None,
|
|
enable_overlap=False,
|
|
enable_overlap_mlx=False,
|
|
server_args=SimpleNamespace(
|
|
disaggregation_decode_enable_offload_kvcache=False,
|
|
enable_hisparse=False,
|
|
),
|
|
model_config=None,
|
|
token_to_kv_pool_allocator=None,
|
|
tree_cache=tree_cache,
|
|
hisparse_coordinator=None,
|
|
req_to_token_pool=None,
|
|
decode_offload_manager=None,
|
|
metrics_collector=None,
|
|
metrics_reporter=None,
|
|
draft_worker=None,
|
|
model_worker=SimpleNamespace(
|
|
prepare_for_kv_cache_release=lambda req: events.append(
|
|
("prepare", req.rid)
|
|
)
|
|
),
|
|
logprob_result_processor=None,
|
|
output_streamer=None,
|
|
abort_request=lambda req: None,
|
|
)
|
|
req = SimpleNamespace(
|
|
rid="r0",
|
|
finished=lambda: True,
|
|
multimodal_inputs=None,
|
|
session=None,
|
|
return_routed_experts=False,
|
|
time_stats=SimpleNamespace(
|
|
set_completion_time=lambda: events.append(("completion", "r0"))
|
|
),
|
|
)
|
|
original_release = batch_result_processor_module.release_kv_cache
|
|
original_get_indexer = batch_result_processor_module.get_global_indexer_capturer
|
|
|
|
def fake_release_kv_cache(release_req, tree_cache):
|
|
events.append(("release", release_req.rid))
|
|
self.assertIs(tree_cache, processor.tree_cache)
|
|
|
|
batch_result_processor_module.release_kv_cache = fake_release_kv_cache
|
|
batch_result_processor_module.get_global_indexer_capturer = lambda: None
|
|
try:
|
|
SchedulerBatchResultProcessor._handle_finished_req(
|
|
processor, req, 0, SimpleNamespace(customized_info=None)
|
|
)
|
|
finally:
|
|
batch_result_processor_module.release_kv_cache = original_release
|
|
batch_result_processor_module.get_global_indexer_capturer = (
|
|
original_get_indexer
|
|
)
|
|
|
|
self.assertEqual(
|
|
events,
|
|
[
|
|
("prepare", "r0"),
|
|
("release", "r0"),
|
|
("completion", "r0"),
|
|
],
|
|
)
|
|
|
|
|
|
if _HAS_MLX:
|
|
|
|
class FakeProjection(nn.Module):
|
|
def __init__(self, out_dim: int = 4):
|
|
super().__init__()
|
|
self.weight = mx.zeros((out_dim, 4), dtype=mx.float32)
|
|
|
|
def __call__(self, x):
|
|
shape = (*x.shape[:-1], self.weight.shape[0])
|
|
return mx.zeros(shape, dtype=x.dtype)
|
|
|
|
class FakeAttention(nn.Module):
|
|
def __init__(self, use_aliases: bool = False):
|
|
super().__init__()
|
|
if use_aliases:
|
|
self.num_attention_heads = 2
|
|
self.num_key_value_heads = 1
|
|
else:
|
|
self.n_heads = 2
|
|
self.n_kv_heads = 1
|
|
self.head_dim = 2
|
|
self.scale = self.head_dim**-0.5
|
|
self.q_proj = FakeProjection(4)
|
|
self.k_proj = FakeProjection(2)
|
|
self.v_proj = FakeProjection(2)
|
|
self.o_proj = FakeProjection(4)
|
|
self.rope = lambda x, offset=None: x
|
|
|
|
class ProjectionOnlyMixer(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.n_heads = 2
|
|
self.n_kv_heads = 1
|
|
self.q_proj = FakeProjection(4)
|
|
self.k_proj = FakeProjection(2)
|
|
self.v_proj = FakeProjection(2)
|
|
self.o_proj = FakeProjection(4)
|
|
|
|
class FakeLayer(nn.Module):
|
|
def __init__(self, attr_name: str, module: nn.Module):
|
|
super().__init__()
|
|
setattr(self, attr_name, module)
|
|
|
|
class FakeModel(nn.Module):
|
|
def __init__(self, layers):
|
|
super().__init__()
|
|
self.layers = layers
|
|
|
|
class IdentityNorm(nn.Module):
|
|
def __call__(self, x):
|
|
return x
|
|
|
|
class IdentityRope:
|
|
def __call__(self, x, offset=None):
|
|
return x
|
|
|
|
class CapturingOutput(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.last_input_shape = None
|
|
|
|
def __call__(self, x):
|
|
self.last_input_shape = x.shape
|
|
return x
|
|
|
|
class FakeGatedAttention(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.num_attention_heads = 2
|
|
self.num_key_value_heads = 1
|
|
self.head_dim = 2
|
|
self.scale = self.head_dim**-0.5
|
|
self.q_proj = FakeProjection(8)
|
|
self.k_proj = FakeProjection(2)
|
|
self.v_proj = FakeProjection(2)
|
|
self.o_proj = CapturingOutput()
|
|
self.q_norm = IdentityNorm()
|
|
self.k_norm = IdentityNorm()
|
|
self.rope = IdentityRope()
|
|
|
|
class RecordingIdentity(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.seen_shapes = []
|
|
|
|
def __call__(self, x):
|
|
self.seen_shapes.append(x.shape)
|
|
return x
|
|
|
|
class FakeMergeableLinearAttention(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.seen_shapes = []
|
|
self.cache_type = None
|
|
|
|
def __call__(self, x, mask=None, cache=None):
|
|
self.seen_shapes.append(x.shape)
|
|
self.cache_type = type(cache).__name__
|
|
cache[0] = mx.arange(x.shape[0], dtype=mx.float32).reshape(x.shape[0], 1)
|
|
return x + 1
|
|
|
|
class FakeBatchableAuxiliaryLayer(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.is_linear = True
|
|
self.input_layernorm = RecordingIdentity()
|
|
self.linear_attn = FakeMergeableLinearAttention()
|
|
self.post_attention_layernorm = RecordingIdentity()
|
|
self.mlp = RecordingIdentity()
|
|
|
|
class FakeMergeableAuxiliaryCache:
|
|
def __init__(self, state=None, tag="init", extra_metadata=None):
|
|
self.cache = [state]
|
|
self.tag = tag
|
|
self.extra_metadata = extra_metadata or {}
|
|
|
|
def __getitem__(self, idx):
|
|
return self.cache[idx]
|
|
|
|
def __setitem__(self, idx, value):
|
|
self.cache[idx] = value
|
|
|
|
@classmethod
|
|
def merge(cls, caches):
|
|
merged = cls(tag="merged")
|
|
values = [cache[0] for cache in caches]
|
|
if all(value is None for value in values):
|
|
return merged
|
|
merged[0] = mx.concatenate(
|
|
[
|
|
(
|
|
value
|
|
if value is not None
|
|
else mx.zeros_like(next(v for v in values if v is not None))
|
|
)
|
|
for value in values
|
|
],
|
|
axis=0,
|
|
)
|
|
return merged
|
|
|
|
def extract(self, idx):
|
|
return type(self)(
|
|
self.cache[0][idx : idx + 1],
|
|
tag=f"split-{idx}",
|
|
extra_metadata={"idx": idx},
|
|
)
|
|
|
|
class FakeNativeCache:
|
|
def __init__(self, value=None, meta_state=None):
|
|
self._state = [
|
|
value if value is not None else mx.array([0.0], dtype=mx.float32)
|
|
]
|
|
if meta_state is not None:
|
|
self.meta_state = meta_state
|
|
self.lengths = None
|
|
self.left_padding = None
|
|
|
|
@property
|
|
def state(self):
|
|
return self._state
|
|
|
|
@state.setter
|
|
def state(self, value):
|
|
self._state = value
|
|
|
|
class FakeAuxiliaryStateModel:
|
|
def __init__(self):
|
|
self.seen_inputs = []
|
|
self.seen_auxiliary_states = []
|
|
|
|
def make_cache(self):
|
|
return [FakeNativeCache(), FakeNativeCache()]
|
|
|
|
def __call__(self, inputs, cache=None):
|
|
self.seen_inputs.append(inputs.tolist())
|
|
if cache is not None:
|
|
self.seen_auxiliary_states.append(cache[0].state[0].tolist())
|
|
cache[0].state = [mx.array([float(inputs.shape[1])], dtype=mx.float32)]
|
|
keys = mx.ones((1, 1, inputs.shape[1], 2), dtype=mx.float32)
|
|
values = keys * 2
|
|
cache[1].update_and_fetch(keys, values)
|
|
return mx.zeros((1, inputs.shape[1], 4), dtype=mx.float32)
|
|
|
|
class FakeDenseModel:
|
|
def __init__(self, num_layers):
|
|
self.num_layers = num_layers
|
|
self.seen_inputs = []
|
|
self.seen_offsets = []
|
|
|
|
def __call__(self, inputs, cache=None):
|
|
self.seen_inputs.append(inputs.tolist())
|
|
if cache is not None:
|
|
offsets = []
|
|
for layer_idx in range(self.num_layers):
|
|
offsets.append(cache[layer_idx].offset)
|
|
scale = float(layer_idx + 1)
|
|
keys = mx.ones((1, 1, inputs.shape[1], 2), dtype=mx.float32) * scale
|
|
cache[layer_idx].update_and_fetch(keys, keys * 2)
|
|
self.seen_offsets.append(offsets)
|
|
return mx.zeros((1, inputs.shape[1], 4), dtype=mx.float32)
|
|
|
|
class FakeWrappedAttentionModel:
|
|
def __init__(self):
|
|
self.attn = MLXAttentionWrapper(FakeAttention(), layer_idx=0)
|
|
self.seen_inputs = []
|
|
self.seen_cache_types = []
|
|
|
|
def __call__(self, inputs, cache=None):
|
|
self.seen_inputs.append(inputs.tolist())
|
|
self.seen_cache_types.append([type(c).__name__ for c in cache])
|
|
hidden = mx.zeros((*inputs.shape, 4), dtype=mx.float32)
|
|
self.attn(hidden, cache=cache[0])
|
|
return mx.zeros((*inputs.shape, 8), dtype=mx.float32)
|
|
|
|
class FakeNestedStateCache:
|
|
@property
|
|
def state(self):
|
|
return [
|
|
mx.array([1.0], dtype=mx.float32),
|
|
None,
|
|
{"nested": (mx.array([2.0], dtype=mx.float32),)},
|
|
]
|
|
|
|
class FakeRequest:
|
|
def __init__(self):
|
|
self.req_pool_idx = None
|
|
self.mamba_pool_idx = None
|
|
self.inflight_middle_chunks = 0
|
|
self.kv_committed_len = 0
|
|
|
|
class FakeTpWorker:
|
|
def __init__(self, next_token_ids):
|
|
self.result = GenerationBatchResult(next_token_ids=next_token_ids)
|
|
self.calls = []
|
|
|
|
def finalize_mlx_result(self, *args):
|
|
self.calls.append(args)
|
|
return self.result
|
|
|
|
class FakeOverlapScheduler(SchedulerMlxOverlapMixin):
|
|
def __init__(self, next_token_ids):
|
|
self.tp_worker = FakeTpWorker(next_token_ids)
|
|
self.last_batch = None
|
|
self.processed_batch = None
|
|
self.processed_result = None
|
|
|
|
def process_batch_result(self, batch, result):
|
|
self.processed_batch = batch
|
|
self.processed_result = result
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|