From efb99506755eeb5dbd328e569b58d6aa3adb6bb4 Mon Sep 17 00:00:00 2001 From: Aidan Foster Date: Mon, 14 Sep 2026 16:35:46 -0700 Subject: [PATCH 1/3] feat: expose multi-view rollout layout options Signed-off-by: Aidan Foster --- .../flashdreams/core/attention/kvcache.py | 63 +++++++++++++++++-- .../core/attention/multiview/rollout.py | 10 ++- flashdreams/tests/test_multiview_cache.py | 30 +++++++++ flashdreams/tests/test_multiview_rollout.py | 8 +-- 4 files changed, 100 insertions(+), 11 deletions(-) diff --git a/flashdreams/flashdreams/core/attention/kvcache.py b/flashdreams/flashdreams/core/attention/kvcache.py index 9b4c6ce5c..9ddc901ed 100644 --- a/flashdreams/flashdreams/core/attention/kvcache.py +++ b/flashdreams/flashdreams/core/attention/kvcache.py @@ -87,6 +87,7 @@ def __init__( capacity: int, regions: Sequence[SlotRegion], seq_dim: int = 2, + storage: Sequence[LayerKV] | None = None, ) -> None: """Allocate storage and copy the prefilled key/value tensors. @@ -95,6 +96,10 @@ def __init__( capacity: Total tokens reserved along ``seq_dim``. regions: Independently rotating physical slot regions. seq_dim: Sequence dimension in every key/value tensor. + storage: Optional caller-owned K/V buffers whose sequence dimension + is exactly ``capacity``. Writes retain these buffers and their + storage addresses, which lets accelerated adapters embed the + cache in a larger attention arena. Raises: ValueError: Inputs or slot regions cannot describe a valid cache. @@ -129,14 +134,32 @@ def __init__( self._k: list[Tensor] = [] self._v: list[Tensor] = [] + if storage is not None and len(storage) != len(prefilled): + raise ValueError( + f"{len(storage)} storage layers for {len(prefilled)} prefill layers." + ) + for layer, (key, value) in enumerate(prefilled): self._validate_prefill_layer(key, value, layer) - key_shape = list(key.shape) - value_shape = list(value.shape) - key_shape[self._seq_dim] = capacity - value_shape[self._seq_dim] = capacity - key_buffer = key.new_zeros(key_shape) - value_buffer = value.new_zeros(value_shape) + if storage is None: + key_shape = list(key.shape) + value_shape = list(value.shape) + key_shape[self._seq_dim] = capacity + value_shape[self._seq_dim] = capacity + key_buffer = key.new_zeros(key_shape) + value_buffer = value.new_zeros(value_shape) + else: + key_buffer, value_buffer = storage[layer] + self._validate_storage_layer( + key, key_buffer, capacity=capacity, layer=layer, kind="key" + ) + self._validate_storage_layer( + value, + value_buffer, + capacity=capacity, + layer=layer, + kind="value", + ) prefix = self._seq_slice(0, self._initial_length, tensor_dim) with torch.no_grad(): key_buffer[prefix] = key @@ -144,6 +167,34 @@ def __init__( self._k.append(key_buffer) self._v.append(value_buffer) + def _validate_storage_layer( + self, + source: Tensor, + storage: Tensor, + *, + capacity: int, + layer: int, + kind: str, + ) -> None: + """Validate caller-owned storage for one prefilled tensor.""" + if storage.ndim != source.ndim: + raise ValueError( + f"storage layer {layer} {kind} rank does not match its prefill." + ) + for dim, (source_size, storage_size) in enumerate( + zip(source.shape, storage.shape, strict=True) + ): + expected = capacity if dim == self._seq_dim else source_size + if storage_size != expected: + raise ValueError( + f"storage layer {layer} {kind} dimension {dim} is " + f"{storage_size}; expected {expected}." + ) + if storage.dtype is not source.dtype or storage.device != source.device: + raise ValueError( + f"storage layer {layer} {kind} must match prefill dtype and device." + ) + @property def capacity(self) -> int: """Return the allocated token capacity.""" diff --git a/flashdreams/flashdreams/core/attention/multiview/rollout.py b/flashdreams/flashdreams/core/attention/multiview/rollout.py index 85bd74af4..184df1a7d 100644 --- a/flashdreams/flashdreams/core/attention/multiview/rollout.py +++ b/flashdreams/flashdreams/core/attention/multiview/rollout.py @@ -188,6 +188,7 @@ def __init__( history_frames: int | None = None, token_frames: int | None = None, use_block_mask: bool = False, + mask_block_size: int | tuple[int, int] = 128, ) -> None: """Validate rollout geometry, allocate output, and prefill model state.""" plan = ar_chunk_plan(geometry) @@ -252,6 +253,7 @@ def __init__( self._sample_type = sample_type self._history_slots = history_slots self._use_block_mask = use_block_mask + self._mask_block_size = mask_block_size self._offset = vision_temporal_offset(self._state.num_text_tokens) minimum_token_frames = geometry.condition_frames + geometry.frames_per_chunk @@ -458,7 +460,11 @@ def _mask( pass_kind=pass_kind, device=self._model.device, ) - return metadata.block_mask() if self._use_block_mask else metadata.mask() + return ( + metadata.block_mask(block_size=self._mask_block_size) + if self._use_block_mask + else metadata.mask() + ) def latent_for(self, start: int, end: int) -> Tensor: """Return decoder-ready latents for an available frame range.""" @@ -513,6 +519,7 @@ def run_rollout( sample_type: Literal["sde", "ode"] = "sde", history_frames: int | None = None, use_block_mask: bool = False, + mask_block_size: int | tuple[int, int] = 128, on_chunk: Callable[[ChunkTrace, Tensor], None] | None = None, ) -> Rollout: """Build and drain a model-independent multi-view rollout.""" @@ -528,6 +535,7 @@ def run_rollout( sample_type=sample_type, history_frames=history_frames, use_block_mask=use_block_mask, + mask_block_size=mask_block_size, ) while not rollout.is_finished: trace, chunk = rollout.step() diff --git a/flashdreams/tests/test_multiview_cache.py b/flashdreams/tests/test_multiview_cache.py index 049014a04..499e4d97d 100644 --- a/flashdreams/tests/test_multiview_cache.py +++ b/flashdreams/tests/test_multiview_cache.py @@ -55,6 +55,36 @@ def cache() -> FixedSlotKVCache: ) +def test_cache_can_use_caller_owned_storage() -> None: + """Embed fixed-slot storage in an allocation owned by an adapter.""" + prefilled = [layer(5, 1.0), layer(5, 2.0)] + storage = [layer(11, 0.0), layer(11, 0.0)] + memory = FixedSlotKVCache( + prefilled, + capacity=11, + regions=[ + SlotRegion( + name="history", + start=5, + slots=2, + slot_tokens=3, + extends_length=True, + ) + ], + storage=storage, + ) + + assert [ + (key.data_ptr(), value.data_ptr()) + for key, value in zip(memory._k, memory._v, strict=True) + ] == [(key.data_ptr(), value.data_ptr()) for key, value in storage] + for (key, value), (prefill_key, prefill_value) in zip( + memory.layers(), prefilled, strict=True + ): + assert torch.equal(key, prefill_key) + assert torch.equal(value, prefill_value) + + def chunks(tokens: int, value: float) -> list[LayerKV]: """Build a two-layer cache write.""" return [layer(tokens, value), layer(tokens, value + 1.0)] diff --git a/flashdreams/tests/test_multiview_rollout.py b/flashdreams/tests/test_multiview_rollout.py index cfc8d809f..1356f10bd 100644 --- a/flashdreams/tests/test_multiview_rollout.py +++ b/flashdreams/tests/test_multiview_rollout.py @@ -238,22 +238,22 @@ def tracked_chunk_mrope_ids( def test_block_mask_selection_crosses_the_protocol_boundary() -> None: - """Supply the same visibility contract in block-sparse form.""" + """Supply the same visibility contract with the requested block geometry.""" model = _Model() controls, text_ids, condition = _inputs() - rollout = ChunkRollout( + run_rollout( model, geometry=GEOMETRY, controls=controls, text_ids=text_ids, condition_tokens=condition, use_block_mask=True, + mask_block_size=(16, 32), ) - rollout.step() - assert model.state.masks assert all(isinstance(mask, BlockMask) for mask in model.state.masks) + assert all(mask.BLOCK_SIZE == (16, 32) for mask in model.state.masks) def test_seeded_rollouts_are_reproducible_through_an_adapter() -> None: From 0b3d84d239e46b21b6d5d654e7665a535495ec09 Mon Sep 17 00:00:00 2001 From: Aidan Foster Date: Tue, 15 Sep 2026 10:38:17 -0700 Subject: [PATCH 2/3] fix: initialize caller-owned cache storage atomically Signed-off-by: Aidan Foster --- .../flashdreams/core/attention/kvcache.py | 27 ++++++++++++------- flashdreams/tests/test_multiview_cache.py | 25 ++++++++++++++++- flashdreams/tests/test_multiview_rollout.py | 5 ++-- 3 files changed, 45 insertions(+), 12 deletions(-) diff --git a/flashdreams/flashdreams/core/attention/kvcache.py b/flashdreams/flashdreams/core/attention/kvcache.py index 9ddc901ed..152026642 100644 --- a/flashdreams/flashdreams/core/attention/kvcache.py +++ b/flashdreams/flashdreams/core/attention/kvcache.py @@ -141,15 +141,10 @@ def __init__( for layer, (key, value) in enumerate(prefilled): self._validate_prefill_layer(key, value, layer) - if storage is None: - key_shape = list(key.shape) - value_shape = list(value.shape) - key_shape[self._seq_dim] = capacity - value_shape[self._seq_dim] = capacity - key_buffer = key.new_zeros(key_shape) - value_buffer = value.new_zeros(value_shape) - else: - key_buffer, value_buffer = storage[layer] + if storage is not None: + for layer, ((key, value), (key_buffer, value_buffer)) in enumerate( + zip(prefilled, storage, strict=True) + ): self._validate_storage_layer( key, key_buffer, capacity=capacity, layer=layer, kind="key" ) @@ -160,8 +155,22 @@ def __init__( layer=layer, kind="value", ) + + for layer, (key, value) in enumerate(prefilled): + if storage is None: + key_shape = list(key.shape) + value_shape = list(value.shape) + key_shape[self._seq_dim] = capacity + value_shape[self._seq_dim] = capacity + key_buffer = key.new_zeros(key_shape) + value_buffer = value.new_zeros(value_shape) + else: + key_buffer, value_buffer = storage[layer] prefix = self._seq_slice(0, self._initial_length, tensor_dim) + suffix = self._seq_slice(self._initial_length, capacity, tensor_dim) with torch.no_grad(): + key_buffer[suffix].zero_() + value_buffer[suffix].zero_() key_buffer[prefix] = key value_buffer[prefix] = value self._k.append(key_buffer) diff --git a/flashdreams/tests/test_multiview_cache.py b/flashdreams/tests/test_multiview_cache.py index 499e4d97d..4c9b58f37 100644 --- a/flashdreams/tests/test_multiview_cache.py +++ b/flashdreams/tests/test_multiview_cache.py @@ -58,7 +58,7 @@ def cache() -> FixedSlotKVCache: def test_cache_can_use_caller_owned_storage() -> None: """Embed fixed-slot storage in an allocation owned by an adapter.""" prefilled = [layer(5, 1.0), layer(5, 2.0)] - storage = [layer(11, 0.0), layer(11, 0.0)] + storage = [layer(11, -9.0), layer(11, -9.0)] memory = FixedSlotKVCache( prefilled, capacity=11, @@ -83,6 +83,29 @@ def test_cache_can_use_caller_owned_storage() -> None: ): assert torch.equal(key, prefill_key) assert torch.equal(value, prefill_value) + for key, value in storage: + assert torch.count_nonzero(key[:, :, 5:]) == 0 + assert torch.count_nonzero(value[:, :, 5:]) == 0 + + +def test_invalid_caller_storage_does_not_mutate_any_layer() -> None: + """Validate every supplied buffer before clearing or copying any of them.""" + prefilled = [layer(5, 1.0), layer(5, 2.0)] + storage = [layer(11, -9.0), layer(11, -8.0)] + storage[1] = (torch.zeros(2, 2, 11, 3), storage[1][1]) + before = [(key.clone(), value.clone()) for key, value in storage] + + with pytest.raises(ValueError, match="storage layer 1 key dimension 0"): + FixedSlotKVCache( + prefilled, + capacity=11, + regions=[], + storage=storage, + ) + + for (key, value), (old_key, old_value) in zip(storage, before, strict=True): + assert torch.equal(key, old_key) + assert torch.equal(value, old_value) def chunks(tokens: int, value: float) -> list[LayerKV]: diff --git a/flashdreams/tests/test_multiview_rollout.py b/flashdreams/tests/test_multiview_rollout.py index 1356f10bd..a3f8e1c28 100644 --- a/flashdreams/tests/test_multiview_rollout.py +++ b/flashdreams/tests/test_multiview_rollout.py @@ -252,8 +252,9 @@ def test_block_mask_selection_crosses_the_protocol_boundary() -> None: ) assert model.state.masks - assert all(isinstance(mask, BlockMask) for mask in model.state.masks) - assert all(mask.BLOCK_SIZE == (16, 32) for mask in model.state.masks) + for mask in model.state.masks: + assert isinstance(mask, BlockMask) + assert mask.BLOCK_SIZE == (16, 32) def test_seeded_rollouts_are_reproducible_through_an_adapter() -> None: From f5d0d1dcd9a34b1e0823b7be7c4ea7c73a55a2c5 Mon Sep 17 00:00:00 2001 From: Aidan Foster Date: Tue, 15 Sep 2026 10:45:48 -0700 Subject: [PATCH 3/3] fix: preserve aliased cache prefill data Signed-off-by: Aidan Foster --- .../flashdreams/core/attention/kvcache.py | 18 +++++++++++----- flashdreams/tests/test_multiview_cache.py | 21 +++++++++++++++++++ 2 files changed, 34 insertions(+), 5 deletions(-) diff --git a/flashdreams/flashdreams/core/attention/kvcache.py b/flashdreams/flashdreams/core/attention/kvcache.py index 152026642..47aff1159 100644 --- a/flashdreams/flashdreams/core/attention/kvcache.py +++ b/flashdreams/flashdreams/core/attention/kvcache.py @@ -156,6 +156,7 @@ def __init__( kind="value", ) + buffers: list[LayerKV] = [] for layer, (key, value) in enumerate(prefilled): if storage is None: key_shape = list(key.shape) @@ -166,13 +167,20 @@ def __init__( value_buffer = value.new_zeros(value_shape) else: key_buffer, value_buffer = storage[layer] - prefix = self._seq_slice(0, self._initial_length, tensor_dim) - suffix = self._seq_slice(self._initial_length, capacity, tensor_dim) - with torch.no_grad(): - key_buffer[suffix].zero_() - value_buffer[suffix].zero_() + buffers.append((key_buffer, value_buffer)) + + prefix = self._seq_slice(0, self._initial_length, tensor_dim) + suffix = self._seq_slice(self._initial_length, capacity, tensor_dim) + with torch.no_grad(): + for (key, value), (key_buffer, value_buffer) in zip( + prefilled, buffers, strict=True + ): key_buffer[prefix] = key value_buffer[prefix] = value + for key_buffer, value_buffer in buffers: + key_buffer[suffix].zero_() + value_buffer[suffix].zero_() + for key_buffer, value_buffer in buffers: self._k.append(key_buffer) self._v.append(value_buffer) diff --git a/flashdreams/tests/test_multiview_cache.py b/flashdreams/tests/test_multiview_cache.py index 4c9b58f37..adfaf900c 100644 --- a/flashdreams/tests/test_multiview_cache.py +++ b/flashdreams/tests/test_multiview_cache.py @@ -108,6 +108,27 @@ def test_invalid_caller_storage_does_not_mutate_any_layer() -> None: assert torch.equal(value, old_value) +def test_prefill_can_alias_unused_caller_storage() -> None: + """Copy aliased prefill values before clearing the unused suffix.""" + storage = [layer(11, -9.0)] + key_buffer, value_buffer = storage[0] + prefilled = [(key_buffer[:, :, 6:], value_buffer[:, :, 6:])] + expected = [(key.clone(), value.clone()) for key, value in prefilled] + + memory = FixedSlotKVCache( + prefilled, + capacity=11, + regions=[], + storage=storage, + ) + + key, value = memory.layers()[0] + assert torch.equal(key, expected[0][0]) + assert torch.equal(value, expected[0][1]) + assert torch.count_nonzero(key_buffer[:, :, 5:]) == 0 + assert torch.count_nonzero(value_buffer[:, :, 5:]) == 0 + + def chunks(tokens: int, value: float) -> list[LayerKV]: """Build a two-layer cache write.""" return [layer(tokens, value), layer(tokens, value + 1.0)]