Spaces:
Running on Zero
Running on Zero
Download space/inference.py from AlayaLab/FloodDiffusion2-Live: direct link, hf CLI and curl.
- Browser
- Download file 19.3 kB
-
https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/space/inference.py
- Command line
-
hf download hf://spaces/AlayaLab/FloodDiffusion2-Live/space/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/space/inference.py
19.3 kB
| """SEED Path 300k EMA streaming inference for a single ZeroGPU session. | |
| Import ``spaces`` before calling load_model at app-module scope. Loading is on | |
| CPU, followed by .to('cuda') so ZeroGPU can virtualize the model. Call start, | |
| step, and close inside the same GPU allocation. Genuine text features can be | |
| added from a session's CPU feature file before generation; the three supplied | |
| fallback prompts and unconditional branch are retained. The official motion | |
| network, scheduler, normalization, stream buffers, and text masks are unchanged. | |
| """ | |
| from collections import OrderedDict | |
| from contextlib import contextmanager | |
| import copy | |
| import importlib | |
| import math | |
| import os | |
| from pathlib import Path | |
| import sys | |
| import numpy as np | |
| import torch | |
| import yaml | |
| PROMPTS = { | |
| "walk": "A person walks forward.", | |
| "jog": "A person jogs forward.", | |
| "run": "A person runs forward.", | |
| } | |
| FPS = 30 | |
| def _repository(repo_root): | |
| root = Path(repo_root).resolve() | |
| if not (root / "models/diffusion_forcing_position_wan.py").is_file(): | |
| raise FileNotFoundError(f"Public model source is missing under {root}") | |
| if str(root) not in sys.path: | |
| sys.path.insert(0, str(root)) | |
| previous = Path.cwd() | |
| os.chdir(root) | |
| try: | |
| yield root | |
| finally: | |
| os.chdir(previous) | |
| def _feature_bank(path): | |
| required = {"", *PROMPTS.values()} | |
| with np.load(path, allow_pickle=False) as archive: | |
| prompts = archive["prompts"].tolist() | |
| if len(prompts) != len(set(prompts)): | |
| raise ValueError("Cached prompt names must be unique") | |
| missing = required - set(prompts) | |
| if missing: | |
| raise ValueError(f"Real T5 cache is missing prompts: {sorted(missing)}") | |
| bank = {} | |
| for index, text in enumerate(prompts): | |
| if text not in required: | |
| continue | |
| value = np.array(archive[f"feature_{index}"], dtype=np.float32, copy=True) | |
| if value.ndim != 2 or value.shape[1] != 4096 or not 0 < len(value) <= 512: | |
| raise ValueError(f"Invalid T5 feature shape for {text!r}: {value.shape}") | |
| if not np.isfinite(value).all(): | |
| raise ValueError(f"Non-finite T5 features for {text!r}") | |
| bank[text] = torch.from_numpy(value) | |
| return bank | |
| def _map_tensors(value, fn): | |
| if isinstance(value, torch.Tensor): | |
| return fn(value) | |
| if isinstance(value, dict): | |
| return {key: _map_tensors(item, fn) for key, item in value.items()} | |
| if isinstance(value, (list, tuple)): | |
| return type(value)(_map_tensors(item, fn) for item in value) | |
| return value | |
| def _signature(value): | |
| if isinstance(value, torch.Tensor): | |
| key = ("tensor", tuple(value.shape), value.dtype, value.device) | |
| return key + (tuple(value.flatten().tolist()),) if value.device.type == "cpu" else key | |
| if isinstance(value, dict): | |
| return (dict, tuple((key, _signature(item)) for key, item in sorted(value.items()))) | |
| if isinstance(value, (list, tuple)): | |
| return (type(value), tuple(_signature(item) for item in value)) | |
| return (type(value), value) | |
| def _copy_inputs(destination, source): | |
| if isinstance(source, torch.Tensor): | |
| if source.device.type == "cuda": | |
| destination.copy_(source) | |
| elif isinstance(source, dict): | |
| for key in source: | |
| _copy_inputs(destination[key], source[key]) | |
| elif isinstance(source, (list, tuple)): | |
| for dst, src in zip(destination, source): | |
| _copy_inputs(dst, src) | |
| class _GraphForward: | |
| """Bounded graph cache adapted from the existing standalone demo wrapper. | |
| Only the unchanged denoiser forward is captured. Startup remains eager. | |
| Capture errors are explicit; a caller can restart with use_graph=False. | |
| Graph objects are created only in start/step, inside the GPU allocation. | |
| """ | |
| def __init__(self, forward, *, enabled, min_tokens, max_entries=4): | |
| self.forward = forward | |
| self.enabled = bool(enabled) | |
| self.min_tokens = min_tokens | |
| self.max_entries = max_entries | |
| self.entries = OrderedDict() | |
| self.captures = self.replays = self.eager_calls = 0 | |
| def __call__(self, *args, **kwargs): | |
| x = args[0] if args else kwargs["x"] | |
| short = any(t.shape[1:].numel() < self.min_tokens for t in x) | |
| if not self.enabled or torch.is_grad_enabled() or short: | |
| self.eager_calls += 1 | |
| return self.forward(*args, **kwargs) | |
| inputs = (args, kwargs) | |
| key = (_signature(inputs), torch.is_autocast_enabled("cuda"), | |
| torch.get_autocast_dtype("cuda")) | |
| entry = self.entries.get(key) | |
| if entry is None: | |
| if len(self.entries) >= self.max_entries: | |
| torch.cuda.synchronize() | |
| self.entries.popitem(last=False) | |
| static = _map_tensors(inputs, lambda t: t.detach().clone().contiguous()) | |
| static_args, static_kwargs = static | |
| stream = torch.cuda.Stream() | |
| stream.wait_stream(torch.cuda.current_stream()) | |
| with torch.cuda.stream(stream): | |
| for _ in range(3): | |
| reference = self.forward(*static_args, **static_kwargs) | |
| torch.cuda.current_stream().wait_stream(stream) | |
| torch.cuda.synchronize() | |
| graph = torch.cuda.CUDAGraph() | |
| with torch.cuda.graph(graph, stream=stream): | |
| output = self.forward(*static_args, **static_kwargs) | |
| graph.replay() | |
| torch.testing.assert_close(output, reference, rtol=1e-5, atol=1e-5) | |
| entry = (graph, static, output) | |
| self.entries[key] = entry | |
| self.captures += 1 | |
| else: | |
| self.entries.move_to_end(key) | |
| _copy_inputs(entry[1], inputs) | |
| graph, _, output = entry | |
| graph.replay() | |
| self.replays += 1 | |
| return _map_tensors(output, lambda t: t.clone()) | |
| def close(self): | |
| if self.entries: | |
| torch.cuda.synchronize() | |
| self.entries.clear() | |
| def status(self): | |
| return {"enabled": self.enabled, "cached_shapes": len(self.entries), | |
| "captures": self.captures, "replays": self.replays, | |
| "eager_calls": self.eager_calls} | |
| class Runtime: | |
| """One active generation session; serialize calls or use a separate runtime. | |
| Recover consumes only NEW raw frames in their original order. Passing a | |
| cumulative frame history repeatedly would advance the recovery twice. | |
| Positive turn_rate turns toward +X from initial +Z (right in Y-up space). | |
| A control/text change enters the noisy window immediately; its committed | |
| output follows the official 29-step lookahead rather than resetting motion. | |
| """ | |
| def __init__(self, model, recovery_type, metadata): | |
| self.model = model | |
| self._recovery_type = recovery_type | |
| self.metadata = metadata | |
| self._schedule = copy.deepcopy(model.schedule_config) | |
| self._forward = model.model.forward | |
| self._graph = None | |
| self._ready = False | |
| self._steps = self._frames = self._recovered = 0 | |
| self._action = None | |
| self._prompt = None | |
| self._history = None | |
| def add_prompt_features(self, bank): | |
| """Copy genuine encoded features from CPU storage into this process. | |
| Call this inside the generation GPU worker after reading the session's | |
| NPZ. A cache changed in a different ZeroGPU fork is not shared here. | |
| This only adds conditioning; it does not reset stream/recovery state. | |
| """ | |
| if not isinstance(bank, dict): | |
| raise TypeError("Prompt features must be a dict mapping text to CPU arrays") | |
| validated = {} | |
| for prompt, value in bank.items(): | |
| if not isinstance(prompt, str): | |
| raise TypeError("Prompt keys must be strings") | |
| if isinstance(value, torch.Tensor): | |
| if value.device.type != "cpu": | |
| raise ValueError("Prompt feature tensors must be on CPU") | |
| feature = value.detach().to(dtype=torch.float32).clone().contiguous() | |
| else: | |
| feature = torch.from_numpy(np.array(value, dtype=np.float32, copy=True)) | |
| if feature.ndim != 2 or feature.shape[1] != 4096 or not 0 < len(feature) <= 512: | |
| raise ValueError(f"Invalid T5 feature shape for {prompt!r}: {tuple(feature.shape)}") | |
| if not torch.isfinite(feature).all(): | |
| raise ValueError(f"Non-finite T5 features for {prompt!r}") | |
| validated[prompt] = feature | |
| self.model.text_module.text_cache.update(validated) | |
| return len(validated) | |
| def start(self, seed=0, history=120, use_graph=True): | |
| """Reset only at explicit session start, inside a real GPU allocation.""" | |
| history = int(history) | |
| if history not in (120, 150): | |
| raise ValueError("history must be 120 or 150 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._graph = _GraphForward(self._forward, enabled=use_graph, | |
| min_tokens=history * self.model.spatial_size) | |
| self.model.model.forward = self._graph | |
| self._recovery = self._recovery_type(smoothing_alpha=1.0, fps=FPS) | |
| self._steps = self._frames = self._recovered = 0 | |
| self._action = None | |
| self._prompt = None | |
| self._history = history | |
| self._ready = True | |
| return self | |
| def step(self, action: str, speed: float, turn_rate: float): | |
| """Backward-compatible walk/jog/run interface.""" | |
| if action not in PROMPTS: | |
| raise ValueError(f"action must be one of {tuple(PROMPTS)}") | |
| frame = self.step_text(PROMPTS[action], speed, turn_rate) | |
| self._action = action | |
| return frame | |
| def step_text(self, prompt: str, speed: float, turn_rate: float, strafe: float = 0.): | |
| """Push encoded text and local velocity; return a completed (138,) or None. | |
| Speed/strafe are local +Z/+X metres per second. Positive turn_rate is | |
| degrees per second rotating +Z toward +X; no heading is inferred from | |
| the text or pelvis orientation. Text changes preserve motion history. | |
| """ | |
| speed, turn_rate, strafe = float(speed), float(turn_rate), float(strafe) | |
| if not all(math.isfinite(v) for v in (speed, turn_rate, strafe)): | |
| raise ValueError("Speed, turn rate, and strafe must be finite") | |
| return self.step_native(prompt, [math.radians(turn_rate) / FPS, | |
| strafe / FPS, speed / FPS]) | |
| def step_native(self, prompt: str, row): | |
| """Push native [yaw radians/frame, local X/frame, local Z/frame]. | |
| The first row is zeroed because frame zero establishes the initial | |
| state in MEI138. Every subsequent native row is used unchanged. | |
| """ | |
| if not self._ready: | |
| raise RuntimeError("Call start inside the GPU allocation first") | |
| if not isinstance(prompt, str) or prompt not in self.model.text_module.text_cache: | |
| raise ValueError("Prompt has no encoded features; call add_prompt_features first") | |
| row = np.asarray(row, dtype=np.float32) | |
| if row.shape != (3,) or not np.isfinite(row).all(): | |
| raise ValueError("Native root control must be a finite (3,) row") | |
| # Native MEI138 increments are radians/frame, local X/frame, local Z/frame. | |
| # Frame 0 is the initial state; the official decoder ignores its deltas. | |
| if self._steps == 0: | |
| row = np.zeros(3, dtype=np.float32) | |
| inputs = {self.model.input_keys["text"]: [prompt], | |
| "position": torch.as_tensor(row, device="cuda").unsqueeze(0)} | |
| result = self.model.stream_generate_step(inputs)["generated"][0] | |
| self._steps += 1 | |
| self._action = None | |
| self._prompt = prompt | |
| if len(result) == 0: | |
| return None | |
| if tuple(result.shape) != (1, 138): | |
| raise RuntimeError(f"Expected one committed MEI138 frame, got {tuple(result.shape)}") | |
| frame = result[0].float().cpu().numpy().copy() | |
| if not np.isfinite(frame).all(): | |
| raise RuntimeError("Model produced a non-finite motion frame") | |
| self._frames += 1 | |
| return frame | |
| def recover(self, raw_frames): | |
| """Consume a NEW chunk and preserve world heading/XZ between chunks.""" | |
| if not self._ready: | |
| raise RuntimeError("Call start before recovering frames") | |
| frames = np.asarray(raw_frames, dtype=np.float64) | |
| if frames.size == 0: | |
| return np.empty((0, 22, 3), dtype=np.float32) | |
| if frames.ndim == 1: | |
| frames = frames[None, :] | |
| if frames.ndim != 2 or frames.shape[1] != 138 or not np.isfinite(frames).all(): | |
| raise ValueError("Recovery requires finite (T,138) NEW frames") | |
| joints = np.stack([self._recovery.process_frame(frame) for frame in frames]) | |
| self._recovered += len(frames) | |
| return joints.astype(np.float32, copy=False) | |
| def render_rotations(self, raw_frame): | |
| """World root and local body rotations for the just-recovered frame.""" | |
| from visualization.tools.rotations import rotation_6d_to_matrix | |
| if not self._ready or self._recovered != self._frames: | |
| raise RuntimeError('Recover the current completed frame before rendering it.') | |
| frame = np.asarray(raw_frame, dtype=np.float64) | |
| c, s = np.cos(self._recovery.heading), np.sin(self._recovery.heading) | |
| heading = np.array([[c, 0., s], [0., 1., 0.], [-s, 0., c]]) | |
| root = heading @ rotation_6d_to_matrix(frame[3:9][None])[0] | |
| body = rotation_6d_to_matrix(frame[12:138].reshape(21, 6)) | |
| return np.concatenate([root[None], body]).astype(np.float32) | |
| def status(self): | |
| return {"ready": self._ready, "steps": self._steps, "frames": self._frames, | |
| "recovered_frames": self._recovered, "action": self._action, | |
| "prompt": self._prompt, | |
| "history": self._history, "fps": FPS, "cfg_scale": 4.0, | |
| "startup_lookahead_frames": 29, | |
| "pending_condition_frames": self._steps - self._frames, | |
| "graph": self._graph.status() if self._graph else | |
| {"enabled": False, "cached_shapes": 0, "captures": 0, | |
| "replays": 0, "eager_calls": 0}} | |
| def close(self): | |
| """Release session graphs before leaving its GPU allocation.""" | |
| if self._graph is not None: | |
| self._graph.close() | |
| self.model.model.forward = self._forward | |
| self._ready = False | |
| def load_model(repo_root, checkpoint_path, features_path): | |
| """Load the official 300k EMA on CPU, then virtualize with .to('cuda').""" | |
| checkpoint_path = Path(checkpoint_path).resolve() | |
| features_path = Path(features_path).resolve() | |
| bank = _feature_bank(features_path) | |
| with _repository(repo_root) as root: | |
| source = importlib.import_module("models.diffusion_forcing_position_wan") | |
| if Path(source.__file__).resolve() != root / "models/diffusion_forcing_position_wan.py": | |
| raise RuntimeError("A different models package is already imported; use a clean process") | |
| from visualization.MEI138.recovery import StreamJointRecovery | |
| config = yaml.safe_load((root / "configs/df_seed_138_path.yaml").read_text()) | |
| params = copy.deepcopy(config["model"]["params"]) | |
| params.update(mean_path=None, std_path=None, loss_W=None) | |
| params["cfg_config"] = {"text_scale": 4.0, "null_scale": -3.0} | |
| original_text = source.T5TextCrossModule | |
| class CachedText(original_text): | |
| # Reuse the official encode/null/stream/trim methods and masks. | |
| # Session prompts are added as genuine CPU features before use. | |
| def __init__(self, **kwargs): | |
| torch.nn.Module.__init__(self) | |
| self.len, self.dim = int(kwargs.get("len", 512)), int(kwargs.get("dim", 4096)) | |
| self.cross_attn_norm = True | |
| self.cross_rope = bool(kwargs.get("cross_rope", False)) | |
| self.drop_out = 0.0 | |
| self.input_keys = kwargs.get("input_keys", {"text": "text", "text_end": "text_end"}) | |
| self.cache_encoded = True | |
| self.text_cache = dict(bank) | |
| # Constructor substitution is limited to cached T5 conditioning, restored | |
| # immediately. The actual motion model is instantiated from public code. | |
| source.T5TextCrossModule = CachedText | |
| try: | |
| model = source.DiffForcingPositionWanModel(**params).eval() | |
| finally: | |
| source.T5TextCrossModule = original_text | |
| # The pinned official release also contains OmegaConf training metadata. | |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False, mmap=True) | |
| if int(checkpoint.get("global_step", -1)) != 300000: | |
| raise ValueError("This demo requires the verified SEED Path 300000-step checkpoint") | |
| state = checkpoint["state_dict"] | |
| if "loss_matrix" in state and model.loss_matrix is None: | |
| delattr(model, "loss_matrix") | |
| model.register_buffer("loss_matrix", torch.empty_like(state["loss_matrix"])) | |
| model.load_state_dict(state, strict=True) | |
| ema = checkpoint["ema_state"]["shadow_params"] | |
| parameters = list(model.named_parameters()) | |
| if len(parameters) != len(ema): | |
| raise ValueError("EMA parameter count differs from the official model") | |
| with torch.no_grad(): | |
| for (name, parameter), value in zip(parameters, ema): | |
| if parameter.shape != value.shape: | |
| raise ValueError(f"EMA parameter shape mismatch for {name}") | |
| parameter.copy_(value) | |
| for name in ("mean", "std", "position_mean", "position_std"): | |
| value = getattr(model, name) | |
| if not torch.isfinite(value).all() or ("std" in name and not (value > 0).all()): | |
| raise ValueError(f"Invalid checkpoint normalization: {name}") | |
| metadata = {"checkpoint": str(checkpoint_path), "step": 300000, "weights": "EMA", | |
| "strict_load": True, "cfg_scale": 4.0, "source": str(root), | |
| "features": str(features_path), "ema_parameters": len(ema), | |
| "precision": "FP32 parameters/latent; unchanged internal BF16 SDPA"} | |
| del checkpoint, state, ema | |
| model.requires_grad_(False) | |
| model.to("cuda") | |
| return Runtime(model, StreamJointRecovery, metadata) | |