"""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