""" Frox AI Morph 1.1 — KV Cache Tests Run with: pytest tests/test_kv_cache.py -v Includes a dedicated regression test for a real bug caught during development: `MorphSessionCache.chat()`-style usage originally called `cache.step(0)` after writing new tokens, which is a no-op — `_seq_len` never advanced, so every write silently landed back at offset 0 and overwrote whatever was written before it. `test_sequential_updates_append_not_overwrite` below fails immediately on that bug and passes with the fix. """ from __future__ import annotations import pytest import torch from inference.cache.kv_cache import MorphKVCache, MorphSessionCache NUM_LAYERS = 2 NUM_KV_HEADS = 2 HEAD_DIM = 8 BLOCK_SIZE = 4 # small on purpose to force multiple blocks quickly @pytest.fixture def cache() -> MorphKVCache: return MorphKVCache( num_layers=NUM_LAYERS, num_kv_heads=NUM_KV_HEADS, head_dim=HEAD_DIM, block_size=BLOCK_SIZE, dtype=torch.float32, device=torch.device("cpu"), ) def _rand_kv(seq_len: int): k = torch.randn(1, NUM_KV_HEADS, seq_len, HEAD_DIM) v = torch.randn(1, NUM_KV_HEADS, seq_len, HEAD_DIM) return k, v class TestBasicReadWrite: def test_single_write_round_trips(self, cache): k, v = _rand_kv(3) cache.update(0, k, v) cache.step(3) k_read, v_read = cache.get(0) assert torch.allclose(k_read, k) assert torch.allclose(v_read, v) assert cache.seq_len == 3 def test_empty_cache_returns_zero_length(self, cache): k, v = cache.get(0) assert k.shape[2] == 0 assert v.shape[2] == 0 assert cache.seq_len == 0 class TestSequentialAppend: """ The core regression test: writing in several separate calls (as happens once per generated token in `MorphInferenceEngine.chat()`) must APPEND, never overwrite. This is exactly the bug that was caught: calling `step(0)` instead of `step(delta_len)` made every write land at offset 0 again. """ def test_sequential_updates_append_not_overwrite(self, cache): chunks = [_rand_kv(1) for _ in range(6)] # simulate 6 generated tokens, one at a time for k, v in chunks: cache.update(0, k, v) cache.step(1) # <- the fix: must advance by the actual delta length assert cache.seq_len == 6 k_read, v_read = cache.get(0) assert k_read.shape[2] == 6, ( f"Expected 6 cached tokens after 6 sequential single-token writes, " f"got {k_read.shape[2]} — writes are overwriting instead of appending " f"(this is exactly the step(0) vs step(delta_len) bug)." ) expected_k = torch.cat([c[0] for c in chunks], dim=2) expected_v = torch.cat([c[1] for c in chunks], dim=2) assert torch.allclose(k_read, expected_k), "Cached content doesn't match write order" assert torch.allclose(v_read, expected_v) def test_mixed_batch_and_single_writes(self, cache): """A multi-token prefill followed by several single-token decode steps.""" prefill_k, prefill_v = _rand_kv(5) cache.update(0, prefill_k, prefill_v) cache.step(5) decode_chunks = [_rand_kv(1) for _ in range(4)] for k, v in decode_chunks: cache.update(0, k, v) cache.step(1) assert cache.seq_len == 9 k_read, _ = cache.get(0) assert k_read.shape[2] == 9 expected = torch.cat([prefill_k] + [c[0] for c in decode_chunks], dim=2) assert torch.allclose(k_read, expected) def test_all_layers_stay_in_sync(self, cache): """ step() is called once per forward pass (not once per layer), so every layer must end up with the same seq_len even though update() is called once per layer per step. """ for _ in range(5): for layer_idx in range(NUM_LAYERS): k, v = _rand_kv(1) cache.update(layer_idx, k, v) cache.step(1) # called once, after all layers for this step for layer_idx in range(NUM_LAYERS): k, v = cache.get(layer_idx) assert k.shape[2] == 5, f"Layer {layer_idx} has {k.shape[2]} tokens, expected 5" class TestPagedAllocation: def test_grows_across_block_boundary(self, cache): """block_size=4 in the fixture; write more than one block's worth.""" k, v = _rand_kv(10) # spans 3 blocks (4+4+2) cache.update(0, k, v) cache.step(10) k_read, v_read = cache.get(0) assert k_read.shape[2] == 10 assert torch.allclose(k_read, k) def test_memory_scales_with_allocated_blocks_not_max_seq_len(self, cache): """A short sequence shouldn't allocate memory for a hypothetical long one.""" k, v = _rand_kv(1) cache.update(0, k, v) cache.step(1) # Only layer 0 was written; layer 1 should have zero blocks allocated mem = cache.memory_mb() assert mem > 0 # Rough sanity: memory should be proportional to ~1 block for 1 layer, # not proportional to some large fixed max_seq_len. bytes_per_block = NUM_KV_HEADS * BLOCK_SIZE * HEAD_DIM * 4 * 2 # k+v, float32 expected_mb = bytes_per_block / (1024 ** 2) assert mem <= expected_mb * 1.5 # generous tolerance class TestUtilization: def test_utilization_reflects_fill_ratio(self, cache): k, v = _rand_kv(2) # half of one 4-token block cache.update(0, k, v) cache.step(2) assert cache.utilization() == pytest.approx(0.5, abs=0.01) class TestReset: def test_reset_clears_everything(self, cache): k, v = _rand_kv(5) cache.update(0, k, v) cache.step(5) assert cache.seq_len == 5 cache.reset() assert cache.seq_len == 0 k_read, _ = cache.get(0) assert k_read.shape[2] == 0 class TestSessionCache: def test_sessions_are_isolated(self): sessions = MorphSessionCache( num_layers=NUM_LAYERS, num_kv_heads=NUM_KV_HEADS, head_dim=HEAD_DIM, block_size=BLOCK_SIZE, device=torch.device("cpu"), max_sessions=4, ) cache_a = sessions.get_or_create("session-a") cache_b = sessions.get_or_create("session-b") k, v = _rand_kv(3) cache_a.update(0, k, v) cache_a.step(3) assert cache_a.seq_len == 3 assert cache_b.seq_len == 0, "Writing to session A must not affect session B" def test_lru_eviction(self): sessions = MorphSessionCache( num_layers=NUM_LAYERS, num_kv_heads=NUM_KV_HEADS, head_dim=HEAD_DIM, block_size=BLOCK_SIZE, device=torch.device("cpu"), max_sessions=2, ) sessions.get_or_create("a") sessions.get_or_create("b") sessions.get_or_create("a") # touch "a" again — "b" is now the LRU sessions.get_or_create("c") # should evict "b", not "a" assert "a" in sessions._caches assert "b" not in sessions._caches assert "c" in sessions._caches def test_reset_session_preserves_other_sessions(self): sessions = MorphSessionCache( num_layers=NUM_LAYERS, num_kv_heads=NUM_KV_HEADS, head_dim=HEAD_DIM, block_size=BLOCK_SIZE, device=torch.device("cpu"), ) cache_a = sessions.get_or_create("a") cache_b = sessions.get_or_create("b") for c in (cache_a, cache_b): k, v = _rand_kv(3) c.update(0, k, v) c.step(3) sessions.reset_session("a") assert sessions.get_or_create("a").seq_len == 0 assert sessions.get_or_create("b").seq_len == 3