"""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 @contextmanager 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) @torch.inference_mode() 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 @torch.inference_mode() 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 @torch.inference_mode() 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]) @torch.inference_mode() 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)