Spaces:
Running on Zero
Running on Zero
Download fdanyone_app.py from rerun/4danyone-rerun: direct link, hf CLI and curl.
- Browser
- Download file 45.9 kB
-
https://huggingface.co/spaces/rerun/4danyone-rerun/resolve/4c983fd1a9ca8bfaeec577c88d117b66c3e3ae59/fdanyone_app.py
- Command line
-
hf download hf://spaces/rerun/4danyone-rerun@4c983fd1a9ca8bfaeec577c88d117b66c3e3ae59/fdanyone_app.py
-
curl -L -o fdanyone_app.py https://huggingface.co/spaces/rerun/4danyone-rerun/resolve/4c983fd1a9ca8bfaeec577c88d117b66c3e3ae59/fdanyone_app.py
45.9 kB
| """4DAnyone on ZeroGPU: one monocular clip becomes six synchronized novel views. | |
| The run is a four-link Gradio event chain so that every phase is visible while | |
| it happens: | |
| begin -> prepare_cpu -> run_gpu -> publish_cpu | |
| ``prepare_run`` and ``generate_run`` are single blocking calls that report | |
| progress through synchronous hooks, so the GPU link runs its pipeline calls on a | |
| worker thread and the decorated generator yields the recording's bytes as the | |
| hooks fill it. That is the only mechanism that streams: yielding from inside a | |
| hook is impossible, and yielding only after the call returns would show nothing | |
| for minutes. The hooks write through an explicit ``RecordingStream``, which is | |
| safe to use from any thread, so no thread-local recording is involved. | |
| Everything else about the shape follows from one fact: ``@spaces.GPU`` runs its | |
| callback in a *forked child process* — one worker per decorated function, kept | |
| alive and reused between requests. A link therefore shares nothing with the next | |
| except its arguments, which are pickled into the worker. Two consequences run | |
| through this module: | |
| * A ``RecordingStream`` cannot cross the fork. The SDK refuses to flush one | |
| whose pid has changed ("Fork detected during flush"), so every link opens its | |
| own stream keyed by the run's token, and the viewer merges same-token streams | |
| into one recording. Blueprints ride the application id, so a layout an early | |
| link sent still governs a later link's data. This costs the viewer nothing, | |
| because ``BinaryStream.read`` already returns a complete RRD document, magic | |
| bytes and manifest included: what a run sends has always been a concatenation | |
| of them, and a stream per link only changes how many carry the same store id. | |
| * Run state cannot cross it either. The chain carries a picklable ``RunSpec`` | |
| rather than a key into a process-local table, and the two GPU phases share one | |
| allocation because ``PreparedRun`` holds decoded frames and a live completion | |
| barrier, neither of which can be pickled from one worker to another. | |
| """ | |
| from __future__ import annotations | |
| import colorsys | |
| import logging | |
| import os | |
| import queue | |
| import threading | |
| import time | |
| import uuid | |
| from collections.abc import Callable, Iterator | |
| import dataclasses | |
| from dataclasses import dataclass, field | |
| from fractions import Fraction | |
| from pathlib import Path | |
| from typing import Any, TypeAlias, TypeVar | |
| import gradio as gr | |
| import numpy as np | |
| import rerun as rr | |
| import rerun.blueprint as rrb | |
| import spaces | |
| import torch | |
| from gradio_rerun import Rerun | |
| from jaxtyping import Float, UInt8 | |
| import fdanyone.motion.body as body_module | |
| from fdanyone.config import INFERENCE, MODES, ModeSettings | |
| from fdanyone.errors import FourDAnyoneError | |
| from fdanyone.model.inference import _tensor_frames | |
| from fdanyone.model.tiny_decoder import decode_tiny_target_video, load_tiny_wan_decoder | |
| from fdanyone.motion.body import BodyMotion, load_body_motion | |
| from fdanyone.pipeline import PreparedRun, generate_run, prepare_run, release_run | |
| from fdanyone.runs import camera_records, read_cameras | |
| from fdanyone.video import choose_canonical_fps | |
| from fdanyone.viz import FRAME_TIMELINE, TIME_TIMELINE | |
| LOGGER: logging.Logger = logging.getLogger("fdanyone.app") | |
| APPLICATION_ID: str = "4danyone-rerun-v2" | |
| """Versioned per blueprint change: the viewer persists blueprints by app id, | |
| so a stale layout from an earlier deploy would otherwise shadow a new one.""" | |
| # --------------------------------------------------------------------------- | |
| # Fixed policy | |
| # --------------------------------------------------------------------------- | |
| MODEL_DIR: Path = Path(os.environ.get("FDANYONE_MODEL_DIR", "models")).expanduser().resolve() | |
| """Root holding the 4DAnyone, GVHMR, BiRefNet, SMPL-X, and turbo assets.""" | |
| DATA_DIR: Path = Path(os.environ.get("FDANYONE_DATA_DIR", "data")).expanduser().resolve() | |
| """Root holding the GVHMR checkout and one scratch tree per run.""" | |
| GVHMR_ROOT: Path = DATA_DIR / "GVHMR" | |
| """GVHMR source checkout that ``download_assets.py`` clones at boot.""" | |
| PROMPT_EMBEDDING: Path = MODEL_DIR / "4danyone" / "prompt_embedding.safetensors" | |
| """Exported prompt context; supplying it keeps the 11 GB T5 encoder out.""" | |
| TINY_DECODER_CHECKPOINT: Path = MODEL_DIR / "fps-assets" / "taehv" / "taew2_2.pth" | |
| """TAEW2.2 weights, shared by the pipeline's decode and this app's previews.""" | |
| EXAMPLE_VIDEO: Path = Path(__file__).parent / "examples" / "jump-rope.mp4" | |
| """The only bundled example, kept small enough to live in the Space repository.""" | |
| SETTINGS: ModeSettings = dataclasses.replace( | |
| MODES["turbo"], nvdec_skeletons=False, regional_compile=False, fp8_w8a8=False | |
| ) | |
| """Turbo policy minus three ZeroGPU incompatibilities: no NVCUVID on the GPU | |
| slices (skeletons decode with PyAV); no torch.compile (dynamo collides with | |
| the ``spaces`` package's patched torch.cuda internals; AOT is the supported | |
| route); no FP8 W8A8 (torchao's float8 path trips an NVML assert in torch | |
| 2.12's allocator under the slice's restricted NVML). The DiT runs bf16.""" | |
| VIEWS: int = 6 | |
| """Novel views generated per run. Six or fewer skips the RCP proposal stage.""" | |
| LAYER_PITCHES: list[int] = [15] | |
| """A single camera ring, pitched fifteen degrees above the subject.""" | |
| DIFFUSION_TIMELINE: str = "diffusion_step" | |
| """Sequence timeline carrying one point per denoising step.""" | |
| GPU_DURATION: int = 420 | |
| """Seconds requested for the one allocation the motion and generation phases share. | |
| Padded, like the two allocations it replaces (180 + 540) were, until a real | |
| ZeroGPU run gives honest timings. Everything the run can do on the CPU — the | |
| input probe, the source transcode, and the result transcode — sits outside it.""" | |
| LATENT_PREVIEW_STRIDE: int = 4 | |
| """Every fourth latent frame is decoded for preview, giving eight per view.""" | |
| TEMPORAL_COMPRESSION: int = 4 | |
| """Pixel frames per latent frame in the Wan 2.2 VAE.""" | |
| GOP_SIZE: int = 25 | |
| """Frames between forced keyframes in every stream this app logs. | |
| About one second at the 24-30 FPS the app canonicalizes to, which is what a | |
| browser needs to resync its decoder after a scrub or a reset.""" | |
| BOX_COLOR: tuple[int, int, int] = (116, 192, 252) | |
| KEYPOINT_COLOR: tuple[int, int, int] = (248, 129, 81) | |
| JOINT_COLOR: tuple[int, int, int] = (255, 212, 59) | |
| BONE_COLOR: tuple[int, int, int] = (116, 192, 252) | |
| # Body evaluation resolves SMPL-X relative to the vendored package's repository | |
| # root, which is not where a Space keeps its models. An absolute root wins the | |
| # join that builds each candidate path. | |
| body_module.SMPLX_MODEL_ROOTS = (MODEL_DIR,) | |
| # --------------------------------------------------------------------------- | |
| # Pure helpers | |
| # --------------------------------------------------------------------------- | |
| RgbFrame: TypeAlias = UInt8[np.ndarray, "height width 3"] | |
| """One decoded preview frame, ready for ``rr.Image``.""" | |
| class ClipInfo: | |
| """What a cheap container probe can say about a candidate input clip.""" | |
| fps: Fraction | |
| """Canonical frame rate the pipeline will resample the input onto.""" | |
| duration_seconds: float | |
| """Decodable duration of the video stream.""" | |
| def required_seconds(self) -> float: | |
| """Seconds of video the frozen 121-frame contract needs after the start.""" | |
| return float(Fraction(INFERENCE.num_frames - 1, 1) / self.fps) | |
| def probe_clip(video_path: Path, start_time: float) -> ClipInfo: | |
| """Reject an input that cannot yield 121 canonical frames, before any GPU work. | |
| The pipeline repeats this check exactly during its real decode; doing it | |
| here turns a mid-run failure into an immediate, actionable message. | |
| """ | |
| import av | |
| from fdanyone.video import _stream_rate | |
| if not video_path.is_file(): | |
| raise FourDAnyoneError(f"Input video does not exist: {video_path}") | |
| if not (start_time >= 0.0): | |
| raise FourDAnyoneError(f"Start time must be zero or positive, got {start_time}.") | |
| with av.open(str(video_path), mode="r") as container: | |
| if not container.streams.video: | |
| raise FourDAnyoneError(f"Input has no video stream: {video_path.name}") | |
| stream = container.streams.video[0] | |
| fps: Fraction = choose_canonical_fps(_stream_rate(stream)) | |
| raw_duration: int | None = stream.duration or container.duration | |
| if raw_duration is None: | |
| raise FourDAnyoneError(f"Input reports no duration: {video_path.name}") | |
| time_base: Fraction = Fraction(stream.time_base) if stream.duration else Fraction(1, 1_000_000) | |
| duration: float = float(raw_duration * time_base) | |
| info: ClipInfo = ClipInfo(fps=fps, duration_seconds=duration) | |
| if duration < start_time + info.required_seconds: | |
| raise FourDAnyoneError( | |
| f"{video_path.name} is {duration:.2f}s long, but {INFERENCE.num_frames} frames at " | |
| f"{float(fps):.3f} FPS from start_time={start_time:.2f}s need " | |
| f"{start_time + info.required_seconds:.2f}s. Pick an earlier start time or a longer clip." | |
| ) | |
| return info | |
| def preview_slice_plan( | |
| num_latent_frames: int, stride: int = LATENT_PREVIEW_STRIDE | |
| ) -> tuple[tuple[int, int], ...]: | |
| """Pair each previewed latent frame with the source frame it decodes to. | |
| The Wan 2.2 VAE compresses four pixel frames into one latent frame, so | |
| latent frame ``i`` is pixel frame ``4 * i`` of the generated 121-frame video. | |
| """ | |
| if num_latent_frames <= 0 or stride <= 0: | |
| raise ValueError(f"num_latent_frames and stride must be positive, got {num_latent_frames}, {stride}.") | |
| return tuple( | |
| (latent_index, latent_index * TEMPORAL_COMPRESSION) | |
| for latent_index in range(0, num_latent_frames, stride) | |
| ) | |
| def decode_preview_frame( | |
| decoder: torch.nn.Module, | |
| latent_slice: Float[torch.Tensor, "1 48 1 latent_h latent_w"], | |
| ) -> RgbFrame: | |
| """Decode one latent frame to the uint8 raster the pipeline would write. | |
| ``_tensor_frames`` is the pipeline's own float-to-uint8 truncation, reused | |
| so a preview and the final MP4 disagree only through the tiny decoder. | |
| """ | |
| video: Float[torch.Tensor, "1 3 frames height width"] = decode_tiny_target_video( | |
| decoder, latent_slice.to(dtype=torch.float16) | |
| ) | |
| return next(iter(_tensor_frames(video[0]))) | |
| # --------------------------------------------------------------------------- | |
| # Model loading at import | |
| # --------------------------------------------------------------------------- | |
| def _load_preview_decoder() -> torch.nn.Module | None: | |
| """Put the tiny decoder on CUDA once, at import, as ZeroGPU expects. | |
| ``FDANYONE_SKIP_LOAD=1`` exists only so the Blocks-construction smoke test | |
| can run on a machine whose GPU another process owns. | |
| """ | |
| if os.environ.get("FDANYONE_SKIP_LOAD") == "1": | |
| LOGGER.warning("FDANYONE_SKIP_LOAD=1: the preview decoder is not loaded.") | |
| return None | |
| return load_tiny_wan_decoder(TINY_DECODER_CHECKPOINT).to("cuda") | |
| PREVIEW_DECODER: torch.nn.Module | None = _load_preview_decoder() | |
| # --------------------------------------------------------------------------- | |
| # Rerun logging | |
| # --------------------------------------------------------------------------- | |
| def _set_frame(recording: rr.RecordingStream, index: int, fps: Fraction) -> None: | |
| """Stamp the next log calls on the shared frame and seconds timelines.""" | |
| recording.set_time(FRAME_TIMELINE, sequence=index) | |
| recording.set_time(TIME_TIMELINE, duration=float(Fraction(index, 1) / fps)) | |
| def _status(recording: rr.RecordingStream, message: str) -> None: | |
| """Append one line to the recording's text log.""" | |
| recording.log("log", rr.TextLog(message, level="INFO")) | |
| def _log_video_stream(recording: rr.RecordingStream, entity: str, video: Path) -> None: | |
| """Stream an MP4's encoded samples onto the duration timeline. | |
| ``VideoStream`` beats ``AssetVideo`` here: samples decode as they arrive | |
| instead of after the whole file, which is what a streamed run needs. | |
| Every clip goes through FFmpeg on the way in, because the web viewer decodes | |
| through WebCodecs and a browser is far pickier than the native decoder. | |
| B-frames are the known hazard (rerun#10090), and both the pipeline's own | |
| outputs and a typical phone upload carry them (``has_b_frames=2``). The | |
| reader already drops them on its own: it detects a B-framed source and | |
| re-encodes it before emitting samples. What it leaves behind is a whole clip | |
| as one 121-frame GOP — a single keyframe, at sample zero. That is the second | |
| half of the same failure. Every decoder reset a browser makes, on a scrub, a | |
| throttled tab, or seven streams contending for one hardware decoder, then has | |
| nowhere to resync short of the very beginning, and WebCodecs answers with | |
| "a key frame is required after configure() or flush()". | |
| ``GOP_SIZE`` buys a recovery point every second instead: the transcode emits | |
| a real IDR with its own SPS/PPS at each boundary, so a reset costs at most a | |
| second of video. Naming the codec alongside it makes the re-encode | |
| unconditional, so an upload in any codec the reader accepts lands as the same | |
| browser-safe H.264 the generated views do, rather than riding on the reader's | |
| own judgement of what needs fixing. | |
| It costs FFmpeg on the CPU at log time — about 2.9 s per 704x1280 generated | |
| view and 13.5 s for a 4K upload — and stays on the CPU (``try_gpu`` off) | |
| because logging runs outside the ZeroGPU allocation the pipeline holds. | |
| """ | |
| from rerun.components import VideoCodec | |
| from rerun.experimental import Mp4Reader, Mp4TranscodeOptions, send_chunks | |
| transcode: Mp4TranscodeOptions = Mp4TranscodeOptions(output_codec=VideoCodec.H264, gop_size=GOP_SIZE) | |
| reader = Mp4Reader( | |
| video, | |
| timeline_name=TIME_TIMELINE, | |
| timeline_type="duration", | |
| entity_path=entity, | |
| transcode=transcode, | |
| ) | |
| send_chunks(list(reader.stream()), recording=recording) | |
| def _log_body(recording: rr.RecordingStream, body: BodyMotion, fps: Fraction) -> None: | |
| """Log the posed SMPL-X joints and skeleton over the whole clip.""" | |
| bones: tuple[tuple[int, int], ...] = tuple( | |
| (index, parent) for index, parent in enumerate(body.parents) if parent >= 0 | |
| ) | |
| recording.log("world/body/joints", rr.Points3D.from_fields(radii=0.015, colors=JOINT_COLOR), static=True) | |
| recording.log( | |
| "world/body/skeleton", rr.LineStrips3D.from_fields(radii=0.008, colors=BONE_COLOR), static=True | |
| ) | |
| for index in range(min(len(body.joints), INFERENCE.num_frames)): | |
| _set_frame(recording, index, fps) | |
| joints: Float[np.ndarray, "joints 3"] = body.joints[index] | |
| recording.log("world/body/joints", rr.Points3D.from_fields(positions=joints)) | |
| recording.log( | |
| "world/body/skeleton", | |
| rr.LineStrips3D.from_fields( | |
| strips=np.stack([joints[[child, parent]] for child, parent in bones], axis=0) | |
| ), | |
| ) | |
| def _log_result(recording: rr.RecordingStream, result_dir: Path) -> None: | |
| """Place the camera rig in the canonical world and attach each output video.""" | |
| for record in camera_records(read_cameras(result_dir)): | |
| camera_id: int = int(record["camera_id"]) | |
| entity: str = f"world/cameras/dense/{camera_id:02d}" | |
| # Hue by yaw, so a frustum in the 3D view and its grid pane match. | |
| red, green, blue = colorsys.hsv_to_rgb((float(record["yaw"]) % 360.0) / 360.0, 0.65, 1.0) | |
| camera_to_world: Float[np.ndarray, "4 4"] = np.asarray(record["camera_to_world"], dtype=np.float64) | |
| recording.log( | |
| entity, | |
| rr.Transform3D(mat3x3=camera_to_world[:3, :3], translation=camera_to_world[:3, 3]), | |
| static=True, | |
| ) | |
| recording.log( | |
| entity, | |
| rr.Pinhole( | |
| image_from_camera=np.asarray(record["K"], dtype=np.float64), | |
| resolution=(int(record["image_width"]), int(record["image_height"])), | |
| camera_xyz=rr.ViewCoordinates.RDF, | |
| image_plane_distance=0.35, | |
| color=(int(red * 255), int(green * 255), int(blue * 255)), | |
| ), | |
| static=True, | |
| ) | |
| video: Path = result_dir / str(record["video"]) | |
| _log_video_stream(recording, f"{entity}/image", video) | |
| # --------------------------------------------------------------------------- | |
| # Blueprints | |
| # --------------------------------------------------------------------------- | |
| def motion_blueprint() -> rrb.Blueprint: | |
| """Source clip beside its detections, with the progress log.""" | |
| return rrb.Blueprint( | |
| rrb.Horizontal( | |
| rrb.Spatial2DView(origin="source", contents=["source/video"], name="Source"), | |
| rrb.Spatial2DView(origin="source", name="Detections"), | |
| rrb.TextLogView(origin="log", name="Progress"), | |
| column_shares=[1.0, 1.0, 1.0], | |
| ), | |
| rrb.TimePanel(timeline=TIME_TIMELINE, play_state="following", state="collapsed"), | |
| ) | |
| def diffusion_blueprint() -> rrb.Blueprint: | |
| """A grid of per-view previews that fills in as the denoiser steps.""" | |
| return rrb.Blueprint( | |
| rrb.Horizontal( | |
| rrb.Grid( | |
| *(rrb.Spatial2DView(origin=f"views/{index:02d}", name=f"View {index:02d}") for index in range(VIEWS)), | |
| grid_columns=3, | |
| ), | |
| rrb.Vertical( | |
| rrb.Spatial2DView(origin="source", contents=["source/preview"], name="Source"), | |
| rrb.TextLogView(origin="log", name="Progress"), | |
| ), | |
| column_shares=[2.0, 1.0], | |
| ), | |
| rrb.TimePanel(timeline=DIFFUSION_TIMELINE, play_state="following", state="collapsed"), | |
| ) | |
| def result_blueprint(fps: Fraction) -> rrb.Blueprint: | |
| """The finished rig above the six generated videos, playing on a loop.""" | |
| return rrb.Blueprint( | |
| rrb.Vertical( | |
| rrb.Horizontal( | |
| rrb.Spatial3DView(origin="world", name="Canonical world"), | |
| rrb.Spatial2DView(origin="source", contents=["source/video"], name="Source"), | |
| ), | |
| rrb.Grid( | |
| *( | |
| rrb.Spatial2DView(origin=f"world/cameras/dense/{index:02d}/image", name=f"View {index:02d}") | |
| for index in range(VIEWS) | |
| ), | |
| grid_columns=3, | |
| ), | |
| row_shares=[1.0, 1.0], | |
| ), | |
| rrb.TimePanel( | |
| timeline=TIME_TIMELINE, | |
| play_state="playing", | |
| loop_mode="all", | |
| state="collapsed", | |
| ), | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Session state and streaming plumbing | |
| # --------------------------------------------------------------------------- | |
| class RunSpec: | |
| """All one link needs to rebuild a run's context, and nothing that cannot travel. | |
| This is what the chain passes between links, as a Gradio ``State``. Every | |
| field is picklable, because ZeroGPU pickles a GPU callback's arguments into | |
| its forked worker; a key into a process-local table would not survive, since | |
| a fresh worker forks from a parent that never saw the entry and a reused one | |
| still holds some earlier run's memory. | |
| """ | |
| token: str | |
| """Run identifier, and the ``recording_id`` every link's stream carries.""" | |
| video_path: Path | |
| """Uploaded source clip.""" | |
| start_time: float | |
| """Seconds into the source clip where the canonical 121 frames begin.""" | |
| seed: int | |
| """Generation seed.""" | |
| fps: Fraction | |
| """Canonical frame rate shared by the source clip and every output video.""" | |
| data_dir: Path | |
| """Fresh per-run data root; the pipeline refuses to overwrite a result.""" | |
| class Link: | |
| """One chain link's own recording, and the sink its bytes are drained from.""" | |
| recording: rr.RecordingStream | |
| """Explicit stream this link and its pipeline hooks log through.""" | |
| stream: Any | None = None | |
| """Binary sink feeding the Gradio viewer; ``None`` when the sink is a file.""" | |
| def read(self) -> bytes | None: | |
| """Bytes buffered since the last read, or ``None`` without a binary sink.""" | |
| return None if self.stream is None else self.stream.read() | |
| def open_link(token: str) -> Link: | |
| """Open this callback's own recording, in whatever process the callback runs in. | |
| Called first thing in every link, GPU ones included, where "this callback" | |
| means the forked child. A stream inherited across the fork cannot be used: | |
| the SDK compares pids on flush and raises "Fork detected during flush". A | |
| stream built here belongs to this process, and sharing ``token`` as the | |
| recording id is what makes the viewer treat all four links as one recording. | |
| ``cleanup_if_forked_child`` drops whatever streams were inherited. The SDK | |
| registers it with ``os.register_at_fork`` already, and ``rr.init`` calls it | |
| outright for the same reason, so this is belt and braces and a no-op outside | |
| a child. It is called unconditionally: guarding it on the name existing | |
| would turn a rename in the SDK into a silent return of the very bug this | |
| function exists to prevent. | |
| """ | |
| rr.cleanup_if_forked_child() | |
| recording: rr.RecordingStream = rr.RecordingStream(APPLICATION_ID, recording_id=token) | |
| return Link(recording=recording, stream=recording.binary_stream()) | |
| class Session: | |
| """One link's runtime state: what its phases mutate and never send onward.""" | |
| spec: RunSpec | |
| """The run this link was rebuilt from.""" | |
| prepared: PreparedRun | None = None | |
| """Result of the motion phase, consumed by the generation phase.""" | |
| summary: dict | None = None | |
| """Published metadata the generation phase returns, set once it succeeds.""" | |
| source_frames: dict[int, RgbFrame] = field(default_factory=dict) | |
| """Decoded source stills for the diffusion pane, keyed by canonical frame.""" | |
| events: queue.Queue[str | None] = field(default_factory=queue.Queue) | |
| """Stage labels the pipeline hooks push from the worker thread.""" | |
| class RunCancelled(FourDAnyoneError): | |
| """Raised inside a pipeline hook once the run's stop sentinel appears.""" | |
| def _stop_path(spec: RunSpec) -> Path: | |
| return spec.data_dir / "stop-requested" | |
| def _check_stop(spec: RunSpec) -> None: | |
| """Abort the pipeline from inside one of its hooks after Stop is pressed. | |
| Gradio's ``cancels`` only closes the request's generator; the pipeline keeps | |
| running — on a worker thread here, in a forked ZeroGPU child on the Space. | |
| A file under the run's own directory is the one channel that reaches both, | |
| and the hooks are the only code of ours the pipeline re-enters, so they are | |
| where the run can be unwound. | |
| """ | |
| if _stop_path(spec).exists(): | |
| raise RunCancelled("Stopped.") | |
| T = TypeVar("T") | |
| def _pump(session: Session, work: Callable[[], T]) -> Iterator[str]: | |
| """Run ``work`` on a worker thread, yielding each label its hooks queue. | |
| The generator's return value is ``work``'s result, so a caller writes | |
| ``result = yield from _pump(...)``. A failure inside the thread is re-raised | |
| here, on the request's own thread, where Gradio can report it. | |
| """ | |
| outcome: list[T] = [] | |
| failure: list[BaseException] = [] | |
| def target() -> None: | |
| try: | |
| outcome.append(work()) | |
| except BaseException as exc: # noqa: BLE001 - re-raised below, unchanged | |
| failure.append(exc) | |
| finally: | |
| session.events.put(None) | |
| worker: threading.Thread = threading.Thread(target=target, daemon=True) | |
| worker.start() | |
| while True: | |
| label: str | None = session.events.get() | |
| if label is None: | |
| break | |
| yield label | |
| worker.join() | |
| if failure: | |
| raise failure[0] | |
| return outcome[0] | |
| # --------------------------------------------------------------------------- | |
| # Pipeline hooks | |
| # --------------------------------------------------------------------------- | |
| def _motion_hook(link: Link, session: Session) -> Callable[[str, dict[str, object]], None]: | |
| """Build the ``on_motion_stage`` hook that draws GVHMR's intermediates.""" | |
| recording: rr.RecordingStream = link.recording | |
| def on_stage(stage: str, payload: dict[str, object]) -> None: | |
| _check_stop(session.spec) | |
| if stage == "bboxes": | |
| boxes: Float[torch.Tensor, "frames 4"] = payload["bbx_xyxy"] # pyrefly: ignore | |
| for index, box in enumerate(boxes.numpy()[: INFERENCE.num_frames]): | |
| _set_frame(recording, index, session.spec.fps) | |
| recording.log( | |
| "source/bboxes", | |
| rr.Boxes2D( | |
| array=box.reshape(1, 4), array_format=rr.Box2DFormat.XYXY, colors=BOX_COLOR | |
| ), | |
| ) | |
| elif stage == "keypoints_2d": | |
| keypoints: Float[torch.Tensor, "frames joints 3"] = payload["kp2d"] # pyrefly: ignore | |
| for index, frame in enumerate(keypoints.numpy()[: INFERENCE.num_frames]): | |
| _set_frame(recording, index, session.spec.fps) | |
| recording.log( | |
| "source/keypoints", | |
| rr.Points2D(positions=frame[:, :2], radii=6.0, colors=KEYPOINT_COLOR), | |
| ) | |
| _status(recording, f"motion: {stage}") | |
| session.events.put(stage) | |
| return on_stage | |
| def _denoise_hook(link: Link, session: Session) -> Callable[[int, tuple[int, ...], torch.Tensor], None]: | |
| """Build the ``on_denoise_step`` hook that decodes and logs previews.""" | |
| recording: rr.RecordingStream = link.recording | |
| def on_step( | |
| step_index: int, | |
| view_indices: tuple[int, ...], | |
| x0_hat: Float[torch.Tensor, "views 48 latent_t latent_h latent_w"], | |
| ) -> None: | |
| _check_stop(session.spec) | |
| if PREVIEW_DECODER is None: | |
| return | |
| plan: tuple[tuple[int, int], ...] = preview_slice_plan(x0_hat.shape[2]) | |
| recording.set_time(DIFFUSION_TIMELINE, sequence=step_index) | |
| with torch.inference_mode(): | |
| for latent_index, source_frame in plan: | |
| _set_frame(recording, source_frame, session.spec.fps) | |
| # The source clip is indexed on `frame` alone, so repeat its | |
| # reference here or the source pane stays empty while the active | |
| # timeline is `diffusion_step`. | |
| preview_frame: RgbFrame | None = session.source_frames.get(source_frame) | |
| if preview_frame is not None: | |
| recording.log("source/preview", rr.Image(preview_frame).compress(jpeg_quality=85)) | |
| for group_index, view in enumerate(view_indices): | |
| frame: RgbFrame = decode_preview_frame( | |
| PREVIEW_DECODER, | |
| x0_hat[group_index : group_index + 1, :, latent_index : latent_index + 1], | |
| ) | |
| recording.log(f"views/{view:02d}", rr.Image(frame).compress(jpeg_quality=85)) | |
| _status(recording, f"denoise step {step_index + 1}/{SETTINGS.num_inference_steps}") | |
| session.events.put(f"step {step_index + 1}") | |
| return on_step | |
| # --------------------------------------------------------------------------- | |
| # Sink-agnostic phases — one implementation for the native CLI and Gradio | |
| # --------------------------------------------------------------------------- | |
| def new_spec(video_path: Path, start_time: float, seed: int) -> RunSpec: | |
| """Probe the clip and mint the run this whole chain will be rebuilt from.""" | |
| info: ClipInfo = probe_clip(video_path, start_time) | |
| token: str = uuid.uuid4().hex | |
| return RunSpec( | |
| token=token, | |
| video_path=video_path, | |
| start_time=start_time, | |
| seed=seed, | |
| fps=info.fps, | |
| data_dir=DATA_DIR / "runs" / token, | |
| ) | |
| def begin_phase(link: Link, spec: RunSpec) -> None: | |
| """Send the first blueprint and the run's opening status line. | |
| The caller binds the recording's sink (binary stream, file, or a live | |
| viewer) BEFORE calling this, so the blueprint and every later row reach it. | |
| """ | |
| link.recording.send_blueprint(motion_blueprint(), make_active=True) | |
| link.recording.log("world", rr.ViewCoordinates.RUB, static=True) | |
| _status( | |
| link.recording, | |
| f"begin: {spec.video_path.name} at {float(spec.fps):.3f} FPS from {spec.start_time:.2f}s", | |
| ) | |
| def _decode_source_stills(spec: RunSpec) -> dict[int, RgbFrame]: | |
| """Decode the preview-plan source frames once; a few stills, CPU only.""" | |
| import av | |
| wanted: set[int] = {source for _, source in preview_slice_plan(31)} | |
| stills: dict[int, RgbFrame] = {} | |
| with av.open(str(spec.video_path)) as container: | |
| stream = container.streams.video[0] | |
| offset: float = spec.start_time | |
| index: int = 0 | |
| for frame in container.decode(stream): | |
| if frame.time is None or frame.time < offset: | |
| continue | |
| if index in wanted: | |
| stills[index] = frame.to_ndarray(format="rgb24") | |
| index += 1 | |
| if index > max(wanted): | |
| break | |
| return stills | |
| def _log_smplx_params(recording: rr.RecordingStream, motion_dir: Path, fps: Fraction) -> None: | |
| """Log the raw SMPL-X parameters so the recording carries the capture itself. | |
| Shape is constant per subject, so betas go static; the per-frame pose, | |
| orientation, and translation land on both timelines like everything else. | |
| """ | |
| from fdanyone.motion.result import MotionResult | |
| motion: MotionResult = MotionResult.load(motion_dir) | |
| params: dict[str, torch.Tensor] = motion.smpl_params_global | |
| betas: Float[torch.Tensor, "frames 10"] = params["betas"] | |
| recording.log("world/body/params/betas", rr.Tensor(betas[0].numpy()), static=True) | |
| num_frames: int = min(int(params["body_pose"].shape[0]), INFERENCE.num_frames) | |
| for index in range(num_frames): | |
| _set_frame(recording, index, fps) | |
| recording.log("world/body/params/body_pose", rr.Tensor(params["body_pose"][index].numpy())) | |
| recording.log("world/body/params/global_orient", rr.Tensor(params["global_orient"][index].numpy())) | |
| recording.log( | |
| "world/body/params/transl", | |
| rr.Scalars(params["transl"][index].numpy()), | |
| ) | |
| def source_phase(link: Link, spec: RunSpec) -> str: | |
| """Log the source clip on the frame timeline; CPU only.""" | |
| _log_video_stream(link.recording, "source/video", spec.video_path) | |
| _status(link.recording, "prepare: source clip logged") | |
| return "Recovering motion with GVHMR." | |
| MOTION_LABELS: dict[str, str] = { | |
| "bboxes": "Tracking the subject.", | |
| "keypoints_2d": "Estimating 2D keypoints.", | |
| "features": "Extracting motion features.", | |
| "smplx": "Fitting the SMPL-X body.", | |
| } | |
| def motion_phase(link: Link, session: Session) -> Iterator[str]: | |
| """Run GVHMR + conditioning, yielding a label as each stage lands.""" | |
| def work() -> PreparedRun: | |
| return prepare_run( | |
| settings=SETTINGS, | |
| video_path=str(session.spec.video_path), | |
| data_dir=str(session.spec.data_dir), | |
| model_dir=str(MODEL_DIR), | |
| checkpoint_path=None, | |
| mhr70_regressor_path=None, | |
| gvhmr_root=str(GVHMR_ROOT), | |
| device="cuda", | |
| start_time=session.spec.start_time, | |
| target_fps="auto", | |
| views_per_layer=VIEWS, | |
| layer_pitches=LAYER_PITCHES, | |
| start_yaw=0, | |
| yaw_span=360, | |
| views_per_group="auto", | |
| enable_rcp=True, | |
| enable_tcr=True, | |
| # ZeroGPU allocates to this process; a forked worker would lose it. | |
| inline_workers=True, | |
| on_motion_stage=_motion_hook(link, session), | |
| prompt_embedding_path=PROMPT_EMBEDDING, | |
| ) | |
| stages: Iterator[str] = _pump(session, work) | |
| while True: | |
| try: | |
| stage: str = next(stages) | |
| except StopIteration as done: | |
| session.prepared = done.value | |
| break | |
| yield MOTION_LABELS.get(stage, stage) | |
| body: BodyMotion = load_body_motion(session.prepared.motion_dir, device="cpu") | |
| _log_body(link.recording, body, session.spec.fps) | |
| _log_smplx_params(link.recording, session.prepared.motion_dir, session.spec.fps) | |
| _status(link.recording, "motion: SMPL-X body and parameters logged") | |
| yield "Generating six views." | |
| def generate_phase(link: Link, session: Session) -> Iterator[str]: | |
| """Denoise with per-step previews, leaving the result on disk for publishing. | |
| The source stills the preview pane needs are decoded here rather than in the | |
| CPU link before it, because the two links no longer share a process. | |
| """ | |
| if session.prepared is None: | |
| raise FourDAnyoneError("The motion phase did not finish.") | |
| prepared: PreparedRun = session.prepared | |
| session.source_frames = _decode_source_stills(session.spec) | |
| link.recording.send_blueprint(diffusion_blueprint(), make_active=True) | |
| yield "Generating six views." | |
| def work() -> dict: | |
| # The worker owns the scratch, so it must also settle it: a Stop closes | |
| # the outer generator, which then never sees the pipeline call end. | |
| try: | |
| return generate_run( | |
| prepared, seed=session.spec.seed, on_denoise_step=_denoise_hook(link, session) | |
| ) | |
| finally: | |
| release_run(prepared) | |
| steps: Iterator[str] = _pump(session, work) | |
| while True: | |
| try: | |
| label: str = next(steps) | |
| except StopIteration as done: | |
| session.summary = done.value | |
| break | |
| yield f"Denoising: {label} of {SETTINGS.num_inference_steps}." | |
| _status(link.recording, "generate: six views published") | |
| yield "Publishing the result." | |
| def publish_phase(link: Link, spec: RunSpec, summary: dict) -> str: | |
| """Place the finished rig and its six videos, then switch to the final layout. | |
| Everything this reads is a file on disk, so it runs outside the GPU | |
| allocation — which matters, because logging each view re-encodes it. | |
| """ | |
| link.recording.reset_time() | |
| _log_result(link.recording, Path(summary["result_dir"])) | |
| _status(link.recording, "done") | |
| link.recording.send_blueprint(result_blueprint(spec.fps), make_active=True) | |
| elapsed: float = float(summary["total_pipeline_elapsed_seconds"]) | |
| return f"Done in {elapsed:.1f}s. Six views at {float(spec.fps):.3f} FPS." | |
| # --------------------------------------------------------------------------- | |
| # Callbacks | |
| # --------------------------------------------------------------------------- | |
| STREAM_SMOKE: bool = os.environ.get("FDANYONE_STREAM_SMOKE") == "1" | |
| """GPU-free debug mode: Run streams synthetic data through the real machinery.""" | |
| SMOKE_DELAY: float = float(os.environ.get("FDANYONE_SMOKE_DELAY", "0.1")) | |
| """Seconds between smoke yields; raise it to eyeball each phase in a browser.""" | |
| SMOKE_SUMMARY: dict = {"result_dir": "", "total_pipeline_elapsed_seconds": 0.0} | |
| """Stand-in for the pipeline's published metadata; the smoke run writes no files.""" | |
| def smoke_motion_phase(link: Link, session: Session) -> Iterator[str]: | |
| """Synthetic motion phase: a moving detection over the real source clip.""" | |
| rng: np.random.Generator = np.random.default_rng(0) | |
| for index in range(0, 121, 10): | |
| _check_stop(session.spec) | |
| _set_frame(link.recording, index, session.spec.fps) | |
| x0: float = 180.0 + 2.0 * index | |
| link.recording.log( | |
| "source/bboxes", | |
| rr.Boxes2D(array=[[x0, 300.0, x0 + 300.0, 1100.0]], array_format=rr.Box2DFormat.XYXY), | |
| ) | |
| keypoints: Float[np.ndarray, "17 2"] = rng.uniform((x0, 350.0), (x0 + 300.0, 1050.0), (17, 2)) | |
| link.recording.log("source/keypoints", rr.Points2D(keypoints)) | |
| _status(link.recording, f"[smoke] motion frame {index}") | |
| time.sleep(SMOKE_DELAY) | |
| yield f"[smoke] motion frame {index}" | |
| def smoke_generate_phase(link: Link, session: Session) -> Iterator[str]: | |
| """Synthetic diffusion phase: gradient frames sharpening per denoising step.""" | |
| rng: np.random.Generator = np.random.default_rng(1) | |
| link.recording.send_blueprint(diffusion_blueprint(), make_active=True) | |
| yield "[smoke] diffusion" | |
| for step in range(4): | |
| _check_stop(session.spec) | |
| link.recording.set_time(DIFFUSION_TIMELINE, sequence=step) | |
| for view in range(VIEWS): | |
| noise: UInt8[np.ndarray, "160 88 3"] = rng.integers( | |
| 0, 256 // (step + 1), (160, 88, 3), dtype=np.uint8 | |
| ) | |
| base: UInt8[np.ndarray, "160 88 3"] = np.full((160, 88, 3), 60 * step, dtype=np.uint8) | |
| link.recording.log(f"views/{view:02d}", rr.Image(base + noise)) | |
| _status(link.recording, f"[smoke] denoise step {step + 1}/4") | |
| time.sleep(3.0 * SMOKE_DELAY) | |
| yield f"[smoke] denoise step {step + 1}/4" | |
| session.summary = SMOKE_SUMMARY | |
| yield "[smoke] publishing" | |
| def smoke_publish_phase(link: Link, spec: RunSpec) -> str: | |
| """Synthetic result phase: flat shades where the six generated videos go.""" | |
| link.recording.reset_time() | |
| for index in range(0, 121, 10): | |
| _set_frame(link.recording, index, spec.fps) | |
| for view in range(VIEWS): | |
| shade: UInt8[np.ndarray, "160 88 3"] = np.full((160, 88, 3), 40 + index, dtype=np.uint8) | |
| link.recording.log(f"world/cameras/dense/{view:02d}/image", rr.Image(shade)) | |
| _status(link.recording, "[smoke] done") | |
| link.recording.send_blueprint(result_blueprint(spec.fps), make_active=True) | |
| return "[smoke] done — final blueprint sent" | |
| def begin( | |
| video: str | None, start_time: float, seed: int | |
| ) -> Iterator[tuple[RunSpec, bytes | None, str, Any]]: | |
| """Validate the input on CPU, open a recording, and switch to the outputs.""" | |
| if video is None: | |
| raise gr.Error("Upload a video, or pick the bundled example.") | |
| try: | |
| spec: RunSpec = new_spec(Path(video), float(start_time), int(seed)) | |
| except FourDAnyoneError as exc: | |
| raise gr.Error(str(exc)) from None | |
| link: Link = open_link(spec.token) | |
| begin_phase(link, spec) | |
| yield spec, link.read(), "Preparing the source clip.", gr.Tabs(selected="outputs") | |
| def prepare_cpu(spec: RunSpec) -> Iterator[tuple[bytes | None, str]]: | |
| """Put the source clip on the frame timeline before any GPU work starts.""" | |
| link: Link = open_link(spec.token) | |
| label: str = source_phase(link, spec) | |
| yield link.read(), label | |
| def run_gpu(spec: RunSpec) -> Iterator[tuple[dict | None, bytes | None, str]]: | |
| """Recover the motion and generate the six views, inside one allocation. | |
| The two phases have to share a process. ``PreparedRun`` carries the decoded | |
| clip, the open conditioning artifacts, and a completion barrier for the | |
| skeleton renderer, so it cannot be pickled from one ZeroGPU worker into | |
| another — and every ``@spaces.GPU`` function gets a worker of its own. | |
| Everything in here runs in the forked child, including ``open_link``: this | |
| is the only process that may build the recording it logs through. | |
| """ | |
| link: Link = open_link(spec.token) | |
| session: Session = Session(spec) | |
| try: | |
| motion: Iterator[str] = ( | |
| smoke_motion_phase(link, session) if STREAM_SMOKE else motion_phase(link, session) | |
| ) | |
| for label in motion: | |
| yield None, link.read(), label | |
| generate: Iterator[str] = ( | |
| smoke_generate_phase(link, session) if STREAM_SMOKE else generate_phase(link, session) | |
| ) | |
| for label in generate: | |
| yield None, link.read(), label | |
| except FourDAnyoneError as exc: | |
| raise gr.Error(str(exc)) from None | |
| yield session.summary, link.read(), "Publishing the result." | |
| def request_stop(spec: RunSpec | None) -> str: | |
| """Write the run's stop sentinel; the pipeline unwinds at its next hook.""" | |
| if spec is None: | |
| return "Nothing is running." | |
| spec.data_dir.mkdir(parents=True, exist_ok=True) | |
| _stop_path(spec).touch() | |
| return "Stopping — the pipeline halts at its next step." | |
| def publish_cpu(spec: RunSpec, summary: dict | None) -> Iterator[tuple[bytes | None, str]]: | |
| """Attach the finished rig and its six videos, off the GPU allocation.""" | |
| if summary is None: | |
| raise gr.Error("The generation phase did not publish a result.") | |
| link: Link = open_link(spec.token) | |
| label: str = smoke_publish_phase(link, spec) if STREAM_SMOKE else publish_phase(link, spec, summary) | |
| yield link.read(), label | |
| # --------------------------------------------------------------------------- | |
| # Interface | |
| # --------------------------------------------------------------------------- | |
| DESCRIPTION: str = """ | |
| # 4DAnyone × Rerun | |
| One monocular clip of a person becomes six synchronized novel views. GVHMR | |
| recovers the SMPL-X motion, a Wan 2.2 diffusion transformer generates every | |
| view, and each phase streams into the Rerun viewer as it happens. | |
| """ | |
| SOURCE_CLIP_HEIGHT: int = 360 | |
| """Display height of the source preview. A portrait clip scaled to the column | |
| width is taller than the fold, which pushes every control below it out of sight.""" | |
| STATUS_CSS: str = """ | |
| #run-status { | |
| display: flex; | |
| flex-direction: column; | |
| justify-content: center; | |
| /* Tall enough for the progress overlay Gradio draws over a running event. */ | |
| min-height: 5rem; | |
| /* Gradio writes `overflow: auto` inline on every block. */ | |
| overflow: visible !important; | |
| padding: 0.75rem 1rem; | |
| border-radius: var(--radius-lg); | |
| background: var(--background-fill-secondary); | |
| } | |
| #run-status p { | |
| font-size: 1.15rem; | |
| line-height: 1.5; | |
| margin: 0; | |
| } | |
| """ | |
| """Banner styling for the run status, the only readout of a multi-minute run. | |
| Gradio otherwise gives a Markdown block the body font and a height its own text | |
| overflows, which turns the line into a scrolling sliver. Gradio 6 takes ``css`` | |
| on ``launch``, not on the ``Blocks`` constructor.""" | |
| def build_demo() -> gr.Blocks: | |
| """Assemble the persistent viewer, the inputs, and the run chain. | |
| Controls sit in a narrow column and the viewer fills a wide one beside it, | |
| the same split every other Rerun Space here uses. Stacking them instead | |
| pushes the viewer off the fold on a laptop, and the viewer is the demo. It | |
| stays outside the tabs for the same reason it is created once: a tab switch | |
| would tear down the component and drop the stream mid-run. | |
| The chain is joined by ``success`` rather than ``then``: each link's inputs | |
| are the previous link's outputs, so a link that failed leaves the next one | |
| nothing to work with, and running it anyway would bury the real error under | |
| a second, meaningless one. | |
| """ | |
| with gr.Blocks(title="4DAnyone × Rerun") as demo: | |
| gr.Markdown(DESCRIPTION) | |
| spec: gr.State = gr.State(None) | |
| summary: gr.State = gr.State(None) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| with gr.Tabs() as tabs: | |
| with gr.Tab("Input", id="input"): | |
| video: gr.Video = gr.Video( | |
| label="Source clip", sources=["upload"], height=SOURCE_CLIP_HEIGHT | |
| ) | |
| start_time: gr.Number = gr.Number( | |
| value=0.0, label="Start time (seconds)", minimum=0.0 | |
| ) | |
| seed: gr.Number = gr.Number(value=0, label="Seed", precision=0, minimum=0) | |
| gr.Examples(examples=[[str(EXAMPLE_VIDEO)]], inputs=[video], cache_examples=False) | |
| with gr.Tab("Outputs", id="outputs"): | |
| gr.Markdown( | |
| "The viewer beside this switches layout with the run: detections on the " | |
| "source clip, then a grid of per-step previews, then the finished rig " | |
| "playing on a loop." | |
| ) | |
| with gr.Row(): | |
| run_button: gr.Button = gr.Button("Run", variant="primary") | |
| stop_button: gr.Button = gr.Button("Stop", variant="stop") | |
| status: gr.Markdown = gr.Markdown( | |
| "Upload a clip, or pick the example, then press Run.", elem_id="run-status" | |
| ) | |
| with gr.Column(scale=3): | |
| viewer: Rerun = Rerun( | |
| streaming=True, | |
| height=760, | |
| panel_states={"time": "collapsed", "blueprint": "hidden", "selection": "hidden"}, | |
| ) | |
| started = run_button.click(begin, [video, start_time, seed], [spec, viewer, status, tabs]) | |
| prepared = started.success(prepare_cpu, spec, [viewer, status]) | |
| generated = prepared.success(run_gpu, spec, [summary, viewer, status]) | |
| published = generated.success(publish_cpu, [spec, summary], [viewer, status]) | |
| # `cancels` detaches the UI at once, but the pipeline itself only halts | |
| # when it reads the sentinel `request_stop` writes: the worker thread | |
| # (and on ZeroGPU, the forked child) never sees a Gradio cancellation. | |
| stop_button.click( | |
| request_stop, spec, status, cancels=[started, prepared, generated, published] | |
| ) | |
| return demo | |
| if __name__ == "__main__": | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") | |
| build_demo().queue(default_concurrency_limit=1).launch(ssr_mode=False, css=STATUS_CSS) | |