Skip to content
Closed
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
23 changes: 21 additions & 2 deletions src/memos/memories/activation/kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,8 +232,27 @@ def _concat_caches(self, caches: list[DynamicCache]) -> DynamicCache:
keys = [c.layers[layer].keys for c in caches]
vals = [c.layers[layer].values for c in caches]
# single concat per layer
merged.layers[layer].keys = torch.cat(keys, dim=-2)
merged.layers[layer].values = torch.cat(vals, dim=-2)
concat_keys = torch.cat(keys, dim=-2)
concat_vals = torch.cat(vals, dim=-2)
merged_layer = merged.layers[layer]
# From transformers>=4.57, DynamicLayer.get_seq_length() and
# DynamicLayer.update() gate on the `is_initialized` flag,
# which is only flipped inside `lazy_initialization`. Bare
# `layer_cls()` construction plus direct attribute assignment
# leaves the flag False, so the merged cache reports length 0
# and DynamicLayer.update() silently overwrites `.keys` /
# `.values` with empty tensors on the first forward pass
# (see issue #2313). Route through the public
# lazy_initialization path so dtype / device / is_initialized
# are set exactly as upstream expects, then overwrite the
# tensors with the merged content. The hasattr guard leaves
# older layer classes (without lazy_initialization) untouched.
if hasattr(merged_layer, "lazy_initialization") and not getattr(
merged_layer, "is_initialized", False
):
merged_layer.lazy_initialization(concat_keys)
merged_layer.keys = concat_keys
merged_layer.values = concat_vals

# Check for old structure (key_cache)
elif hasattr(caches[0], "key_cache"):
Expand Down
66 changes: 59 additions & 7 deletions tests/memories/activation/test_kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,26 @@ def kv_memory(dummy_config):
yield KVCacheMemory(dummy_config)


def make_filled_cache():
# Create a DynamicCache with at least one dummy tensor layer
def make_filled_cache(seq_len: int = 3):
"""Build a DynamicCache with one populated layer.

Works against both the current transformers layout (``layers`` list of
``DynamicLayer``) and the legacy layout (``key_cache`` /
``value_cache`` lists). Tensors are 4-D
``[batch, num_heads, seq_len, head_dim]`` as expected by
``DynamicLayer.update``.
"""
cache = DynamicCache()
cache.key_cache.append(torch.zeros(1, 2, 3))
cache.value_cache.append(torch.zeros(1, 2, 3))
keys = torch.zeros(1, 2, seq_len, 4)
vals = torch.zeros(1, 2, seq_len, 4)
if hasattr(cache, "layers"):
# transformers >= 4.57: DynamicCache exposes `layers` and
# `.update(...)` routes through DynamicLayer.lazy_initialization,
# so `is_initialized` becomes True.
cache.update(keys, vals, layer_idx=0)
else: # pragma: no cover - legacy path retained for older transformers
cache.key_cache.append(keys)
cache.value_cache.append(vals)
return cache


Expand All @@ -58,9 +73,46 @@ def test_get_cache_merge(kv_memory):
kv_memory.add([item1, item2])
merged = kv_memory.get_cache([item1.id, item2.id])
assert isinstance(merged, DynamicCache)
# Check the number of layers in merged key/value cache
assert len(merged.key_cache) == 1
assert len(merged.value_cache) == 1
# Check that the merged cache exposes at least one populated layer,
# regardless of the transformers layout at runtime.
if hasattr(merged, "layers"):
assert len(merged.layers) == 1
assert merged.layers[0].keys.numel() > 0
assert merged.layers[0].values.numel() > 0
else: # pragma: no cover - legacy transformers layout
assert len(merged.key_cache) == 1
assert len(merged.value_cache) == 1


def test_concat_layer_reports_initialized_and_full_seq_length(kv_memory):
"""Regression for issue #2313.

On transformers >= 4.57, DynamicLayer.get_seq_length() short-circuits to
0 when `is_initialized` is False. `_concat_caches` used to build layers
via bare `layer_cls()` + direct `.keys` / `.values` assignment, which
bypasses `lazy_initialization` and leaves the flag False. The
consequence: merged caches reported length 0, and the first forward
pass silently discarded every merged token. This test asserts the
invariants of that fix and fails on the pre-fix implementation.
"""
seq_len = 5
item1 = KVCacheItem(memory=make_filled_cache(seq_len=seq_len))
item2 = KVCacheItem(memory=make_filled_cache(seq_len=seq_len))
kv_memory.add([item1, item2])

merged = kv_memory.get_cache([item1.id, item2.id])
assert merged is not None

# This path is only exercised on transformers >= 4.57. The bug lives
# here — get_seq_length() must reflect the true merged length.
if hasattr(merged, "layers"):
assert len(merged.layers) == 1, f"expected 1 layer, got {len(merged.layers)}"
assert merged.layers[0].keys.shape[-2] == 2 * seq_len
assert merged.layers[0].values.shape[-2] == 2 * seq_len
# `is_initialized` is set inside DynamicLayer.lazy_initialization —
# bypassing that path is what triggered the silent data loss.
assert getattr(merged.layers[0], "is_initialized", True) is True
assert merged.get_seq_length() == 2 * seq_len


def test_delete_and_get_all(kv_memory):
Expand Down
Loading