Spaces:
Running on Zero
Running on Zero
File size: 14,293 Bytes
9a25493 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 | """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
|