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
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 20 additions & 9 deletions gmlx/serve/patches/apc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 []
Expand Down
2 changes: 2 additions & 0 deletions gmlx/upstream/seams.json
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
4 changes: 4 additions & 0 deletions gmlx/upstream/seams.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
5 changes: 4 additions & 1 deletion tests/container/test_container_launch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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 "
Expand Down
125 changes: 125 additions & 0 deletions tests/serve/test_apc_quantized_harvest.py
Original file line number Diff line number Diff line change
@@ -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)
6 changes: 3 additions & 3 deletions tests/serve/test_server_patches.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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 == []
Expand Down
Loading