Spaces:
Running on Zero
Running on Zero
Download fast_window_runtime.py from AlayaLab/FloodDiffusion2-Live: direct link, hf CLI and curl.
- Browser
- Download file 25.5 kB
-
https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/fast_window_runtime.py
- Command line
-
hf download hf://spaces/AlayaLab/FloodDiffusion2-Live/fast_window_runtime.py
-
curl -L -o fast_window_runtime.py https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/fast_window_runtime.py
25.5 kB
| """Private SEED path runtime: complete GPU sampling graphs and mixed precision. | |
| This preserves the official finite motion window. Motion K/V is recomputed: | |
| removing an old prefix changes deeper-layer history representations. Text K/V | |
| is safely cached because its projections depend only on the current genuine | |
| text features and fixed weights, with cross-text RoPE disabled in this release. | |
| No public/vendor files are changed; close restores original methods/FP32 weights. | |
| """ | |
| from __future__ import annotations | |
| from contextlib import contextmanager | |
| import copy | |
| import math | |
| import time | |
| import torch | |
| from fixed_text_runtime import FixedPromptBank, _feature_key | |
| from space.inference import FPS | |
| from window_runtime import WindowRuntime | |
| DEFAULT_TEXT_BUCKETS = (32, 64, 128, 256, 512, 1024) | |
| def _motion_attention_mask(indices, valid, target_start, history_cutoff): | |
| """Keep the partial mask, isolating new queries from an old prompt epoch. | |
| All arguments can stay on the GPU with fixed shapes during graph replay. | |
| Old queries keep their normal causal keys, so the cutoff introduces no | |
| fully masked valid rows. New queries cannot read them in any model layer. | |
| A zero cutoff is exactly the original mask. | |
| """ | |
| partial = (valid[:, None] & valid[None, :] & | |
| ((indices[:, None] >= target_start) | (indices[None, :] <= indices[:, None]))) | |
| epoch = ((indices[:, None] < history_cutoff) | (indices[None, :] >= history_cutoff)) | |
| return partial & epoch | |
| class _MixedPrecisionWeights: | |
| def __init__(self, model, precision): | |
| if precision not in ("bf16", "fp32"): | |
| raise ValueError("precision must be 'bf16' or 'fp32'") | |
| self.precision = precision | |
| self.originals = [] | |
| if precision == "fp32": | |
| return | |
| backbone = model.model | |
| modules = [backbone.patch_embedding, backbone.text_embedding] | |
| for block in backbone.blocks: | |
| modules.append(block.ffn) | |
| for attention in (block.self_attn, block.cross_attn): | |
| modules.extend((attention.q, attention.k, attention.v, attention.o)) | |
| seen = set() | |
| try: | |
| for module in modules: | |
| for parameter in module.parameters(): | |
| if id(parameter) in seen: | |
| continue | |
| seen.add(id(parameter)) | |
| original = parameter.data | |
| self.originals.append((parameter, original)) | |
| parameter.data = original.to(dtype=torch.bfloat16) | |
| except BaseException: | |
| # Construction can fail before the caller receives this handle. | |
| # Restore already converted storages here, including on OOM. | |
| self.restore() | |
| raise | |
| def restore(self): | |
| """Restore the actual original storages, without BF16 round-trip loss.""" | |
| for parameter, original in self.originals: | |
| parameter.data = original | |
| self.originals.clear() | |
| def configure_mixed_precision(model, precision="bf16"): | |
| """Match the old benchmark's BF16 projections; return a .restore() handle. | |
| Time embeddings, normalization, modulation, output head, latent state and | |
| integration remain FP32. Reference forwards must also use CUDA BF16 | |
| autocast with cache_enabled=False when precision is bf16. | |
| """ | |
| return _MixedPrecisionWeights(model, precision) | |
| class _SamplingCore: | |
| """One text-capacity graph sharing the runtime's actual ring and step id.""" | |
| def __init__(self, runtime, capacity, null_length): | |
| self.runtime, self.model = runtime, runtime.model | |
| self.capacity, self.null_length = int(capacity), int(null_length) | |
| self.history = int(runtime._history) | |
| self.precision = runtime._fast_precision | |
| self.device = self.model.generated[0].device | |
| self.state = self.model.generated[0] | |
| self.roots = self.model._stream_position | |
| self.step_id = runtime._fast_step_id | |
| self.history_cutoff = runtime._fast_history_cutoff | |
| self.time_table = runtime._fast_time_table | |
| self.rows = torch.arange(self.history, device=self.device, dtype=torch.long) | |
| self.text_length = max(self.capacity, self.null_length) | |
| width = self.model.text_module.dim | |
| dim = self.model.model.dim | |
| dtype = torch.bfloat16 if self.precision == "bf16" else torch.float32 | |
| self.context = torch.zeros(self.capacity, width, device=self.device) | |
| self.null_context = torch.zeros(self.null_length, width, device=self.device) | |
| self.cross_mask = torch.zeros(self.history, self.capacity, dtype=torch.bool, device=self.device) | |
| self.projected_text = torch.zeros(2, self.text_length, dim, device=self.device, dtype=dtype) | |
| # WanRMSNorm restores FP32 through its FP32 learned scale. | |
| self.text_k = [torch.zeros(2, self.text_length, dim, device=self.device) | |
| for _ in self.model.model.blocks] | |
| self.text_v = [torch.zeros(2, self.text_length, dim, device=self.device, dtype=dtype) | |
| for _ in self.model.model.blocks] | |
| self.output = torch.zeros(self.model.model_input_dim, device=self.device) | |
| self.graph = None | |
| self.text_revision = None | |
| self.text_refreshes = 0 | |
| def _autocast(self): | |
| return torch.autocast("cuda", enabled=self.precision == "bf16", | |
| dtype=torch.bfloat16, cache_enabled=False) | |
| def prepare_text(self, context, null_context, mask, revision=None): | |
| """Update only mutable inputs; no capture or model-state reset.""" | |
| self.cross_mask.zero_() | |
| self.cross_mask[:mask.shape[0]].copy_(mask) | |
| if revision is not None and revision == self.text_revision: | |
| return | |
| self.context.copy_(context) | |
| self.null_context.copy_(null_context) | |
| padded = self.context.new_zeros((2, self.text_length, self.context.shape[1])) | |
| padded[0, :self.capacity].copy_(self.context) | |
| padded[1, :self.null_length].copy_(self.null_context) | |
| # This work happens only for new/changed prompt features or a moved | |
| # bank span. It is independent of pose, timesteps and motion history. | |
| with self._autocast(): | |
| projected = self.model.model.text_embedding(padded) | |
| self.projected_text.copy_(projected) | |
| for block, keys, values in zip(self.model.model.blocks, self.text_k, self.text_v): | |
| attention = block.cross_attn | |
| keys.copy_(attention.norm_k(attention.k(projected))) | |
| values.copy_(attention.v(projected)) | |
| self.text_revision = revision | |
| self.text_refreshes += 1 | |
| def _cached_text_projections(self): | |
| """Bind cached tensors during eager/capture; replay needs no patching.""" | |
| saved = [] | |
| def replace(module, forward): | |
| saved.append((module, module.forward)) | |
| module.forward = forward | |
| replace(self.model.model.text_embedding, lambda _value: self.projected_text) | |
| for block, keys, values in zip(self.model.model.blocks, self.text_k, self.text_v): | |
| replace(block.cross_attn.k, lambda _value, cached=keys: cached) | |
| replace(block.cross_attn.norm_k, lambda value: value) | |
| replace(block.cross_attn.v, lambda _value, cached=values: cached) | |
| try: | |
| yield | |
| finally: | |
| for module, forward in reversed(saved): | |
| module.forward = forward | |
| def step(self): | |
| """GPU schedule, full-window denoising, CFG, Euler, output and tick.""" | |
| # CPU-generated table matches torch.tensor(current_step / 30) exactly. | |
| end = self.step_id + 1 | |
| window_start = (end - self.history).clamp(min=0) | |
| indices = window_start + self.rows | |
| valid = indices < end | |
| target_start = (self.step_id - 29).clamp(min=0) | |
| updating = valid & (indices >= target_start) | |
| latent = self.state.index_select(1, indices) | |
| raw_root = self.roots[0].index_select(0, indices) | |
| normalized_root = (raw_root - self.model.position_mean) / self.model.position_std | |
| conditioned = torch.cat((normalized_root.T[:, :, None, None], latent), dim=0) | |
| conditioned = torch.where(valid[None, :, None, None], conditioned, torch.zeros_like(conditioned)) | |
| level = (-indices.float() / 30 + self.time_table.index_select(0, self.step_id.reshape(1))[0]).clamp(0, 1) | |
| level = level * self.model.time_embedding_scale | |
| self_mask = _motion_attention_mask(indices, valid, target_start, self.history_cutoff) | |
| null_mask = valid[:, None].expand(self.history, self.null_length) | |
| with self._autocast(), self._cached_text_projections(): | |
| predicted = self.runtime._forward( | |
| [conditioned, conditioned], [level, level], [self.context, self.null_context], self.history, | |
| attn_mask=[self_mask, self_mask], text_k_rope_ids=None, text_q_rope_ids=None, | |
| cross_attn_mask=[self.cross_mask, null_mask], text_seq_len=self.text_length, y=None) | |
| # Keep scalar order equal to the official text_scale*pt + null_scale*pn. | |
| velocity = self.model.cfg_config["text_scale"] * predicted[0] + self.model.cfg_config["null_scale"] * predicted[1] | |
| updated = torch.where(updating[None, :, None, None], latent + velocity * (1 / 30), latent) | |
| self.state.index_copy_(1, indices, updated) | |
| committed_index = (end - 30).clamp(min=0).reshape(1) | |
| pose = self.state.index_select(1, committed_index)[:, 0, 0, 0] | |
| root = self.roots[0].index_select(0, committed_index)[0] | |
| self.output.copy_(torch.cat((root, pose * self.model.std + self.model.mean))) | |
| self.step_id.add_(1) | |
| def snapshot(self): | |
| """All mutable GPU data required to compare step() and graph.replay().""" | |
| return dict(state=self.state.clone(), roots=self.roots.clone(), step_id=self.step_id.clone(), | |
| history_cutoff=self.history_cutoff.clone(), | |
| context=self.context.clone(), null_context=self.null_context.clone(), | |
| cross_mask=self.cross_mask.clone(), projected_text=self.projected_text.clone(), | |
| text_k=[value.clone() for value in self.text_k], | |
| text_v=[value.clone() for value in self.text_v], output=self.output.clone()) | |
| def restore(self, snapshot): | |
| for name in ("state", "roots", "step_id", "history_cutoff", "context", "null_context", "cross_mask", "projected_text", "output"): | |
| getattr(self, name).copy_(snapshot[name]) | |
| for name in ("text_k", "text_v"): | |
| for destination, source in zip(getattr(self, name), snapshot[name]): | |
| destination.copy_(source) | |
| def capture(self): | |
| """Prepare once without consuming RNG or retaining probe motion.""" | |
| snapshot = self.snapshot() | |
| stream = torch.cuda.Stream() | |
| stream.wait_stream(torch.cuda.current_stream()) | |
| try: | |
| with torch.cuda.stream(stream): | |
| for _ in range(3): | |
| self.restore(snapshot) | |
| self.step_id.fill_(self.history - 1) | |
| self.step() | |
| torch.cuda.current_stream().wait_stream(stream) | |
| torch.cuda.synchronize() | |
| self.restore(snapshot) | |
| self.step_id.fill_(self.history - 1) | |
| self.step() | |
| reference = {"state": self.state.clone(), "output": self.output.clone(), "step_id": self.step_id.clone()} | |
| self.restore(snapshot) | |
| self.step_id.fill_(self.history - 1) | |
| graph = torch.cuda.CUDAGraph() | |
| with torch.cuda.graph(graph, stream=stream): | |
| self.step() | |
| self.restore(snapshot) | |
| self.step_id.fill_(self.history - 1) | |
| graph.replay() | |
| torch.testing.assert_close(self.state, reference["state"], rtol=1e-5, atol=1e-5) | |
| torch.testing.assert_close(self.output, reference["output"], rtol=1e-5, atol=1e-5) | |
| torch.testing.assert_close(self.step_id, reference["step_id"], rtol=0, atol=0) | |
| self.graph = graph | |
| finally: | |
| self.restore(snapshot) | |
| class FastWindowRuntime(WindowRuntime): | |
| """Compatible local runtime using full-step graphs and a BF16 backbone. | |
| Native path controls, pending prompt/path replacement, 29-frame cold-start | |
| lookahead, finite 90/120/150-frame attention windows and ring rollback are kept. | |
| use_graph=False executes exactly the same optimized step eagerly. | |
| precision='fp32' keeps all original weights for precision comparisons. | |
| reset_history_on_prompt_change masks committed motion from earlier prompt | |
| epochs while preserving pending latents, scheduling, noise and recovery. | |
| All step APIs observe prompt changes; only step_replanned(replan_text=True) | |
| also replaces the text conditions of already pending frames. | |
| """ | |
| def __init__(self, model, recovery_type, metadata): | |
| super().__init__(model, recovery_type, dict(metadata)) | |
| self._fast_original_stream = None | |
| self._precision_handle = None | |
| self._fast_cores = {} | |
| self._overflow_core = None | |
| self._last_fast_core = None | |
| self._fast_text = None | |
| self._fast_step_id = self._fast_time_table = None | |
| self._fast_history_cutoff = None | |
| self._history_reset_enabled = False | |
| self._history_cutoff_absolute = None | |
| self._history_cutoff_local = 0 | |
| self._history_resets = 0 | |
| self._visible_history = 0 | |
| self._fast_precision = "bf16" | |
| self._fast_use_graph = False | |
| self._fast_graph_replays = self._fast_eager_calls = 0 | |
| self._fast_prewarm = None | |
| def _assert_fast_contract(self): | |
| model = self.model | |
| scheduler = model.time_scheduler | |
| if (model.batch_size != 1 or model.input_dim <= 0 or | |
| model.model_input_dim != model.input_dim + model.stream_condition_dim or | |
| tuple(model.spatial_shape) != (1, 1) or tuple(model.model.patch_size) != (1, 1, 1) or | |
| model.stream_condition_dim != 3 or model.attn_type != "partial" or | |
| model.prediction_type != "vel" or scheduler.steps != 30 or scheduler.chunk_size != 30 or | |
| scheduler.noise_type != "linear" or scheduler.sigma_type != "zero" or | |
| model.text_module.cross_rope): | |
| raise RuntimeError("FastWindowRuntime requires pose channels plus root3, 30/30 linear/zero, partial attention") | |
| def _start_fast_session(self, seed, history): | |
| """Base session initialization with local finite-history windows. | |
| The public Runtime intentionally still accepts only its released window | |
| sizes. Keep its seed, scheduler and recovery initialization here, while | |
| the local full-step runtime also supports 60/90 total attention frames. | |
| No forward-only graph wrapper is needed before full-step prewarming. | |
| """ | |
| history = int(history) | |
| if history not in (60, 90, 120, 150): | |
| raise ValueError("history must be 60, 90, 120 or 150 total window frames") | |
| if next(self.model.parameters()).device.type != "cuda": | |
| raise RuntimeError("start must run inside the ZeroGPU CUDA allocation") | |
| self.close() | |
| torch.manual_seed(int(seed)) | |
| self.model.schedule_config.clear() | |
| self.model.schedule_config.update(copy.deepcopy(self._schedule)) | |
| self.model.cfg_config.update(text_scale=4.0, null_scale=-3.0) | |
| self.model.init_generated(history, batch_size=1, | |
| schedule_config=copy.deepcopy(self._schedule)) | |
| self._recovery = self._recovery_type(smoothing_alpha=1.0, fps=getattr(self, "fps", FPS)) | |
| self._steps = self._frames = self._recovered = 0 | |
| self._action = self._prompt = None | |
| self._history = history | |
| self._ready = True | |
| def start(self, seed=0, history=120, use_graph=True, *, precision="bf16", | |
| text_token_buckets=DEFAULT_TEXT_BUCKETS, reset_history_on_prompt_change=False, | |
| cfg_scale=4.0): | |
| if precision not in ("bf16", "fp32"): | |
| raise ValueError("precision must be 'bf16' or 'fp32'") | |
| if not isinstance(reset_history_on_prompt_change, bool): | |
| raise TypeError("reset_history_on_prompt_change must be a bool") | |
| cfg_scale = float(cfg_scale) | |
| if not math.isfinite(cfg_scale) or cfg_scale < 0: | |
| raise ValueError("cfg_scale must be a finite nonnegative number") | |
| buckets = tuple(text_token_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") | |
| began = time.perf_counter() | |
| try: | |
| self._start_fast_session(seed, history) | |
| self.model.cfg_config.update(text_scale=cfg_scale, null_scale=1.0-cfg_scale) | |
| self._assert_fast_contract() | |
| self.model.model.forward = self._forward | |
| self._fast_precision, self._fast_use_graph = precision, bool(use_graph) | |
| self._precision_handle = configure_mixed_precision(self.model, precision) | |
| device = self.model.generated[0].device | |
| self._fast_step_id = torch.zeros((), dtype=torch.long, device=device) | |
| # Every capacity graph reads this same mutable scalar. Prompt | |
| # switches update its value, never the graph shape or capture. | |
| self._fast_history_cutoff = torch.zeros((), dtype=torch.long, device=device) | |
| self._history_reset_enabled = reset_history_on_prompt_change | |
| self._history_cutoff_absolute = None | |
| self._history_cutoff_local = 0 | |
| self._history_resets = self._visible_history = 0 | |
| self._fast_time_table = torch.tensor([index * (1 / 30) for index in range(self.model.buf_len + 1)], | |
| device=device, dtype=torch.float32) | |
| text = self.model.text_module | |
| self._fast_text = FixedPromptBank(text, text.get_stream_context, buckets) | |
| null = text.text_cache[""].to(device=device, dtype=torch.float32) | |
| for capacity in buckets: | |
| core = _SamplingCore(self, capacity, len(null)) | |
| context = torch.zeros_like(core.context) | |
| copied = min(len(null), capacity) | |
| context[:copied].copy_(null[:copied]) | |
| mask = torch.zeros_like(core.cross_mask) | |
| mask[:, :copied] = True | |
| core.prepare_text(context, null, mask) | |
| if use_graph: | |
| core.capture() | |
| self._fast_cores[capacity] = core | |
| self._fast_original_stream = self.model.stream_generate_step | |
| self.model.stream_generate_step = self._stream_generate_step | |
| self._fast_graph_replays = self._fast_eager_calls = 0 | |
| self._fast_prewarm = dict(complete=True, seconds=time.perf_counter() - began, | |
| captures=len(buckets) if use_graph else 0, capacities=list(buckets)) | |
| self.metadata.update(cfg_scale=cfg_scale, | |
| precision=("BF16 projections; FP32 normalization/time/head/latent/Euler" if precision == "bf16" else "FP32 parameters and sampling; original BF16 SDPA"), | |
| sampling_graph="full GPU step", motion_kv=False, text_kv=True) | |
| return self | |
| except BaseException: | |
| self.close() | |
| raise | |
| def _update_prompt_history(self, prompt): | |
| """Run after input validation/condition commit, before the GPU step. | |
| Runtime counters remain absolute while the model counters roll back. | |
| The earliest uncommitted frame is self._frames, not self._steps. | |
| Keep that epoch boundary fixed as new committed history accumulates. | |
| """ | |
| if (self._history_reset_enabled and self._prompt is not None and | |
| prompt != self._prompt): | |
| self._history_cutoff_absolute = self._frames | |
| self._history_resets += 1 | |
| offset = self._steps - self.model.current_step | |
| cutoff = max(0, (self._history_cutoff_absolute or 0) - offset) | |
| if cutoff != self._history_cutoff_local: | |
| self._fast_history_cutoff.fill_(cutoff) | |
| self._history_cutoff_local = cutoff | |
| end = self.model.current_step + 1 | |
| target_start = max(0, end - self.window_frames) | |
| window_start = max(0, end - self._history) | |
| self._visible_history = max(0, target_start - max(window_start, cutoff)) | |
| def _stream_generate_step(self, inputs): | |
| model = self.model | |
| inputs = model._extract_inputs(inputs) | |
| device = model.generated[0].device | |
| model._update_stream_condition(inputs, device) | |
| model.text_module.update_stream(inputs, device, model.param_dtype) | |
| model.condition_frames += 1 | |
| if model.condition_frames > model.buf_len: | |
| model._rollback() | |
| self._fast_step_id.sub_(model.seq_len) | |
| model._commit_stream_condition(model.condition_frames - 1) | |
| # Public step methods have already validated prompt features and root | |
| # rows. In particular, a rejected replan cannot advance the cutoff. | |
| self._update_prompt_history(inputs["text"][0]) | |
| step = model.current_step | |
| end = step + 1 | |
| start = max(0, end - self._history) | |
| contexts, meta = self._fast_text(start, end, device, model.param_dtype) | |
| null_feature = model.text_module.text_cache[""] | |
| null = null_feature.to(device=device, dtype=model.param_dtype) | |
| context, mask = contexts[0], meta["cross_attn_mask"][0] | |
| compatible = self._fast_text.graph_compatible and len(null) == self._fast_cores[self._fast_text.capacity].null_length | |
| if compatible: | |
| core = self._fast_cores[self._fast_text.capacity] | |
| null_key = _feature_key(null_feature) | |
| revision = (id(self._fast_text.bank), self._fast_text.uploads, null_key) | |
| if null_key[1] is None: | |
| revision = None | |
| else: | |
| contract = (len(context), len(null)) | |
| if self._overflow_core is None or (self._overflow_core.capacity, self._overflow_core.null_length) != contract: | |
| self._overflow_core = _SamplingCore(self, *contract) | |
| core, revision = self._overflow_core, None | |
| core.prepare_text(context, null, mask, revision) | |
| self._last_fast_core = core | |
| if self._fast_use_graph and compatible: | |
| core.graph.replay() | |
| self._fast_graph_replays += 1 | |
| else: | |
| core.step() | |
| self._fast_eager_calls += 1 | |
| # The official sigma=0 implementation still draws this exact shape. | |
| # Retain its RNG advancement so newly filled rollback noise matches. | |
| torch.randn_like(model.generated[0][:, max(0, step - 29):end]) | |
| model.current_step += 1 | |
| committable = max(0, model.condition_frames - 29) | |
| if committable > model.current_commit: | |
| model.current_commit = committable | |
| return {"generated": [core.output.unsqueeze(0)]} | |
| return {"generated": [core.output.new_empty((0, model.model_input_dim))]} | |
| def status(self): | |
| status = super().status() | |
| status.update(cfg_scale=float(self.model.cfg_config['text_scale']), | |
| history_reset_enabled=self._history_reset_enabled, | |
| history_cutoff_absolute=self._history_cutoff_absolute, | |
| history_resets=self._history_resets, visible_history=self._visible_history) | |
| status["graph"] = dict(enabled=self._fast_use_graph, cached_shapes=len(self._fast_cores) if self._fast_use_graph else 0, | |
| captures=0, replays=self._fast_graph_replays, eager_calls=self._fast_eager_calls) | |
| status["prewarm"] = dict(self._fast_prewarm) if self._fast_prewarm else None | |
| status["fixed_text"] = self._fast_text.status() if self._fast_text else None | |
| status["fast"] = dict(precision=self._fast_precision, complete_step_graph=True, motion_kv=False, | |
| text_kv=True, finite_window=self._history, | |
| text_refreshes=sum(core.text_refreshes for core in self._fast_cores.values())) | |
| return status | |
| def close(self): | |
| if self._fast_original_stream is not None: | |
| self.model.stream_generate_step = self._fast_original_stream | |
| self._fast_original_stream = None | |
| if self._fast_cores or self._overflow_core is not None: | |
| torch.cuda.synchronize() | |
| for core in self._fast_cores.values(): | |
| core.graph = None | |
| self._fast_cores.clear() | |
| self._overflow_core = None | |
| self._last_fast_core = None | |
| self._fast_text = None | |
| self._fast_step_id = self._fast_time_table = None | |
| self._fast_history_cutoff = None | |
| if self._precision_handle is not None: | |
| self._precision_handle.restore() | |
| self._precision_handle = None | |
| super().close() | |