Add batch support for sliding window cache (#487)

* add batch support for sliding window cache

* fix

* fix masks

* fix cache

* another test

* comment
This commit is contained in:
Awni Hannun
2025-09-23 20:56:47 -07:00
committed by GitHub
parent eaf1748ea5
commit 50012d153d
6 changed files with 418 additions and 36 deletions
+8
View File
@@ -26,8 +26,10 @@ from .models import cache
from .models.cache import (
ArraysCache,
BatchKVCache,
BatchRotatingKVCache,
KVCache,
QuantizedKVCache,
RotatingKVCache,
load_prompt_cache,
)
from .sample_utils import make_sampler
@@ -866,6 +868,12 @@ def _make_cache(model, left_padding):
elif isinstance(c, ArraysCache):
c.left_padding = mx.array(left_padding)
batch_cache.append(c)
elif isinstance(c, RotatingKVCache):
if c.keep > 0:
raise ValueError(
"RotatingKVCache with keep tokens is not supported."
)
batch_cache.append(BatchRotatingKVCache(c.max_size, left_padding))
else:
raise ValueError(f"{type(c)} does not yet support batching")
return batch_cache
+256 -22
View File
@@ -1,5 +1,6 @@
# Copyright © 2023-2024 Apple Inc.
import copy
from typing import Any, Dict, List, Optional
import mlx.core as mx
@@ -73,10 +74,10 @@ def load_prompt_cache(file_name, return_metadata=False):
arrays = tree_unflatten(list(arrays.items()))
cache_metadata = tree_unflatten(list(cache_metadata.items()))
info, metadata, classes = cache_metadata
cache = [globals()[c]() for c in classes]
for c, state, meta_state in zip(cache, arrays, info):
c.state = state
c.meta_state = meta_state
cache = [
globals()[c].from_state(state, meta_state)
for c, state, meta_state in zip(classes, arrays, info)
]
if return_metadata:
return cache, metadata
return cache
@@ -141,6 +142,14 @@ class _BaseCache:
def is_trimmable(self):
return False
@classmethod
def from_state(cls, state, meta_state):
# Create an instance of cls without calling __init__
obj = cls.__new__(cls)
obj.state = state
obj.meta_state = meta_state
return obj
class ConcatenateKVCache(_BaseCache):
"""ConcatenateKVCache the simplest KV cache implementation.
@@ -188,11 +197,12 @@ class ConcatenateKVCache(_BaseCache):
class QuantizedKVCache(_BaseCache):
step = 256
def __init__(self, group_size: int = 64, bits: int = 8):
self.keys = None
self.values = None
self.offset = 0
self.step = 256
self.group_size = group_size
self.bits = bits
@@ -254,11 +264,11 @@ class QuantizedKVCache(_BaseCache):
@property
def meta_state(self):
return tuple(map(str, (self.step, self.offset, self.group_size, self.bits)))
return tuple(map(str, (self.offset, self.group_size, self.bits)))
@meta_state.setter
def meta_state(self, v):
self.step, self.offset, self.group_size, self.bits = map(int, v)
self.offset, self.group_size, self.bits = map(int, v)
def is_trimmable(self):
return True
@@ -273,11 +283,12 @@ class QuantizedKVCache(_BaseCache):
class KVCache(_BaseCache):
step = 256
def __init__(self):
self.keys = None
self.values = None
self.offset = 0
self.step = 256
def update_and_fetch(self, keys, values):
prev = self.offset
@@ -341,14 +352,14 @@ class KVCache(_BaseCache):
class RotatingKVCache(_BaseCache):
step = 256
def __init__(self, max_size=None, keep=0, step=256):
def __init__(self, max_size, keep=0):
self.keep = keep
self.keys = None
self.values = None
self.offset = 0
self.max_size = max_size
self.step = step
self._idx = 0
def _trim(self, trim_size, v, append=None):
@@ -389,9 +400,9 @@ class RotatingKVCache(_BaseCache):
self.keys = self._temporal_order(self.keys)
self.values = self._temporal_order(self.values)
# The largest size is self.max_size + S to ensure
# The largest size is self.max_size + S - 1 to ensure
# every token gets at least self.max_size context
trim_size = self._idx - self.max_size
trim_size = self._idx - self.max_size + 1
self.keys = self._trim(trim_size, self.keys, keys)
self.values = self._trim(trim_size, self.values, values)
self.offset += keys.shape[2]
@@ -459,13 +470,11 @@ class RotatingKVCache(_BaseCache):
@property
def meta_state(self):
return tuple(
map(str, (self.keep, self.max_size, self.step, self.offset, self._idx))
)
return tuple(map(str, (self.keep, self.max_size, self.offset, self._idx)))
@meta_state.setter
def meta_state(self, v):
self.keep, self.max_size, self.step, self.offset, self._idx = map(
self.keep, self.max_size, self.offset, self._idx = map(
int,
v,
)
@@ -487,7 +496,7 @@ class RotatingKVCache(_BaseCache):
):
if N > 1:
window_size = window_size or self.max_size
offset = min(self.max_size, self.offset)
offset = min(self.max_size - 1, self.offset)
if offset + N > window_size or return_array:
return create_causal_mask(N, offset, window_size=window_size)
else:
@@ -500,16 +509,19 @@ class RotatingKVCache(_BaseCache):
idx = self._idx
if idx >= self.max_size:
idx = 0
mask_size = min(self.max_size, self.offset)
if self.offset < self.max_size:
mask_size = self.offset + 1
else:
mask_size = self.max_size
mask = mx.arange(mask_size) >= (mask_size - window_size)
mask = mx.roll(mask, shift=idx + 1)
return mask[:, None]
return mask
class ArraysCache(_BaseCache):
def __init__(self, size, left_padding: Optional[List[int]] = None):
self.cache = [None] * size
self.left_padding = left_padding
self.left_padding = mx.array(left_padding) if left_padding else None
def __setitem__(self, idx, value):
self.cache[idx] = value
@@ -552,7 +564,7 @@ class MambaCache(ArraysCache):
class ChunkedKVCache(KVCache):
def __init__(self, chunk_size=None):
def __init__(self, chunk_size):
super().__init__()
self.chunk_size = chunk_size
self.start_position = 0
@@ -633,6 +645,8 @@ class CacheList(KVCache):
class BatchKVCache(_BaseCache):
step = 256
def __init__(self, left_padding: List[int]):
"""
The BatchKV cache expects inputs to be left-padded.
@@ -657,7 +671,6 @@ class BatchKVCache(_BaseCache):
self.left_padding = mx.array(left_padding)
self.offset = mx.array([-l for l in left_padding])
self._idx = 0
self.step = 256
def update_and_fetch(self, keys, values):
prev = self._idx
@@ -756,3 +769,224 @@ class BatchKVCache(_BaseCache):
mx.concatenate, zip(*(pad(self), pad(other)))
)
self._idx = max_idx
class BatchRotatingKVCache(_BaseCache):
step = 256
def __init__(self, max_size, left_padding: List[int]):
self.keys = None
self.values = None
self.left_padding = mx.array(left_padding)
self.offset = mx.array([-l for l in left_padding])
self.max_size = max_size
self._idx = 0
self._offset = 0
self.rotated = False
def _trim(self, trim_size, v, append=None):
if trim_size > 0:
v = v[..., trim_size:, :]
if append is not None:
return mx.concatenate([v, append], axis=2)
return v
def _temporal_order(self):
"""
Rearrange the cache into temporal order.
"""
if self.rotated:
self.keys = mx.roll(self.keys, -self._idx, axis=2)
self.values = mx.roll(self.values, -self._idx, axis=2)
self._idx = self.keys.shape[2]
self.rotated = False
def _update_concat(self, keys, values):
if self.keys is None:
self.keys = keys
self.values = values
else:
# Put the keys/values in temporal order to
# preserve context
self._temporal_order()
# Slice off the end if needed
if self.keys.shape[2] > self._idx:
self.keys = self.keys[..., : self._idx, :]
self.values = self.values[..., : self._idx, :]
# The largest size is self.max_size + S - 1 to ensure
# every token gets at least self.max_size context
trim_size = self._idx - self.max_size + 1
if trim_size > 0:
self.left_padding -= trim_size
self.keys = self._trim(trim_size, self.keys, keys)
self.values = self._trim(trim_size, self.values, values)
self.offset += keys.shape[2]
self._offset += keys.shape[2]
self._idx = self.keys.shape[2]
return self.keys, self.values
def _update_in_place(self, keys, values):
# May not have hit the max size yet, so potentially
# keep growing the cache
B, n_kv_heads, S, k_head_dim = keys.shape
prev = self._offset
if self.keys is None or (
prev >= self.keys.shape[2] and self.keys.shape[2] < self.max_size
):
v_head_dim = values.shape[3]
new_size = min(self.step, self.max_size - prev)
k_shape = (B, n_kv_heads, new_size, k_head_dim)
v_shape = (B, n_kv_heads, new_size, v_head_dim)
new_k = mx.zeros(k_shape, keys.dtype)
new_v = mx.zeros(v_shape, values.dtype)
if self.keys is not None:
self.keys = mx.concatenate([self.keys, new_k], axis=2)
self.values = mx.concatenate([self.values, new_v], axis=2)
else:
self.keys, self.values = new_k, new_v
self._idx = prev
# Trim if needed
trim_size = self.keys.shape[2] - self.max_size
if trim_size > 0:
self.keys = self._trim(trim_size, self.keys)
self.values = self._trim(trim_size, self.values)
self._idx = self.max_size
self.left_padding -= trim_size
# Rotate
if self._idx == self.max_size:
self.rotated = True
self._idx = 0
if self.rotated:
self.left_padding -= S
# Assign
self.keys[..., self._idx : self._idx + S, :] = keys
self.values[..., self._idx : self._idx + S, :] = values
self._offset += S
self.offset += S
self._idx += S
# If the buffer is not full, slice off the end
if self._offset < self.max_size:
return (
self.keys[..., : self._offset, :],
self.values[..., : self._offset, :],
)
return self.keys, self.values
def update_and_fetch(self, keys, values):
if keys.shape[2] == 1:
return self._update_in_place(keys, values)
return self._update_concat(keys, values)
@property
def state(self):
k, v = self.keys, self.values
if self._offset < k.shape[2]:
k, v = k[..., : self._offset, :], v[..., : self._offset, :]
return k, v, self.offset, self.left_padding
@state.setter
def state(self, v):
self.keys, self.values, self.offset, self.left_padding = v
@property
def meta_state(self):
return tuple(map(str, (self.max_size, self._offset, self._idx, self.rotated)))
@meta_state.setter
def meta_state(self, v):
self.max_size, self._offset, self._idx = map(
int,
v[:3],
)
self.rotated = bool(v[3])
def is_trimmable(self):
return self._offset < self.max_size
def trim(self, n):
n = min(self._offset, n)
self._offset -= n
self._idx -= n
self.offset -= n
return n
def to_quantized(self, group_size: int = 64, bits: int = 4) -> QuantizedKVCache:
raise NotImplementedError("BatchRotatingKVCache Quantization NYI")
def make_mask(
self, N: int, window_size: Optional[int] = None, return_array: bool = False
):
left_padding = self.left_padding
window_size = window_size or self.max_size
offset = min(self.max_size - 1, self._offset)
rinds = mx.arange(offset + N)
linds = mx.arange(offset, offset + N) if offset else rinds
linds = linds[:, None]
rinds = rinds[None]
mask = linds >= rinds
mask &= linds < rinds + window_size
if (trim_size := self._idx - self.max_size + int(N > 1)) > 0:
left_padding = left_padding - trim_size
rotated = N == 1 and (self.rotated or self._idx >= self.max_size)
if rotated:
left_padding = left_padding - 1
mask = mask & (rinds >= mx.expand_dims(left_padding, (1, 2, 3)))
if rotated:
idx = self._idx
if idx >= self.max_size:
idx = 0
mask = mx.roll(mask, shift=idx + 1, axis=-1)
return mask
def filter(self, batch_indices):
"""
In-place filter to keep just the given indices in the cache.
"""
self.keys = self.keys[batch_indices]
self.values = self.values[batch_indices]
self.offset = self.offset[batch_indices]
self.left_padding = self.left_padding[batch_indices]
def extend(self, other):
"""
In-place extend this cache with the other cache.
"""
if (self.rotated != other.rotated) or self._idx != other._idx:
self._temporal_order()
other._temporal_order()
max_idx = max(self._idx, other._idx)
max_size = max(self.keys.shape[2], other.keys.shape[2])
def pad(c):
left = max_idx - c._idx
right = max_size - c.keys.shape[2] - left
k, v = c.keys, c.values
if right < 0:
k = k[..., :right, :]
v = v[..., :right, :]
right = 0
if left != 0 or right != 0:
pad = [(0, 0), (0, 0), (left, right), (0, 0)]
k = mx.pad(k, pad)
v = mx.pad(v, pad)
left_padding = c.left_padding + left
return k, v, c.offset, left_padding
self.keys, self.values, self.offset, self.left_padding = map(
mx.concatenate, zip(*(pad(self), pad(other)))
)
self._idx = max_idx
self._offset = max(self._offset, other._offset)
+1 -6
View File
@@ -87,8 +87,6 @@ class Attention(nn.Module):
keys = self.rope(keys)
# Sliding window
if isinstance(mask, mx.array) and mask.shape[-1] != keys.shape[-2]:
mask = mask[..., -keys.shape[-2] :]
output = scaled_dot_product_attention(
queries, keys, values, cache=cache, scale=self.scale, mask=mask
)
@@ -194,7 +192,6 @@ class Gemma3Model(nn.Module):
cache[0],
window_size=self.window_size,
)
for i, (layer, c) in enumerate(zip(self.layers, cache)):
is_global = (
i % self.sliding_window_pattern == self.sliding_window_pattern - 1
@@ -246,7 +243,5 @@ class Model(nn.Module):
):
caches.append(KVCache())
else:
caches.append(
RotatingKVCache(max_size=self.args.sliding_window, keep=0)
)
caches.append(RotatingKVCache(max_size=self.args.sliding_window))
return caches
+51
View File
@@ -11,6 +11,7 @@ from mlx_lm.generate import (
generate,
stream_generate,
)
from mlx_lm.models.cache import RotatingKVCache
from mlx_lm.sample_utils import make_logits_processors, make_sampler
from mlx_lm.utils import load
@@ -301,6 +302,56 @@ class TestGenerate(unittest.TestCase):
batch_tokens = batch_responses[uids[e]]
self.assertEqual(tokens, batch_tokens)
def test_batch_sliding_window(self):
prompts = [
"Write a story about Einstein",
"Hi",
"What time is it?",
"How tall is Mt Everest?",
]
prompts = [
self.tokenizer.apply_chat_template(
[{"role": "user", "content": p}],
tokenize=True,
add_generation_prompt=True,
)
for p in prompts
]
self.model.make_cache = lambda: [
RotatingKVCache(max_size=4) for _ in self.model.layers
]
batch_gen = BatchGenerator(
self.model,
stop_tokens=self.tokenizer.eos_token_ids,
max_tokens=10,
prefill_batch_size=1,
prefill_step_size=8,
completion_batch_size=2,
)
uids = batch_gen.insert(prompts)
batch_responses = {uid: [] for uid in uids}
while responses := batch_gen.next():
for r in responses:
batch_responses[r.uid].append(r.logprobs)
for e, uid in enumerate(uids):
for i, response in enumerate(
stream_generate(
self.model,
self.tokenizer,
prompts[e],
max_tokens=10,
)
):
batch_logprobs = batch_responses[uid][i]
logprobs = response.logprobs
self.assertTrue(
mx.allclose(batch_logprobs, logprobs, rtol=1e-4, atol=1e-4)
)
del self.model.make_cache
if __name__ == "__main__":
unittest.main()
+3 -3
View File
@@ -37,7 +37,7 @@ class TestModels(unittest.TestCase):
def test_rotating_kv_cache(self):
b, h, d = 1, 2, 32
cache = RotatingKVCache(max_size=8, step=4)
cache = RotatingKVCache(max_size=8)
k = mx.random.uniform(shape=(b, h, 2, d))
v = mx.random.uniform(shape=(b, h, 2, d))
@@ -70,7 +70,7 @@ class TestModels(unittest.TestCase):
idx %= 8
# Try with nonzero keep
cache = RotatingKVCache(max_size=8, step=4, keep=2)
cache = RotatingKVCache(max_size=8, keep=2)
# Check a large update
k = mx.random.uniform(shape=(b, h, 20, d))
@@ -98,7 +98,7 @@ class TestModels(unittest.TestCase):
# alternating prompt/prefill with generation
d = 4
h = 2
cache = RotatingKVCache(max_size=18, step=4)
cache = RotatingKVCache(max_size=18)
x = mx.random.uniform(shape=(1, h, 8, d))
k, v = cache.update_and_fetch(x, x)
+99 -5
View File
@@ -11,6 +11,7 @@ from mlx_lm.generate import generate_step
from mlx_lm.models.base import create_attention_mask, create_causal_mask
from mlx_lm.models.cache import (
BatchKVCache,
BatchRotatingKVCache,
CacheList,
ChunkedKVCache,
KVCache,
@@ -391,7 +392,7 @@ class TestPromptCache(unittest.TestCase):
kv = mx.zeros((1, 1, 10, 32))
cache.update_and_fetch(kv, kv)
mask = cache.make_mask(3, window_size=5)
self.assertEqual(mask.shape, (3, 11))
self.assertEqual(mask.shape, (3, 10))
self.assertTrue(mx.all(mask.sum(axis=-1) == 5))
for i in range(3):
s = 11 - 3 + i
@@ -405,7 +406,7 @@ class TestPromptCache(unittest.TestCase):
self.assertEqual(mask, None)
mask = cache.make_mask(1, window_size=5)
self.assertEqual(mask.squeeze(1).tolist(), [True] + [False] * 3 + [True] * 4)
self.assertEqual(mask.tolist(), [True] + [False] * 3 + [True] * 4)
cmask = create_attention_mask(mx.zeros((1, 1, 32)), cache, window_size=5)
self.assertTrue(mx.array_equal(cmask, mask))
@@ -413,9 +414,7 @@ class TestPromptCache(unittest.TestCase):
cache.update_and_fetch(kv, kv)
mask = cache.make_mask(1, window_size=5)
self.assertEqual(
mask.squeeze(1).tolist(), [True] * 2 + [False] * 3 + [True] * 3
)
self.assertEqual(mask.tolist(), [True] * 2 + [False] * 3 + [True] * 3)
cmask = create_attention_mask(mx.zeros((1, 1, 32)), cache, window_size=5)
self.assertTrue(mx.array_equal(cmask, mask))
@@ -460,6 +459,101 @@ class TestPromptCache(unittest.TestCase):
self.assertEqual(cache_a.offset.tolist(), [6, 7, 6, 1, 4])
self.assertEqual(cache_a.left_padding.tolist(), [2, 1, 2, 7, 4])
def test_batch_rotating_kv_cache(self):
cache = BatchRotatingKVCache(max_size=4, left_padding=[2, 0])
mask = cache.make_mask(4)
self.assertFalse(mx.any(mask[0, 0, 0, :]))
self.assertTrue(
mx.array_equal(mask[1, 0, 0, :], mx.array([True, False, False, False]))
)
# Batch update works
k, v = mx.zeros((2, 1, 4, 8)), mx.zeros((2, 1, 4, 8))
k, v = cache.update_and_fetch(k, v)
mask = cache.make_mask(4)
k, v = mx.zeros((2, 1, 4, 8)), mx.zeros((2, 1, 4, 8))
k, v = cache.update_and_fetch(k, v)
self.assertEqual(mask.shape[-2:], (4, k.shape[2]))
self.assertEqual(
mask[0, 0, 0, :].tolist(), [False, True, True, True, False, False, False]
)
# Single query update works
cache = BatchRotatingKVCache(max_size=4, left_padding=[2, 0])
k, v = mx.zeros((2, 1, 4, 8)), mx.zeros((2, 1, 4, 8))
k, v = cache.update_and_fetch(k, v)
mask = cache.make_mask(1)
k, v = mx.zeros((2, 1, 1, 8)), mx.zeros((2, 1, 1, 8))
k, v = cache.update_and_fetch(k, v)
self.assertEqual(mask.shape[-2:], (1, k.shape[2]))
self.assertEqual(mask[0, 0, 0].tolist(), [True, False, True, True])
self.assertEqual(mask[1, 0, 0].tolist(), [True, True, True, True])
# Check filtering
cache = BatchRotatingKVCache(max_size=4, left_padding=[2, 0, 3])
k, v = mx.zeros((3, 1, 3, 8)), mx.zeros((3, 1, 3, 8))
cache.update_and_fetch(k, v)
cache.filter(mx.array([1]))
self.assertEqual(cache.keys.shape, (1, 1, 3, 8))
# Check extend
cache = BatchRotatingKVCache(max_size=4, left_padding=[2, 1])
other = BatchRotatingKVCache(max_size=4, left_padding=[2, 2])
k, v = mx.zeros((2, 1, 5, 8)), mx.zeros((2, 1, 5, 8))
cache.update_and_fetch(k, v)
other.update_and_fetch(k, v)
k, v = mx.zeros((2, 1, 1, 8)), mx.zeros((2, 1, 1, 8))
cache.update_and_fetch(k, v)
cache.extend(other)
# Check mask when going from prompt -> extend -> prompt
cache = BatchRotatingKVCache(max_size=8, left_padding=[4])
k, v = mx.zeros((1, 1, 8, 8)), mx.zeros((1, 1, 8, 8))
cache.update_and_fetch(k, v)
mask = cache.make_mask(1)
self.assertEqual(
mask.squeeze().tolist(), [True, False, False, False, True, True, True, True]
)
k, v = mx.zeros((1, 1, 1, 8)), mx.zeros((1, 1, 1, 8))
cache.update_and_fetch(k, v)
mask = cache.make_mask(2)
expected = mx.array(
[
[False, False, False, True, True, True, True, True, False],
[False, False, False, True, True, True, True, True, True],
]
)
self.assertTrue(mx.array_equal(mask.squeeze(), expected))
def test_save_load_batch_caches(self):
cache_file = os.path.join(self.test_dir, "prompt_cache.safetensors")
cache = [
MambaCache(left_padding=[1, 2]),
BatchKVCache(left_padding=[1, 2]),
BatchRotatingKVCache(max_size=10, left_padding=[1, 2]),
]
for c in cache:
if isinstance(c, MambaCache):
c[0] = mx.random.uniform(shape=(4, 4, 4))
c[1] = mx.random.uniform(shape=(4, 4, 4))
else:
x = mx.random.uniform(shape=(4, 4, 7, 4))
y = mx.random.uniform(shape=(4, 4, 7, 4))
c.update_and_fetch(x, y)
save_prompt_cache(cache_file, cache)
loaded_cache = load_prompt_cache(cache_file)
left_padding = mx.array([1, 2])
for c, lc in zip(cache, loaded_cache):
self.assertTrue(mx.array_equal(c.left_padding, left_padding))
if __name__ == "__main__":
unittest.main()