FloodDiffusion2-Live / fixed_text_runtime.py
caiyiyi1998's picture
Initial commit
9a25493
Raw History Blame Contribute Delete
14.3 kB
"""Local runtime with mutable, fixed-shape text conditioning for CUDA graphs.
For the released SEED model (text cross-RoPE disabled), each query sees exactly
one prompt segment. Repeated segments can therefore share that prompt's token
span. A bounded GPU token bank keeps forward shapes fixed while arbitrary
encoded prompts enter and leave the visible history. No model source changes.
"""
from __future__ import annotations
import torch
from space.inference import _signature
from window_runtime import WindowRuntime
def _feature_key(feature):
try:
version = feature._version
except RuntimeError: # Inference tensors have no version counter.
version = None
return id(feature), version, tuple(feature.shape), feature.dtype
class FixedPromptBank:
"""One-stream text-context adapter; all mutations use the caller's stream.
Eviction follows the actual attention window, not the longer retained ring
history. Slot moves only change masks and values, never the bank shape.
RoPE-enabled or oversized contexts retain the official implementation and
are explicitly marked for eager execution, without on-switch graph capture.
"""
def __init__(self, text_module, original_context, buckets=(32, 64, 128, 256, 512, 1024)):
buckets = tuple(buckets)
if (not buckets or any(isinstance(n, bool) or not isinstance(n, int) or n <= 0 for n in buckets)
or tuple(sorted(set(buckets))) != buckets):
raise ValueError("Text token buckets must be increasing positive integers")
self.text = text_module
self.original_context = original_context
self.buckets = buckets
self.capacity = buckets[0]
self.bank = None
self.spans = {}
self.keys = {}
self._states = {}
self.uploads = self.compactions = self.fallback_calls = 0
self.max_active_tokens = self.active_tokens = 0
self.last_mode = "uninitialized"
self.null_length = int(self.text.text_cache[""].shape[0])
@property
def graph_compatible(self):
return self.last_mode == "fixed_bank"
def _fallback(self, reason, start, end, device, dtype):
self.last_mode = reason
self.fallback_calls += 1
return self.original_context(start, end, device, dtype)
def _first_free_span(self, length):
cursor = 0
for begin, size in sorted(self.spans.values()):
if begin - cursor >= length:
return cursor
cursor = begin + size
return cursor if self.capacity - cursor >= length else None
def _upload(self, name, begin, length, device, dtype):
feature = self.text.text_cache[name]
self.bank[begin:begin + length].copy_(feature.to(device=device, dtype=dtype))
self.spans[name] = (begin, length)
self.keys[name] = _feature_key(feature)
self.uploads += 1
def __call__(self, start, end, device, dtype):
if self.text.cross_rope:
return self._fallback("cross_rope_eager", start, end, device, dtype)
if len(self.text.stream_segments) != 1:
return self._fallback("batch_eager", start, end, device, dtype)
if int(self.text.text_cache[""].shape[0]) != self.null_length:
return self._fallback("changed_null_length_eager", start, end, device, dtype)
segments = self.text.stream_segments[0]
ends = [index for _, index in segments[1:]] + [self.text.stream_frames[0]]
active = [(name, begin, finish) for (name, begin), finish in zip(segments, ends)
if begin < end and finish > start]
if not active:
return self._fallback("empty_window_eager", start, end, device, dtype)
names = list(dict.fromkeys(name for name, _, _ in active))
lengths = {name: int(self.text.text_cache[name].shape[0]) for name in names}
self.active_tokens = sum(lengths.values())
self.max_active_tokens = max(self.max_active_tokens, self.active_tokens)
capacity = next((size for size in self.buckets if size >= self.active_tokens), None)
if capacity is None:
return self._fallback("capacity_overflow_eager", start, end, device, dtype)
width = int(self.text.text_cache[names[0]].shape[1])
wanted_shape = (capacity, width)
device = torch.device(device)
if capacity not in self._states:
self._states[capacity] = (torch.zeros(wanted_shape, device=device, dtype=dtype), {}, {})
bank, spans, keys = self._states[capacity]
if bank.shape != wanted_shape or bank.device != device or bank.dtype != dtype:
return self._fallback("changed_tensor_contract_eager", start, end, device, dtype)
self.capacity, self.bank, self.spans, self.keys = capacity, bank, spans, keys
# Drop only unused or resized spans. Stale values are harmless because
# every query's mask selects exactly its current prompt's valid span.
for name in list(self.spans):
if name not in lengths or self.spans[name][1] != lengths[name]:
del self.spans[name]
del self.keys[name]
for name in names:
length = lengths[name]
key = _feature_key(self.text.text_cache[name])
if name in self.spans:
# Unversioned inference tensors are refreshed so an in-place
# feature edit cannot leave captured conditioning stale.
if key != self.keys[name] or key[1] is None:
self._upload(name, self.spans[name][0], length, device, dtype)
continue
begin = self._first_free_span(length)
if begin is None:
# Fragmentation, not overflow: stage a complete compact layout
# from genuine cached features while preserving bank identity.
self.spans.clear()
self.keys.clear()
cursor = 0
for active_name in names:
self._upload(active_name, cursor, lengths[active_name], device, dtype)
cursor += lengths[active_name]
self.compactions += 1
break
self._upload(name, begin, length, device, dtype)
mask = torch.zeros(end - start, self.capacity, dtype=torch.bool, device=device)
for name, begin, finish in active:
offset, length = self.spans[name]
mask[max(0, begin - start):min(end, finish) - start, offset:offset + length] = True
self.last_mode = "fixed_bank"
return [self.bank], {"cross_attn_mask": [mask]}
def status(self):
return dict(buckets=list(self.buckets), capacity=self.capacity, active_tokens=self.active_tokens,
max_active_tokens=self.max_active_tokens, resident_prompts=len(self.spans),
feature_uploads=self.uploads, compactions=self.compactions,
fallback_calls=self.fallback_calls, mode=self.last_mode)
def _pad_motion_inputs(args, kwargs, max_motion_frames):
"""Canonicalize a short 1x1 masked motion window for a prepared graph.
The WAN already pads embedded states to seq_len. Raw zero padding changes
only invalid states (e.g. the embedding bias); explicit self-attention
masks prevent every valid query from reading those added keys. Valid
temporal RoPE positions remain 0..N-1. This helper declines other contracts.
"""
if max_motion_frames is None or len(args) != 4 or args[3] != max_motion_frames:
return args, kwargs, None
if any(kwargs.get(name) is not None for name in ("y", "rope_ids", "text_k_rope_ids", "text_q_rope_ids")):
return args, kwargs, None
values, times, contexts = args[:3]
self_masks, cross_masks = kwargs.get("attn_mask"), kwargs.get("cross_attn_mask")
sequences = (values, times, contexts, self_masks, cross_masks)
if any(not isinstance(sequence, (list, tuple)) for sequence in sequences):
return args, kwargs, None
count = len(values)
if not count or any(len(sequence) != count for sequence in sequences):
return args, kwargs, None
lengths = []
for value, time, context, self_mask, cross_mask in zip(*sequences):
if (value.ndim != 4 or value.shape[-2:] != (1, 1) or
not 0 < value.shape[1] <= max_motion_frames):
return args, kwargs, None
length = value.shape[1]
if (time.ndim != 1 or time.shape[0] != length or context.ndim != 2 or
self_mask.shape != (length, length) or self_mask.dtype != torch.bool or
cross_mask.shape != (length, context.shape[0]) or cross_mask.dtype != torch.bool):
return args, kwargs, None
lengths.append(length)
if all(length == max_motion_frames for length in lengths):
return args, kwargs, None
padded_values, padded_times, padded_self, padded_cross = [], [], [], []
for value, time, _, self_mask, cross_mask in zip(*sequences):
length = value.shape[1]
padded = value.new_zeros((value.shape[0], max_motion_frames, 1, 1))
padded[:, :length].copy_(value)
padded_values.append(padded)
padded = time.new_zeros(max_motion_frames)
padded[:length].copy_(time)
padded_times.append(padded)
padded = self_mask.new_zeros((max_motion_frames, max_motion_frames))
padded[:length, :length].copy_(self_mask)
padded_self.append(padded)
padded = cross_mask.new_zeros((max_motion_frames, cross_mask.shape[1]))
padded[:length].copy_(cross_mask)
padded_cross.append(padded)
changed_args = (padded_values, padded_times, contexts, args[3])
changed_kwargs = dict(kwargs, attn_mask=padded_self, cross_attn_mask=padded_cross)
return changed_args, changed_kwargs, lengths
class _BankGraphGate:
def __init__(self, graph, bank, max_motion_frames=None):
self.graph, self.bank = graph, bank
self.max_motion_frames = max_motion_frames
self.allow_capture = True
self.uncached_eager_calls = 0
self.startup_padded_calls = 0
def __call__(self, *args, **kwargs):
if self.bank.graph_compatible:
if not self.graph.enabled:
return self.graph(*args, **kwargs)
graph_args, graph_kwargs, lengths = _pad_motion_inputs(args, kwargs, self.max_motion_frames)
inputs = (graph_args, graph_kwargs)
key = (_signature(inputs), torch.is_autocast_enabled("cuda"),
torch.get_autocast_dtype("cuda"))
if self.allow_capture or key in self.graph.entries:
output = self.graph(*graph_args, **graph_kwargs)
if lengths is not None:
self.startup_padded_calls += 1
return [value[:, :length] for value, length in zip(output, lengths)]
return output
self.uncached_eager_calls += 1
# A supported prompt is never truncated/rejected to fit the fast path.
# Preserve existing graph entries for when a large history ages out.
self.graph.eager_calls += 1
return self.graph.forward(*args, **kwargs)
class FixedTextWindowRuntime(WindowRuntime):
"""WindowRuntime plus a dynamic prompt bank with fixed capacity buckets.
Add genuine prompt features with the inherited add_prompt_features even
during a session. New prompts fitting the active token budget reuse the
same captured shapes. Bucket capacities are explicit and configurable at
start. Overflow is correct eager inference, reported in status(). Prewarm
all bucket graphs then freeze_graph_shapes() to prohibit mid-session capture.
"""
def __init__(self, model, recovery_type, metadata):
super().__init__(model, recovery_type, metadata)
self._original_text_context = None
self._fixed_text_bank = None
self._bank_graph_gate = None
@torch.inference_mode()
def start(self, seed=0, history=120, use_graph=True, *, text_token_buckets=(32, 64, 128, 256, 512, 1024)):
buckets = tuple(text_token_buckets)
# Validate before resetting a running session.
if (not buckets or any(isinstance(n, bool) or not isinstance(n, int) or n <= 0 for n in buckets)
or tuple(sorted(set(buckets))) != buckets):
raise ValueError("Text token buckets must be increasing positive integers")
super().start(seed=seed, history=history, use_graph=use_graph)
text = self.model.text_module
self._original_text_context = text.get_stream_context
self._fixed_text_bank = FixedPromptBank(text, self._original_text_context, buckets)
text.get_stream_context = self._fixed_text_bank
self._graph.max_entries = max(self._graph.max_entries, len(buckets))
pad_history = (history if tuple(self.model.spatial_shape) == (1, 1) and
tuple(self.model.model.patch_size) == (1, 1, 1) else None)
self._bank_graph_gate = _BankGraphGate(self._graph, self._fixed_text_bank, pad_history)
self.model.model.forward = self._bank_graph_gate
return self
def freeze_graph_shapes(self):
"""Forbid new captures after prewarming; unknown shapes execute eagerly."""
if not self._ready or self._bank_graph_gate is None:
raise RuntimeError("Call start before freezing graph shapes")
self._bank_graph_gate.allow_capture = False
def close(self):
if self._original_text_context is not None:
self.model.text_module.get_stream_context = self._original_text_context
self._original_text_context = None
super().close()
def status(self):
status = super().status()
status["fixed_text"] = self._fixed_text_bank.status() if self._fixed_text_bank else None
if self._bank_graph_gate is not None:
status["fixed_text"].update(capture_frozen=not self._bank_graph_gate.allow_capture,
uncached_eager_calls=self._bank_graph_gate.uncached_eager_calls,
startup_padded_calls=self._bank_graph_gate.startup_padded_calls)
return status