Skip to content
Merged
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
84 changes: 76 additions & 8 deletions flashdreams/flashdreams/core/attention/kvcache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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.
Expand Down Expand Up @@ -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."""
Expand Down
10 changes: 9 additions & 1 deletion flashdreams/flashdreams/core/attention/multiview/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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."""
Expand All @@ -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()
Expand Down
74 changes: 74 additions & 0 deletions flashdreams/tests/test_multiview_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand Down
11 changes: 6 additions & 5 deletions flashdreams/tests/test_multiview_rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading