diff --git a/src/memos/llms/hf.py b/src/memos/llms/hf.py index 0dd841c1a..99c7a394d 100644 --- a/src/memos/llms/hf.py +++ b/src/memos/llms/hf.py @@ -82,7 +82,12 @@ def generate( if past_key_values is None: return self._generate_full(prompt, **kwargs) else: - return self._generate_with_cache(prompt, past_key_values, **kwargs) + from memos.memories.activation.kv import clone_dynamic_cache + + # The model appends new K/V tensors to the cache it receives, so + # hand it a clone and keep the caller's cache (e.g. a stored + # activation memory) unchanged by this call. + return self._generate_with_cache(prompt, clone_dynamic_cache(past_key_values), **kwargs) def generate_stream( self, messages: MessageList, past_key_values: DynamicCache | None = None, **kwargs @@ -102,7 +107,11 @@ def generate_stream( if past_key_values is None: yield from self._generate_full_stream(prompt) else: - yield from self._generate_with_cache_stream(prompt, past_key_values) + from memos.memories.activation.kv import clone_dynamic_cache + + yield from self._generate_with_cache_stream( + prompt, clone_dynamic_cache(past_key_values) + ) def _generate_full(self, prompt: str, **kwargs) -> str: """ diff --git a/src/memos/memories/activation/kv.py b/src/memos/memories/activation/kv.py index 1981b958f..73f98901e 100644 --- a/src/memos/memories/activation/kv.py +++ b/src/memos/memories/activation/kv.py @@ -1,3 +1,4 @@ +import copy import os import pickle @@ -206,7 +207,10 @@ def _concat_caches(self, caches: list[DynamicCache]) -> DynamicCache: assert caches, "Need at least one cache" if len(caches) == 1: - return caches[0] + # Return a copy: the stored cache must never be handed out by + # reference, because generation appends new K/V tensors to the + # cache object it receives and would grow the store every turn. + return clone_dynamic_cache(caches[0]) merged = DynamicCache() @@ -248,13 +252,89 @@ def _concat_caches(self, caches: list[DynamicCache]) -> DynamicCache: merged.value_cache.append(torch.cat(vals, dim=-2)) else: - raise AttributeError( - "DynamicCache object has neither 'layers' nor 'key_cache' attributes" - ) + raise TypeError("DynamicCache object has neither 'layers' nor 'key_cache' attributes") return merged +def clone_dynamic_cache(cache: DynamicCache) -> DynamicCache: + """ + Return an independent copy of a DynamicCache with cloned K/V tensors. + + Generation mutates the cache object it receives in place, so a stored cache + must never be handed to a model by reference — hand out a clone instead. + Compatible with both old (key_cache/value_cache) and new (layers) structures. + """ + import torch + + cloned = DynamicCache() + + if hasattr(cache, "layers"): + cloned.layers = [] + for layer in cache.layers: + # Avoid invoking a layer constructor: modern transformers layers + # such as DynamicSlidingWindowLayer require constructor metadata. + new_layer = copy.copy(layer) + layer_attrs = vars(layer) + # Preserve layer state and clone every tensor, including K/V tensors. + kv_attrs = {"keys", "values", "key_cache", "value_cache"} + for attr, value in layer_attrs.items(): + if attr in kv_attrs: + continue + setattr( + new_layer, + attr, + value.clone() if isinstance(value, torch.Tensor) else copy.deepcopy(value), + ) + # transformers>=4.56 layers expose keys/values, but some versions + # instead carry per-layer key_cache/value_cache (see + # move_dynamic_cache_htod); a clone that skips one shape would + # silently return a content-empty layer. + # Select one naming scheme, matching move_dynamic_cache_htod's + # precedence, while retaining independent guards for asymmetric + # test doubles and cache layers. + has_per_layer_cache = any( + getattr(layer, name, None) is not None for name in ("key_cache", "value_cache") + ) + if has_per_layer_cache: + if "keys" in layer_attrs: + new_layer.keys = None + if "values" in layer_attrs: + new_layer.values = None + if getattr(layer, "key_cache", None) is not None: + new_layer.key_cache = layer.key_cache.clone() + if getattr(layer, "value_cache", None) is not None: + new_layer.value_cache = layer.value_cache.clone() + else: + if "key_cache" in layer_attrs: + new_layer.key_cache = None + if "value_cache" in layer_attrs: + new_layer.value_cache = None + if getattr(layer, "keys", None) is not None: + new_layer.keys = layer.keys.clone() + if getattr(layer, "values", None) is not None: + new_layer.values = layer.values.clone() + cloned.layers.append(new_layer) + elif hasattr(cache, "key_cache"): + # Legacy DynamicCache keeps generation state such as _seen_tokens on + # the cache itself. Keep that state independent of the stored cache; + # key/value lists are populated from cloned tensors below. + for attr, value in vars(cache).items(): + if attr not in {"key_cache", "value_cache"}: + setattr( + cloned, + attr, + value.clone() if isinstance(value, torch.Tensor) else copy.deepcopy(value), + ) + for keys, values in zip(cache.key_cache, cache.value_cache, strict=True): + cloned.key_cache.append(keys.clone() if keys is not None else None) + cloned.value_cache.append(values.clone() if values is not None else None) + else: + raise TypeError("DynamicCache object has neither 'layers' nor 'key_cache' attributes") + + return cloned + + def move_dynamic_cache_htod(dynamic_cache: DynamicCache, device: str) -> DynamicCache: """ Move DynamicCache from CPU to GPU device. diff --git a/tests/cache_helpers.py b/tests/cache_helpers.py new file mode 100644 index 000000000..69e07f620 --- /dev/null +++ b/tests/cache_helpers.py @@ -0,0 +1,63 @@ +import pytest +import torch + +from transformers import DynamicCache + + +def make_filled_cache(): + cache = DynamicCache() + keys = torch.zeros(1, 2, 3, 4) if hasattr(cache, "layers") else torch.zeros(1, 2, 3) + values = torch.zeros_like(keys) + cache.update(keys, values, layer_idx=0) + return cache + + +def cache_keys(cache, layer_idx=0): + if hasattr(cache, "layers"): + return cache.layers[layer_idx].keys + return cache.key_cache[layer_idx] + + +def cache_values(cache, layer_idx=0): + if hasattr(cache, "layers"): + return cache.layers[layer_idx].values + return cache.value_cache[layer_idx] + + +def set_cache_keys(cache, value, layer_idx=0): + if hasattr(cache, "layers"): + cache.layers[layer_idx].keys = value + else: + cache.key_cache[layer_idx] = value + + +def cache_layer_count(cache): + if hasattr(cache, "layers"): + return len(cache.layers) + return len(cache.key_cache) + + +def make_real_hybrid_cache(populate=True): + if not hasattr(DynamicCache(), "layers"): + pytest.skip("requires transformers >=4.56") + + class HybridConfig: + num_hidden_layers = 2 + sliding_window = 4 + + def __init__(self): + self.layer_types = ["full_attention", "sliding_attention"] + + def get_text_config(self): + return self + + try: + cache = DynamicCache(config=HybridConfig()) + if populate: + keys = torch.zeros(1, 2, 3, 4) + values = torch.zeros(1, 2, 3, 4) + cache.update(keys, values, layer_idx=0) + cache.update(keys, values, layer_idx=1) + except TypeError: + pytest.skip("DynamicCache(config=...) is not supported") + return cache diff --git a/tests/llms/test_hf.py b/tests/llms/test_hf.py index 375bf2247..bb5390d7e 100644 --- a/tests/llms/test_hf.py +++ b/tests/llms/test_hf.py @@ -9,6 +9,9 @@ from memos.configs.llm import HFLLMConfig, LLMConfigFactory from memos.llms.factory import LLMFactory from memos.llms.hf import HFLLM +from tests.cache_helpers import cache_keys as _cache_keys +from tests.cache_helpers import cache_values as _cache_values +from tests.cache_helpers import make_filled_cache as _make_filled_cache @patch("transformers.AutoModelForCausalLM", MagicMock()) @@ -182,3 +185,56 @@ def test_kv_cache_generation_with_sampling(self): kv_cache = DynamicCache() resp = llm.generate([{"role": "user", "content": "Sampling"}], past_key_values=kv_cache) self.assertEqual(resp, self.standard_response) + + def test_generate_with_cache_does_not_mutate_caller_cache(self): + """Regression for issue #2301: generation must not append K/V tensors + into the caller's stored cache (activation memory grew every turn).""" + config = HFLLMConfig( + model_name_or_path="qwen3:0.6b", + temperature=0.7, + max_tokens=3, + do_sample=True, + add_generation_prompt=True, + ) + llm = self._create_llm(config) + + kv_cache = _make_filled_cache() + original_key_shape = _cache_keys(kv_cache).shape + original_value_shape = _cache_values(kv_cache).shape + captured = {} + + def forward(*args, **kwargs): + # transformers appends the new tokens' K/V to the cache in place. + # _prefill always passes the cache by keyword; .get keeps the mock + # resilient to an explicit-None caller without inventing a + # positional call shape. + kv = kwargs.get("past_key_values") + self.assertIsNotNone(kv, "forward() called without past_key_values") + captured["kv"] = kv + if hasattr(kv, "layers"): + kv.layers[0].keys = torch.cat([kv.layers[0].keys, torch.ones(1, 2, 1, 4)], dim=-2) + kv.layers[0].values = torch.cat( + [kv.layers[0].values, torch.ones(1, 2, 1, 4)], dim=-2 + ) + else: + kv.key_cache[0] = torch.cat([kv.key_cache[0], torch.ones(1, 1, 3)], dim=-2) + kv.value_cache[0] = torch.cat([kv.value_cache[0], torch.ones(1, 1, 3)], dim=-2) + out = MagicMock() + # Deterministic non-EOS argmax so the loop runs all max_tokens turns + # instead of sometimes sampling eos_token_id (2) on the first step. + logits = torch.full((1, 1, 100), -1e9) + logits[0, 0, 10] = 0.0 + out.logits = logits + out.past_key_values = kv + return out + + self.mock_model.side_effect = forward + try: + llm.generate([{"role": "user", "content": "Hi"}], past_key_values=kv_cache) + finally: + self.mock_model.side_effect = None + + self.assertEqual(_cache_keys(kv_cache).shape, original_key_shape) + self.assertEqual(_cache_values(kv_cache).shape, original_value_shape) + self.assertIsNotNone(captured.get("kv")) + self.assertIsNot(captured["kv"], kv_cache) diff --git a/tests/memories/activation/test_kv.py b/tests/memories/activation/test_kv.py index 6490d687f..6244c08fd 100644 --- a/tests/memories/activation/test_kv.py +++ b/tests/memories/activation/test_kv.py @@ -6,8 +6,18 @@ from transformers import DynamicCache from memos.configs.memory import KVCacheMemoryConfig +from memos.memories.activation import kv as kv_module from memos.memories.activation.item import KVCacheItem -from memos.memories.activation.kv import KVCacheMemory +from memos.memories.activation.kv import KVCacheMemory, clone_dynamic_cache +from tests import cache_helpers +from tests.cache_helpers import ( + cache_keys, + cache_layer_count, + cache_values, + make_filled_cache, + make_real_hybrid_cache, + set_cache_keys, +) @pytest.fixture @@ -33,14 +43,6 @@ def kv_memory(dummy_config): yield KVCacheMemory(dummy_config) -def make_filled_cache(): - # Create a DynamicCache with at least one dummy tensor layer - cache = DynamicCache() - cache.key_cache.append(torch.zeros(1, 2, 3)) - cache.value_cache.append(torch.zeros(1, 2, 3)) - return cache - - def test_extract_and_add_and_get(kv_memory): # Test extract, add, and get functionality item = kv_memory.extract("hello world") @@ -59,8 +61,22 @@ def test_get_cache_merge(kv_memory): 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 + assert cache_layer_count(merged) == 1 + assert cache_values(merged) is not None + + +def test_make_real_hybrid_cache_skips_update_typeerror(monkeypatch): + class IncompatibleCache: + def __init__(self, *args, **kwargs): + self.layers = [] + + def update(self, *args, **kwargs): + raise TypeError("hybrid update signature is unsupported") + + monkeypatch.setattr(cache_helpers, "DynamicCache", IncompatibleCache) + + with pytest.raises(pytest.skip.Exception, match=r"DynamicCache\(config=\.\.\.\)"): + cache_helpers.make_real_hybrid_cache() def test_delete_and_get_all(kv_memory): @@ -84,3 +100,372 @@ class DummyTextualMemory: item = kv_memory.from_textual_memory(DummyTextualMemory()) assert isinstance(item, KVCacheItem) assert item.metadata["bar"] == 1 + + +def test_get_cache_single_item_returns_independent_copy(kv_memory): + # Regression for issue #2301: with a single cache, get_cache used to hand + # out the stored object, so generation appended new K/V tensors into the + # store and the activation memory grew every turn. + item = KVCacheItem(memory=make_filled_cache()) + kv_memory.add([item]) + + merged = kv_memory.get_cache([item.id]) + assert merged is not item.memory + original_shape = cache_keys(item.memory).shape + + # In-place mutation must not leak either: verify storage independence + # before replacing the list slot with generation's appended tensor. + merged_keys = cache_keys(merged) + merged_keys.fill_(99.0) + assert not torch.all(cache_keys(item.memory) == 99.0), "get_cache shares storage with store" + merged_keys.zero_() + + # Simulate generation appending to the handed-out cache. + appended = torch.ones((*merged_keys.shape[:-2], 1, merged_keys.shape[-1])) + set_cache_keys(merged, torch.cat([merged_keys, appended], dim=-2)) + assert cache_keys(item.memory).shape == original_shape + + +def test_get_cache_multi_item_merge_does_not_alias_inputs(kv_memory): + item1 = KVCacheItem(memory=make_filled_cache()) + item2 = KVCacheItem(memory=make_filled_cache()) + kv_memory.add([item1, item2]) + + merged = kv_memory.get_cache([item1.id, item2.id]) + assert merged is not item1.memory + assert merged is not item2.memory + + +def test_clone_dynamic_cache_copies_legacy_tensors(): + cache = make_filled_cache() + original_shape = cache_keys(cache).shape + cloned = clone_dynamic_cache(cache) + + assert cloned is not cache + assert cache_keys(cloned) is not cache_keys(cache) + assert torch.equal(cache_keys(cloned), cache_keys(cache)) + + # In-place mutation must not leak either: verify storage independence + # before replacing the list slot. + cloned_keys = cache_keys(cloned) + cloned_keys.fill_(99.0) + assert not torch.all(cache_keys(cache) == 99.0), "clone shares storage with original" + cloned_keys.zero_() + + replacement = torch.ones((*cloned_keys.shape[:-2], 5, cloned_keys.shape[-1])) + set_cache_keys(cloned, replacement) + assert cache_keys(cache).shape == original_shape + + +def test_clone_dynamic_cache_preserves_legacy_cache_state(): + cache = make_filled_cache() + if not hasattr(cache, "_seen_tokens"): + pytest.skip("_seen_tokens is not present in this transformers version") + cache._seen_tokens = 2 + + cloned = clone_dynamic_cache(cache) + + assert cloned._seen_tokens == 2 + cloned.update(torch.ones(1, 1, 3), torch.ones(1, 1, 3), layer_idx=0) + assert cloned._seen_tokens == 3 + assert cache._seen_tokens == 2 + + +@pytest.mark.skipif( + hasattr(DynamicCache(), "layers"), reason="requires the legacy DynamicCache API" +) +def test_clone_dynamic_cache_copies_legacy_tensor_state(): + cache = make_filled_cache() + cache._cos_cached = torch.arange(3) + + cloned = clone_dynamic_cache(cache) + + assert torch.equal(cloned._cos_cached, cache._cos_cached) + assert cloned._cos_cached is not cache._cos_cached + cloned._cos_cached[0] = 99 + assert cache._cos_cached[0] == 0 + + +def test_clone_dynamic_cache_rejects_mismatched_legacy_layers(): + class LegacyCache: + def __init__(self): + self.key_cache = [torch.zeros(1, 2, 3)] + self.value_cache = [] + + cache = LegacyCache() + + with pytest.raises(ValueError): + clone_dynamic_cache(cache) + + +def test_clone_dynamic_cache_handles_layers_structure(): + # transformers >= 4.56 exposes DynamicCache.layers with per-layer keys/values. + class FakeLayer: + def __init__(self): + self.keys = None + self.values = None + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + cache.layers = [FakeLayer()] + cache.layers[0].keys = torch.zeros(1, 2, 3) + cache.layers[0].values = torch.zeros(1, 2, 4) + + cloned = clone_dynamic_cache(cache) + assert isinstance(cloned, DynamicCache) + assert len(cloned.layers) == 1 + assert cloned.layers[0].keys is not cache.layers[0].keys + assert torch.equal(cloned.layers[0].keys, cache.layers[0].keys) + + # In-place mutation must not leak either: verify storage independence + # before replacing the layer attribute. + cloned.layers[0].keys.fill_(99.0) + assert not torch.all(cache.layers[0].keys == 99.0), "clone shares tensor storage with original" + cloned.layers[0].keys.zero_() + + cloned.layers[0].keys = torch.ones(2, 2, 3) + assert cache.layers[0].keys.shape == (1, 2, 3) + + +def test_clone_dynamic_cache_preserves_real_hybrid_layers(): + cache = make_real_hybrid_cache() + + cloned = clone_dynamic_cache(cache) + + assert [type(layer) for layer in cloned.layers] == [type(layer) for layer in cache.layers] + assert cloned.layers[1].sliding_window == 4 + assert cloned.layers[1].cumulative_length == cache.layers[1].cumulative_length + assert torch.equal(cloned.layers[0].keys, cache.layers[0].keys) + assert torch.equal(cloned.layers[1].values, cache.layers[1].values) + assert cloned.layers[1].keys is not cache.layers[1].keys + + cloned.layers[1].update(torch.ones(1, 2, 1, 4), torch.ones(1, 2, 1, 4)) + assert cloned.layers[1].cumulative_length == 4 + assert cloned.layers[1].keys.shape[-2] == 3 + assert cache.layers[1].keys.shape[-2] == 3 + assert cache.layers[1].cumulative_length == 3 + + +def test_clone_dynamic_cache_preserves_uninitialized_real_hybrid_layers(): + cache = make_real_hybrid_cache(populate=False) + + cloned = clone_dynamic_cache(cache) + + assert [type(layer) for layer in cloned.layers] == [type(layer) for layer in cache.layers] + assert cloned.layers[0].keys is None + assert cloned.layers[0].values is None + assert cloned.layers[1].keys is None + assert cloned.layers[1].values is None + assert cloned.layers[1].sliding_window == cache.layers[1].sliding_window + + +def test_clone_dynamic_cache_layers_guard_keys_and_values_independently(): + # A layer may legitimately have only one side populated; the clone must + # not crash on the missing side nor fabricate a value for it. + class FakeLayer: + def __init__(self): + self.keys = None + self.values = None + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + keys_only = FakeLayer() + keys_only.keys = torch.zeros(1, 2, 3) + values_only = FakeLayer() + values_only.values = torch.zeros(1, 2, 4) + cache.layers = [keys_only, values_only] + + cloned = clone_dynamic_cache(cache) + assert torch.equal(cloned.layers[0].keys, keys_only.keys) + assert cloned.layers[0].values is None + assert cloned.layers[1].keys is None + assert torch.equal(cloned.layers[1].values, values_only.values) + + +def test_clone_dynamic_cache_handles_per_layer_key_value_cache(): + # Some transformers versions carry per-layer key_cache/value_cache + # instead of keys/values (mirrors move_dynamic_cache_htod); the clone + # must copy those tensors too instead of returning an empty layer. + class FakeLayer: + pass + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + layer = FakeLayer() + layer.key_cache = torch.zeros(1, 2, 3) + layer.value_cache = torch.zeros(1, 2, 3) + cache.layers = [layer] + + cloned = clone_dynamic_cache(cache) + assert torch.equal(cloned.layers[0].key_cache, layer.key_cache) + assert torch.equal(cloned.layers[0].value_cache, layer.value_cache) + + cloned.layers[0].key_cache.fill_(99.0) + assert not torch.all(layer.key_cache == 99.0), "clone shares tensor storage with original" + cloned.layers[0].value_cache.fill_(99.0) + assert not torch.all(layer.value_cache == 99.0), ( + "clone shares value_cache tensor storage with original" + ) + + +def test_clone_dynamic_cache_replaces_preexisting_destination_layers(monkeypatch): + class FakeLayer: + pass + + class FakeLayeredCache: + pass + + class DestinationCache: + def __init__(self): + self.layers = [object()] + + source = FakeLayeredCache() + layer = FakeLayer() + layer.keys = torch.zeros(1, 2, 3) + layer.values = torch.zeros(1, 2, 3) + source.layers = [layer] + monkeypatch.setattr(kv_module, "DynamicCache", DestinationCache) + + cloned = clone_dynamic_cache(source) + + assert len(cloned.layers) == 1 + assert cloned.layers[0].keys is not layer.keys + assert cloned.layers[0].values is not layer.values + + +def test_clone_dynamic_cache_clones_layer_kv_tensors_once(): + class CloneCountingTensor(torch.Tensor): + clone_count = 0 + + def clone(self, *args, **kwargs): + type(self).clone_count += 1 + return super().clone(*args, **kwargs) + + class FakeLayer: + pass + + class FakeLayeredCache: + pass + + def counting_tensor(value): + return torch.full((1, 2, 3), value).as_subclass(CloneCountingTensor) + + cache = FakeLayeredCache() + layer = FakeLayer() + layer.keys = counting_tensor(0) + layer.values = counting_tensor(0) + layer.key_cache = counting_tensor(1) + layer.value_cache = counting_tensor(1) + cache.layers = [layer] + + CloneCountingTensor.clone_count = 0 + cloned = clone_dynamic_cache(cache) + + assert CloneCountingTensor.clone_count == 2 + assert cloned.layers[0].key_cache is not layer.key_cache + assert cloned.layers[0].value_cache is not layer.value_cache + + +def test_clone_dynamic_cache_preserves_layer_state(): + # DynamicLayer.update() uses these flags to decide whether to append to or + # replace the existing history on its first update. + class StatefulLayer: + def __init__(self): + self.is_initialized = False + self._seen_tokens = 0 + self.keys = None + self.values = None + + def update(self, keys, values): + if self.is_initialized: + self.keys = torch.cat([self.keys, keys], dim=-2) + self.values = torch.cat([self.values, values], dim=-2) + else: + self.keys = keys + self.values = values + self.is_initialized = True + self._seen_tokens += keys.shape[-2] + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + layer = StatefulLayer() + layer.keys = torch.zeros(1, 2, 3) + layer.values = torch.zeros(1, 2, 3) + layer.is_initialized = True + layer._seen_tokens = 2 + cache.layers = [layer] + + cloned = clone_dynamic_cache(cache) + + assert cloned.layers[0].is_initialized is True + assert cloned.layers[0]._seen_tokens == 2 + cloned.layers[0].update(torch.ones(1, 1, 3), torch.ones(1, 1, 3)) + assert cloned.layers[0].keys.shape == (1, 3, 3) + assert cloned.layers[0].values.shape == (1, 3, 3) + + +def test_clone_dynamic_cache_copies_mutable_layer_state(): + class FakeLayer: + def __init__(self): + self.keys = None + self.values = None + self.metadata = {"history": ["original"]} + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + cache.layers = [FakeLayer()] + + cloned = clone_dynamic_cache(cache) + cloned.layers[0].metadata["history"].append("clone") + + assert cloned.layers[0].metadata == {"history": ["original", "clone"]} + assert cache.layers[0].metadata == {"history": ["original"]} + assert cloned.layers[0].metadata is not cache.layers[0].metadata + assert cloned.layers[0].metadata["history"] is not cache.layers[0].metadata["history"] + + +def test_clone_dynamic_cache_rejects_unknown_shape(): + class UnknownCache: + pass + + with pytest.raises(TypeError, match="neither 'layers' nor 'key_cache'"): + clone_dynamic_cache(UnknownCache()) + + +def test_clone_dynamic_cache_prefers_per_layer_cache_attributes(): + # A layer exposing both naming schemes must follow the same precedence as + # move_dynamic_cache_htod: key_cache/value_cache take priority over keys/values. + class FakeLayer: + def __init__(self): + self.keys = None + self.values = None + self.key_cache = None + self.value_cache = None + + class FakeLayeredCache: + pass + + cache = FakeLayeredCache() + layer = FakeLayer() + layer.keys = torch.zeros(1, 2, 3) + layer.values = torch.zeros(1, 2, 3) + layer.key_cache = torch.ones(1, 2, 3) + layer.value_cache = torch.ones(1, 2, 3) + cache.layers = [layer] + + cloned = clone_dynamic_cache(cache) + + assert torch.equal(cloned.layers[0].key_cache, layer.key_cache) + assert torch.equal(cloned.layers[0].value_cache, layer.value_cache) + assert cloned.layers[0].keys is None + assert cloned.layers[0].values is None