Spaces:
Running on Zero
Running on Zero
Download fixed_text_runtime.py from AlayaLab/FloodDiffusion2-Live: direct link, hf CLI and curl.
- Browser
- Download file 14.3 kB
-
https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/fixed_text_runtime.py
- Command line
-
hf download hf://spaces/AlayaLab/FloodDiffusion2-Live/fixed_text_runtime.py
-
curl -L -o fixed_text_runtime.py https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/fixed_text_runtime.py
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]) | |
| 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 | |
| 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 | |