diff --git a/flashdreams/flashdreams/core/attention/kvcache.py b/flashdreams/flashdreams/core/attention/kvcache.py index 9b4c6ce5c..47aff1159 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,21 +134,84 @@ 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) - prefix = self._seq_slice(0, self._initial_length, tensor_dim) - with torch.no_grad(): + 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" + ) + self._validate_storage_layer( + value, + value_buffer, + capacity=capacity, + layer=layer, + kind="value", + ) + + buffers: list[LayerKV] = [] + 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] + 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) + 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..adfaf900c 100644 --- a/flashdreams/tests/test_multiview_cache.py +++ b/flashdreams/tests/test_multiview_cache.py @@ -55,6 +55,80 @@ 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, -9.0), layer(11, -9.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) + 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 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)] diff --git a/flashdreams/tests/test_multiview_rollout.py b/flashdreams/tests/test_multiview_rollout.py index cfc8d809f..a3f8e1c28 100644 --- a/flashdreams/tests/test_multiview_rollout.py +++ b/flashdreams/tests/test_multiview_rollout.py @@ -238,22 +238,23 @@ 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) + 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: