fix: make write_token dynamic (#29271)

This commit is contained in:
Chang Min Bark
2026-06-30 22:55:57 -07:00
committed by GitHub
parent 546d6e23ad
commit 41779d56fd
3 changed files with 31 additions and 7 deletions
@@ -112,9 +112,12 @@ class ContiguousAttentionKVCache:
def write_token(self, k: mx.array, v: mx.array) -> None:
"""Write one token. k, v shape: (1, n_kv_heads, 1, head_dim)."""
self.keys[:, :, self.offset : self.offset + 1, :] = k
self.values[:, :, self.offset : self.offset + 1, :] = v
self.offset += 1
end = self.offset + 1
if end > self.max_seq_len:
self._grow(end)
self.keys[:, :, self.offset : end, :] = k
self.values[:, :, self.offset : end, :] = v
self.offset = end
def get_kv(self) -> tuple[mx.array, mx.array]:
"""Return valid K/V: (1, n_kv_heads, offset, head_dim)."""
@@ -1198,10 +1198,6 @@ class MlxModelRunner:
"""
caches = prev.caches
# TODO (changminbark): Need to fix
# ContiguousAttentionKVCache.write_token to accommodate dynamic growing
# like ContiguousAttentionKVCache.update_and_fetch.
# After prev's graph ran, each attention KV cache offset was
# bumped by one per layer - attention wrapper's `write_token`
# mutates the Python offset synchronously at graph-build time.