From 547d203350f9863b0346744a688c2e335f14d7ee Mon Sep 17 00:00:00 2001 From: Asher Feldman <59994+asher@users.noreply.github.com> Date: Sun, 4 Oct 2026 17:28:40 -0700 Subject: [PATCH 1/2] fix(serve): affine kv-bits caches store prompt cache blocks after batched and lone requests --- CHANGELOG.md | 6 ++ gmlx/serve/patches/apc.py | 29 +++-- gmlx/upstream/seams.json | 2 + gmlx/upstream/seams.py | 4 + tests/serve/test_apc_quantized_harvest.py | 125 ++++++++++++++++++++++ tests/serve/test_server_patches.py | 6 +- 6 files changed, 160 insertions(+), 12 deletions(-) create mode 100644 tests/serve/test_apc_quantized_harvest.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 52a7e171..e087e85d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,12 @@ adhere to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). ## [Unreleased] +### Fixed + +- The server with affine `--kv-bits` stores prompt cache blocks, so a later + prompt reuses the part it shares. It stored none, and each batched prefill + logged `APC harvest failed`. + ## [0.4.20] - 2026-10-04 ### Added diff --git a/gmlx/serve/patches/apc.py b/gmlx/serve/patches/apc.py index 089f3e39..b4f96552 100644 --- a/gmlx/serve/patches/apc.py +++ b/gmlx/serve/patches/apc.py @@ -34,14 +34,13 @@ def install_apc_lone_harvest() -> None: This generalizes the harvest: when ``_idx`` is absent, fall back to ``offset`` over the single row (no left padding), which is exactly the slice the stock - code would take on a batched cache. A quantized KV cache keeps ``keys`` as a - tuple, which neither path can slice - skip it as before. ``ar.py`` calls the - function by module attribute (``_apc.harvest_blocks_from_batch_cache``), so - replacing it on the module is picked up at the call site without a fork. - - Against the pinned mlx-vlm 0.6.15 the replacement's behavioral delta - is the quantized (tuple ``keys``) decline: stock dequantizes and - stores, this keeps the block tier out of quantized-KV serving. + code would take on a batched cache. An affine cache (``--kv-bits``) keeps + ``keys`` as a tuple, which neither path can slice: it goes through stock + ``layer_kv_for_apc``, which dequantizes it, so affine keeps block reuse. + A batch row is extracted first, so only that row is dequantized. ``ar.py`` + calls the function by module attribute + (``_apc.harvest_blocks_from_batch_cache``), so replacing it on the module + is picked up at the call site without a fork. Signature contract: matches the 0.6.15 harvest exactly -- ``full_token_ids`` third positional, ``batch_idx`` keyword-only @@ -64,11 +63,23 @@ def harvest_blocks_from_batch_cache( if keys is None or values is None: return [] idx = getattr(c, "_idx", None) + if isinstance(keys, tuple): + # Affine cache: dequantize as stock does, one row only. + if idx is not None and hasattr(c, "extract"): + k, v = apc.layer_kv_for_apc(c.extract(row)) + else: + k, v = apc.layer_kv_for_apc( + c, batch_idx=None if idx is None else row) + if k is None or v is None: + return [] + layer_keys.append(k) + layer_values.append(v) + continue left_padding = getattr(c, "left_padding", None) if idx is None: # Lone-request fast path: a plain KVCache has a scalar `offset` # and no `_idx`/`left_padding`. Harvest its one row over - # [0, offset). Skip a quantized cache (tuple `keys`). + # [0, offset). offset = getattr(c, "offset", None) if offset is None or not isinstance(keys, mx.array): return [] diff --git a/gmlx/upstream/seams.json b/gmlx/upstream/seams.json index 576d7dbe..83314414 100644 --- a/gmlx/upstream/seams.json +++ b/gmlx/upstream/seams.json @@ -38,6 +38,7 @@ "mlx_vlm.apc:_safetensors_dtype_info": "f2d094d0b1ac7fa757bfe67ae7152a44d1351f717c98ed62082c2f3b4fb7fc63", "mlx_vlm.apc:_sequence_hash": "fd3c2cde19f0140e8c00d03618b5f26675d49a31aa7be834f8aedc63f5d7868e", "mlx_vlm.apc:harvest_blocks_from_batch_cache": "90ffae03efc603b07605490aa69aca8df969131b45c814c0ef311526c19211ba", + "mlx_vlm.apc:layer_kv_for_apc": "237ce309de2a0b7c5ca78b7bbab49e21c1f8278b6038bc0015182c5675f9b228", "mlx_vlm.apc:model_apc_mode": "259d55e1d4d86aff2f0adce50a0f8ae007c10452bdd1f3e4e3677f8e725c372a", "mlx_vlm.apc:multimodal_token_ids_from_config": "1714e1247cabf0a32d8e02361ecb41f4d89c6101751a63a8adecc3a24c041dfd", "mlx_vlm.generate.ar:BatchGenerator.__init__": "92d2de43ca109cf29cf1bc313d71ba478df0037aee9d293760ca070eac54b392", @@ -74,6 +75,7 @@ "mlx_vlm.models.cache:BatchKVCache": "11beb6c23464729eae4e1794d5c3133831dfd977315956bba991847ec6fb9c3c", "mlx_vlm.models.cache:BatchKVCache.extract": "11f78479e10a60b2ac960d0ec1b88fda272cde44cb877bd2318ad1f5efc2a502", "mlx_vlm.models.cache:BatchKVCache.filter": "9bf1a403703f244f75723125d4903034349bcfef545acf80cdef6804fdc60574", + "mlx_vlm.models.cache:BatchQuantizedKVCache.extract": "45ad86b4d80ca213e78bce9ee32c9b374beca1932ac2e278e2a581e3e25ac875", "mlx_vlm.models.cache:BatchQuantizedKVCache.make_mask": "17cab343ff407f6c461bcd3cc94753c012d51a07863fe765ed46b9913211a525", "mlx_vlm.models.cache:BatchQuantizedKVCache.update_and_fetch": "4411abff7fa4bc0a9273ce01483741fdc44d7ffc0f5ba60168c7268f00cc2550", "mlx_vlm.models.cache:BatchRotatingKVCache": "3cc47579ef4c8c4cf3326151565c6eebac6eb023fb8d9441e3c1831f29c00bb3", diff --git a/gmlx/upstream/seams.py b/gmlx/upstream/seams.py index 00bedc05..b024a3f4 100644 --- a/gmlx/upstream/seams.py +++ b/gmlx/upstream/seams.py @@ -356,6 +356,10 @@ class Seam: # --- APC internals (lone-harvest patch, gmlx manager subclass, apc_pooling) --- Seam("mlx_vlm.apc", "harvest_blocks_from_batch_cache", "server_patches.install_apc_lone_harvest", critical=True), + Seam("mlx_vlm.apc", "layer_kv_for_apc", + "server_patches.install_apc_lone_harvest (affine cache)"), + Seam("mlx_vlm.models.cache", "BatchQuantizedKVCache.extract", + "server_patches.install_apc_lone_harvest (affine batch row)"), Seam("mlx_vlm.apc", "_clone_layer_major_kv_cache_for_apc", "apc_manager.GmlxAPCManager.store_kv_blocks", critical=True), Seam("mlx_vlm.apc", "_sequence_hash", diff --git a/tests/serve/test_apc_quantized_harvest.py b/tests/serve/test_apc_quantized_harvest.py new file mode 100644 index 00000000..f4b47631 --- /dev/null +++ b/tests/serve/test_apc_quantized_harvest.py @@ -0,0 +1,125 @@ +"""APC block harvest on an affine (``--kv-bits``) cache. + +The cache a batched prefill commits under the affine scheme is mlx-vlm's +``BatchQuantizedKVCache``, whose keys are a tuple. The lone path holds a +``QuantizedKVCache``. The installed harvest must store the same blocks as +mlx-vlm's stock harvest for both, so that affine keeps partial prefix reuse. +""" +from __future__ import annotations + +import importlib + +import pytest + +pytest.importorskip("mlx_vlm") + +import mlx.core as mx # noqa: E402 + +from gmlx.serve.patches.apc import install_apc_lone_harvest # noqa: E402 + +B, H, D, T, BLOCK = 2, 2, 64, 64, 16 +LEFT_PAD = [0, 16] +N_LAYERS = 4 + + +class _Layers: + """A full-attention stack without make_cache, so _make_cache builds one + batch cache per layer.""" + layers = [object()] * N_LAYERS + + +def _kv(batch, seed): + mx.random.seed(seed) + k = mx.random.normal((batch, H, T, D)).astype(mx.float16) + v = mx.random.normal((batch, H, T, D)).astype(mx.float16) + return k, v + + +def _affine_batch_cache(): + ar = importlib.import_module("mlx_vlm.generate.ar") + caches = ar._make_cache(_Layers(), LEFT_PAD, kv_bits=8, kv_group_size=64, + kv_quant_scheme="uniform") + for i, c in enumerate(caches): + c.update_and_fetch(*_kv(B, i)) + mx.eval([c.state for c in caches]) + return caches + + +def _affine_lone_cache(): + cache_mod = importlib.import_module("mlx_vlm.models.cache") + caches = [cache_mod.QuantizedKVCache(group_size=64, bits=8) + for _ in range(N_LAYERS)] + for i, c in enumerate(caches): + c.update_and_fetch(*_kv(1, i)) + mx.eval([c.state for c in caches]) + return caches + + +def _manager(): + apc = importlib.import_module("mlx_vlm.apc") + return apc.APCManager(num_blocks=64, block_size=BLOCK) + + +@pytest.fixture +def harvests(monkeypatch): + """(stock, installed) harvest functions; monkeypatch restores the + module after the test.""" + apc = importlib.import_module("mlx_vlm.apc") + stock = apc.harvest_blocks_from_batch_cache + assert stock.__module__ == "mlx_vlm.apc" + monkeypatch.setattr(apc, "harvest_blocks_from_batch_cache", stock) + monkeypatch.setattr(apc, "_kq_lone_harvest", False, raising=False) + install_apc_lone_harvest() + installed = apc.harvest_blocks_from_batch_cache + assert installed is not stock + return stock, installed + + +def _assert_same_blocks(got, want): + assert len(got) == len(want) > 0 + for g, w in zip(got, want): + assert g.block_hash == w.block_hash + assert g.token_ids == w.token_ids + assert len(g.keys) == len(w.keys) == N_LAYERS + for gk, wk, gv, wv in zip(g.keys, w.keys, g.values, w.values): + assert gk.shape == wk.shape == (1, H, BLOCK, D) + assert gk.dtype == wk.dtype + assert mx.array_equal(gk, wk).item() + assert mx.array_equal(gv, wv).item() + + +@pytest.mark.parametrize("row", [0, 1]) +def test_batched_affine_harvest_matches_stock(harvests, row): + stock, installed = harvests + caches = _affine_batch_cache() + # mlx-vlm keeps the last layer fp16, so the stack mixes both kinds. + assert sum(isinstance(c.keys, tuple) for c in caches) == N_LAYERS - 1 + ids = list(range(T - LEFT_PAD[row])) + want = stock(_manager(), caches, ids, batch_idx=row) + got = installed(_manager(), caches, ids, batch_idx=row) + assert len(got) == len(ids) // BLOCK + _assert_same_blocks(got, want) + + +@pytest.mark.parametrize("row", [0, 1]) +def test_commit_after_batched_affine_prefill_stores_blocks(harvests, row): + """ar.py commits through the module global after a batched prefill and + logs 'APC harvest failed during batched prefill' on any exception.""" + stock, _ = harvests + apc = importlib.import_module("mlx_vlm.apc") + caches = _affine_batch_cache() + ids = list(range(T - LEFT_PAD[row])) + want = stock(_manager(), caches, ids, batch_idx=row) + got = apc.commit_prefix_blocks(_manager(), caches, ids, batch_idx=row) + _assert_same_blocks(got, want) + + +def test_lone_affine_harvest_matches_stock(harvests): + stock, installed = harvests + caches = _affine_lone_cache() + assert all(isinstance(c.keys, tuple) for c in caches) + ids = list(range(T)) + want = stock(_manager(), caches, ids) + got = installed(_manager(), caches, ids) + assert len(got) == T // BLOCK + _assert_same_blocks(got, want) diff --git a/tests/serve/test_server_patches.py b/tests/serve/test_server_patches.py index 2133a1b1..f2efe7bd 100644 --- a/tests/serve/test_server_patches.py +++ b/tests/serve/test_server_patches.py @@ -1229,8 +1229,8 @@ def test_lone_harvest_preserves_batched_path(): def test_lone_harvest_skips_unsupported_cache(): - """A cache with neither _idx nor offset (or quantized tuple keys) is declined, - not crashed.""" + """A cache with neither _idx nor offset, or with tuple keys and no + dequantize_for_apc, is declined, not crashed.""" import mlx.core as mx apc = importlib.import_module("mlx_vlm.apc") sp.install_apc_lone_harvest() @@ -1240,7 +1240,7 @@ def test_lone_harvest_skips_unsupported_cache(): nocache = [_FakeKVCache(mx.zeros((1, 2, 8, 8)), mx.zeros((1, 2, 8, 8)), None)] assert apc.harvest_blocks_from_batch_cache(mgr, nocache, list(range(8))) == [] - # quantized: keys is a tuple, not an mx.array + # tuple keys that nothing can dequantize quant = [_FakeKVCache((mx.zeros((1, 2, 8, 4)),), (mx.zeros((1, 2, 8, 4)),), 8)] assert apc.harvest_blocks_from_batch_cache(mgr, quant, list(range(8))) == [] assert mgr.calls == [] From 9f4618db6e3af495c15c2dab1af4ec5bd2b37b55 Mon Sep 17 00:00:00 2001 From: Asher Feldman <59994+asher@users.noreply.github.com> Date: Sun, 4 Oct 2026 17:28:40 -0700 Subject: [PATCH 2/2] test(container): launch output tests clear the caller's PYTHONPATH --- tests/container/test_container_launch.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/container/test_container_launch.py b/tests/container/test_container_launch.py index d8841d9f..42ddb000 100644 --- a/tests/container/test_container_launch.py +++ b/tests/container/test_container_launch.py @@ -994,6 +994,7 @@ def test_the_dry_run_shows_the_context_window_claude_code_gets(env, capsys, monk entry, line): name = "CLAUDE_CODE_MAX_CONTEXT_TOKENS" monkeypatch.delenv(name, raising=False) + monkeypatch.delenv("PYTHONPATH", raising=False) # no launch warning if entry: _user_config(env.home, "launch:\n container:\n clients:\n claude-code:\n" f" env: [\"{entry.replace('NAME', name)}\"]\n") @@ -4555,7 +4556,9 @@ def test_the_model_comes_from_the_flag_the_setting_or_the_server(env, capsys, mo "--model, or set launch.agents.bot.model.") in capsys.readouterr().out.splitlines() -def test_a_global_env_entry_that_launch_sets_is_named_once_with_its_block(env, capsys): +def test_a_global_env_entry_that_launch_sets_is_named_once_with_its_block(env, capsys, + monkeypatch): + monkeypatch.delenv("PYTHONPATH", raising=False) # no launch warning _user_config(env.home, "launch:\n container:\n open_browser: false\n" " env: [IS_SANDBOX=1, UV_CACHE_DIR=/c]\n") line = ("[launch] the entry IS_SANDBOX in launch.container.env has no effect for "