FloodDiffusion2-Live / space /inference.py
caiyiyi1998's picture
Initial commit
9a25493
Raw History Blame Contribute Delete
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
@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)