diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..53aee21b9a286bf9d1904e9527c3f0b047f4f0ea 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +*.mp4 filter=lfs diff=lfs merge=lfs -text diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..60bf7cdb5d42ec8ed6587ef023f658aacca95f74 --- /dev/null +++ b/.gitignore @@ -0,0 +1,10 @@ +# Pixi's materialized environments; pixi.lock is the source of truth. +.pixi/ + +# Everything download_assets.py fetches at boot, plus per-run scratch. +models/ +data/ + +__pycache__/ +*.py[cod] +.pytest_cache/ diff --git a/PROVENANCE.md b/PROVENANCE.md new file mode 100644 index 0000000000000000000000000000000000000000..e2faeaed9348dd5ca89392e455a251c8ddfb49b4 --- /dev/null +++ b/PROVENANCE.md @@ -0,0 +1,36 @@ +# Provenance + +`fdanyone/` is a copy, not a fork. Regenerate it with `./sync_vendor.sh`. + +| Item | Value | +| --- | --- | +| Source repository | | +| Branch | `space-streaming` | +| Commit | `0cc334c2b260ff19b23d423ac8123d318ed602ff` | +| Synced | 2026-08-27T05:15:32Z | + +## Excluded from the copy + +- `fdanyone/nerfstudio/`, `fdanyone/freetimegs/`, `fdanyone/vendor/freetimegs/` — 3DGS + and 4DGS reconstruction, which this Space does not run. +- `fdanyone/cli.py` — the Fire shim for `scripts/`, which is not copied either. +- `__pycache__/`. + +## Patched in the copy + +- `fdanyone/assets.py`: `MODEL_FILES` drops `models_t5_umt5-xxl-enc-bf16.pth`, + `4danyone/umt5-xxl/`, and the perceptual VGG-19. `prepare_run` calls + `ensure_models`, which downloads every missing entry — inside the ZeroGPU + allocation. The Space passes `prompt_embedding_path`, so the 11 GB encoder is + never loaded, and it must never be fetched either. + +## GVHMR + +GVHMR is a git submodule of the source repository and is deliberately absent +here. `download_assets.py` clones it into the ephemeral disk at boot, at the +pinned revision below, and `fdanyone_app.py` passes that path as `gvhmr_root`. + +| Item | Value | +| --- | --- | +| Repository | | +| Revision | `6ec3ca39336c50492c0fae65fba2fb831fc7d866` | diff --git a/fdanyone/__init__.py b/fdanyone/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..523aa32c1ef902ece03f75a6bbb86caa3800275f --- /dev/null +++ b/fdanyone/__init__.py @@ -0,0 +1 @@ +"""4DAnyone inference.""" diff --git a/fdanyone/assets.py b/fdanyone/assets.py new file mode 100644 index 0000000000000000000000000000000000000000..62d799b9138673bec36fe105928248bfe10342f7 --- /dev/null +++ b/fdanyone/assets.py @@ -0,0 +1,191 @@ +"""Locate the model files used by 4DAnyone. + +Every published file is anchored by one immutable Hugging Face revision and +downloaded on demand. ``fdanyone.download`` fetches missing files; +the resolvers here only locate them. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import TYPE_CHECKING + +from fdanyone.errors import AssetError + +if TYPE_CHECKING: + from fdanyone.config import ModeSettings + +HF_REPO_ID = "AntResearch/4DAnyone" +HF_REVISION = "7850985888b56aabf09e69480b73248f1a76bcbe" + +BIREFNET_REPO_ID = "ZhengPeng7/BiRefNet" +BIREFNET_REVISION = "e2bf8e4460fc8fa32bba5ea4d94b3233d367b0e4" +BIREFNET_DIR = "birefnet" +BIREFNET_FILES = ( + "BiRefNet_config.py", + "birefnet.py", + "config.json", + "model.safetensors", +) + +CHECKPOINT = "4danyone/model.safetensors" +MHR70_REGRESSOR = "4danyone/smplx_to_goliath70.pt" +WAN_VAE = "4danyone/Wan2.2_VAE.pth" +TEXT_ENCODER = "4danyone/models_t5_umt5-xxl-enc-bf16.pth" +TOKENIZER_DIR = "4danyone/umt5-xxl" +TOKENIZER_FILES = tuple( + f"{TOKENIZER_DIR}/{name}" + for name in ("special_tokens_map.json", "spiece.model", "tokenizer.json", "tokenizer_config.json") +) + +GVHMR_CHECKPOINT = "gvhmr/gvhmr_siga24_release.ckpt" +HMR2_CHECKPOINT = "gvhmr/epoch=10-step=25000.ckpt" +VITPOSE_CHECKPOINT = "gvhmr/vitpose-h-multi-coco.pth" +YOLO_CHECKPOINT = "gvhmr/yolov8x.pt" +PERCEPTUAL_VGG19 = "perceptual/imagenet-vgg-verydeep-19-conv.safetensors" +TAEW2_2 = "fps-assets/taehv/taew2_2.pth" +TAEW2_2_URL = ( + "https://raw.githubusercontent.com/madebyollin/taehv/" + "e743234f3217ab3d1570f65642ab06596d1bd7c5/taew2_2.pth" +) +TAEW2_2_SHA256 = "d053e216ca50e2bb837bbcd79b85f0366bea00e5938025572382a773b74c559a" +TURBO_LORA = ( + "fps-assets/turbo-lora/LoRAs/Wan22-Turbo/" + "Wan22_TI2V_5B_Turbo_lora_rank_64_fp16.safetensors" +) +TURBO_LORA_REPO_ID = "Kijai/WanVideo_comfy" +TURBO_LORA_REVISION = "86c2b0442e01eeee630b48fd7efc0cd37af03252" +TURBO_LORA_REPO_FILE = ( + "LoRAs/Wan22-Turbo/Wan22_TI2V_5B_Turbo_lora_rank_64_fp16.safetensors" +) +TURBO_LORA_SHA256 = "0ace5244e3d1256f884662c261b017249796cf5b95f05d5ed93cc02a478967b8" + +SMPLX_MODEL = "body_models/smplx/SMPLX_NEUTRAL.npz" + +# Space patch (sync_vendor.sh): the exported prompt embedding replaces the +# UMT5-XXL encoder and its tokenizer, and reconstruction never runs here. +MODEL_FILES = ( + CHECKPOINT, + MHR70_REGRESSOR, + WAN_VAE, + GVHMR_CHECKPOINT, + HMR2_CHECKPOINT, + VITPOSE_CHECKPOINT, + YOLO_CHECKPOINT, +) + +EXAMPLE_FILES = ( + "data/source/pexels/10331522-uhd_2160_4096_25fps.mp4", + "data/source/pexels/2785536-uhd_2160_3840_25fps.mp4", + "data/source/pexels/5435720-uhd_2160_4096_25fps.mp4", + "data/source/pexels/5885633-hd_1080_1920_25fps.mp4", + "data/source/pexels/5999210-uhd_2160_4096_25fps.mp4", + "data/source/pexels/6980035-uhd_2160_4096_30fps.mp4", + "data/source/pexels/7080903-hd_1080_1920_30fps.mp4", + "data/source/pexels/7480858-uhd_2160_3840_25fps.mp4", +) + +# Upstream GVHMR resolves its model files relative to its own checkout, so the +# install commands link each downloaded file to the location GVHMR expects. +GVHMR_LINKS = ( + (GVHMR_CHECKPOINT, "inputs/checkpoints/gvhmr/gvhmr_siga24_release.ckpt"), + (HMR2_CHECKPOINT, "inputs/checkpoints/hmr2/epoch=10-step=25000.ckpt"), + (VITPOSE_CHECKPOINT, "inputs/checkpoints/vitpose/vitpose-h-multi-coco.pth"), + (YOLO_CHECKPOINT, "inputs/checkpoints/yolo/yolov8x.pt"), + (SMPLX_MODEL, "inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz"), +) + + +@dataclass(frozen=True, slots=True) +class BaseAssets: + vae: Path + """Wan VAE checkpoint.""" + text_encoder: Path | None + """UMT5 text-encoder checkpoint; ``None`` when a prompt embedding replaces it.""" + tokenizer: Path | None + """UMT5 tokenizer directory; ``None`` when a prompt embedding replaces it.""" + tiny_decoder: Path | None + """Pinned TAEW2.2 checkpoint when turbo mode needs it.""" + turbo_lora: Path | None + """Pinned Wan2.2 Turbo-LoRA when turbo mode needs it.""" + + +def _require_file(path: Path, label: str, command: str) -> Path: + resolved = path.expanduser().resolve() + if not resolved.is_file(): + raise AssetError(f"{label} does not exist: {resolved}. Run `python {command}` to install it.") + return resolved + + +def resolve_checkpoint(path: str | Path | None = None, model_dir: str | Path = "models") -> Path: + if path is not None: + resolved = Path(path).expanduser().resolve() + if not resolved.is_file(): + raise AssetError(f"Checkpoint override does not exist: {resolved}") + return resolved + return _require_file(Path(model_dir) / CHECKPOINT, "Checkpoint", "scripts/download_model.py") + + +def resolve_regressor(path: str | Path | None = None, model_dir: str | Path = "models") -> Path: + if path is not None: + resolved = Path(path).expanduser().resolve() + if not resolved.is_file(): + raise AssetError(f"MHR70 regressor override does not exist: {resolved}") + return resolved + return _require_file(Path(model_dir) / MHR70_REGRESSOR, "MHR70 regressor", "scripts/download_model.py") + + +def resolve_foreground_model(model_dir: str | Path = "models") -> Path: + root = Path(model_dir).expanduser() / BIREFNET_DIR + for relative in BIREFNET_FILES: + _require_file(root / relative, "BiRefNet file", "scripts/download_model.py") + return root.resolve() + + +def resolve_perceptual_vgg19(model_dir: str | Path = "models") -> Path: + """Resolve the converted VGG-19 weights used by perceptual reconstruction.""" + + return _require_file( + Path(model_dir) / PERCEPTUAL_VGG19, + "Perceptual VGG-19 weights", + "scripts/download_model.py", + ) + + +def resolve_base_assets( + model_dir: str | Path, + settings: ModeSettings, + *, + have_prompt_embedding: bool = False, +) -> BaseAssets: + """Resolve the local VAE, T5 encoder, and tokenizer. + + The encoder and its tokenizer serve only the single fixed prompt, so a + caller that already holds the exported embedding needs neither and must + not be asked for the 11 GB checkpoint. + """ + + root = Path(model_dir).expanduser() + text_encoder: Path | None = None + tokenizer: Path | None = None + if not have_prompt_embedding: + for relative in TOKENIZER_FILES: + _require_file(root / relative, "Tokenizer file", "scripts/download_model.py") + text_encoder = _require_file(root / TEXT_ENCODER, "Text encoder", "scripts/download_model.py") + tokenizer = (root / TOKENIZER_DIR).expanduser().resolve() + return BaseAssets( + vae=_require_file(root / WAN_VAE, "VAE", "scripts/download_model.py"), + text_encoder=text_encoder, + tokenizer=tokenizer, + tiny_decoder=( + _require_file(root / TAEW2_2, "TAEW2.2 checkpoint", "scripts/download_model.py") + if settings.tiny_decoders + else None + ), + turbo_lora=( + _require_file(root / TURBO_LORA, "Turbo-LoRA", "scripts/download_model.py") + if settings.turbo_lora + else None + ), + ) diff --git a/fdanyone/config.py b/fdanyone/config.py new file mode 100644 index 0000000000000000000000000000000000000000..341f1bcae776655c674947041103ef89602bcf02 --- /dev/null +++ b/fdanyone/config.py @@ -0,0 +1,200 @@ +"""Fixed model and preprocessing settings used by the released method. + +Only reader-useful choices live in the CLI. These values describe the trained +model and therefore stay together here instead of being exposed as knobs. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import Literal + +ModeName = Literal["turbo", "reference"] + +MVS_ATTENTION_RANGE: tuple[int | None, int | None, int | None] = (0, None, 1) +USE_VIEWPACK: bool = True +USE_POSE_ENCODER: bool = True +POSE_ENCODER_TYPE: str = "rgb" + + +@dataclass(frozen=True, slots=True) +class InferenceConstants: + """Architecture and media constants shared by both inference modes.""" + + num_frames: int = 121 + """Frames generated for every view.""" + height: int = 1280 + """Output height in pixels.""" + width: int = 704 + """Output width in pixels.""" + prompt: str = "视频中的人在做动作" + """Fixed positive prompt used by the released model.""" + auto_downsample_fps: tuple[tuple[int, int], ...] = ( + (24, 1), + (24000, 1001), + (25, 1), + (30, 1), + (30000, 1001), + ) + """Input rates that may be reduced to their exact supported divisor.""" + temporal_sampling_policy: str = "nearest_source_pts_on_zero_based_cfr_clock" + """Canonical clip sampling policy.""" + rcp_jpeg_quality: int = 85 + """JPEG quality at the proposal-to-target boundary.""" + skeleton_h264_crf: int = 17 + """CRF used for skeleton conditioning videos.""" + target_h264_crf: int = 18 + """CRF used for generated target videos.""" + h264_preset: str = "medium" + """H.264 encoder preset.""" + skeleton_max_dimension: int = 2048 + """Largest skeleton-render canvas dimension.""" + denoising_strength: float = 1.0 + """Flow-match denoising strength.""" + + +@dataclass(frozen=True, slots=True) +class ModeSettings: + """One complete, immutable inference policy.""" + + mode: ModeName + """Public CLI and metadata name.""" + num_inference_steps: int + """Number of flow-matching denoising steps.""" + scheduler_shift: float + """Flow-match sigma shift.""" + stream_dit_weights: bool + """Whether all wrapped DiT weights stream from host memory.""" + fp8_w8a8: bool + """Whether safe interior projections use dynamic per-tensor FP8.""" + regional_compile: bool + """Whether repeated DiT blocks use the fixed turbo compile policy.""" + overlap_target_skeletons: bool = False + """Whether target skeleton rendering overlaps the proposal stage.""" + bf16_block_glue: bool = False + """Run transformer normalization, modulation, and gates in BF16.""" + direct_rcp_latent_handoff: bool = False + """Feed generated RCP latents to target conditioning without re-encoding.""" + async_video_encode: bool = False + """Overlap CPU x264 encoding with the next per-view GPU VAE decode.""" + tiny_decoders: bool = False + """Whether proposal and target views use the pinned TAEW2.2 decoder.""" + nvdec_skeletons: bool = False + """Whether skeleton videos use fail-closed CUDA decoding.""" + turbo_lora: bool = False + """Whether to merge the pinned Wan2.2 Turbo-LoRA before quantization.""" + + @property + def dit_pose_batch_size(self) -> int | None: + """Return the validated FP8 pose-activation batch.""" + + return 4 if self.fp8_w8a8 else None + + @property + def exact_attention(self) -> bool: + """Return whether this mode requires exact SDPA attention.""" + + return self.mode == "reference" + + @property + def skeleton_video_decoder(self) -> str: + """Return the concrete skeleton-video backend.""" + + return "torchcodec_cuda" if self.nvdec_skeletons else "pyav" + + +MODES: Mapping[ModeName, ModeSettings] = MappingProxyType( + { + "turbo": ModeSettings( + mode="turbo", + num_inference_steps=4, + scheduler_shift=17.0, + stream_dit_weights=False, + fp8_w8a8=True, + regional_compile=True, + overlap_target_skeletons=True, + bf16_block_glue=True, + direct_rcp_latent_handoff=True, + async_video_encode=True, + tiny_decoders=True, + nvdec_skeletons=True, + turbo_lora=True, + ), + "reference": ModeSettings( + mode="reference", + num_inference_steps=24, + scheduler_shift=5.0, + stream_dit_weights=True, + fp8_w8a8=False, + regional_compile=False, + ), + } +) + + +@dataclass(frozen=True) +class CameraConfig: + count: int = 24 + pitch_degrees: float = 15.0 + + def __post_init__(self) -> None: + if self.count <= 0: + raise ValueError("Camera count must be positive.") + + +@dataclass(frozen=True) +class ForegroundConfig: + """Pinned standard BiRefNet inference contract.""" + + image_size: tuple[int, int] = (1024, 1024) + batch_size: int = 4 + + +@dataclass(frozen=True) +class FramingConfig: + """Sequence-level camera solve matching the current GVHMR demo.""" + + reference_radius: float = 3.0 + reference_target_height: float = 1.0 + reference_focal_normalized: float = 1664.0 / 1280.0 + height_target_ratio: float = 0.80 + height_percentile: float = 95.0 + width_target_ratio: float = 0.90 + width_percentile: float = 80.0 + min_radius: float = 1.5 + max_radius: float = 8.0 + input_min_confidence: float = 0.55 + max_focal_normalized: float = 4.0 + cutoff_target_ratio: float = 0.99 + cutoff_percentile: float = 80.0 + + +@dataclass(frozen=True) +class CropConfig: + """Source-mask crop; generated cameras use a plain center aspect crop.""" + + margin_top: float = 0.04 + margin_right: float = 0.04 + margin_bottom: float = 0.04 + margin_left: float = 0.04 + allow_upscale: bool = True + mask_threshold: float = 0.05 + + @property + def margins(self) -> tuple[float, float, float, float]: + return (self.margin_top, self.margin_right, self.margin_bottom, self.margin_left) + + +@dataclass(frozen=True) +class SkeletonConfig: + draw_body_reference_px: float = 640.0 + + +INFERENCE = InferenceConstants() +CAMERA = CameraConfig() +FOREGROUND = ForegroundConfig() +FRAMING = FramingConfig() +CROP = CropConfig() +SKELETON = SkeletonConfig() diff --git a/fdanyone/device.py b/fdanyone/device.py new file mode 100644 index 0000000000000000000000000000000000000000..73bd828c3cc4e361ae562f2610bc9274f2f94204 --- /dev/null +++ b/fdanyone/device.py @@ -0,0 +1,25 @@ +"""CUDA device selection shared by pipeline and isolated workers.""" + +from __future__ import annotations + +from fdanyone.errors import ConfigurationError + + +def select_cuda_device(device: str) -> tuple[str, int]: + """Validate, select, and normalize one CUDA device.""" + + import torch + + try: + requested = torch.device(device) + except (RuntimeError, TypeError, ValueError) as exc: + raise ConfigurationError(f"Invalid CUDA device {device!r}.") from exc + if requested.type != "cuda" or not torch.cuda.is_available(): + raise ConfigurationError(f"4DAnyone requires an available CUDA device, got {device!r}.") + index = torch.cuda.current_device() if requested.index is None else requested.index + if index < 0 or index >= torch.cuda.device_count(): + raise ConfigurationError( + f"CUDA device index {index} is unavailable; visible device count is {torch.cuda.device_count()}." + ) + torch.cuda.set_device(index) + return f"cuda:{index}", index diff --git a/fdanyone/download.py b/fdanyone/download.py new file mode 100644 index 0000000000000000000000000000000000000000..d9933fad66caec4b04831887cf0ea3279ca32158 --- /dev/null +++ b/fdanyone/download.py @@ -0,0 +1,427 @@ +"""Download the published 4DAnyone assets from Hugging Face. + +Missing model checkpoints and bundled example clips are fetched automatically +when inference needs them; the scripts under ``scripts/`` pre-fetch the same +files. SMPL-X is licensed separately, so first-run inference starts its +interactive installer only when a terminal is available. +""" + +from __future__ import annotations + +import getpass +import logging +import os +import shlex +import shutil +import sys +import tempfile +import urllib.error +import urllib.parse +import urllib.request +import zipfile +from pathlib import Path, PurePosixPath + +from fdanyone.assets import ( + BIREFNET_DIR, + BIREFNET_FILES, + BIREFNET_REPO_ID, + BIREFNET_REVISION, + EXAMPLE_FILES, + GVHMR_LINKS, + HF_REPO_ID, + HF_REVISION, + MODEL_FILES, + PERCEPTUAL_VGG19, + SMPLX_MODEL, + TAEW2_2, + TAEW2_2_SHA256, + TAEW2_2_URL, + TURBO_LORA, + TURBO_LORA_REPO_FILE, + TURBO_LORA_REPO_ID, + TURBO_LORA_REVISION, + TURBO_LORA_SHA256, + resolve_perceptual_vgg19, +) +from fdanyone.errors import AssetError + +LOGGER = logging.getLogger("fdanyone") + +SMPLX_HOME = "https://smpl-x.is.tue.mpg.de/" +SMPLX_DOWNLOAD_URL = "https://download.is.tue.mpg.de/download.php?domain=smplx&sfile=models_smplx_v1_1.zip" +SMPLX_ARCHIVE_MEMBER = ("models", "smplx", "SMPLX_NEUTRAL.npz") + + +def _snapshot( + allow_patterns: list[str], + local_dir: Path, + *, + repo_id: str = HF_REPO_ID, + revision: str = HF_REVISION, +) -> None: + try: + from huggingface_hub import snapshot_download + except ImportError as exc: + raise AssetError("Install requirements.txt before downloading assets.") from exc + + try: + snapshot_download( + repo_id=repo_id, + revision=revision, + allow_patterns=allow_patterns, + local_dir=local_dir, + ) + except Exception as exc: + raise AssetError( + f"Could not download {repo_id}@{revision}. Check the network connection and Hugging Face access." + ) from exc + + +def ensure_foreground_model(model_dir: str | Path = "models") -> Path: + root = Path(model_dir).expanduser().resolve() / BIREFNET_DIR + missing = [relative for relative in BIREFNET_FILES if not (root / relative).is_file()] + if missing: + LOGGER.info("Downloading BiRefNet foreground model (first run only)") + _snapshot( + missing, + root, + repo_id=BIREFNET_REPO_ID, + revision=BIREFNET_REVISION, + ) + return root + + +def require_gvhmr_checkout(gvhmr_root: str | Path) -> Path: + root = Path(gvhmr_root).expanduser().resolve() + if not (root / "hmr4d/__init__.py").is_file(): + raise AssetError( + f"GVHMR is not initialized at {root}. Run `git submodule update --init third_party/GVHMR` first." + ) + return root + + +def _ensure_link(source: Path, destination: Path) -> None: + source = source.expanduser().resolve() + destination.parent.mkdir(parents=True, exist_ok=True) + if destination.is_symlink(): + try: + if destination.resolve(strict=True).samefile(source): + return + except FileNotFoundError: + pass + destination.unlink() + elif destination.exists(): + if destination.samefile(source): + return + raise AssetError( + f"GVHMR asset location is occupied by an unrelated file: {destination}. " + f"Move it away so the downloaded {source.name} can be linked." + ) + relative = os.path.relpath(source, start=destination.parent) + destination.symlink_to(relative) + + +def create_classic_gvhmr_links( + model_dir: str | Path = "models", + gvhmr_root: str | Path = "third_party/GVHMR", + *, + require_models: bool = True, + require_smplx: bool = True, +) -> Path: + """Create the ignored compatibility links expected by upstream GVHMR.""" + + root = require_gvhmr_checkout(gvhmr_root) + models = Path(model_dir).expanduser().resolve() + for relative, target in GVHMR_LINKS: + source = models / relative + if not source.is_file(): + required = require_smplx if relative == SMPLX_MODEL else require_models + if required: + command = "scripts/download_smplx.py" if relative == SMPLX_MODEL else "scripts/download_model.py" + raise AssetError(f"Model file is missing: {source}. Run `python {command}` first.") + continue + _ensure_link(source, root / target) + return root + + +def ensure_models( + model_dir: str | Path = "models", + gvhmr_root: str | Path = "third_party/GVHMR", +) -> Path: + """Download any missing published model file and refresh the GVHMR links.""" + + require_gvhmr_checkout(gvhmr_root) + models = Path(model_dir).expanduser().resolve() + missing = [relative for relative in MODEL_FILES if not (models / relative).is_file()] + if missing: + LOGGER.info("Downloading %d model files from %s (first run only)", len(missing), HF_REPO_ID) + # Repository paths match the local layout, so download straight into + # place; huggingface_hub stages and resumes partial files itself. + _snapshot(missing, models) + ensure_foreground_model(models) + create_classic_gvhmr_links(models, gvhmr_root, require_smplx=False) + return models + + +def _verify_sha256(path: Path, expected: str, label: str) -> None: + import hashlib + + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1 << 20), b""): + digest.update(chunk) + if digest.hexdigest() != expected: + path.unlink(missing_ok=True) + raise AssetError(f"{label} failed its SHA-256 check; the download was removed. Re-run it.") + + +def ensure_turbo_assets(model_dir: str | Path = "models") -> Path: + """Download the pinned turbo-mode checkpoints (TAEW2.2 decoder, Turbo-LoRA).""" + + models = Path(model_dir).expanduser().resolve() + taew = models / TAEW2_2 + if not taew.is_file(): + LOGGER.info("Downloading the TAEW2.2 tiny decoder (first run only)") + taew.parent.mkdir(parents=True, exist_ok=True) + staged = taew.with_suffix(".part") + try: + urllib.request.urlretrieve(TAEW2_2_URL, staged) + except urllib.error.URLError as exc: + raise AssetError(f"Could not download the TAEW2.2 decoder from {TAEW2_2_URL}.") from exc + _verify_sha256(staged, TAEW2_2_SHA256, "TAEW2.2 decoder") + staged.replace(taew) + lora = models / TURBO_LORA + if not lora.is_file(): + LOGGER.info( + "Downloading the Wan2.2 Turbo-LoRA from %s (first run only)", TURBO_LORA_REPO_ID + ) + _snapshot( + [TURBO_LORA_REPO_FILE], + models / "fps-assets" / "turbo-lora", + repo_id=TURBO_LORA_REPO_ID, + revision=TURBO_LORA_REVISION, + ) + _verify_sha256(lora, TURBO_LORA_SHA256, "Turbo-LoRA") + return models + + +def ensure_perceptual_vgg19(model_dir: str | Path = "models") -> Path: + """Download only the optional VGG-19 reconstruction asset when missing.""" + + models = Path(model_dir).expanduser().resolve() + destination = models / PERCEPTUAL_VGG19 + if not destination.is_file(): + LOGGER.info("Downloading the perceptual VGG-19 model (first use only)") + _snapshot([PERCEPTUAL_VGG19], models) + return resolve_perceptual_vgg19(models) + + +def download_model( + model_dir: str = "models", + gvhmr_root: str = "third_party/GVHMR", +) -> dict[str, str]: + """Download the published model checkpoints.""" + + models = ensure_models(model_dir, gvhmr_root) + ensure_turbo_assets(models) + return { + "models": str(models), + "revision": HF_REVISION, + "foreground_revision": BIREFNET_REVISION, + "turbo_lora_revision": TURBO_LORA_REVISION, + } + + +def download_example(data_dir: str = "data") -> dict[str, str]: + """Download the bundled example clips.""" + + data = Path(data_dir).expanduser().resolve() + destinations = {relative: data / Path(relative).relative_to("data") for relative in EXAMPLE_FILES} + missing = [relative for relative, destination in destinations.items() if not destination.is_file()] + if missing: + # Repository paths carry a leading ``data/`` prefix while --data_dir is + # the local root itself, so stage the snapshot and move each file. + staging = data / ".download" + _snapshot(missing, staging) + for relative in missing: + destination = destinations[relative] + destination.parent.mkdir(parents=True, exist_ok=True) + (staging / relative).replace(destination) + shutil.rmtree(staging) + return {"examples": str(data / "source/pexels"), "revision": HF_REVISION} + + +def ensure_example_video(video_path: str | Path) -> Path: + """Fetch a bundled example clip when its expected file is missing.""" + + path = Path(video_path).expanduser() + if path.is_file(): + return path + matches = [relative for relative in EXAMPLE_FILES if PurePosixPath(relative).name == path.name] + if not matches: + raise AssetError(f"Input video does not exist: {path.resolve()}") + LOGGER.info("Downloading the bundled example clip %s", path.name) + path.parent.mkdir(parents=True, exist_ok=True) + # Stage beside the destination so the final rename stays on one filesystem. + staging = path.parent / ".download" + _snapshot(matches[:1], staging) + (staging / matches[0]).replace(path) + shutil.rmtree(staging) + return path + + +def _parse_interactive_path(value: str) -> Path: + try: + parts = shlex.split(value.strip()) + except ValueError as exc: + raise AssetError(f"Could not parse the archive path: {exc}") from None + if len(parts) != 1: + raise AssetError("Enter one ZIP or SMPLX_NEUTRAL.npz path.") + return Path(parts[0]).expanduser() + + +def _copy_model_from_source(source: Path, destination: Path) -> None: + if source.name == "SMPLX_NEUTRAL.npz": + shutil.copyfile(source, destination) + return + if zipfile.is_zipfile(source): + try: + with zipfile.ZipFile(source) as archive: + candidates = [ + info + for info in archive.infolist() + if not info.is_dir() and PurePosixPath(info.filename).parts[-3:] == SMPLX_ARCHIVE_MEMBER + ] + if len(candidates) != 1: + raise AssetError( + "The archive must contain exactly one models/smplx/SMPLX_NEUTRAL.npz file. " + "Download models_smplx_v1_1.zip from the official SMPL-X website." + ) + with archive.open(candidates[0]) as model, destination.open("wb") as output: + shutil.copyfileobj(model, output, length=8 * 1024 * 1024) + except zipfile.BadZipFile: + raise AssetError(f"SMPL-X archive is invalid: {source}") from None + return + raise AssetError("Select models_smplx_v1_1.zip or SMPLX_NEUTRAL.npz.") + + +def install_smplx( + source_path: str | Path, + model_dir: str | Path = "models", + gvhmr_root: str | Path = "third_party/GVHMR", +) -> Path: + """Install a user-provided official ZIP or neutral NPZ.""" + + source = Path(source_path).expanduser().resolve() + if not source.is_file(): + raise AssetError(f"SMPL-X source does not exist: {source}") + + target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL + target.parent.mkdir(parents=True, exist_ok=True) + temporary = target.parent / f".{target.name}.download-{os.getpid()}" + try: + _copy_model_from_source(source, temporary) + temporary.replace(target) + finally: + temporary.unlink(missing_ok=True) + create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True) + return target + + +def _download_official(username: str, password: str, destination: Path) -> None: + payload = urllib.parse.urlencode({"username": username, "password": password}).encode() + request = urllib.request.Request( + SMPLX_DOWNLOAD_URL, + data=payload, + headers={"User-Agent": "4DAnyone SMPL-X installer"}, + method="POST", + ) + try: + with urllib.request.urlopen(request, timeout=120) as response, destination.open("wb") as output: + shutil.copyfileobj(response, output, length=8 * 1024 * 1024) + except (OSError, urllib.error.URLError) as exc: + raise AssetError(f"Official SMPL-X download failed: {exc}") from None + if not zipfile.is_zipfile(destination): + raise AssetError( + "The SMPL-X website did not return a ZIP archive. Check the account, license acceptance, or website." + ) + + +def _prompt_for_archive(model_dir: str, gvhmr_root: str) -> dict[str, str] | None: + print(f"Download models_smplx_v1_1.zip from:\n {SMPLX_DOWNLOAD_URL}") + while True: + try: + value = input("Archive path (drag the downloaded ZIP here): ").strip() + except EOFError: + value = "" + if not value: + print("SMPL-X setup cancelled; the downloaded ZIP was not modified.") + return None + try: + installed = install_smplx(_parse_interactive_path(value), model_dir, gvhmr_root) + except AssetError as exc: + print(f"error: {exc}") + continue + return {"installed": str(installed)} + + +def download_smplx( + archive_path: str | None = None, + model_dir: str = "models", + gvhmr_root: str = "third_party/GVHMR", +) -> dict[str, str] | None: + """Install the separately licensed SMPL-X neutral body model.""" + + target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL + if target.is_file(): + create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True) + return {"installed": str(target)} + + if archive_path is not None: + return {"installed": str(install_smplx(archive_path, model_dir, gvhmr_root))} + + print(f"SMPL-X requires a free account and license acceptance at {SMPLX_HOME}") + try: + accepted = input("Have you registered and accepted the SMPL-X license? [y/N]: ").strip().lower() + except EOFError: + accepted = "" + if accepted in {"y", "yes"}: + username = input("SMPL-X username or email: ").strip() + password = getpass.getpass("SMPL-X password: ") + if username and password: + with tempfile.TemporaryDirectory(prefix="fdanyone-smplx-") as temporary_dir: + archive = Path(temporary_dir) / "models_smplx_v1_1.zip" + try: + _download_official(username, password, archive) + installed = install_smplx(archive, model_dir, gvhmr_root) + except AssetError as exc: + print(f"Automatic download was unavailable: {exc}") + else: + return {"installed": str(installed)} + return _prompt_for_archive(model_dir, gvhmr_root) + + +def ensure_smplx( + model_dir: str | Path = "models", + gvhmr_root: str | Path = "third_party/GVHMR", +) -> Path: + """Install SMPL-X interactively on first use, without blocking jobs.""" + + require_gvhmr_checkout(gvhmr_root) + target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL + if target.is_file(): + create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True) + return target + + if not getattr(sys.stdin, "isatty", lambda: False)(): + raise AssetError( + f"SMPL-X is not installed at {target}, and inference has no interactive terminal. " + "Run `python scripts/download_smplx.py` before starting this job." + ) + + print("SMPL-X is required and has not been installed; starting its licensed setup.") + result = download_smplx(model_dir=str(model_dir), gvhmr_root=str(gvhmr_root)) + if result is None or not target.is_file(): + raise AssetError("SMPL-X setup was cancelled; inference cannot continue.") + LOGGER.info("SMPL-X installed; continuing inference") + return target diff --git a/fdanyone/errors.py b/fdanyone/errors.py new file mode 100644 index 0000000000000000000000000000000000000000..3093effd8132fd6bc0aa5b4c61e48032c57f8fe8 --- /dev/null +++ b/fdanyone/errors.py @@ -0,0 +1,17 @@ +"""Project-specific errors with actionable user-facing messages.""" + + +class FourDAnyoneError(RuntimeError): + """Base class for expected pipeline failures.""" + + +class ConfigurationError(FourDAnyoneError): + """Raised when a frozen inference contract is violated.""" + + +class AssetError(FourDAnyoneError): + """Raised when a model or gated asset is missing or invalid.""" + + +class VideoContractError(FourDAnyoneError): + """Raised when an input cannot provide the canonical 121-frame clip.""" diff --git a/fdanyone/foreground.py b/fdanyone/foreground.py new file mode 100644 index 0000000000000000000000000000000000000000..5ea44de1a4f6f364d68a6f63e1c8d5efea45ddc0 --- /dev/null +++ b/fdanyone/foreground.py @@ -0,0 +1,65 @@ +"""Pinned BiRefNet inference over the canonical source clip.""" + +from __future__ import annotations + +import gc +from pathlib import Path + +import numpy as np +from PIL import Image + +from fdanyone.config import FOREGROUND + + +def predict_foreground_masks( + frames: tuple[np.ndarray, ...], + model_path: str | Path, + device: str, + *, + batch_size: int = FOREGROUND.batch_size, +) -> np.ndarray: + """Return full-raster 8-bit foreground masks for the canonical clip.""" + + import torch + from torchvision import transforms + from torchvision.transforms.functional import to_pil_image + from transformers import AutoModelForImageSegmentation + + if not frames: + raise ValueError("Foreground inference requires at least one frame.") + if batch_size <= 0: + raise ValueError("batch_size must be positive.") + shape = frames[0].shape + if any(frame.dtype != np.uint8 or frame.shape != shape for frame in frames): + raise ValueError("Foreground frames must share one RGB uint8 raster.") + + model = AutoModelForImageSegmentation.from_pretrained( + str(Path(model_path).expanduser().resolve()), + local_files_only=True, + trust_remote_code=True, + ) + model = model.eval().half().to(device) + transform = transforms.Compose( + [ + transforms.Resize(FOREGROUND.image_size), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ) + output: list[np.ndarray] = [] + try: + for start in range(0, len(frames), batch_size): + images = [Image.fromarray(frame, mode="RGB") for frame in frames[start : start + batch_size]] + inputs = torch.stack([transform(image) for image in images]).to(device=device, dtype=torch.float16) + with torch.inference_mode(): + predictions = model(inputs)[-1].sigmoid().cpu() + for image, prediction in zip(images, predictions, strict=True): + mask = to_pil_image(prediction).resize(image.size).convert("L") + output.append(np.asarray(mask, dtype=np.uint8).copy()) + del inputs, predictions + finally: + del model + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return np.stack(output) diff --git a/fdanyone/geometry/__init__.py b/fdanyone/geometry/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..7c99e8acfc0f643f2b2d1f21f33eeb97440d7077 --- /dev/null +++ b/fdanyone/geometry/__init__.py @@ -0,0 +1 @@ +"""Camera and crop geometry.""" diff --git a/fdanyone/geometry/cameras.py b/fdanyone/geometry/cameras.py new file mode 100644 index 0000000000000000000000000000000000000000..2333f41fc07024a1d4b9067d614979a50284ed18 --- /dev/null +++ b/fdanyone/geometry/cameras.py @@ -0,0 +1,250 @@ +"""Canonical uniform camera rings and camera serialization.""" + +from __future__ import annotations + +from dataclasses import asdict, dataclass + +import numpy as np + +from fdanyone.config import CAMERA, FRAMING, CameraConfig + +WORLD_FRAME = { + "name": "canonical_human_world", + "handedness": "right", + "axes": { + "x": "right in the source-facing ring camera (the subject's anatomical left)", + "y": "up", + "z": "front; the subject initially faces +z and the source-facing camera lies on the +z side", + }, + "origin": "initial root projected to the ground plane", + "units": "meters", +} + +CAMERA_FRAME = { + "name": "opencv_camera", + "handedness": "right", + "axes": {"x": "image right", "y": "image down", "z": "forward from camera into the scene"}, + "matrix_convention": "column vectors: x_camera = world_to_camera @ x_world_homogeneous", + "intrinsics_convention": "pixels with origin at the top-left", +} + + +def _normalize(vector: np.ndarray) -> np.ndarray: + norm = float(np.linalg.norm(vector)) + if norm <= 1e-12: + raise ValueError("Cannot normalize a zero-length direction vector.") + return vector / norm + + +def look_at_pytorch3d(position: np.ndarray, target: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + """Return the column-vector w2c used by PyTorch3D's look_at_rotation.""" + + position = np.asarray(position, dtype=np.float64) + target = np.asarray(target, dtype=np.float64) + up = np.array([0.0, 1.0, 0.0], dtype=np.float64) + z_axis = _normalize(target - position) + x_axis = _normalize(np.cross(up, z_axis)) + y_axis = _normalize(np.cross(z_axis, x_axis)) + rotation = np.stack([x_axis, y_axis, z_axis], axis=0) + translation = -(rotation @ position) + return rotation, translation + + +def pytorch3d_to_opencv(rotation: np.ndarray, translation: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + rotation_cv = np.asarray(rotation, dtype=np.float64).copy() + translation_cv = np.asarray(translation, dtype=np.float64).copy() + rotation_cv[:2] *= -1.0 + translation_cv[:2] *= -1.0 + return rotation_cv, translation_cv + + +def homogeneous_w2c(rotation: np.ndarray, translation: np.ndarray) -> np.ndarray: + matrix = np.eye(4, dtype=np.float64) + matrix[:3, :3] = rotation + matrix[:3, 3] = translation + return matrix + + +@dataclass(frozen=True) +class Camera: + camera_id: int + layer_index: int + yaw_degrees: float + azimuth_degrees: float + pitch_degrees: float + position: tuple[float, float, float] + K: tuple[tuple[float, float, float], ...] + world_to_camera: tuple[tuple[float, float, float, float], ...] + camera_to_world: tuple[tuple[float, float, float, float], ...] + image_width: int + image_height: int + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, payload: dict) -> Camera: + """Invert ``to_dict``, coercing JSON lists back to the tuple fields.""" + + return cls( + camera_id=int(payload["camera_id"]), + layer_index=int(payload["layer_index"]), + yaw_degrees=float(payload["yaw_degrees"]), + azimuth_degrees=float(payload["azimuth_degrees"]), + pitch_degrees=float(payload["pitch_degrees"]), + position=tuple(float(value) for value in payload["position"]), + K=tuple(tuple(float(value) for value in row) for row in payload["K"]), + world_to_camera=tuple( + tuple(float(value) for value in row) for row in payload["world_to_camera"] + ), + camera_to_world=tuple( + tuple(float(value) for value in row) for row in payload["camera_to_world"] + ), + image_width=int(payload["image_width"]), + image_height=int(payload["image_height"]), + ) + + +def reference_intrinsics( + image_height: int, + image_width: int, + max_render_height: int = 1280, + *, + focal_normalized: float = FRAMING.reference_focal_normalized, +) -> np.ndarray: + """Reproduce the renderer's downscale-then-rescale intrinsic construction.""" + + divisor = 2 + while image_height / divisor > max_render_height: + divisor += 1 + render_height = image_height // divisor + render_width = image_width // divisor + scale = image_height / render_height + if focal_normalized <= 0: + raise ValueError("focal_normalized must be positive.") + focal = focal_normalized * render_height + intrinsic = np.array( + [[focal, 0.0, render_width / 2.0], [0.0, focal, render_height / 2.0], [0.0, 0.0, 1.0]], + dtype=np.float64, + ) + intrinsic *= scale + intrinsic[2, 2] = 1.0 + return intrinsic + + +def camera_ring( + *, + center: np.ndarray, + front_direction: np.ndarray, + K: np.ndarray, + image_height: int, + image_width: int, + radius: float = FRAMING.reference_radius, + target_height: float = FRAMING.reference_target_height, + spec: CameraConfig = CAMERA, + start_yaw_degrees: float = 0.0, + yaw_span_degrees: float = 360.0, + layer_index: int = 0, + camera_id_offset: int = 0, +) -> tuple[Camera, ...]: + center = np.asarray(center, dtype=np.float64).copy() + front_direction = np.asarray(front_direction, dtype=np.float64) + front_xz = front_direction[[0, 2]] + front_azimuth = np.arctan2(front_xz[1], front_xz[0]) + # Relative yaw zero faces the person's front; start_yaw chooses where the + # first camera lies and IDs then advance uniformly through the span. + azimuth_start = front_azimuth + np.pi + np.deg2rad(start_yaw_degrees) + target = center.copy() + if radius <= 0: + raise ValueError("radius must be positive.") + target[1] = target_height + camera_height = target_height + radius * np.tan(np.deg2rad(spec.pitch_degrees)) + + cameras: list[Camera] = [] + yaw_span_radians = np.deg2rad(yaw_span_degrees) + for view_index in range(spec.count): + camera_id = camera_id_offset + view_index + yaw_degrees = start_yaw_degrees + view_index / spec.count * yaw_span_degrees + azimuth = azimuth_start + view_index / spec.count * yaw_span_radians + position = np.array( + [ + center[0] + radius * np.cos(azimuth), + camera_height, + center[2] + radius * np.sin(azimuth), + ], + dtype=np.float64, + ) + rotation_p3d, translation_p3d = look_at_pytorch3d(position, target) + rotation_cv, translation_cv = pytorch3d_to_opencv(rotation_p3d, translation_p3d) + w2c = homogeneous_w2c(rotation_cv, translation_cv) + c2w = np.linalg.inv(w2c) + cameras.append( + Camera( + camera_id=camera_id, + layer_index=layer_index, + yaw_degrees=float(yaw_degrees), + azimuth_degrees=float(np.rad2deg(azimuth) % 360.0), + pitch_degrees=spec.pitch_degrees, + position=tuple(float(value) for value in position), + K=tuple(tuple(float(value) for value in row) for row in K), + world_to_camera=tuple(tuple(float(value) for value in row) for row in w2c), + camera_to_world=tuple(tuple(float(value) for value in row) for row in c2w), + image_width=image_width, + image_height=image_height, + ) + ) + return tuple(cameras) + + +def camera_grid( + *, + center: np.ndarray, + front_direction: np.ndarray, + K: np.ndarray, + image_height: int, + image_width: int, + views_per_layer: int, + layer_pitches: tuple[int, ...], + start_yaw: int, + yaw_span: int, + radius: float = FRAMING.reference_radius, + target_height: float = FRAMING.reference_target_height, +) -> tuple[Camera, ...]: + """Create a layer-major target camera grid.""" + + return tuple( + camera + for layer_index, pitch in enumerate(layer_pitches) + for camera in camera_ring( + center=center, + front_direction=front_direction, + K=K, + image_height=image_height, + image_width=image_width, + radius=radius, + target_height=target_height, + spec=CameraConfig(count=views_per_layer, pitch_degrees=float(pitch)), + start_yaw_degrees=float(start_yaw), + yaw_span_degrees=float(yaw_span), + layer_index=layer_index, + camera_id_offset=layer_index * views_per_layer, + ) + ) + + +def project_points(points_world: np.ndarray, camera: Camera) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + points = np.asarray(points_world, dtype=np.float64) + w2c = np.asarray(camera.world_to_camera) + K = np.asarray(camera.K) + points_h = np.concatenate([points, np.ones((points.shape[0], 1), dtype=points.dtype)], axis=1) + points_camera = points_h @ w2c[:3].T + homogeneous = points_camera @ K.T + xy = homogeneous[:, :2] / np.maximum(homogeneous[:, 2:3], 1e-8) + valid = ( + (points_camera[:, 2] > 0.0) + & (xy[:, 0] >= 0.0) + & (xy[:, 0] < camera.image_width) + & (xy[:, 1] >= 0.0) + & (xy[:, 1] < camera.image_height) + ) + return xy.astype(np.float32), points_camera[:, 2].astype(np.float32), valid diff --git a/fdanyone/geometry/crop.py b/fdanyone/geometry/crop.py new file mode 100644 index 0000000000000000000000000000000000000000..08f501c2bdc311cc737aa8f202f049de725c3d7f --- /dev/null +++ b/fdanyone/geometry/crop.py @@ -0,0 +1,144 @@ +"""Deterministic source-mask and center-aspect crops.""" + +from __future__ import annotations + +import math +from dataclasses import dataclass + +import numpy as np + + +@dataclass(frozen=True) +class Crop: + top: int + left: int + height: int + width: int + original_height: int + original_width: int + output_height: int + output_width: int + + @property + def scale_y(self) -> float: + return self.output_height / self.height + + @property + def scale_x(self) -> float: + return self.output_width / self.width + + +def center_crop(image_height: int, image_width: int, output_height: int, output_width: int) -> Crop: + aspect_ratio = output_height / output_width + if image_height / image_width >= aspect_ratio: + crop_width = image_width + crop_height = min(image_height, max(1, int(round(crop_width * aspect_ratio)))) + else: + crop_height = image_height + crop_width = min(image_width, max(1, int(round(crop_height / aspect_ratio)))) + return Crop( + top=(image_height - crop_height) // 2, + left=(image_width - crop_width) // 2, + height=crop_height, + width=crop_width, + original_height=image_height, + original_width=image_width, + output_height=output_height, + output_width=output_width, + ) + + +def mask_bounds(masks: np.ndarray, threshold: float) -> tuple[int, int, int, int] | None: + """Return union bounds as left/top inclusive and right/bottom exclusive.""" + + values = np.asarray(masks) + if values.ndim == 4 and values.shape[1] == 1: + values = values[:, 0] + if values.ndim not in (2, 3): + raise ValueError(f"Expected masks [H,W] or [F,H,W], got {values.shape}.") + if not 0 <= threshold <= 1: + raise ValueError("threshold must be in [0, 1].") + cutoff = threshold * 255.0 if np.issubdtype(values.dtype, np.integer) else threshold + union = values.max(axis=0) if values.ndim == 3 else values + foreground = np.asarray(union > cutoff) + rows = np.flatnonzero(foreground.any(axis=1)) + columns = np.flatnonzero(foreground.any(axis=0)) + if rows.size == 0 or columns.size == 0: + return None + return int(columns[0]), int(rows[0]), int(columns[-1] + 1), int(rows[-1] + 1) + + +def expand_bounds(bounds, margins, image_height: int, image_width: int): + if bounds is None: + return None + xmin, ymin, xmax, ymax = (float(value) for value in bounds) + top, right, bottom, left = (float(value) for value in margins) + box_width, box_height = xmax - xmin, ymax - ymin + return ( + max(0, int(math.floor(xmin - left * box_width))), + max(0, int(math.floor(ymin - top * box_height))), + min(image_width, int(math.ceil(xmax + right * box_width))), + min(image_height, int(math.ceil(ymax + bottom * box_height))), + ) + + +def crop_from_bounds( + *, + bounds: tuple[int, int, int, int] | None, + image_height: int, + image_width: int, + output_height: int, + output_width: int, + margins: tuple[float, float, float, float], + allow_upscale: bool = True, +) -> Crop: + bounds = expand_bounds(bounds, margins, image_height, image_width) + if bounds is None: + return center_crop(image_height, image_width, output_height, output_width) + + xmin, ymin, xmax, ymax = bounds + required_height = float(ymax - ymin) + required_width = float(xmax - xmin) + if not allow_upscale: + required_height = max(required_height, min(output_height, image_height)) + required_width = max(required_width, min(output_width, image_width)) + + aspect_ratio = output_height / output_width + crop_height = max(required_height, required_width * aspect_ratio) + crop_width = crop_height / aspect_ratio + max_height = min(float(image_height), float(image_width) * aspect_ratio) + crop_height = max(1, min(image_height, int(math.ceil(min(crop_height, max_height))))) + crop_width = max(1, min(image_width, int(math.ceil(crop_height / aspect_ratio)))) + + left_low = max(0.0, xmax - crop_width) + left_high = min(float(xmin), image_width - crop_width) + top_low = max(0.0, ymax - crop_height) + top_high = min(float(ymin), image_height - crop_height) + if left_low <= left_high: + left = int(math.floor((left_low + left_high) / 2.0)) + else: + left = max(0, min(int(math.floor((xmin + xmax - crop_width) / 2.0)), image_width - crop_width)) + if top_low <= top_high: + top = int(math.floor((top_low + top_high) / 2.0)) + else: + top = max(0, min(int(math.floor((ymin + ymax - crop_height) / 2.0)), image_height - crop_height)) + return Crop( + top=top, + left=left, + height=crop_height, + width=crop_width, + original_height=image_height, + original_width=image_width, + output_height=output_height, + output_width=output_width, + ) + + +def transform_intrinsics(K: np.ndarray, crop: Crop) -> np.ndarray: + transformed = np.asarray(K, dtype=np.float64).copy() + transformed[0, 2] -= crop.left + transformed[1, 2] -= crop.top + transformed[0] *= crop.scale_x + transformed[1] *= crop.scale_y + transformed[2, 2] = 1.0 + return transformed diff --git a/fdanyone/geometry/framing.py b/fdanyone/geometry/framing.py new file mode 100644 index 0000000000000000000000000000000000000000..fad54a10b32d170b4a71f320de2c1b2c8c406069 --- /dev/null +++ b/fdanyone/geometry/framing.py @@ -0,0 +1,579 @@ +"""Source-aware static-camera framing for the canonical camera ring.""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from dataclasses import asdict, dataclass + +import cv2 +import numpy as np + +from fdanyone.config import FRAMING, FramingConfig +from fdanyone.geometry.cameras import Camera + +FULL_BODY_BOTTOM = 0.82 +CLOSE_UP_BOTTOM = 0.35 +HALF_BODY_BOTTOM = 0.42 +_ANATOMY_COORDS = (0.0, 0.18, 0.45, 0.72, 0.94, 1.0) +_FINGER_TOKENS = ("thumb", "index", "middle", "ring", "pinky") + + +@dataclass(frozen=True) +class InputFraming: + label: str + visible_body_bottom: float + visible_body_bottom_p50: float + visible_body_bottom_p80: float + visible_height_ratio: float + wrist_out_ratio_x: float + wrist_out_ratio_y: float + valid_frame_ratio: float + torso_valid_ratio: float + projection_alignment_error_ratio: float | None + confidence: float + + def to_dict(self) -> dict[str, object]: + return {**asdict(self), "fmask_used": True} + + +@dataclass(frozen=True) +class AdaptiveThresholds: + closeup_strength: float + height_target_ratio: float + height_percentile: float + width_target_ratio: float + width_percentile: float + + +@dataclass(frozen=True) +class RadiusSolve: + radius: float + height_ratio: float + width_ratio: float + bound: str | None + limiting_constraint: str + + +@dataclass(frozen=True) +class FocalSolve: + focal_normalized: float + height_ratio: float + width_ratio: float + bound: str | None + limiting_constraint: str + + +@dataclass(frozen=True) +class SequenceFraming: + radius: float + target_height: float + focal_normalized: float + input: InputFraming + input_applied: bool + radius_solve: RadiusSolve + adaptive_thresholds: AdaptiveThresholds | None = None + focal_solve: FocalSolve | None = None + cutoff_ratio: float | None = None + target_bound: str | None = None + + def to_dict(self) -> dict[str, object]: + return { + "method": "sequence_input_profile_static_radius_target_focal", + "radius": self.radius, + "target_height": self.target_height, + "focal_normalized": self.focal_normalized, + "input_framing_applied": self.input_applied, + "input_framing": self.input.to_dict(), + "adaptive_thresholds": (None if self.adaptive_thresholds is None else asdict(self.adaptive_thresholds)), + "radius_solver": asdict(self.radius_solve), + "focal_solver": None if self.focal_solve is None else asdict(self.focal_solve), + "cutoff_ratio": self.cutoff_ratio, + "target_bound": self.target_bound, + } + + +def _name_index(names: Sequence[str]) -> dict[str, int]: + normalized = [str(name).strip().lower().replace("_", "-") for name in names] + mapping = dict(zip(normalized, range(len(normalized)), strict=True)) + if len(mapping) != len(normalized): + raise ValueError("Keypoint names must be unique after normalization.") + return mapping + + +def _required_indices(names: Sequence[str]) -> dict[str, int]: + required = ["nose", "left-eye", "right-eye", "left-ear", "right-ear", "neck"] + for side in ("left", "right"): + required.extend( + f"{side}-{part}" + for part in ( + "shoulder", + "hip", + "knee", + "ankle", + "big-toe-tip", + "small-toe-tip", + "heel", + ) + ) + mapping = _name_index(names) + missing = [name for name in required if name not in mapping] + if missing: + raise ValueError(f"Missing framing keypoints: {missing}.") + return {name: mapping[name] for name in required} + + +def _sample_path(anchors: Sequence[np.ndarray], samples_per_segment: int) -> tuple[np.ndarray, np.ndarray]: + point_chunks: list[np.ndarray] = [] + coordinate_chunks: list[np.ndarray] = [] + for index, (start, end) in enumerate(zip(anchors, anchors[1:], strict=False)): + endpoint = index == len(anchors) - 2 + weights = np.linspace( + 0.0, + 1.0, + samples_per_segment + int(endpoint), + endpoint=endpoint, + dtype=np.float64, + ) + point_chunks.append(start[:, None] * (1.0 - weights[None, :, None]) + end[:, None] * weights[None, :, None]) + coordinate_chunks.append(_ANATOMY_COORDS[index] * (1.0 - weights) + _ANATOMY_COORDS[index + 1] * weights) + return np.concatenate(point_chunks, axis=1), np.concatenate(coordinate_chunks) + + +def anatomy_samples( + keypoints: np.ndarray, + names: Sequence[str], + samples_per_segment: int = 8, +) -> tuple[np.ndarray, np.ndarray]: + points = np.asarray(keypoints, dtype=np.float64) + if points.ndim != 3 or points.shape[1:] != (len(names), 3) or not np.isfinite(points).all(): + raise ValueError(f"Expected finite keypoints [frames,{len(names)},3], got {points.shape}.") + if samples_per_segment < 2: + raise ValueError("samples_per_segment must be at least two.") + ids = _required_indices(names) + face = np.mean( + points[:, [ids[name] for name in ("nose", "left-eye", "right-eye", "left-ear", "right-ear")]], axis=1 + ) + head_top = face + 0.65 * (face - points[:, ids["neck"]]) + paths: list[np.ndarray] = [] + coordinates: list[np.ndarray] = [] + for side in ("left", "right"): + foot = np.mean( + points[:, [ids[f"{side}-{part}"] for part in ("big-toe-tip", "small-toe-tip", "heel")]], + axis=1, + ) + path, path_coordinates = _sample_path( + [ + head_top, + points[:, ids[f"{side}-shoulder"]], + points[:, ids[f"{side}-hip"]], + points[:, ids[f"{side}-knee"]], + points[:, ids[f"{side}-ankle"]], + foot, + ], + samples_per_segment, + ) + paths.append(path) + coordinates.append(path_coordinates) + return np.concatenate(paths, axis=1), np.concatenate(coordinates) + + +def project_incam(points: np.ndarray, intrinsics: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + points = np.asarray(points, dtype=np.float64) + cameras = np.asarray(intrinsics, dtype=np.float64) + if points.ndim != 3 or points.shape[-1] != 3: + raise ValueError(f"Expected points [frames,keypoints,3], got {points.shape}.") + if cameras.shape == (3, 3): + cameras = np.broadcast_to(cameras, (points.shape[0], 3, 3)) + if cameras.shape == (1, 3, 3): + cameras = np.broadcast_to(cameras, (points.shape[0], 3, 3)) + if cameras.shape != (points.shape[0], 3, 3): + raise ValueError(f"Expected intrinsics [{points.shape[0]},3,3], got {cameras.shape}.") + if not np.isfinite(points).all() or not np.isfinite(cameras).all(): + raise ValueError("Projection inputs must be finite.") + homogeneous = np.einsum("fij,fkj->fki", cameras, points) + depths = points[..., 2] + xy = homogeneous[..., :2] / np.maximum(homogeneous[..., 2:3], 1e-8) + return xy, depths + + +def _inside(xy: np.ndarray, depths: np.ndarray, width: int, height: int) -> np.ndarray: + return (depths > 1e-6) & (xy[..., 0] >= 0) & (xy[..., 0] < width) & (xy[..., 1] >= 0) & (xy[..., 1] < height) + + +def _mask_support(xy: np.ndarray, inside: np.ndarray, masks: np.ndarray) -> np.ndarray: + masks = np.asarray(masks) + if masks.ndim != 3 or masks.shape[0] != xy.shape[0]: + raise ValueError(f"Expected masks [frames,height,width], got {masks.shape}.") + height, width = masks.shape[1:] + patch_radius = max(3, int(round(min(width, height) * 0.005))) + kernel = np.ones((2 * patch_radius + 1,) * 2, dtype=np.uint8) + support = np.zeros_like(inside) + for frame_index, mask in enumerate(masks): + dilated = cv2.dilate((mask >= 64).astype(np.uint8), kernel) + candidates = np.flatnonzero(inside[frame_index]) + if candidates.size: + pixels = np.rint(xy[frame_index, candidates]).astype(np.int64) + pixels[:, 0] = np.clip(pixels[:, 0], 0, width - 1) + pixels[:, 1] = np.clip(pixels[:, 1], 0, height - 1) + support[frame_index, candidates] = dilated[pixels[:, 1], pixels[:, 0]] > 0 + return support + + +def _torso_valid(vitpose: np.ndarray, width: int, height: int) -> np.ndarray: + detector = np.asarray(vitpose, dtype=np.float64) + if detector.ndim != 3 or detector.shape[1:] != (17, 3): + raise ValueError(f"Expected VitPose [frames,17,3], got {detector.shape}.") + valid = ( + (detector[..., 2] >= 0.3) + & (detector[..., 0] >= 0) + & (detector[..., 0] < width) + & (detector[..., 1] >= 0) + & (detector[..., 1] < height) + ) + shoulders = valid[:, 5] & valid[:, 6] + torso = shoulders & (valid[:, 11] | valid[:, 12]) + return shoulders if float(torso.mean()) < 0.5 else torso + + +def _alignment_error( + projected: np.ndarray, + names: Sequence[str], + vitpose: np.ndarray, + width: int, + height: int, +) -> float | None: + coco_names = ( + "nose", + "left-eye", + "right-eye", + "left-ear", + "right-ear", + "left-shoulder", + "right-shoulder", + "left-elbow", + "right-elbow", + "left-wrist", + "right-wrist", + "left-hip", + "right-hip", + "left-knee", + "right-knee", + "left-ankle", + "right-ankle", + ) + mapping = _name_index(names) + if any(name not in mapping for name in coco_names): + return None + detector = np.asarray(vitpose, dtype=np.float64) + valid = ( + (detector[..., 2] >= 0.3) + & (detector[..., 0] >= 0) + & (detector[..., 0] < width) + & (detector[..., 1] >= 0) + & (detector[..., 1] < height) + ) + if not valid.any(): + return None + errors = np.linalg.norm(projected[:, [mapping[name] for name in coco_names]] - detector[..., :2], axis=-1) + return float(np.median(errors[valid]) / height) + + +def analyze_input_framing( + incam_keypoints: np.ndarray, + names: Sequence[str], + intrinsics: np.ndarray, + vitpose: np.ndarray, + masks: np.ndarray, +) -> InputFraming: + masks = np.asarray(masks) + if masks.ndim != 3: + raise ValueError(f"Expected masks [frames,height,width], got {masks.shape}.") + frame_count, height, width = masks.shape + if np.asarray(incam_keypoints).shape[0] != frame_count or np.asarray(vitpose).shape[0] != frame_count: + raise ValueError("Input-framing arrays do not share one frame count.") + + anatomy, coordinates = anatomy_samples(incam_keypoints, names) + anatomy_xy, anatomy_depth = project_incam(anatomy, intrinsics) + support = _mask_support(anatomy_xy, _inside(anatomy_xy, anatomy_depth, width, height), masks) + torso_valid = _torso_valid(vitpose, width, height) + bottoms = np.full(frame_count, np.nan) + visible_heights = np.full(frame_count, np.nan) + for frame_index in range(frame_count): + valid = np.flatnonzero(support[frame_index]) + if valid.size and torso_valid[frame_index]: + bottoms[frame_index] = float(coordinates[valid].max()) + y = anatomy_xy[frame_index, valid, 1] + visible_heights[frame_index] = float(np.clip((y.max() - y.min()) / height, 0.0, 1.0)) + valid_frames = np.isfinite(bottoms) + if not valid_frames.any(): + raise ValueError("No valid frames for input framing analysis.") + + projected, depths = project_incam(incam_keypoints, intrinsics) + mapping = _name_index(names) + wrists = [mapping[name] for name in ("left-wrist", "right-wrist")] + wrist_xy = projected[:, wrists] + wrist_depth = depths[:, wrists] + wrist_out_x = (wrist_depth <= 1e-6) | (wrist_xy[..., 0] < 0) | (wrist_xy[..., 0] >= width) + wrist_out_y = (wrist_depth <= 1e-6) | (wrist_xy[..., 1] < 0) | (wrist_xy[..., 1] >= height) + alignment = _alignment_error(projected, names, vitpose, width, height) + valid_ratio = float(valid_frames.mean()) + torso_ratio = float(torso_valid.mean()) + alignment_score = 0.6 if alignment is None else float(np.clip(1.0 - alignment / 0.15, 0.0, 1.0)) + confidence = float(np.clip(0.45 * valid_ratio + 0.35 * torso_ratio + 0.20 * alignment_score, 0.0, 1.0)) + bottom = float(np.percentile(bottoms[valid_frames], 20.0)) + if bottom >= FULL_BODY_BOTTOM: + label = "full_body" + elif bottom >= HALF_BODY_BOTTOM: + label = "half_body" + else: + label = "close_up" + finite_heights = visible_heights[np.isfinite(visible_heights)] + return InputFraming( + label=label, + visible_body_bottom=round(bottom, 6), + visible_body_bottom_p50=round(float(np.percentile(bottoms[valid_frames], 50.0)), 6), + visible_body_bottom_p80=round(float(np.percentile(bottoms[valid_frames], 80.0)), 6), + visible_height_ratio=round(float(np.percentile(finite_heights, 50.0)) if finite_heights.size else 0.0, 6), + wrist_out_ratio_x=round(float(np.any(wrist_out_x, axis=1).mean()), 6), + wrist_out_ratio_y=round(float(np.any(wrist_out_y, axis=1).mean()), 6), + valid_frame_ratio=round(valid_ratio, 6), + torso_valid_ratio=round(torso_ratio, 6), + projection_alignment_error_ratio=None if alignment is None else round(alignment, 6), + confidence=round(confidence, 6), + ) + + +def _camera_coordinates(points: np.ndarray, cameras: Sequence[Camera]) -> np.ndarray: + world = np.asarray(points, dtype=np.float64) + points_h = np.concatenate([world, np.ones((*world.shape[:-1], 1), dtype=np.float64)], axis=-1) + w2c = np.asarray([camera.world_to_camera for camera in cameras], dtype=np.float64) + return np.einsum("cij,fkj->cfki", w2c[:, :3], points_h) + + +def projected_axis_ratios( + points: np.ndarray, + cameras: Sequence[Camera], + focal_normalized: float, + *, + axis: int, + axis_scale: float = 1.0, +) -> np.ndarray: + camera_points = _camera_coordinates(points, cameras) + depths = camera_points[..., 2] + coordinates = focal_normalized * camera_points[..., axis] / np.maximum(depths, 1e-6) / axis_scale + ratios = coordinates.max(axis=2) - coordinates.min(axis=2) + ratios[np.any(depths <= 1e-6, axis=2)] = np.inf + return ratios.max(axis=0) + + +def _selected(points: np.ndarray, names: Sequence[str], *, exclude_hands: bool, exclude_fingers: bool) -> np.ndarray: + excluded = _FINGER_TOKENS + if exclude_hands: + excluded = ("wrist", *_FINGER_TOKENS) + elif not exclude_fingers: + excluded = () + ids = [index for index, name in enumerate(names) if not any(token in name.lower() for token in excluded)] + return np.asarray(points)[:, ids] + + +def solve_radius( + points: np.ndarray, + names: Sequence[str], + camera_factory: Callable[[float, float], Sequence[Camera]], + aspect_ratio: float, + spec: FramingConfig = FRAMING, +) -> RadiusSolve: + height_points = _selected(points, names, exclude_hands=True, exclude_fingers=True) + width_points = _selected(points, names, exclude_hands=False, exclude_fingers=True) + + def evaluate(radius: float) -> tuple[float, float, float]: + cameras = camera_factory(radius, spec.reference_target_height) + heights = projected_axis_ratios(height_points, cameras, spec.reference_focal_normalized, axis=1) + widths = projected_axis_ratios( + width_points, cameras, spec.reference_focal_normalized, axis=0, axis_scale=aspect_ratio + ) + height = float(np.percentile(heights, spec.height_percentile)) + width = float(np.percentile(widths, spec.width_percentile)) + return max(height / spec.height_target_ratio, width / spec.width_target_ratio), height, width + + def result(radius: float, values: tuple[float, float, float], bound: str | None) -> RadiusSolve: + _, height, width = values + limiting = "width" if width / spec.width_target_ratio > height / spec.height_target_ratio else "height" + return RadiusSolve(radius, height, width, bound, limiting) + + lower_values = evaluate(spec.min_radius) + if lower_values[0] <= 1.0: + return result(spec.min_radius, lower_values, "min") + upper_values = evaluate(spec.max_radius) + if upper_values[0] > 1.0: + return result(spec.max_radius, upper_values, "max") + lower, upper, best = spec.min_radius, spec.max_radius, upper_values + for _ in range(24): + midpoint = (lower + upper) / 2 + values = evaluate(midpoint) + if values[0] > 1.0: + lower = midpoint + else: + upper, best = midpoint, values + if values[0] <= 1.0 and abs(values[0] - 1.0) <= 1e-4: + break + return result(upper, best, None) + + +def adaptive_thresholds(profile: InputFraming) -> AdaptiveThresholds: + strength = float( + np.clip((FULL_BODY_BOTTOM - profile.visible_body_bottom) / (FULL_BODY_BOTTOM - CLOSE_UP_BOTTOM), 0.0, 1.0) + ) + return AdaptiveThresholds(strength, 0.80 + 0.12 * strength, 95.0, 0.90 + 0.20 * strength, 80.0 - 30.0 * strength) + + +def _cutoff_points(points: np.ndarray, coordinates: np.ndarray, cutoff: float) -> np.ndarray: + path_length = points.shape[1] // 2 + output = [] + for offset in (0, path_length): + local_coordinates = coordinates[offset : offset + path_length] + local_points = points[:, offset : offset + path_length] + after = int(np.searchsorted(local_coordinates, cutoff, side="right")) + if after == 0: + output.append(local_points[:, 0]) + elif after >= len(local_coordinates): + output.append(local_points[:, -1]) + else: + before = after - 1 + weight = (cutoff - local_coordinates[before]) / (local_coordinates[after] - local_coordinates[before]) + output.append(local_points[:, before] * (1.0 - weight) + local_points[:, after] * weight) + return np.stack(output, axis=1) + + +def _visible_anatomy( + points: np.ndarray, + names: Sequence[str], + bottom: float, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + samples, coordinates = anatomy_samples(points, names) + core = samples[:, coordinates <= bottom + 1e-8] + cutoff = _cutoff_points(samples, coordinates, bottom) + if not np.any(np.isclose(coordinates, bottom, atol=1e-8)): + core = np.concatenate([core, cutoff], axis=1) + mapping = _name_index(names) + arm_tokens = ("shoulder", "acromion", "elbow", "olecranon", "cubital-fossa", "wrist") + arms = np.asarray(points)[ + :, [index for name, index in mapping.items() if any(token in name for token in arm_tokens)] + ] + return ( + np.asarray(core, dtype=np.float32), + np.asarray(np.concatenate([core, arms], axis=1), dtype=np.float32), + np.asarray(cutoff, dtype=np.float32), + ) + + +def _solve_focal( + core: np.ndarray, + width_points: np.ndarray, + cameras: Sequence[Camera], + thresholds: AdaptiveThresholds, + aspect_ratio: float, + spec: FramingConfig, +) -> tuple[FocalSolve, np.ndarray]: + unit_heights = projected_axis_ratios(core, cameras, 1.0, axis=1) + unit_widths = projected_axis_ratios(width_points, cameras, 1.0, axis=0, axis_scale=aspect_ratio) + unit_height = float(np.percentile(unit_heights, thresholds.height_percentile)) + unit_width = float(np.percentile(unit_widths, thresholds.width_percentile)) + height_focal = thresholds.height_target_ratio / unit_height + width_focal = thresholds.width_target_ratio / unit_width if unit_width > 0 else np.inf + unconstrained = min(height_focal, width_focal) + focal = float(np.clip(unconstrained, spec.reference_focal_normalized, spec.max_focal_normalized)) + bound = ( + "min" + if unconstrained < spec.reference_focal_normalized + else "max" + if unconstrained > spec.max_focal_normalized + else None + ) + return ( + FocalSolve( + focal, + float(np.percentile(unit_heights * focal, thresholds.height_percentile)), + float(np.percentile(unit_widths * focal, thresholds.width_percentile)), + bound, + "width" if width_focal < height_focal else "height", + ), + unit_heights, + ) + + +def _cutoff_ratio(points: np.ndarray, cameras: Sequence[Camera], focal: float, percentile: float) -> float: + camera_points = _camera_coordinates(points, cameras) + positions = 0.5 + focal * camera_points[..., 1] / np.maximum(camera_points[..., 2], 1e-6) + positions[camera_points[..., 2] <= 1e-6] = np.nan + return float(np.percentile(positions[np.isfinite(positions)], percentile)) + + +def solve_sequence_framing( + keypoints: np.ndarray, + names: Sequence[str], + profile: InputFraming, + camera_factory: Callable[[float, float], Sequence[Camera]], + aspect_ratio: float, + spec: FramingConfig = FRAMING, +) -> SequenceFraming: + radius_result = solve_radius(keypoints, names, camera_factory, aspect_ratio, spec) + applied = profile.confidence >= spec.input_min_confidence + thresholds = adaptive_thresholds(profile) if applied else None + if not applied or thresholds is None or thresholds.closeup_strength <= 0: + return SequenceFraming( + radius_result.radius, + spec.reference_target_height, + spec.reference_focal_normalized, + profile, + applied, + radius_result, + thresholds, + ) + + core, width_points, cutoff = _visible_anatomy(keypoints, names, profile.visible_body_bottom) + centers = (core[..., 1].min(axis=1) + core[..., 1].max(axis=1)) / 2 + anatomical_target = float(np.median(centers)) + alignment = min(1.0, 2.0 * thresholds.closeup_strength) + initial_target = spec.reference_target_height * (1.0 - alignment) + anatomical_target * alignment + + def evaluate(target_height: float) -> tuple[FocalSolve, float]: + cameras = camera_factory(radius_result.radius, target_height) + focal, _ = _solve_focal(core, width_points, cameras, thresholds, aspect_ratio, spec) + return focal, _cutoff_ratio(cutoff, cameras, focal.focal_normalized, spec.cutoff_percentile) + + lower, upper = initial_target - 0.5, initial_target + 0.5 + lower_value, upper_value = evaluate(lower), evaluate(upper) + increasing = upper_value[1] > lower_value[1] + if not min(lower_value[1], upper_value[1]) <= spec.cutoff_target_ratio <= max(lower_value[1], upper_value[1]): + if abs(lower_value[1] - spec.cutoff_target_ratio) <= abs(upper_value[1] - spec.cutoff_target_ratio): + target, value, target_bound = lower, lower_value, "min" + else: + target, value, target_bound = upper, upper_value, "max" + else: + target, value, target_bound = lower, lower_value, None + for _ in range(20): + midpoint = (lower + upper) / 2 + candidate = evaluate(midpoint) + error = candidate[1] - spec.cutoff_target_ratio + if abs(error) < abs(value[1] - spec.cutoff_target_ratio): + target, value = midpoint, candidate + if abs(error) <= 1e-4: + break + if (error < 0) == increasing: + lower = midpoint + else: + upper = midpoint + + return SequenceFraming( + radius_result.radius, + target, + value[0].focal_normalized, + profile, + True, + radius_result, + thresholds, + value[0], + value[1], + target_bound, + ) diff --git a/fdanyone/io.py b/fdanyone/io.py new file mode 100644 index 0000000000000000000000000000000000000000..6544384dc4786d08f5b6ea151595c583999aacc0 --- /dev/null +++ b/fdanyone/io.py @@ -0,0 +1,115 @@ +"""Filesystem helpers for crash-safe result publication.""" + +from __future__ import annotations + +import errno +import json +import os +import shutil +import time +import uuid +from contextlib import AbstractContextManager +from pathlib import Path +from typing import Any + +from fdanyone.errors import FourDAnyoneError + +_RETRYABLE_TREE_ERRORS = {errno.EBUSY, errno.ENOTEMPTY, errno.ESTALE} + + +def write_json(path: str | Path, value: object, *, sort_keys: bool = True) -> None: + target = Path(path) + target.parent.mkdir(parents=True, exist_ok=True) + temporary = target.with_name(f".{target.name}.{uuid.uuid4().hex}.tmp") + temporary.write_text(json.dumps(value, indent=2, sort_keys=sort_keys) + "\n") + os.replace(temporary, target) + + +def read_json(path: str | Path | None) -> dict[str, Any] | None: + """Read a JSON object, returning None when there is no such file.""" + + if path is None: + return None + target = Path(path) + if not target.is_file(): + return None + try: + payload = json.loads(target.read_text()) + except json.JSONDecodeError as exc: + raise FourDAnyoneError(f"Cannot parse {target}: {exc}") from exc + if not isinstance(payload, dict): + raise FourDAnyoneError(f"{target} must contain a JSON object.") + return payload + + +def remove_tree( + path: str | Path, + *, + attempts: int = 8, + initial_delay_seconds: float = 0.1, + ignore_errors: bool = False, +) -> None: + """Remove a tree, tolerating short directory-entry lag on network filesystems.""" + + target = Path(path) + if attempts <= 0: + raise ValueError(f"attempts must be positive, got {attempts}.") + for attempt in range(attempts): + try: + shutil.rmtree(target) + return + except FileNotFoundError: + return + except OSError as exc: + retryable = exc.errno in _RETRYABLE_TREE_ERRORS and attempt + 1 < attempts + if not retryable: + if ignore_errors: + return + raise + time.sleep(initial_delay_seconds * (2**attempt)) + + +class AtomicResultDirectory(AbstractContextManager[Path]): + """Build beside the destination and rename only after all validation passes.""" + + def __init__(self, destination: str | Path): + expanded = Path(destination).expanduser() + # Resolve the parent for a stable absolute location, but preserve the + # leaf itself so a dangling output symlink cannot be followed and + # mistaken for a nonexistent destination. + self.destination = expanded.parent.resolve() / expanded.name + self.working = self.destination.with_name(f".{self.destination.name}.work-{uuid.uuid4().hex[:10]}") + self._committed = False + + def _destination_exists(self) -> bool: + return os.path.lexists(self.destination) + + def __enter__(self) -> Path: + if self._destination_exists(): + raise FourDAnyoneError( + f"Output directory already exists: {self.destination}. Choose a new directory to avoid mixed runs." + ) + self.working.mkdir(parents=True) + return self.working + + def commit(self) -> Path: + if self._committed: + return self.destination + if self._destination_exists(): + raise FourDAnyoneError( + f"Output directory appeared during inference: {self.destination}. Refusing to overwrite it." + ) + os.replace(self.working, self.destination) + self._committed = True + return self.destination + + def __exit__(self, exc_type, exc_value, traceback) -> bool: + if exc_type is None: + try: + self.commit() + except BaseException: + remove_tree(self.working, ignore_errors=True) + raise + else: + remove_tree(self.working, ignore_errors=True) + return False diff --git a/fdanyone/model/__init__.py b/fdanyone/model/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..1a13461abad48298b8a4d34beb304f2f3b26251a --- /dev/null +++ b/fdanyone/model/__init__.py @@ -0,0 +1 @@ +"""4DAnyone model loading and stage inference.""" diff --git a/fdanyone/model/inference.py b/fdanyone/model/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..1f9551ca8e358d8f6dcc0985f72b2a07b1d6d490 --- /dev/null +++ b/fdanyone/model/inference.py @@ -0,0 +1,689 @@ +"""Hydra/Lightning-free multi-view generation. + +RCP and final target generation share one source encoding and prompt embedding. +Target groups execute sequentially on one GPU; TCR optionally shifts their +membership between denoising steps. +""" + +from __future__ import annotations + +import gc +import logging +from collections.abc import Callable, Iterable +from concurrent.futures import Future, ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import TYPE_CHECKING, TypeAlias + +import numpy as np +from PIL import Image + +from fdanyone.assets import BaseAssets +from fdanyone.config import INFERENCE, ModeSettings +from fdanyone.errors import FourDAnyoneError +from fdanyone.model.loader import load_pipeline, offload_all, onload_all +from fdanyone.model.prepared import forward_dynamic, prepare_static_conditioning +from fdanyone.model.profiling import CudaStageTimer, StageClock, profile_dit_step +from fdanyone.model.routing import routing_steps +from fdanyone.model.tiny_decoder import decode_tiny_target_video +from fdanyone.skeleton.pipeline import Conditioning, SkeletonVideo +from fdanyone.video import CanonicalClip, write_video, write_video_async +from fdanyone.views import ViewPlan + +LOGGER = logging.getLogger("fdanyone") + +if TYPE_CHECKING: + import torch + +DenoiseStepHook: TypeAlias = Callable[[int, tuple[int, ...], "torch.Tensor"], None] +"""Receives ``(step_index, view_indices, x0_hat)`` for one denoised group.""" + + +@dataclass(frozen=True) +class GeneratedViews: + """Paths and measurements produced by one resolved view plan.""" + + rcp_videos: tuple[Path, ...] + target_videos: tuple[Path, ...] + view_plan: ViewPlan + seed: int + device: str + elapsed_seconds: dict[str, float] + peak_vram_allocated_bytes: int + peak_vram_reserved_bytes: int + + +def _empty_cuda_cache() -> None: + import torch + + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +def _move_model(model, device: str) -> None: + if getattr(model, "vram_management_enabled", False): + onload_all(model) + else: + model.to(device) + _empty_cuda_cache() + + +def _offload_model(model) -> None: + if getattr(model, "vram_management_enabled", False): + offload_all(model) + else: + model.to("cpu") + _empty_cuda_cache() + + +def _bf16_autocast(): + """Match Lightning's ``bf16-mixed`` inference context without Lightning.""" + + import torch + + return torch.autocast(device_type="cuda", dtype=torch.bfloat16) + + +def _channels_last_source_layout(video): + """Preserve the frozen source tensor's VFHWC-backed VCFHW layout.""" + + import torch + + if video.ndim != 5: + raise FourDAnyoneError(f"Expected a 5D video tensor, got shape {tuple(video.shape)}.") + return video.contiguous(memory_format=torch.channels_last_3d) + + +def _encode_prompt(pipe, device: str) -> dict[str, object]: + import torch + + _move_model(pipe.text_encoder, device) + with torch.inference_mode(), _bf16_autocast(): + context = pipe.prompter.encode_prompt(INFERENCE.prompt, positive=True, device=device) + prompt = {"context": context.detach().to("cpu")} + _offload_model(pipe.text_encoder) + return prompt + + +def _load_prompt_embedding(path: Path) -> dict[str, object]: + """Read an exported prompt context, refusing one built for another prompt.""" + + from safetensors import safe_open + + with safe_open(str(path), framework="pt", device="cpu") as handle: + metadata: dict[str, str] = handle.metadata() or {} + stored_prompt: str | None = metadata.get("prompt") + if stored_prompt != INFERENCE.prompt: + raise FourDAnyoneError( + f"Prompt embedding {path} was exported for {stored_prompt!r}, " + f"not the configured prompt {INFERENCE.prompt!r}." + ) + return {"context": handle.get_tensor("context")} + + +def _encode_videos(pipe, videos, device: str): + import torch + + videos = videos.to(dtype=pipe.torch_dtype, device=device) + with torch.inference_mode(), _bf16_autocast(): + latents = pipe.encode_video(videos) + return latents.detach().to("cpu") + + +def _noise(pipe, num_views: int, num_frames: int, seed: int, device: str): + import torch + + shape = ( + num_views, + pipe.vae.model.z_dim, + (num_frames - 1) // 4 + 1, + INFERENCE.height // pipe.vae.upsampling_factor, + INFERENCE.width // pipe.vae.upsampling_factor, + ) + generator = torch.Generator("cpu").manual_seed(seed) + return torch.randn(shape, generator=generator, device="cpu", dtype=torch.float32).to( + dtype=pipe.torch_dtype, device=device + ) + + +def _load_skeleton_cache( + conditioning: Conditioning, + skeletons: Iterable[SkeletonVideo], + cache: dict[SkeletonVideo, object], + device: str, +) -> None: + import torch + + for skeleton in skeletons: + if skeleton in cache: + continue + LOGGER.info("Loading skeleton conditioning from %s", skeleton.path.name) + tensor = conditioning.load_skeleton_tensor( + [skeleton], + device=device, + ).to(dtype=torch.bfloat16, device="cpu") + cache[skeleton] = tensor.contiguous() + + +def _skeleton_group( + cache: dict[SkeletonVideo, object], + skeletons: Iterable[SkeletonVideo], + device: str, + *, + channels_last: bool = False, +): + import torch + + items = tuple(skeletons) + first = cache[items[0]] + memory_format = torch.channels_last_3d if channels_last else torch.contiguous_format + group = torch.empty( + (len(items), *first.shape[1:]), + dtype=first.dtype, + device=device, + memory_format=memory_format, + ) + for output_index, skeleton in enumerate(items): + group[output_index].copy_(cache[skeleton][0]) + return group + + +def _denoise( + pipe, + src_latents, + prompt, + skeletons: tuple[SkeletonVideo, ...], + skeleton_cache: dict[SkeletonVideo, object], + view_plan: ViewPlan | None, + device: str, + seed: int, + stage_timer: CudaStageTimer, + settings: ModeSettings, + *, + on_step: DenoiseStepHook | None = None, +): + """Denoise either one full proposal group or routed target groups.""" + + import torch + from tqdm.auto import tqdm + + num_views: int = len(skeletons) + if view_plan is not None and num_views != view_plan.num_target_views: + raise FourDAnyoneError( + f"Target generation requires {view_plan.num_target_views} skeleton views, got {num_views}." + ) + pipe.scheduler.set_timesteps( + settings.num_inference_steps, + denoising_strength=INFERENCE.denoising_strength, + shift=settings.scheduler_shift, + ) + latents = _noise(pipe, num_views, INFERENCE.num_frames, seed, device) + source = src_latents.to(dtype=pipe.torch_dtype, device=device) + context = {name: value.to(dtype=pipe.torch_dtype, device=device) for name, value in prompt.items()} + if view_plan is None: + group_size: int = num_views + routes = routing_steps( + views_per_layer=num_views, + num_layers=1, + group_size=num_views, + num_steps=settings.num_inference_steps, + enable_tcr=False, + circular=False, + ) + # Preserve this RCP channels_last_3d payload as data. Its layout selects + # the banked CUDA kernel and must not be normalized by the merged path. + prepared_skeletons = _skeleton_group( + skeleton_cache, + skeletons, + device, + channels_last=True, + ) + pose_batch_size: int | None = None + description: str = f"RCP 1-to-{num_views}" + else: + group_size = view_plan.views_per_group + routes = routing_steps( + views_per_layer=view_plan.views_per_layer, + num_layers=view_plan.num_layers, + group_size=group_size, + num_steps=settings.num_inference_steps, + enable_tcr=view_plan.tcr_active, + circular=view_plan.closed_yaw, + ) + prepared_skeletons = tuple(skeleton_cache[skeleton] for skeleton in skeletons) + pose_batch_size = settings.dit_pose_batch_size + description = f"Generate {num_views} target views" + with torch.inference_mode(), _bf16_autocast(): + prepared = prepare_static_conditioning( + pipe.dit, + x_src=source, + context=context["context"], + skeletons=prepared_skeletons, + timesteps=pipe.scheduler.timesteps, + group_size=group_size, + pose_batch_size=pose_batch_size, + stage_timer=stage_timer, + ) + del prepared_skeletons + + profile_next: bool = view_plan is not None + with torch.inference_mode(), _bf16_autocast(): + for step_index, groups in enumerate(tqdm(routes, desc=description)): + timestep = pipe.scheduler.timesteps[step_index] + for view_indices in groups: + with torch.profiler.record_function("dit.route_copies"): + index = torch.tensor(view_indices, dtype=torch.long, device=device) + local_latents = torch.index_select(latents, 0, index) + with profile_dit_step(enabled=profile_next): + prediction = forward_dynamic( + pipe.dit, + x=local_latents, + prepared=prepared, + view_indices=index, + step_index=step_index, + ) + profile_next = False + if on_step is not None: + # Flow matching predicts the noise-to-data velocity, so the + # clean estimate is one full sigma step along it. + sigma = pipe.scheduler.sigmas[step_index].to(local_latents.device) + on_step(step_index, view_indices, local_latents - sigma * prediction) + local_latents = pipe.scheduler.step(prediction, timestep, local_latents) + with torch.profiler.record_function("dit.route_copies"): + latents.index_copy_(0, index, local_latents) + del local_latents, prediction + return latents.detach().to("cpu") + + +def _tensor_frames(video) -> Iterable[np.ndarray]: + """Match DiffSynth's float-to-uint8 truncation exactly.""" + + frames = video.detach().float().add_(1.0).mul_(127.5).clamp_(0.0, 255.0).to("cpu") + for frame_index in range(frames.shape[1]): + yield frames[:, frame_index].permute(1, 2, 0).numpy().astype(np.uint8) + + +def _save_rcp_jpegs(video, camera_id: int, root: Path) -> Path: + import torchvision.transforms.functional as transform + + frame_dir = root / f"{camera_id:06d}" + frame_dir.mkdir(parents=True, exist_ok=False) + normalized = video.detach().float().mul(0.5).add_(0.5).clamp_(0.0, 1.0).to("cpu") + for frame_index in range(normalized.shape[1]): + image = transform.to_pil_image(normalized[:, frame_index]) + image.save(frame_dir / f"{frame_index:06d}.jpg", quality=INFERENCE.rcp_jpeg_quality) + return frame_dir + + +def _rcp_reference_video_layout(frame_first_video): + """Match ``prepare_batch`` for JPEG-backed ``[V,F,C,H,W]`` data.""" + + if frame_first_video.ndim != 5: + raise FourDAnyoneError(f"Expected a 5D frame-first video tensor, got shape {tuple(frame_first_video.shape)}.") + return frame_first_video.permute(0, 2, 1, 3, 4) + + +def _load_rcp_reference_videos(frame_dirs: Iterable[Path], num_frames: int): + """Decode RCP references exactly like the frozen JPEG-backed input path.""" + + import torch + import torchvision.transforms.functional as transform + + videos = [] + for frame_dir in frame_dirs: + frames = [] + for frame_index in range(num_frames): + path = frame_dir / f"{frame_index:06d}.jpg" + with Image.open(path) as image: + frames.append(transform.to_tensor(image.convert("RGB"))) + videos.append(torch.stack(frames, dim=0)) + frame_first_video = torch.stack(videos, dim=0).mul_(2.0).sub_(1.0) + return _rcp_reference_video_layout(frame_first_video) + + +def _decode_view(pipe, latents, device: str, *, decoder=None): + """Decode one full- or tiny-decoder latent view.""" + + import torch + + if decoder is None: + return pipe.decode_video(latents.to(dtype=pipe.torch_dtype, device=device))[0] + return decode_tiny_target_video( + decoder, + latents.to(dtype=torch.float16, device=device), + )[0] + + +def _emit_decoded_videos( + decoded_views: Iterable[tuple[object, Path]], + clip: CanonicalClip, + settings: ModeSettings, +) -> tuple[Path, ...]: + """Write decoded views serially or through the bounded encoder pool.""" + + outputs: list[Path] = [] + encode_futures: list[Future[Path]] = [] + executor: ThreadPoolExecutor | None = ( + ThreadPoolExecutor(max_workers=2) if settings.async_video_encode else None + ) + try: + for video, path in decoded_views: + if executor is None: + outputs.append( + write_video( + _tensor_frames(video), + path, + clip.fps, + crf=INFERENCE.target_h264_crf, + preset=INFERENCE.h264_preset, + ) + ) + else: + outputs.append(path) + encode_futures.append( + write_video_async( + executor, + _tensor_frames(video), + path, + clip.fps, + crf=INFERENCE.target_h264_crf, + preset=INFERENCE.h264_preset, + ) + ) + for future in encode_futures: + future.result() + finally: + if executor is not None: + executor.shutdown(wait=True, cancel_futures=True) + return tuple(outputs) + + +def _decode_rcp( + pipe, + latents, + camera_ids: tuple[int, ...], + output_dir: Path, + clip: CanonicalClip, + device: str, + settings: ModeSettings, + *, + rcp_decoder=None, +) -> tuple[tuple[Path, ...], tuple[Path, ...]]: + import torch + + if latents.shape[0] != len(camera_ids): + raise FourDAnyoneError(f"RCP decode expected {len(camera_ids)} latent views, got {latents.shape[0]}.") + frame_root = output_dir / "frames" + video_root = output_dir / "videos" + frame_root.mkdir(parents=True, exist_ok=False) + video_root.mkdir(parents=True, exist_ok=False) + frame_outputs: list[Path] = [] + + def decoded_views() -> Iterable[tuple[object, Path]]: + with torch.inference_mode(), _bf16_autocast(): + for latent_index, camera_id in enumerate(camera_ids): + LOGGER.info("Decoding RCP camera %02d", camera_id) + camera_latents = latents[latent_index : latent_index + 1] + video = _decode_view( + pipe, + camera_latents, + device, + decoder=rcp_decoder, + ) + if not settings.direct_rcp_latent_handoff: + video_for_output = video.detach().to("cpu") + frame_outputs.append( + _save_rcp_jpegs(video_for_output, camera_id, frame_root) + ) + else: + video_for_output = video + video_path: Path = video_root / f"{camera_id:02d}.mp4" + yield video_for_output, video_path + del video + torch.cuda.empty_cache() + + video_outputs: tuple[Path, ...] = _emit_decoded_videos( + decoded_views(), + clip, + settings, + ) + return tuple(frame_outputs), video_outputs + + +def _decode_targets( + pipe, + latents, + output_dir: Path, + clip: CanonicalClip, + device: str, + settings: ModeSettings, + *, + target_decoder=None, +) -> tuple[Path, ...]: + import torch + + video_root = output_dir / "videos" + video_root.mkdir(parents=True, exist_ok=False) + + def decoded_views() -> Iterable[tuple[object, Path]]: + with torch.inference_mode(), _bf16_autocast(): + for camera_id in range(latents.shape[0]): + LOGGER.info("Decoding target camera %02d", camera_id) + camera_latents = latents[camera_id : camera_id + 1] + video = _decode_view( + pipe, + camera_latents, + device, + decoder=target_decoder, + ) + path: Path = video_root / f"{camera_id:02d}.mp4" + yield video, path + del video + torch.cuda.empty_cache() + + return _emit_decoded_videos(decoded_views(), clip, settings) + + +def generate_views( + *, + clip: CanonicalClip, + conditioning: Conditioning, + checkpoint_path: str | Path, + assets: BaseAssets, + output_dir: str | Path, + device: str, + seed: int, + settings: ModeSettings, + on_denoise_step: DenoiseStepHook | None = None, + prompt_embedding_path: Path | None = None, +) -> GeneratedViews: + """Generate the proposal (when enabled) and the requested target views.""" + + import torch + + if conditioning.num_frames != INFERENCE.num_frames or len(clip.frames) != INFERENCE.num_frames: + raise FourDAnyoneError("Generation requires the frozen 121-frame contract.") + if seed < 0: + raise FourDAnyoneError(f"seed must be non-negative, got {seed}.") + device_index = int(device.removeprefix("cuda:")) + LOGGER.info("Using %s (%s)", device, torch.cuda.get_device_name(device_index)) + + root = Path(output_dir).expanduser().resolve() + root.mkdir(parents=True, exist_ok=False) + view_plan = conditioning.view_plan + if len(conditioning.target_skeletons) != view_plan.num_target_views: + raise FourDAnyoneError("Target skeleton count does not match the resolved view plan.") + if len(conditioning.rcp_skeletons) != len(view_plan.rcp_camera_ids): + raise FourDAnyoneError("RCP skeleton count does not match the resolved view plan.") + loaded = load_pipeline( + checkpoint_path=checkpoint_path, + assets=assets, + device=device, + settings=settings, + load_text_encoder=prompt_embedding_path is None, + ) + pipe = loaded.pipe + clock = StageClock(cuda=CudaStageTimer(device=device)) + skeleton_cache: dict[SkeletonVideo, object] = {} + torch.cuda.reset_peak_memory_stats(device_index) + + with clock.stage("prompt_t5"): + if prompt_embedding_path is None: + prompt = _encode_prompt(pipe, device) + else: + prompt = _load_prompt_embedding(prompt_embedding_path) + loaded.release_text_encoder() + + with clock.stage("source_vae_encode"): + _move_model(pipe.vae, device) + source_video = _channels_last_source_layout(conditioning.load_source_tensor()) + source_latents = _encode_videos(pipe, source_video, device) + del source_video + _offload_model(pipe.vae) + + rcp_videos: tuple[Path, ...] = () + target_sources = source_latents + if view_plan.enable_rcp: + with clock.stage("rcp_prepare"): + _load_skeleton_cache( + conditioning, + conditioning.rcp_skeletons, + skeleton_cache, + device, + ) + _move_model(pipe.dit, device) + with clock.stage("rcp_dit"): + rcp_latents = _denoise( + pipe, + source_latents, + prompt, + conditioning.rcp_skeletons, + skeleton_cache, + None, + device, + seed, + clock.cuda, + settings, + ) + clock.elapsed["rcp_pose_encoder"] = clock.cuda.elapsed_seconds("pose_encoder") + with clock.stage("rcp_offload"): + _offload_model(pipe.dit) + + with clock.stage("rcp_decode_jpeg_reencode"): + rcp_root = root / "rcp" + rcp_root.mkdir() + rcp_decoder = ( + loaded.target_decoder if settings.tiny_decoders else None + ) + decode_model = pipe.vae if rcp_decoder is None else rcp_decoder + _move_model(decode_model, device) + frame_dirs, rcp_videos = _decode_rcp( + pipe, + rcp_latents, + view_plan.rcp_camera_ids, + rcp_root, + clip, + device, + settings, + rcp_decoder=rcp_decoder, + ) + if rcp_decoder is not None: + _offload_model(rcp_decoder) + if settings.direct_rcp_latent_handoff: + selected_rcp_latents = rcp_latents + target_sources = torch.cat( + [source_latents, selected_rcp_latents], + dim=0, + ) + else: + if rcp_decoder is not None: + _move_model(pipe.vae, device) + rcp_reference_videos = _load_rcp_reference_videos( + frame_dirs[:4], + INFERENCE.num_frames, + ) + rcp_reference_latents = _encode_videos( + pipe, + rcp_reference_videos, + device, + ) + selected_rcp_latents = rcp_reference_latents + # VAE38 encodes batch elements independently, so the existing + # source encoding matches re-encoding source plus references. + target_sources = torch.cat( + [source_latents, selected_rcp_latents], + dim=0, + ) + del rcp_reference_videos, rcp_reference_latents + del rcp_latents + _offload_model(pipe.vae) + + conditioning.wait_for_target_skeletons() + with clock.stage("target_prepare"): + _load_skeleton_cache( + conditioning, + conditioning.target_skeletons, + skeleton_cache, + device, + ) + _move_model(pipe.dit, device) + with clock.stage("target_dit"): + target_latents = _denoise( + pipe, + target_sources, + prompt, + conditioning.target_skeletons, + skeleton_cache, + view_plan, + device, + seed, + clock.cuda, + settings, + on_step=on_denoise_step, + ) + # One timer spans both stages, so the target share is what RCP did not spend. + pose_encoder_seconds = clock.cuda.elapsed_seconds("pose_encoder") + clock.elapsed["target_pose_encoder"] = pose_encoder_seconds - clock.elapsed.get("rcp_pose_encoder", 0.0) + clock.elapsed["pose_encoder"] = pose_encoder_seconds + with clock.stage("target_offload"): + _offload_model(pipe.dit) + # The decoded skeleton tensors (~650 MB per view) have no reader past the + # target DiT stage; release them before the decode/encode phase allocates. + skeleton_cache.clear() + + with clock.stage("target_vae_decode"): + target_root = root / "target" + target_root.mkdir() + tiny_target_decoder = ( + loaded.target_decoder if settings.tiny_decoders else None + ) + target_decoder = ( + pipe.vae if tiny_target_decoder is None else tiny_target_decoder + ) + _move_model(target_decoder, device) + target_videos = _decode_targets( + pipe, + target_latents, + target_root, + clip, + device, + settings, + target_decoder=tiny_target_decoder, + ) + _offload_model(target_decoder) + peak_vram_allocated = int(torch.cuda.max_memory_allocated(device_index)) + peak_vram_reserved = int(torch.cuda.max_memory_reserved(device_index)) + + del target_latents, target_sources, source_latents, prompt, skeleton_cache, loaded, pipe + _empty_cuda_cache() + return GeneratedViews( + rcp_videos=rcp_videos, + target_videos=target_videos, + view_plan=view_plan, + seed=seed, + device=device, + elapsed_seconds=clock.elapsed, + peak_vram_allocated_bytes=peak_vram_allocated, + peak_vram_reserved_bytes=peak_vram_reserved, + ) diff --git a/fdanyone/model/loader.py b/fdanyone/model/loader.py new file mode 100644 index 0000000000000000000000000000000000000000..0e6e90016fd3d70bf3964cff634129b5049c4297 --- /dev/null +++ b/fdanyone/model/loader.py @@ -0,0 +1,426 @@ +"""Direct, registry-free loading of the frozen Wan/SpaTem inference stack.""" + +from __future__ import annotations + +import gc +import logging +import os +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from fdanyone.assets import BaseAssets +from fdanyone.config import ( + MVS_ATTENTION_RANGE, + POSE_ENCODER_TYPE, + USE_POSE_ENCODER, + USE_VIEWPACK, + ModeSettings, +) +from fdanyone.errors import AssetError, ConfigurationError + +LOGGER = logging.getLogger("fdanyone") + +WAN22_TI2V_5B_CONFIG = { + "has_image_input": False, + "patch_size": (1, 2, 2), + "in_dim": 48, + "dim": 3072, + "ffn_dim": 14336, + "freq_dim": 256, + "text_dim": 4096, + "out_dim": 48, + "num_heads": 24, + "num_layers": 30, + "eps": 1e-6, + "seperated_timestep": True, + "require_clip_embedding": False, + "require_vae_embedding": False, + "fuse_vae_embedding_in_latents": True, +} + + +@dataclass +class LoadedPipeline: + pipe: object + target_decoder: object | None = None + + def release_text_encoder(self) -> None: + """Release T5 after the single fixed prompt has been encoded.""" + + import torch + + self.pipe.text_encoder = None + self.pipe.prompter.text_encoder = None + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +def _load_checkpoint(path: Path): + try: + from safetensors.torch import load_file + except ImportError as exc: + raise AssetError("safetensors is required to load the 4DAnyone checkpoint.") from exc + return load_file(str(path), device="cpu") + + +def _strict_assign(module, state_dict: dict, label: str) -> None: + """Load into a meta-initialized module without a second parameter copy.""" + + try: + incompatible = module.load_state_dict(state_dict, strict=True, assign=True) + except TypeError as exc: + raise ConfigurationError("4DAnyone requires PyTorch >=2.8 for assign-based model loading.") from exc + except RuntimeError as exc: + raise AssetError(f"{label} is incompatible with the released architecture: {exc}") from exc + if incompatible.missing_keys or incompatible.unexpected_keys: + raise AssetError( + f"{label} strict load failed; missing={incompatible.missing_keys}, " + f"unexpected={incompatible.unexpected_keys}" + ) + + +def _load_dit( + checkpoint_path: Path, + dtype, + *, + settings: ModeSettings, + turbo_lora_path: Path | None, +): + import torch + + from fdanyone.vendor.diffsynth.models.wan_video_dit import ( + AttentionModule, + DiTBlock, + WanModel, + precompute_freqs_cis_3d, + ) + from fdanyone.vendor.diffsynth.pipelines.wan_video_spatem import ( + WanVideoSpaTemPipeline, + ) + + with torch.device("meta"): + dit = WanModel(**WAN22_TI2V_5B_CONFIG) + shell = WanVideoSpaTemPipeline(device="cpu", torch_dtype=dtype) + shell.dit = dit + shell.init_spatem_modules( + use_mvs_attn=True, + range_mvs_attn=MVS_ATTENTION_RANGE, + # ``ViewPack`` is the upstream module name for RCP references. + use_viewpack=USE_VIEWPACK, + use_pose_encoder=USE_POSE_ENCODER, + pose_encoder_type=POSE_ENCODER_TYPE, + ) + state_dict = _load_checkpoint(checkpoint_path) + _strict_assign(dit, state_dict, "4DAnyone DiT checkpoint") + del state_dict + for module in dit.modules(): + if isinstance(module, DiTBlock): + module.use_bf16_block_glue = settings.bf16_block_glue + if isinstance(module, AttentionModule): + module.exact = settings.exact_attention + # ``freqs`` is a derived, non-persistent tensor and therefore is not in the + # state dict populated above. + dit.freqs = precompute_freqs_cis_3d(WAN22_TI2V_5B_CONFIG["dim"] // WAN22_TI2V_5B_CONFIG["num_heads"]) + legacy_backend: str | None = os.environ.get("FDANYONE_ATTENTION_BACKEND") + if legacy_backend: + assert legacy_backend.lower() in {"sdpa", "sageattention"} + LOGGER.warning("FDANYONE_ATTENTION_BACKEND is obsolete and ignored; --mode owns attention.") + if turbo_lora_path is not None: + from fdanyone.model.turbo_lora import merge_wan_turbo_lora + + report = merge_wan_turbo_lora( + dit, + turbo_lora_path, + ) + LOGGER.info( + "Turbo-LoRA merged: modules=%d direct_parameters=%d max_rank=%d", + report.lora_modules, + report.direct_parameters, + report.max_rank, + ) + return dit.eval().requires_grad_(False) + + +def _load_vae(path: Path, dtype): + import torch + + from fdanyone.vendor.diffsynth.models.wan_video_vae import WanVideoVAE38 + + state_dict = torch.load(path, map_location="cpu", weights_only=True) + state_dict = WanVideoVAE38.state_dict_converter().from_civitai(state_dict) + with torch.device("meta"): + vae = WanVideoVAE38() + _strict_assign(vae, state_dict, "Wan2.2 VAE") + del state_dict + # Wan's latent normalization tensors are plain attributes rather than + # registered buffers, so materialize them after meta initialization. + mean = ( + -0.2289, + -0.0052, + -0.1323, + -0.2339, + -0.2799, + 0.0174, + 0.1838, + 0.1557, + -0.1382, + 0.0542, + 0.2813, + 0.0891, + 0.1570, + -0.0098, + 0.0375, + -0.1825, + -0.2246, + -0.1207, + -0.0698, + 0.5109, + 0.2665, + -0.2108, + -0.2158, + 0.2502, + -0.2055, + -0.0322, + 0.1109, + 0.1567, + -0.0729, + 0.0899, + -0.2799, + -0.1230, + -0.0313, + -0.1649, + 0.0117, + 0.0723, + -0.2839, + -0.2083, + -0.0520, + 0.3748, + 0.0152, + 0.1957, + 0.1433, + -0.2944, + 0.3573, + -0.0548, + -0.1681, + -0.0667, + ) + std = ( + 0.4765, + 1.0364, + 0.4514, + 1.1677, + 0.5313, + 0.4990, + 0.4818, + 0.5013, + 0.8158, + 1.0344, + 0.5894, + 1.0901, + 0.6885, + 0.6165, + 0.8454, + 0.4978, + 0.5759, + 0.3523, + 0.7135, + 0.6804, + 0.5833, + 1.4146, + 0.8986, + 0.5659, + 0.7069, + 0.5338, + 0.4889, + 0.4917, + 0.4069, + 0.4999, + 0.6866, + 0.4093, + 0.5709, + 0.6065, + 0.6415, + 0.4944, + 0.5726, + 1.2042, + 0.5458, + 1.6887, + 0.3971, + 1.0600, + 0.3943, + 0.5537, + 0.5444, + 0.4089, + 0.7468, + 0.7744, + ) + vae.mean = torch.tensor(mean) + vae.std = torch.tensor(std) + vae.scale = [vae.mean, 1.0 / vae.std] + return vae.to(dtype=dtype).eval().requires_grad_(False) + + +def _load_text_encoder(path: Path, dtype): + import torch + + from fdanyone.vendor.diffsynth.models.wan_video_text_encoder import WanTextEncoder + + state_dict = torch.load(path, map_location="cpu", weights_only=True) + state_dict = WanTextEncoder.state_dict_converter().from_civitai(state_dict) + with torch.device("meta"): + text_encoder = WanTextEncoder() + _strict_assign(text_encoder, state_dict, "Wan T5 text encoder") + del state_dict + return text_encoder.to(dtype=dtype).eval().requires_grad_(False) + + +def onload_all(model: Any) -> None: + """Move every VRAM-managed submodule of ``model`` to its compute device.""" + + for module in model.modules(): + if hasattr(module, "onload"): + module.onload() + + +def offload_all(model: Any) -> None: + """Move every VRAM-managed submodule of ``model`` back to host memory.""" + + for module in model.modules(): + if hasattr(module, "offload"): + module.offload() + + +def _enable_dit_streaming( + pipe: Any, + *, + device: str, + persistent_parameters: int, +) -> None: + """Keep a bounded prefix of DiT weights on the GPU and stream the rest.""" + + import torch + + from fdanyone.vendor.diffsynth.models.wan_video_dit import RMSNorm + from fdanyone.vendor.diffsynth.vram_management import ( + AutoWrappedLinear, + AutoWrappedModule, + enable_vram_management, + ) + + # The vendored wrapper only owns registered child modules. Stage the full + # model first so standalone parameters (for example block modulation) stay + # on the compute device when wrapped modules are moved back to the CPU. + pipe.dit.to(device=device) + dtype: torch.dtype = next(iter(pipe.dit.parameters())).dtype + module_config: dict[str, Any] = { + "offload_dtype": dtype, + "offload_device": "cpu", + "onload_dtype": dtype, + "onload_device": device, + "computation_dtype": pipe.torch_dtype, + "computation_device": device, + } + enable_vram_management( + pipe.dit, + module_map={ + torch.nn.Linear: AutoWrappedLinear, + # Every DiT ``Conv3d`` together (patch embedding, view-pack + # projections, pose encoder) is 15.9M parameters -- 0.3% of the + # model, 30 MiB in bf16. Streaming them buys nothing and forces + # callers to read layer metadata such as ``kernel_size`` through + # the wrapper, so they stay resident and unwrapped. + torch.nn.LayerNorm: AutoWrappedModule, + RMSNorm: AutoWrappedModule, + }, + module_config=module_config, + max_num_param=persistent_parameters, + overflow_module_config={**module_config, "onload_device": "cpu"}, + ) + onload_all(pipe.dit) + torch.cuda.empty_cache() + + +def _enable_regional_compile(dit: Any, *, mode: str = "default") -> None: + """Compile repeated transformer blocks in place with a static-shape contract. + + ``Module.compile`` swaps only the block's call implementation, so module + identity, class, and fully qualified parameter names survive. A rebuilt + ``ModuleList`` of wrappers would rename every DiT weight. + """ + + for block in dit.blocks: + block.compile(dynamic=False, mode=mode) + + +def load_pipeline( + *, + checkpoint_path: str | Path, + assets: BaseAssets, + device: str, + settings: ModeSettings, + load_text_encoder: bool = True, +) -> LoadedPipeline: + """Load exactly the runtime models required by the released checkpoint.""" + + import torch + + from fdanyone.vendor.diffsynth.pipelines.wan_video_spatem import ( + WanVideoSpaTemPipeline, + ) + + dtype = torch.bfloat16 + if load_text_encoder and assets.text_encoder is None: + raise AssetError( + "Text encoder does not exist. Supply an exported prompt embedding through " + "prompt_embedding_path, or run `python scripts/download_model.py` to install it." + ) + tokenizer_path: str | None = None if assets.tokenizer is None else str(assets.tokenizer) + pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=dtype, tokenizer_path=tokenizer_path) + pipe.dit = _load_dit( + Path(checkpoint_path), + dtype, + settings=settings, + turbo_lora_path=assets.turbo_lora, + ) + if settings.fp8_w8a8: + from fdanyone.model.quantization import quantize_dit_fp8_w8a8 + + pipe.dit.to(device=device) + w8a8_report = quantize_dit_fp8_w8a8(pipe.dit) + LOGGER.info( + "FP8 W8A8 DiT: modules=%d granularity=PerTensor", + w8a8_report.module_count, + ) + if settings.regional_compile: + _enable_regional_compile( + pipe.dit, + mode="max-autotune-no-cudagraphs", + ) + LOGGER.info( + "Regional DiT compile: blocks=%d dynamic=False mode=%s", + len(pipe.dit.blocks), + "max-autotune-no-cudagraphs", + ) + pipe.vae = _load_vae(assets.vae, dtype) + target_decoder: object | None = None + if settings.tiny_decoders: + from fdanyone.model.tiny_decoder import load_tiny_wan_decoder + + if assets.tiny_decoder is None: + raise ConfigurationError("TAEW2.2 checkpoint path was not resolved.") + target_decoder = load_tiny_wan_decoder(assets.tiny_decoder) + if load_text_encoder: + pipe.text_encoder = _load_text_encoder(assets.text_encoder, dtype) + pipe.prompter.fetch_models(pipe.text_encoder) + pipe.height_division_factor = pipe.vae.upsampling_factor * 2 + pipe.width_division_factor = pipe.vae.upsampling_factor * 2 + persistent_parameters: int | None = 0 if settings.stream_dit_weights else None + if persistent_parameters is not None: + _enable_dit_streaming( + pipe, + device=device, + persistent_parameters=persistent_parameters, + ) + return LoadedPipeline(pipe=pipe, target_decoder=target_decoder) diff --git a/fdanyone/model/prepared.py b/fdanyone/model/prepared.py new file mode 100644 index 0000000000000000000000000000000000000000..835fc58319c67f9992ff311fc25af023c62f6783 --- /dev/null +++ b/fdanyone/model/prepared.py @@ -0,0 +1,247 @@ +"""Stage-static conditioning for the Wan DiT. + +``prepare_static_conditioning`` computes every tensor that does not change +between denoising steps of one generation stage; ``forward_dynamic`` then runs +only the noisy-latent patching, the route gather, the transformer blocks, and +the head. Both are exact refactors of ``WanModel.forward`` for the released +packed-source configuration, and ``tests/test_prepared_parity.py`` pins that +equivalence bit for bit. +""" + +from __future__ import annotations + +from contextlib import nullcontext +from dataclasses import dataclass + +import torch +from einops import rearrange, repeat + +from fdanyone.vendor.diffsynth.models.wan_video_dit import ( + WanModel, + pack_viewpack_tokens, + sinusoidal_embedding_1d, + split_source_views, +) + + +@dataclass(frozen=True) +class PreparedCrossAttention: + """Prompt keys and values projected for one transformer block.""" + + key: torch.Tensor + value: torch.Tensor + + +@dataclass(frozen=True) +class PreparedWanConditioning: + """All stage-static inputs consumed by ``forward_dynamic``.""" + + cross_attention: tuple[PreparedCrossAttention, ...] + source_tokens: torch.Tensor + pose_tokens: torch.Tensor | None + null_pose_tokens: torch.Tensor | None + temporal_freqs: torch.Tensor + multiview_freqs: torch.Tensor + time_embeddings: tuple[torch.Tensor, ...] + time_modulations: tuple[torch.Tensor, ...] + grid_size: tuple[int, int, int] + group_size: int + packed_views: int + + +def _prepare_source_tokens( + model: WanModel, + x_src: torch.Tensor, +) -> tuple[torch.Tensor, tuple[int, int, int], int]: + """Patch and pack fixed source/reference latents once.""" + + x_src, x_src_2x, x_src_4x = split_source_views(x_src) + source_tokens, grid = model.patchify(x_src) + packed = [source_tokens] + # ``WanModel.forward`` crops the packed tile to the noisy-latent grid; + # ``forward_dynamic`` rejects any group whose grid differs from this source + # grid, so the two crops are the same. + grid_size = tuple(int(size) for size in grid) + if model.use_viewpack and x_src_2x is not None: + x_pack = pack_viewpack_tokens( + model.viewpack_embedding, x_src_2x, x_src_4x, grid_size + ) + packed.append(x_pack.to(source_tokens.dtype)) + return torch.cat(packed, dim=0), grid_size, len(packed) + + +def _prepare_prompt_kv( + model: WanModel, + context: torch.Tensor, +) -> tuple[PreparedCrossAttention, ...]: + """Project the fixed prompt into per-block keys and values.""" + + return tuple( + PreparedCrossAttention( + key=block.cross_attn.norm_k(block.cross_attn.k(context)), + value=block.cross_attn.v(context), + ) + for block in model.blocks + ) + + +def _prepare_pose_tokens( + model: WanModel, + *, + skeletons: torch.Tensor | tuple[torch.Tensor, ...], + packed_views: int, + group_size: int, + pose_batch_size: int | None, + device: torch.device, + stage_timer, +) -> tuple[torch.Tensor, torch.Tensor]: + """Encode every skeleton view once, in bounded activation batches. + + ``skeletons`` is either one stacked ``[v, c, f, h, w]`` tensor or a tuple + of per-view ``[1, c, f, h, w]`` tensors; the tuple form lets callers feed + cached host views without materializing the whole set as one copy. + """ + + views = ( + tuple(skeletons.split(1, dim=0)) + if isinstance(skeletons, torch.Tensor) + else tuple(skeletons) + ) + null_view = -torch.ones_like(views[0]) + num_real_views = len(views) + num_pose_views = num_real_views + packed_views + batch_size = group_size + packed_views if pose_batch_size is None else pose_batch_size + encoded_pose = None + timer = nullcontext() if stage_timer is None else stage_timer.measure("pose_encoder") + with torch.profiler.record_function("dit.pose_encoder"), timer: + for start in range(0, num_pose_views, batch_size): + end = min(start + batch_size, num_pose_views) + parts = list(views[start:min(end, num_real_views)]) + parts.extend([null_view] * (end - max(start, num_real_views))) + pose_batch = parts[0] if len(parts) == 1 else torch.cat(parts, dim=0) + encoded_batch = rearrange( + model.pose_encoder(pose_batch.to(device=device)), "v c f h w -> v (f h w) c" + ) + if encoded_pose is None: + encoded_pose = encoded_batch.new_empty( + (num_pose_views, *encoded_batch.shape[1:]) + ) + encoded_pose[start:end].copy_(encoded_batch) + if encoded_pose is None: + raise ValueError("Prepared pose conditioning requires at least one view.") + return encoded_pose[:-packed_views], encoded_pose[-packed_views:] + + +def prepare_static_conditioning( + model: WanModel, + *, + x_src: torch.Tensor, + context: torch.Tensor, + skeletons: torch.Tensor | tuple[torch.Tensor, ...] | None, + timesteps: torch.Tensor, + group_size: int, + pose_batch_size: int | None = None, + stage_timer=None, +) -> PreparedWanConditioning: + """Compute every stage-invariant tensor once for dynamic denoising.""" + + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}.") + if pose_batch_size is not None and pose_batch_size <= 0: + raise ValueError(f"pose_batch_size must be positive or None, got {pose_batch_size}.") + if model.has_image_input: + raise ValueError("Prepared conditioning supports the released packed-source path only.") + + context = model.text_embedding(context) + source_tokens, grid_size, packed_views = _prepare_source_tokens(model, x_src) + f, h, w = grid_size + effective_views = group_size + packed_views + device = source_tokens.device + pose_tokens = None + null_pose_tokens = None + if model.use_pose_encoder: + if skeletons is None: + raise ValueError("Prepared pose conditioning requires skeleton tensors.") + pose_tokens, null_pose_tokens = _prepare_pose_tokens( + model, + skeletons=skeletons, + packed_views=packed_views, + group_size=group_size, + pose_batch_size=pose_batch_size, + device=device, + stage_timer=stage_timer, + ) + + time_embeddings = [] + time_modulations = [] + for timestep in timesteps.flatten(): + batched = timestep.reshape(1).to(dtype=source_tokens.dtype, device=device) + batched = torch.cat([batched] * group_size, dim=0) + batched = torch.cat( + [batched, torch.zeros(packed_views, dtype=batched.dtype, device=batched.device)] + ) + embedding = model.time_embedding( + sinusoidal_embedding_1d(model.freq_dim, batched).to(source_tokens.dtype) + ) + time_embeddings.append(embedding) + time_modulations.append(model.time_projection(embedding).unflatten(1, (6, model.dim))) + + return PreparedWanConditioning( + cross_attention=_prepare_prompt_kv( + model, repeat(context, "1 l c -> v l c", v=effective_views) + ), + source_tokens=source_tokens, + pose_tokens=pose_tokens, + null_pose_tokens=null_pose_tokens, + temporal_freqs=model._rope_table(f, h, w, device), + multiview_freqs=model._rope_table(effective_views, h, w, device), + time_embeddings=tuple(time_embeddings), + time_modulations=tuple(time_modulations), + grid_size=grid_size, + group_size=group_size, + packed_views=packed_views, + ) + + +def forward_dynamic( + model: WanModel, + *, + x: torch.Tensor, + prepared: PreparedWanConditioning, + view_indices: torch.Tensor, + step_index: int, +) -> torch.Tensor: + """Run only noisy-latent patching, route gather, blocks, and head.""" + + if x.shape[0] != prepared.group_size or view_indices.numel() != prepared.group_size: + raise ValueError("Dynamic group shape does not match prepared conditioning.") + x, grid_size = model.patchify(x) + if tuple(int(size) for size in grid_size) != prepared.grid_size: + raise ValueError("Dynamic latent grid does not match prepared conditioning.") + x = torch.cat([x, prepared.source_tokens], dim=0) + if prepared.pose_tokens is not None: + if prepared.null_pose_tokens is None: + raise ValueError("Prepared null-pose tokens are missing.") + x[: prepared.group_size].add_( + torch.index_select(prepared.pose_tokens, 0, view_indices) + ) + x[prepared.group_size :].add_(prepared.null_pose_tokens) + + t = prepared.time_embeddings[step_index] + t_mod = prepared.time_modulations[step_index] + v = x.shape[0] + for block, cross_attention in zip(model.blocks, prepared.cross_attention): + x = block( + x, + None, + t_mod, + prepared.temporal_freqs, + prepared.multiview_freqs, + (v, *prepared.grid_size), + cross_attention_key=cross_attention.key, + cross_attention_value=cross_attention.value, + ) + + x = x[: -prepared.packed_views] + t = t[: -prepared.packed_views] + return model.unpatchify(model.head(x, t), prepared.grid_size) diff --git a/fdanyone/model/profiling.py b/fdanyone/model/profiling.py new file mode 100644 index 0000000000000000000000000000000000000000..73327e6213fa3a8a7f5d1e5ca8301c36a542ec0f --- /dev/null +++ b/fdanyone/model/profiling.py @@ -0,0 +1,147 @@ +"""Low-overhead CUDA stage timing and one-step DiT profiling.""" + +from __future__ import annotations + +import json +import os +import time +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass, field +from pathlib import Path +from typing import Protocol, cast + + +class _CudaEvent(Protocol): + def record(self) -> None: ... + + def elapsed_time(self, end_event: _CudaEvent) -> float: ... + + +@dataclass +class CudaStageTimer: + """Accumulate named CUDA event pairs and synchronize only when read.""" + + device: str + _events: dict[str, list[tuple[_CudaEvent, _CudaEvent]]] = field(default_factory=dict) + + @contextmanager + def measure(self, name: str) -> Iterator[None]: + """Record one asynchronous CUDA interval under ``name``.""" + + import torch + + start = cast(_CudaEvent, torch.cuda.Event(enable_timing=True)) + end = cast(_CudaEvent, torch.cuda.Event(enable_timing=True)) + start.record() + try: + yield + finally: + end.record() + self._events.setdefault(name, []).append((start, end)) + + def elapsed_seconds(self, name: str) -> float: + """Return the sum of all intervals recorded under ``name``.""" + + import torch + + intervals = self._events.get(name, []) + if not intervals: + return 0.0 + torch.cuda.synchronize(self.device) + return sum(start.elapsed_time(end) for start, end in intervals) / 1000.0 + + +@dataclass +class StageClock: + """One run's named wall-clock stage totals beside its CUDA stage timer.""" + + cuda: CudaStageTimer + elapsed: dict[str, float] = field(default_factory=dict) + + @contextmanager + def stage(self, name: str) -> Iterator[None]: + """Record the wall-clock duration of the enclosed stage under ``name``.""" + + started = time.monotonic() + try: + yield + finally: + self.elapsed[name] = time.monotonic() - started + + +_PROFILE_LABELS = ( + "dit.temporal_attention", + "dit.multiview_attention", + "dit.cross_attention", + "dit.ffn", + "dit.rope", + "dit.pose_encoder", + "dit.route_copies", +) + + +def _event_time(event: object, device: bool) -> float: + """Read a profiler event duration in microseconds across torch versions.""" + + candidates = ("device_time_total", "cuda_time_total") if device else ("cpu_time_total",) + for attribute in candidates: + value = getattr(event, attribute, None) + if value is not None: + return float(value) + return 0.0 + + +def _stage_summary(event: object | None) -> dict[str, float | int]: + """Summarize one profiled named range, or report an unrecorded stage.""" + + if event is None: + return {"calls": 0, "cpu_seconds": 0.0, "device_seconds": 0.0} + return { + "calls": int(getattr(event, "count", 0)), + "cpu_seconds": _event_time(event, device=False) / 1_000_000.0, + "device_seconds": _event_time(event, device=True) / 1_000_000.0, + } + + +@contextmanager +def profile_dit_step(*, enabled: bool = True) -> Iterator[None]: + """Profile one call when enabled and ``FDANYONE_DIT_PROFILE_DIR`` is set.""" + + output_dir: str | None = os.environ.get("FDANYONE_DIT_PROFILE_DIR") + if not enabled or output_dir is None: + yield + return + + import torch + + destination = Path(output_dir).expanduser().resolve() + destination.mkdir(parents=True, exist_ok=True) + activities = [torch.profiler.ProfilerActivity.CPU] + if torch.cuda.is_available(): + activities.append(torch.profiler.ProfilerActivity.CUDA) + with torch.profiler.profile( + activities=activities, + record_shapes=True, + profile_memory=True, + with_stack=False, + ) as profile: + yield + + profile.export_chrome_trace(str(destination / "dit-step-trace.json")) + averages = profile.key_averages() + (destination / "dit-step-table.txt").write_text( + averages.table( + sort_by="self_cuda_time_total" if torch.cuda.is_available() else "self_cpu_time_total", + row_limit=-1, + ) + ) + events = list(averages) + by_key = {event.key: event for event in events} + stages = {label: _stage_summary(by_key.get(label)) for label in _PROFILE_LABELS} + (destination / "dit-step-summary.json").write_text( + json.dumps( + {"profiled_steps": 1, "stages": stages}, + indent=2, + ) + ) diff --git a/fdanyone/model/quantization.py b/fdanyone/model/quantization.py new file mode 100644 index 0000000000000000000000000000000000000000..a9ac8f69b2c7df69d641a138a9c5bde94589010c --- /dev/null +++ b/fdanyone/model/quantization.py @@ -0,0 +1,87 @@ +"""Selective quantization policies for the 4DAnyone DiT.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass + +import torch + +_FP8_WEIGHT_ONLY_PATTERN: re.Pattern[str] = re.compile( + r"blocks\.(?P\d+)\." + r"(?:(?:self_attn|self_attn_mvs|cross_attn)\.(?:q|k|v|o)|" + r"ffn\.(?:0|2))" +) + +QUANTIZED_BLOCK_RANGE: tuple[int, int] = (3, 26) +"""Inclusive interior DiT block range the guide approves for FP8 projections.""" + + +def is_quantizable_dit_projection(fqn: str) -> bool: + """Return whether an FQN is on the safe interior-projection allowlist. + + W8A8 uses the same boundary-protected projection set established by the + earlier weight-only experiment; the name describes the surviving policy. + """ + + match: re.Match[str] | None = _FP8_WEIGHT_ONLY_PATTERN.fullmatch(fqn) + if match is None: + return False + block_index: int = int(match.group("block")) + first_block, last_block = QUANTIZED_BLOCK_RANGE + return first_block <= block_index <= last_block + + +@dataclass(frozen=True, slots=True) +class FP8W8A8Report: + """Summary of one TorchAO dynamic W8A8 conversion.""" + + module_count: int + """Number of guide-approved Linear modules passed to TorchAO.""" + + +def _cast_w8a8_activation_to_bf16( + module: torch.nn.Module, + args: tuple[object, ...], +) -> tuple[object, ...] | None: + """Restore the BF16 autocast input contract at TorchAO's subclass seam.""" + + del module + if not args: + return None + activation = args[0] + if ( + not isinstance(activation, torch.Tensor) + or not activation.is_floating_point() + or activation.dtype is torch.bfloat16 + ): + return None + return (activation.to(dtype=torch.bfloat16), *args[1:]) + + +def quantize_dit_fp8_w8a8(dit: torch.nn.Module) -> FP8W8A8Report: + """Apply dynamic per-tensor W8A8 to the approved projections.""" + + from torchao.quantization import ( + Float8DynamicActivationFloat8WeightConfig, + PerTensor, + quantize_, + ) + + selected: set[str] = { + fqn + for fqn, module in dit.named_modules() + if isinstance(module, torch.nn.Linear) and is_quantizable_dit_projection(fqn) + } + quantize_( + dit, + Float8DynamicActivationFloat8WeightConfig( + granularity=PerTensor(), activation_value_ub=None + ), + filter_fn=lambda _, fqn: fqn in selected, + ) + for fqn in selected: + dit.get_submodule(fqn).register_forward_pre_hook( + _cast_w8a8_activation_to_bf16 + ) + return FP8W8A8Report(module_count=len(selected)) diff --git a/fdanyone/model/routing.py b/fdanyone/model/routing.py new file mode 100644 index 0000000000000000000000000000000000000000..c40b8f227ea80190b7e372da76a3e469311f58c9 --- /dev/null +++ b/fdanyone/model/routing.py @@ -0,0 +1,69 @@ +"""Target-context routing across view groups.""" + +from __future__ import annotations + +from itertools import pairwise + + +def _validate_grouping(num_views: int, group_size: int) -> int: + if num_views <= 0 or group_size <= 0 or num_views % group_size: + raise ValueError(f"num_views={num_views} must be divisible by positive group_size={group_size}.") + return num_views // group_size + + +def view_groups( + num_views: int, + group_size: int, + offset: int = 0, + *, + circular: bool = True, +) -> tuple[tuple[int, ...], ...]: + """Partition one camera layer, optionally without joining its endpoints.""" + + num_groups = _validate_grouping(num_views, group_size) + if circular: + return tuple( + tuple((group_index * group_size + offset + local_index) % num_views for local_index in range(group_size)) + for group_index in range(num_groups) + ) + + # Shifting an open sequence creates smaller boundary groups instead of a + # false neighborhood between the two ends of a partial yaw span. + offset %= group_size + boundaries = [0] + if offset: + boundaries.append(offset) + boundaries.extend(range(offset + group_size, num_views, group_size)) + boundaries.append(num_views) + return tuple(tuple(range(start, end)) for start, end in pairwise(boundaries)) + + +def routing_steps( + *, + views_per_layer: int, + num_layers: int, + group_size: int, + num_steps: int, + enable_tcr: bool, + circular: bool, +) -> tuple[tuple[tuple[int, ...], ...], ...]: + """Return layer-local target groups for every denoising step.""" + + if num_steps <= 0: + raise ValueError(f"num_steps must be positive, got {num_steps}.") + if num_layers <= 0: + raise ValueError(f"num_layers must be positive, got {num_layers}.") + _validate_grouping(views_per_layer, group_size) + return tuple( + tuple( + tuple(layer_index * views_per_layer + view_index for view_index in group) + for layer_index in range(num_layers) + for group in view_groups( + views_per_layer, + group_size, + step_index if enable_tcr else 0, + circular=circular, + ) + ) + for step_index in range(num_steps) + ) diff --git a/fdanyone/model/tiny_decoder.py b/fdanyone/model/tiny_decoder.py new file mode 100644 index 0000000000000000000000000000000000000000..4b543e7d8b130a0750111454c104dec0b156afeb --- /dev/null +++ b/fdanyone/model/tiny_decoder.py @@ -0,0 +1,60 @@ +"""Target-only adapter for the pinned TAEW2.2 decoder experiment.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import torch + +from fdanyone.errors import AssetError, ConfigurationError + + +def validate_tiny_wan_decoder(decoder: Any) -> None: + """Reject tiny decoders that do not match Wan 2.2 TI2V-5B latents.""" + + latent_channels: int = int(getattr(decoder, "latent_channels", 0)) + patch_size: int = int(getattr(decoder, "patch_size", 0)) + if latent_channels != 48 or patch_size != 2: + raise ConfigurationError( + "The 4DAnyone Wan 2.2 5B VAE requires 48 latent channels and " + f"patch size 2; got {latent_channels} channels and patch size " + f"{patch_size}. Use taew2_2, not taew2_1." + ) + + +def load_tiny_wan_decoder(checkpoint_path: str | Path) -> torch.nn.Module: + """Load the target-only TAEW2.2 decoder in its documented FP16 dtype.""" + + path: Path = Path(checkpoint_path).expanduser().resolve() + if not path.is_file(): + raise AssetError(f"TAEW2.2 checkpoint does not exist: {path}") + try: + from taehv import TAEHV # pyrefly: ignore [missing-import] + except ImportError as exc: + raise AssetError( + "TAEHV is required by target_decoder='taew2_2'; use the pixi " + "fast environment." + ) from exc + decoder: torch.nn.Module = TAEHV(checkpoint_path=str(path)) + validate_tiny_wan_decoder(decoder) + return decoder.to(dtype=torch.float16).eval().requires_grad_(False) + + +def decode_tiny_target_video( + decoder: Any, + latents: torch.Tensor, +) -> torch.Tensor: + """Decode normalized ``NCTHW`` latents to ``NCTHW`` pixels in [-1, 1].""" + + if latents.ndim != 5 or latents.shape[1] != 48: + raise ConfigurationError( + "TAEW2.2 target decode expects NCTHW latents with 48 channels; " + f"got {tuple(latents.shape)}." + ) + frame_first: torch.Tensor = decoder.decode_video( + latents.transpose(1, 2), + parallel=False, + show_progress_bar=False, + ) + return frame_first.transpose(1, 2).mul(2.0).sub(1.0) diff --git a/fdanyone/model/turbo_lora.py b/fdanyone/model/turbo_lora.py new file mode 100644 index 0000000000000000000000000000000000000000..c67fa87fecb51e01e2619a51394781293192d5f3 --- /dev/null +++ b/fdanyone/model/turbo_lora.py @@ -0,0 +1,125 @@ +"""Strict streaming merge for Kijai's Wan2.2 5B Turbo-LoRA format.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +import torch + +from fdanyone.errors import AssetError + + +@dataclass(frozen=True, slots=True) +class TurboLoraReport: + """Auditable summary of one completed adapter merge.""" + + lora_modules: int + direct_parameters: int + max_rank: int + + +def _target_name(adapter_name: str, suffix: str, parameter: str) -> str: + """Map one publisher key to the corresponding Wan parameter name.""" + + prefix = "diffusion_model." + if not adapter_name.startswith(prefix) or not adapter_name.endswith(suffix): + raise AssetError(f"Unsupported Turbo-LoRA key {adapter_name!r}.") + return adapter_name[len(prefix) : -len(suffix)] + parameter + + +def merge_wan_turbo_lora( + model: torch.nn.Module, + path: str | Path, +) -> TurboLoraReport: + """Merge all low-rank, bias, and normalization deltas without fallback.""" + + checkpoint = Path(path).expanduser().resolve() + if not checkpoint.is_file(): + raise AssetError(f"Turbo-LoRA checkpoint is missing: {checkpoint}") + try: + from safetensors import safe_open + except ImportError as exc: + raise AssetError("safetensors is required for Turbo-LoRA.") from exc + + parameters = dict(model.named_parameters()) + with safe_open(checkpoint, framework="pt", device="cpu") as adapter: + keys = set(adapter.keys()) + down_keys = sorted(key for key in keys if key.endswith(".lora_down.weight")) + up_keys = {key for key in keys if key.endswith(".lora_up.weight")} + direct_keys = sorted( + key for key in keys if key.endswith((".diff", ".diff_b")) + ) + consumed_up: set[str] = set() + max_rank = 0 + + for down_key in down_keys: + up_key = down_key.removesuffix(".lora_down.weight") + ".lora_up.weight" + if up_key not in keys: + raise AssetError(f"Turbo-LoRA has no paired key {up_key!r}.") + consumed_up.add(up_key) + target_name = _target_name(down_key, ".lora_down.weight", ".weight") + if target_name not in parameters: + raise AssetError(f"Turbo-LoRA target is missing: {target_name}") + down_shape = tuple(adapter.get_slice(down_key).get_shape()) + up_shape = tuple(adapter.get_slice(up_key).get_shape()) + target_shape = tuple(parameters[target_name].shape) + if ( + len(down_shape) != 2 + or len(up_shape) != 2 + or up_shape[1] != down_shape[0] + or (up_shape[0], down_shape[1]) != target_shape + ): + raise AssetError( + f"Turbo-LoRA shape mismatch for {target_name}: " + f"down={down_shape}, up={up_shape}, target={target_shape}." + ) + max_rank = max(max_rank, down_shape[0]) + if consumed_up != up_keys: + raise AssetError( + f"Turbo-LoRA has unpaired up keys: {sorted(up_keys - consumed_up)}" + ) + + direct_targets: dict[str, str] = {} + for key in direct_keys: + if key.endswith(".diff_b"): + target_name = _target_name(key, ".diff_b", ".bias") + else: + target_name = _target_name(key, ".diff", ".weight") + if target_name not in parameters: + raise AssetError(f"Turbo-LoRA target is missing: {target_name}") + patch_shape = tuple(adapter.get_slice(key).get_shape()) + if patch_shape != tuple(parameters[target_name].shape): + raise AssetError( + f"Turbo-LoRA shape mismatch for {target_name}: " + f"patch={patch_shape}, target={tuple(parameters[target_name].shape)}." + ) + direct_targets[key] = target_name + + recognized = set(down_keys) | up_keys | set(direct_keys) + if recognized != keys: + raise AssetError(f"Unsupported Turbo-LoRA keys: {sorted(keys - recognized)}") + + with torch.no_grad(): + for down_key in down_keys: + up_key = down_key.removesuffix(".lora_down.weight") + ".lora_up.weight" + target_name = _target_name( + down_key, + ".lora_down.weight", + ".weight", + ) + parameter = parameters[target_name] + down = adapter.get_tensor(down_key).float() + up = adapter.get_tensor(up_key).float() + delta = torch.mm(up, down) + parameter.add_(delta.to(device=parameter.device, dtype=parameter.dtype)) + for key, target_name in direct_targets.items(): + parameter = parameters[target_name] + delta = adapter.get_tensor(key) + parameter.add_(delta.to(device=parameter.device, dtype=parameter.dtype)) + + return TurboLoraReport( + lora_modules=len(down_keys), + direct_parameters=len(direct_keys), + max_rank=max_rank, + ) diff --git a/fdanyone/motion/__init__.py b/fdanyone/motion/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..5fdaf00a22886cf3990a0b8c760681fdc6986835 --- /dev/null +++ b/fdanyone/motion/__init__.py @@ -0,0 +1,5 @@ +"""GVHMR motion recovery.""" + +from fdanyone.motion.result import MotionResult + +__all__ = ["MotionResult"] diff --git a/fdanyone/motion/body.py b/fdanyone/motion/body.py new file mode 100644 index 0000000000000000000000000000000000000000..43a5582a165f35d41faa24372c8be3a6fd3b4786 --- /dev/null +++ b/fdanyone/motion/body.py @@ -0,0 +1,143 @@ +"""Rebuild the posed SMPL-X body of a finished GVHMR motion result. + +The heavy dependencies (``numpy``, ``smplx``, ``torch``) stay function-local so +importing this module costs nothing until a body is actually evaluated. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from pathlib import Path +from typing import TYPE_CHECKING + +from fdanyone.errors import FourDAnyoneError +from fdanyone.runs import discover_run + +if TYPE_CHECKING: + from fractions import Fraction + + import numpy as np + +LOGGER = logging.getLogger("fdanyone.motion.body") + +_REPO_ROOT = Path(__file__).resolve().parents[2] + +# SMPL-X body model roots tried in order, relative to the repository root. +SMPLX_MODEL_ROOTS = ( + Path("models"), + Path("third_party/GVHMR/inputs/checkpoints"), +) + +# SMPL-X keeps SMPL's body-joint order, so the first 55 entries of the joint +# tensor are the body, jaw, and hand skeleton addressed by ``parents``. +NUM_SMPLX_SKELETON_JOINTS = 55 + + +@dataclass(frozen=True) +class BodyMotion: + """SMPL-X geometry in the canonical 4DAnyone human world.""" + + vertices: np.ndarray + joints: np.ndarray + faces: np.ndarray + parents: tuple[int, ...] + keypoints_2d: np.ndarray | None + image_size: tuple[int, int] + fps: Fraction + + +def _smplx_model_path() -> Path: + from fdanyone.assets import SMPLX_MODEL + + for relative in SMPLX_MODEL_ROOTS: + model = _REPO_ROOT / relative / SMPLX_MODEL + if model.is_file(): + return model.parents[1] + raise FourDAnyoneError( + "The licensed SMPL-X body model is missing. Run `python scripts/download_smplx.py`; " + f"expected {SMPLX_MODEL} under one of: {[str(_REPO_ROOT / path) for path in SMPLX_MODEL_ROOTS]}." + ) + + +def _canonical_rotation(joints: np.ndarray) -> np.ndarray: + """Rotate the first frame onto the canonical yaw-zero human world. + + Mirrors GVHMR's ``compute_T_ayfz2ay``: canonical ``+x`` is the subject's + anatomical left at frame zero, ``+y`` is up, and ``+z`` completes the + right-handed frame. + """ + + import numpy as np + + first = joints[0] + left = (first[1, [0, 2]] - first[2, [0, 2]]) + (first[16, [0, 2]] - first[17, [0, 2]]) + norm = float(np.linalg.norm(left)) + if norm <= 1e-4: + LOGGER.warning("Cannot determine the facing direction; leaving the motion world unrotated.") + return np.eye(3, dtype=np.float64) + x_dir = np.array([left[0] / norm, 0.0, left[1] / norm], dtype=np.float64) + y_dir = np.array([0.0, 1.0, 0.0], dtype=np.float64) + z_dir = np.cross(x_dir, y_dir) + return np.stack([x_dir, y_dir, z_dir], axis=-1) + + +def load_body_motion(motion_dir: Path, device: str = "cpu") -> BodyMotion: + """Rebuild SMPL-X vertices and joints from a saved GVHMR motion result.""" + + import numpy as np + import smplx + import torch + + from fdanyone.motion.result import MotionResult + + motion = MotionResult.load(motion_dir) + model_path = _smplx_model_path() + parameters = motion.smpl_params_global + num_frames = motion.num_frames + # GVHMR's "supermotion" body model is plain neutral SMPL-X with ten shape + # coefficients, twelve hand PCA components, and a non-flat hand mean. + body_model = smplx.create( + model_path=str(model_path), + model_type="smplx", + gender="neutral", + num_betas=10, + num_pca_comps=12, + flat_hand_mean=False, + use_pca=True, + batch_size=num_frames, + ).to(device) + with torch.inference_mode(): + output = body_model( + betas=parameters["betas"].to(device), + global_orient=parameters["global_orient"].to(device), + body_pose=parameters["body_pose"].to(device), + transl=parameters["transl"].to(device), + ) + vertices = output.vertices.detach().cpu().numpy().astype(np.float64) + joints = output.joints.detach().cpu().numpy().astype(np.float64)[:, :NUM_SMPLX_SKELETON_JOINTS] + + # Canonicalize exactly like the conditioning stage: drop the first-frame + # root to the ground plane origin, then align the initial facing yaw. + offset = joints[0, 0].copy() + offset[1] = float(vertices[..., 1].min()) + rotation = _canonical_rotation(joints - offset) + vertices = (vertices - offset) @ rotation + joints = (joints - offset) @ rotation + + keypoints = motion.observed_keypoints_2d.detach().cpu().numpy().astype(np.float32) + return BodyMotion( + vertices=vertices.astype(np.float32), + joints=joints.astype(np.float32), + faces=np.asarray(body_model.faces, dtype=np.uint32), + parents=tuple(int(value) for value in body_model.parents.detach().cpu().numpy()), + keypoints_2d=keypoints, + image_size=(motion.image_width, motion.image_height), + fps=motion.fps, + ) + + +def posed_smplx(data_dir: str | Path, clip: str, device: str = "cpu") -> BodyMotion: + """Rebuild the posed body of one finished run, found by clip name.""" + + return load_body_motion(discover_run(Path(data_dir), clip).motion_dir, device=device) diff --git a/fdanyone/motion/gvhmr.py b/fdanyone/motion/gvhmr.py new file mode 100644 index 0000000000000000000000000000000000000000..a48b5751b454910a47b619fef50b59009db17644 --- /dev/null +++ b/fdanyone/motion/gvhmr.py @@ -0,0 +1,332 @@ +"""Classic GVHMR inference used by 4DAnyone. + +The official demo imports training, evaluation, visualization, and +moving-camera modules eagerly. This file keeps the released static-camera path +in one place without exposing Hydra or backend abstractions to 4DAnyone users. +""" + +from __future__ import annotations + +import contextlib +import importlib +import importlib.util +import itertools +import json +import os +import subprocess +import sys +import types +import warnings +from collections.abc import Callable, Iterator +from contextlib import contextmanager +from pathlib import Path +from typing import TypeAlias + +from fdanyone.errors import AssetError, VideoContractError +from fdanyone.motion.result import SMPL_PARAMETER_NAMES, MotionResult +from fdanyone.vendor.pytorch3d_compat import install_if_needed as install_pytorch3d_compat +from fdanyone.video import CanonicalClip + +MotionStageHook: TypeAlias = Callable[[str, dict[str, object]], None] +"""Receives ``(stage_name, payload)`` as each motion stage completes.""" + +GVHMR_ASSETS = ( + "inputs/checkpoints/gvhmr/gvhmr_siga24_release.ckpt", + "inputs/checkpoints/hmr2/epoch=10-step=25000.ckpt", + "inputs/checkpoints/vitpose/vitpose-h-multi-coco.pth", + "inputs/checkpoints/yolo/yolov8x.pt", + "inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz", +) + + +def validate_gvhmr(root: str | Path) -> tuple[Path, str]: + """Locate the GVHMR checkout and files consumed by inference.""" + + path = Path(root).expanduser().resolve() + required = ("hmr4d/__init__.py", "tools/demo/demo.py", *GVHMR_ASSETS) + missing = [relative for relative in required if not (path / relative).is_file()] + if missing: + formatted = "\n - ".join(missing) + raise AssetError( + f"GVHMR is incomplete under {path}. Run `git submodule update --init third_party/GVHMR`, " + f"`python scripts/download_model.py`, and `python scripts/download_smplx.py`; missing:\n" + f" - {formatted}" + ) + try: + revision = subprocess.check_output( + ["git", "-C", str(path), "rev-parse", "HEAD"], + text=True, + stderr=subprocess.DEVNULL, + ).strip() + except (OSError, subprocess.CalledProcessError) as exc: + raise AssetError(f"GVHMR must be a git checkout: {path}") from exc + if len(revision) != 40: + raise AssetError(f"Cannot identify the GVHMR revision at {path}.") + return path, revision + + +def hydra_override(name: str, value: str | Path) -> str: + """Quote a path for the internal GVHMR Hydra config.""" + + if not name.isidentifier(): + raise ValueError(f"Invalid Hydra field name: {name!r}.") + return f"{name}={json.dumps(str(value), ensure_ascii=False)}" + + +@contextmanager +def gvhmr_imports(root: Path) -> Iterator[None]: + """Temporarily import GVHMR as if its checkout were the working tree.""" + + old_cwd = Path.cwd() + root_text = str(root) + already_present = root_text in sys.path + if not already_present: + sys.path.insert(0, root_text) + os.chdir(root) + try: + yield + finally: + os.chdir(old_cwd) + if not already_present: + with contextlib.suppress(ValueError): + sys.path.remove(root_text) + + +@contextmanager +def _legacy_checkpoint_loading(): + """Restore pre-2.6 ``torch.load`` behavior for trusted GVHMR assets.""" + + name = "TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD" + previous = os.environ.get(name) + os.environ[name] = "1" + try: + with warnings.catch_warnings(): + warnings.filterwarnings( + "ignore", + message=r"Environment variable TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD detected.*", + category=UserWarning, + ) + yield + finally: + if previous is None: + os.environ.pop(name, None) + else: + os.environ[name] = previous + + +def _install_optional_import_stubs() -> None: + """Avoid visualization and moving-camera dependencies we never call.""" + + def unavailable(*_args, **_kwargs): + raise RuntimeError("Wis3D visualization is not part of 4DAnyone inference.") + + if importlib.util.find_spec("wis3d") is None: + module = types.ModuleType("hmr4d.utils.wis3d_utils") + module.make_wis3d = unavailable + module.add_motion_as_lines = unavailable + sys.modules[module.__name__] = module + + def moving_camera_unavailable(*_args, **_kwargs): + raise RuntimeError("SimpleVO is unavailable in the static-camera 4DAnyone runtime.") + + module = types.ModuleType("hmr4d.utils.preproc.relpose.simple_vo") + module.SimpleVO = moving_camera_unavailable + sys.modules[module.__name__] = module + + +def _register_inference_store() -> None: + """Register only the Hydra groups referenced by GVHMR's demo config.""" + + _install_optional_import_stubs() + for module in ( + "hmr4d.model.gvhmr.gvhmr_pl_demo", + "hmr4d.model.gvhmr.utils.endecoder", + "hmr4d.network.gvhmr.relative_transformer", + ): + importlib.import_module(module) + + # GVHMR installs a colored handler on the root logger. Remove only that + # duplicate because the public pipeline already owns a handler. + logger_module = sys.modules.get("hmr4d.utils.pylogger") + logger = getattr(logger_module, "Log", None) + handler = getattr(logger_module, "ch", None) + if ( + logger is not None + and handler in logger.handlers + and any(candidate is not handler for candidate in logger.handlers) + ): + logger.removeHandler(handler) + + +def _run_preprocess(cfg, *, on_stage: MotionStageHook | None = None) -> None: + """Run tracker, ViTPose, and image-feature extraction.""" + + import torch + from hmr4d.utils.geo.hmr_cam import get_bbx_xys_from_xyxy + from hmr4d.utils.preproc.tracker import Tracker + from hmr4d.utils.preproc.vitfeat_extractor import Extractor + from hmr4d.utils.preproc.vitpose import VitPoseExtractor + from hmr4d.utils.pylogger import Log + + if not bool(cfg.static_cam): + raise ValueError("4DAnyone requires GVHMR static_cam=true.") + + Log.info("[Preprocess] Start!") + started = Log.time() + video_path = cfg.video_path + paths = cfg.paths + + if not Path(paths.bbx).exists(): + tracker = Tracker() + bbx_xyxy = tracker.get_one_track(video_path).float() + bbx_xys = get_bbx_xys_from_xyxy(bbx_xyxy, base_enlarge=1.2).float() + torch.save({"bbx_xyxy": bbx_xyxy, "bbx_xys": bbx_xys}, paths.bbx) + del tracker + else: + bbx_xys = torch.load(paths.bbx, weights_only=True)["bbx_xys"] + Log.info("[Preprocess] bbx (xyxy, xys) from %s", paths.bbx) + if on_stage is not None: + # Only ``bbx_xys`` reaches the model, so read the boxes back from the + # file both branches guarantee instead of widening the cached branch. + tracked_boxes = torch.load(paths.bbx, weights_only=True)["bbx_xyxy"] + on_stage("bboxes", {"bbx_xyxy": tracked_boxes.detach().cpu()}) + + if not Path(paths.vitpose).exists(): + extractor = VitPoseExtractor() + torch.save(extractor.extract(video_path, bbx_xys), paths.vitpose) + del extractor + else: + Log.info("[Preprocess] vitpose from %s", paths.vitpose) + if on_stage is not None: + keypoints_2d = torch.load(paths.vitpose, weights_only=True) + on_stage("keypoints_2d", {"kp2d": keypoints_2d.detach().cpu()}) + + if not Path(paths.vit_features).exists(): + extractor = Extractor() + torch.save(extractor.extract_video_features(video_path, bbx_xys), paths.vit_features) + del extractor + else: + Log.info("[Preprocess] vit_features from %s", paths.vit_features) + if on_stage is not None: + on_stage("features", {}) + + Log.info("[Preprocess] End. Time elapsed: %.2fs", Log.time() - started) + + +def _load_data(cfg): + """Build the static-camera tensors consumed by GVHMR.""" + + import torch + from hmr4d.utils.geo.hmr_cam import estimate_K + from hmr4d.utils.geo_transform import compute_cam_angvel + from hmr4d.utils.video_io_utils import get_video_lwh + + if not bool(cfg.static_cam): + raise ValueError("4DAnyone requires GVHMR static_cam=true.") + paths = cfg.paths + length, width, height = get_video_lwh(cfg.video_path) + rotation_world_to_camera = torch.eye(3).repeat(length, 1, 1) + intrinsics = estimate_K(width, height).repeat(length, 1, 1) + return { + "length": torch.tensor(length), + "bbx_xys": torch.load(paths.bbx, weights_only=True)["bbx_xys"], + "kp2d": torch.load(paths.vitpose, weights_only=True), + "K_fullimg": intrinsics, + "cam_angvel": compute_cam_angvel(rotation_world_to_camera), + "f_imgseq": torch.load(paths.vit_features, weights_only=True), + } + + +def _verify_gvhmr_decode(clip: CanonicalClip, working_video: Path, reader_factory) -> None: + """Ensure GVHMR's own video reader sees the canonical RGB frames.""" + + import numpy as np + + reader = reader_factory(str(working_video)) + sentinel = object() + try: + for index, (actual, expected) in enumerate(itertools.zip_longest(reader, clip.rgb_frames, fillvalue=sentinel)): + if actual is sentinel or expected is sentinel or not np.array_equal(actual, expected): + raise VideoContractError(f"GVHMR decoded a different canonical frame at index {index}.") + finally: + close = getattr(reader, "close", None) + if close is not None: + close() + + +def run_gvhmr( + *, + clip: CanonicalClip, + working_video: str | Path, + output_dir: str | Path, + gvhmr_root: str | Path, + device: str, + on_stage: MotionStageHook | None = None, +) -> MotionResult: + """Recover static-camera human motion from the canonical source clip.""" + + root, revision = validate_gvhmr(gvhmr_root) + working_video = Path(working_video).expanduser().resolve() + output_root = Path(output_dir).expanduser().resolve() + output_root.mkdir(parents=True, exist_ok=True) + + with gvhmr_imports(root), _legacy_checkpoint_loading(): + install_pytorch3d_compat() + import hydra + import torch + from hmr4d.model.gvhmr.gvhmr_pl_demo import DemoPL + from hmr4d.utils.net_utils import detach_to_cpu + from hmr4d.utils.video_io_utils import get_video_reader + from hydra import compose, initialize_config_module + from omegaconf import open_dict + + _register_inference_store() + with initialize_config_module(version_base="1.3", config_module="hmr4d.configs"): + cfg = compose( + config_name="demo", + overrides=[ + hydra_override("video_name", working_video.stem), + "static_cam=true", + "verbose=false", + "use_dpvo=false", + hydra_override("output_root", output_root), + ], + ) + Path(cfg.output_dir).mkdir(parents=True, exist_ok=True) + Path(cfg.preprocess_dir).mkdir(parents=True, exist_ok=True) + with open_dict(cfg): + cfg.video_path = str(working_video) + + _verify_gvhmr_decode(clip, working_video, get_video_reader) + _run_preprocess(cfg, on_stage=on_stage) + data = _load_data(cfg) + if int(data["length"]) != len(clip.frames): + raise RuntimeError(f"GVHMR decoded {int(data['length'])} frames, expected {len(clip.frames)}.") + observed_keypoints_2d = data["kp2d"].detach().cpu() + model: DemoPL = hydra.utils.instantiate(cfg.model, _recursive_=False) + model.load_pretrained_model(cfg.ckpt_path) + model = model.eval().to(device) + with torch.inference_mode(): + prediction = detach_to_cpu(model.predict(data, static_cam=True)) + del model, data + torch.cuda.empty_cache() + + result = MotionResult( + gvhmr_revision=revision, + fps=clip.fps, + frame_timestamps_sec=tuple(float(frame.canonical_timestamp) for frame in clip.frames), + source_frame_indices=tuple(frame.source_index for frame in clip.frames), + source_pts=tuple(frame.source_pts for frame in clip.frames), + source_size_bytes=clip.source_size_bytes, + source_mtime_ns=clip.source_mtime_ns, + image_height=clip.height, + image_width=clip.width, + smpl_params_global={name: prediction["smpl_params_global"][name] for name in SMPL_PARAMETER_NAMES}, + smpl_params_incam={name: prediction["smpl_params_incam"][name] for name in SMPL_PARAMETER_NAMES}, + K_fullimg=prediction["K_fullimg"], + observed_keypoints_2d=observed_keypoints_2d, + ) + result.validate(expected_frames=len(clip.frames)) + if on_stage is not None: + on_stage("smplx", {"result": result}) + return result diff --git a/fdanyone/motion/result.py b/fdanyone/motion/result.py new file mode 100644 index 0000000000000000000000000000000000000000..85b29958256c5677806990d623d9300971424943 --- /dev/null +++ b/fdanyone/motion/result.py @@ -0,0 +1,188 @@ +"""Reusable GVHMR result stored as JSON plus safetensors.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from fractions import Fraction +from pathlib import Path +from typing import Any + +from fdanyone.errors import FourDAnyoneError + +SMPL_PARAMETER_NAMES = ("body_pose", "betas", "global_orient", "transl") +SMPL_PARAMETER_WIDTHS = {"body_pose": 63, "betas": 10, "global_orient": 3, "transl": 3} + + +@dataclass(frozen=True) +class MotionResult: + gvhmr_revision: str + fps: Fraction + frame_timestamps_sec: tuple[float, ...] + source_frame_indices: tuple[int, ...] + source_pts: tuple[int | None, ...] + source_size_bytes: int + source_mtime_ns: int + image_height: int + image_width: int + smpl_params_global: dict[str, Any] + smpl_params_incam: dict[str, Any] + K_fullimg: Any + observed_keypoints_2d: Any + motion_world: str = "gvhmr_gravity_aligned_y_up" + + @property + def num_frames(self) -> int: + return len(self.frame_timestamps_sec) + + def validate(self, expected_frames: int = 121) -> None: + try: + import torch + except ImportError as exc: + raise FourDAnyoneError("PyTorch is required to validate motion tensors.") from exc + if self.motion_world != "gvhmr_gravity_aligned_y_up": + raise FourDAnyoneError(f"Unknown motion world convention: {self.motion_world!r}.") + if ( + not isinstance(self.gvhmr_revision, str) + or len(self.gvhmr_revision) != 40 + or any(character not in "0123456789abcdef" for character in self.gvhmr_revision.lower()) + ): + raise FourDAnyoneError("MotionResult requires a 40-character GVHMR git revision.") + if self.fps <= 0: + raise FourDAnyoneError(f"MotionResult FPS must be positive, got {self.fps}.") + if self.num_frames != expected_frames: + raise FourDAnyoneError(f"MotionResult has {self.num_frames} frames, expected {expected_frames}.") + for name, values in ( + ("source_frame_indices", self.source_frame_indices), + ("source_pts", self.source_pts), + ): + if len(values) != self.num_frames: + raise FourDAnyoneError(f"MotionResult {name} has {len(values)} values, expected {self.num_frames}.") + expected_timestamps = tuple(float(Fraction(index, 1) / self.fps) for index in range(self.num_frames)) + if self.frame_timestamps_sec != expected_timestamps: + raise FourDAnyoneError("MotionResult timestamps are not the exact zero-based CFR timeline.") + if any(index < 0 for index in self.source_frame_indices) or any( + right < left for left, right in zip(self.source_frame_indices, self.source_frame_indices[1:], strict=False) + ): + raise FourDAnyoneError("MotionResult source-frame indices must be non-negative and monotonic.") + if self.source_size_bytes <= 0 or self.source_mtime_ns <= 0: + raise FourDAnyoneError("MotionResult has an invalid source-file identity.") + if self.image_height <= 0 or self.image_width <= 0: + raise FourDAnyoneError("MotionResult image dimensions must be positive.") + for group_name, parameters in ( + ("smpl_params_global", self.smpl_params_global), + ("smpl_params_incam", self.smpl_params_incam), + ): + if set(parameters) != set(SMPL_PARAMETER_NAMES): + raise FourDAnyoneError( + f"MotionResult {group_name} must contain {SMPL_PARAMETER_NAMES}, got {tuple(parameters)}." + ) + for name, tensor in parameters.items(): + expected_shape = (self.num_frames, SMPL_PARAMETER_WIDTHS[name]) + if not isinstance(tensor, torch.Tensor) or tuple(tensor.shape) != expected_shape: + raise FourDAnyoneError(f"{group_name}.{name} must have shape {expected_shape}.") + if not bool(torch.isfinite(tensor).all()): + raise FourDAnyoneError(f"{group_name}.{name} contains non-finite values.") + if not isinstance(self.K_fullimg, torch.Tensor) or tuple(self.K_fullimg.shape) != (self.num_frames, 3, 3): + raise FourDAnyoneError(f"K_fullimg must have shape ({self.num_frames}, 3, 3).") + if not bool(torch.isfinite(self.K_fullimg).all()): + raise FourDAnyoneError("K_fullimg contains non-finite values.") + expected_keypoint_shape = (self.num_frames, 17, 3) + if ( + not isinstance(self.observed_keypoints_2d, torch.Tensor) + or tuple(self.observed_keypoints_2d.shape) != expected_keypoint_shape + ): + raise FourDAnyoneError(f"observed_keypoints_2d must have shape {expected_keypoint_shape}.") + if not bool(torch.isfinite(self.observed_keypoints_2d).all()): + raise FourDAnyoneError("observed_keypoints_2d contains non-finite values.") + + def validate_against_clip(self, clip) -> None: + """Reject a cached result produced from a different video timeline.""" + + self.validate(expected_frames=len(clip.frames)) + expected_timestamps = tuple(float(frame.canonical_timestamp) for frame in clip.frames) + expected_indices = tuple(frame.source_index for frame in clip.frames) + expected_pts = tuple(frame.source_pts for frame in clip.frames) + if self.fps != clip.fps: + raise FourDAnyoneError(f"Motion FPS {self.fps} does not match canonical FPS {clip.fps}.") + if self.frame_timestamps_sec != expected_timestamps: + raise FourDAnyoneError("Motion timestamps do not match the canonical clip.") + if self.source_frame_indices != expected_indices or self.source_pts != expected_pts: + raise FourDAnyoneError("Motion source-frame identity does not match the canonical clip.") + if (self.source_size_bytes, self.source_mtime_ns) != ( + clip.source_size_bytes, + clip.source_mtime_ns, + ): + raise FourDAnyoneError("Cached GVHMR motion belongs to a different source file.") + + def save(self, directory: str | Path) -> Path: + from safetensors.torch import save_file + + self.validate() + root = Path(directory).expanduser().resolve() + root.mkdir(parents=True, exist_ok=True) + tensor_path = root / "motion.safetensors" + tensors = { + **{ + f"smpl_params_global.{name}": self.smpl_params_global[name].detach().cpu().contiguous() + for name in SMPL_PARAMETER_NAMES + }, + **{ + f"smpl_params_incam.{name}": self.smpl_params_incam[name].detach().cpu().contiguous() + for name in SMPL_PARAMETER_NAMES + }, + "K_fullimg": self.K_fullimg.detach().cpu().contiguous(), + "observed_keypoints_2d": self.observed_keypoints_2d.detach().cpu().contiguous(), + } + save_file(tensors, str(tensor_path)) + metadata = { + "gvhmr_revision": self.gvhmr_revision, + "num_frames": self.num_frames, + "fps_num": self.fps.numerator, + "fps_den": self.fps.denominator, + "frame_timestamps_sec": self.frame_timestamps_sec, + "source_frame_indices": self.source_frame_indices, + "source_pts": self.source_pts, + "source_size_bytes": self.source_size_bytes, + "source_mtime_ns": self.source_mtime_ns, + "image_height": self.image_height, + "image_width": self.image_width, + "motion_world": self.motion_world, + "tensor_file": tensor_path.name, + } + (root / "motion.json").write_text(json.dumps(metadata, indent=2, sort_keys=True) + "\n") + return root + + @classmethod + def load(cls, directory: str | Path) -> MotionResult: + from safetensors.torch import load_file + + root = Path(directory).expanduser().resolve() + metadata = json.loads((root / "motion.json").read_text()) + tensors = load_file(str(root / metadata["tensor_file"]), device="cpu") + # ``safetensors`` may keep the source file memory-mapped for as long as + # returned tensors own that storage. Motion results live until final + # export, which in turn can leave an in-use ``.efc_*`` tombstone on + # CPFS/NFS when scratch cleanup unlinks the backing file. The payload + # is tiny compared with the generated videos, so take ownership here + # and close the file mapping at this explicit process boundary. + owned = {name: tensor.clone() for name, tensor in tensors.items()} + del tensors + result = cls( + gvhmr_revision=str(metadata["gvhmr_revision"]), + fps=Fraction(int(metadata["fps_num"]), int(metadata["fps_den"])), + frame_timestamps_sec=tuple(float(value) for value in metadata["frame_timestamps_sec"]), + source_frame_indices=tuple(int(value) for value in metadata["source_frame_indices"]), + source_pts=tuple(None if value is None else int(value) for value in metadata["source_pts"]), + source_size_bytes=int(metadata["source_size_bytes"]), + source_mtime_ns=int(metadata["source_mtime_ns"]), + image_height=int(metadata["image_height"]), + image_width=int(metadata["image_width"]), + smpl_params_global={name: owned[f"smpl_params_global.{name}"] for name in SMPL_PARAMETER_NAMES}, + smpl_params_incam={name: owned[f"smpl_params_incam.{name}"] for name in SMPL_PARAMETER_NAMES}, + K_fullimg=owned["K_fullimg"], + observed_keypoints_2d=owned["observed_keypoints_2d"], + motion_world=str(metadata["motion_world"]), + ) + result.validate() + return result diff --git a/fdanyone/motion/worker.py b/fdanyone/motion/worker.py new file mode 100644 index 0000000000000000000000000000000000000000..b6bc0089c39087b74387ec757ee02bfdf8d41492 --- /dev/null +++ b/fdanyone/motion/worker.py @@ -0,0 +1,53 @@ +"""Private subprocess entry point for motion inference.""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +from fdanyone.device import select_cuda_device +from fdanyone.motion.gvhmr import MotionStageHook, run_gvhmr +from fdanyone.video import load_canonical_working_clip + + +def run_motion_worker( + *, + gvhmr_root: str | Path, + working_video: str | Path, + clip_metadata: str | Path, + output_dir: str | Path, + result_dir: str | Path, + device: str, + on_stage: MotionStageHook | None = None, +) -> None: + """Recover motion for one request and publish it under ``result_dir``.""" + + device, _ = select_cuda_device(device) + clip = load_canonical_working_clip(working_video, clip_metadata) + common = { + "clip": clip, + "working_video": working_video, + "output_dir": output_dir, + "device": device, + } + result = run_gvhmr(gvhmr_root=gvhmr_root, on_stage=on_stage, **common) + result.save(result_dir) + + +def main(request_path: str) -> None: + request = json.loads(Path(request_path).read_text()) + run_motion_worker( + gvhmr_root=request["gvhmr_root"], + working_video=request["working_video"], + clip_metadata=request["clip_metadata"], + output_dir=request["output_dir"], + result_dir=request["result_dir"], + device=request["device"], + ) + + +if __name__ == "__main__": + if len(sys.argv) != 2: + raise SystemExit("Usage: python -m fdanyone.motion.worker REQUEST.json") + main(sys.argv[1]) diff --git a/fdanyone/output.py b/fdanyone/output.py new file mode 100644 index 0000000000000000000000000000000000000000..96df106e4aac1e05f0baf2f61ecf2fc0a0fb35ab --- /dev/null +++ b/fdanyone/output.py @@ -0,0 +1,218 @@ +"""Publish generated videos and their camera metadata.""" + +from __future__ import annotations + +import json +import platform +import shutil +import sys +import time +from pathlib import Path +from typing import TYPE_CHECKING + +from fdanyone.config import INFERENCE, ModeSettings +from fdanyone.errors import FourDAnyoneError +from fdanyone.io import write_json + +if TYPE_CHECKING: + from fdanyone.model.inference import GeneratedViews + from fdanyone.motion.result import MotionResult + from fdanyone.skeleton.pipeline import Conditioning + from fdanyone.video import CanonicalClip + + +def _copy_file(source: Path, destination: Path) -> Path: + destination.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(source, destination) + return destination + + +def _runtime_metadata(device: str) -> dict: + import torch + + cuda = { + "available": torch.cuda.is_available(), + "torch_cuda": torch.version.cuda, + "cudnn": torch.backends.cudnn.version(), + } + if torch.cuda.is_available(): + torch_device = torch.device(device) + properties = torch.cuda.get_device_properties(torch_device) + cuda.update( + { + "device": device, + "device_name": torch.cuda.get_device_name(torch_device), + "device_capability": list(torch.cuda.get_device_capability(torch_device)), + "device_total_memory_bytes": properties.total_memory, + } + ) + return { + "python": sys.version.split()[0], + "platform": platform.platform(), + "torch": torch.__version__, + "cuda": cuda, + } + + +def _camera_rig_payload(payload: dict, cameras: list[dict]) -> dict: + """Keep the final OpenCV camera rig needed by downstream tools.""" + + records = [] + for camera in cameras: + camera_id = int(camera["camera_id"]) + records.append( + { + "camera_id": camera_id, + "layer_index": int(camera["layer_index"]), + "pitch": int(camera["pitch_degrees"]), + "yaw": float(camera["yaw_degrees"]), + "K": camera["K"], + "camera_to_world": camera["camera_to_world"], + "image_width": int(camera["image_width"]), + "image_height": int(camera["image_height"]), + "video": f"videos/dense/{camera_id:02d}.mp4", + "skeleton_video": f"skeletons/{camera_id:02d}.mp4", + } + ) + return { + "camera_model": "OPENCV", + "world_frame": payload["world_frame"], + "camera_frame": payload["camera_frame"], + "front_camera_ids": payload["front_camera_ids"], + "framing": payload["framing"], + "cameras": records, + } + + +def _target_cameras(payload: object, expected_count: int) -> list[dict]: + """Read the camera records produced by the conditioning stage.""" + + if not isinstance(payload, dict) or payload.get("camera_model") != "OPENCV": + raise FourDAnyoneError("Conditioning did not produce an OpenCV camera rig.") + cameras = payload.get("cameras") + if not isinstance(cameras, list) or len(cameras) != expected_count: + raise FourDAnyoneError(f"Conditioning must contain {expected_count} target cameras.") + if [camera.get("camera_id") for camera in cameras if isinstance(camera, dict)] != list(range(expected_count)): + raise FourDAnyoneError("Target cameras are not in canonical order.") + return cameras + + +def export_result( + *, + clip: CanonicalClip, + conditioning: Conditioning, + generated: GeneratedViews, + destination: str | Path, + motion: MotionResult, + model_identity: dict, + pipeline_started: float, + settings: ModeSettings, +) -> dict: + """Publish proposal, target, skeleton, camera, and metadata artifacts.""" + + root = Path(destination).expanduser().resolve() + attention_backend = "sdpa" if settings.exact_attention else "sageattention" + view_plan = generated.view_plan + if conditioning.view_plan != view_plan: + raise FourDAnyoneError("Conditioning and generation resolved different view plans.") + if len(generated.rcp_videos) != len(view_plan.rcp_camera_ids): + raise FourDAnyoneError( + f"Generation returned {len(generated.rcp_videos)} RCP videos, expected {len(view_plan.rcp_camera_ids)}." + ) + if len(generated.target_videos) != view_plan.num_target_views: + raise FourDAnyoneError( + f"Generation returned {len(generated.target_videos)} target videos, expected {view_plan.num_target_views}." + ) + if len(conditioning.target_skeletons) != view_plan.num_target_views: + raise FourDAnyoneError( + f"Conditioning returned {len(conditioning.target_skeletons)} target skeletons, " + f"expected {view_plan.num_target_views}." + ) + + sparse_root = root / "videos" / "sparse" + dense_root = root / "videos" / "dense" + skeletons_root = root / "skeletons" + dense_root.mkdir(parents=True, exist_ok=False) + skeletons_root.mkdir(exist_ok=False) + if generated.rcp_videos: + sparse_root.mkdir(exist_ok=False) + + output_sparse = tuple( + _copy_file(source, sparse_root / f"{camera_id:02d}.mp4") + for camera_id, source in zip(view_plan.rcp_camera_ids, generated.rcp_videos, strict=True) + ) + output_dense = tuple( + _copy_file(source, dense_root / f"{camera_id:02d}.mp4") + for camera_id, source in enumerate(generated.target_videos) + ) + for camera_id, skeleton in enumerate(conditioning.target_skeletons): + _copy_file(skeleton.path, skeletons_root / f"{camera_id:02d}.mp4") + + camera_payload = json.loads((conditioning.root / "cameras.json").read_text()) + conditioning_metadata = json.loads((conditioning.root / "metadata.json").read_text()) + camera_records = _target_cameras(camera_payload, view_plan.num_target_views) + total_elapsed = time.monotonic() - pipeline_started + + metadata = { + "input": { + "filename": clip.source_path.name, + "fps": f"{clip.fps_num}/{clip.fps_den}", + "start_time_seconds": float(clip.start_time), + "num_frames": len(clip.frames), + "width": clip.width, + "height": clip.height, + }, + "motion": { + "method": "GVHMR", + "revision": motion.gvhmr_revision, + }, + "preprocessing": { + "source_crop_policy": conditioning_metadata["source_crop_policy"], + "foreground_model": conditioning_metadata["foreground_model"], + "framing": conditioning_metadata["framing"], + "skeleton_draw_scale": conditioning_metadata["skeleton_draw_scale"], + "target_render_deferred": bool( + conditioning_metadata.get("target_render_deferred", False) + ), + "target_render_overlap": conditioning_metadata.get("target_render_overlap"), + }, + "model": dict(model_identity), + "generation": { + "mode": settings.mode, + "seed": generated.seed, + "view_plan": { + **view_plan.to_dict(), + "num_layers": view_plan.num_layers, + "num_target_views": view_plan.num_target_views, + "groups_per_layer": view_plan.groups_per_layer, + "tcr_active": view_plan.tcr_active, + "routing_topology": "circular" if view_plan.closed_yaw else "open", + }, + "attention_backend": attention_backend, + "inference_steps": settings.num_inference_steps, + "elapsed_seconds": generated.elapsed_seconds, + "total_elapsed_seconds": total_elapsed, + "peak_vram_allocated_bytes": generated.peak_vram_allocated_bytes, + "peak_vram_reserved_bytes": generated.peak_vram_reserved_bytes, + }, + "output": { + "rcp_views": len(output_sparse), + "target_views": len(output_dense), + "frames_per_video": INFERENCE.num_frames, + "width": INFERENCE.width, + "height": INFERENCE.height, + "fps": f"{clip.fps_num}/{clip.fps_den}", + }, + "runtime": _runtime_metadata(generated.device), + } + write_json(root / "cameras.json", _camera_rig_payload(camera_payload, camera_records)) + write_json(root / "metadata.json", metadata) + return { + "attention_backend": attention_backend, + "num_rcp_videos": len(output_sparse), + "num_target_videos": len(output_dense), + "fps": f"{clip.fps_num}/{clip.fps_den}", + "peak_vram_allocated_bytes": generated.peak_vram_allocated_bytes, + "peak_vram_reserved_bytes": generated.peak_vram_reserved_bytes, + "total_pipeline_elapsed_seconds": total_elapsed, + } diff --git a/fdanyone/pipeline.py b/fdanyone/pipeline.py new file mode 100644 index 0000000000000000000000000000000000000000..459f1701f992b554aa4e9a25b8404f42696e308b --- /dev/null +++ b/fdanyone/pipeline.py @@ -0,0 +1,591 @@ +"""Top-level inference orchestration.""" + +from __future__ import annotations + +import json +import logging +import os +import subprocess +import sys +import tempfile +import time +from dataclasses import dataclass +from pathlib import Path +from typing import TYPE_CHECKING + +from fdanyone.assets import ( + CHECKPOINT, + HF_REPO_ID, + HF_REVISION, + BaseAssets, + resolve_base_assets, + resolve_checkpoint, + resolve_foreground_model, + resolve_regressor, +) +from fdanyone.config import INFERENCE, ModeSettings +from fdanyone.device import select_cuda_device +from fdanyone.download import ensure_example_video, ensure_models, ensure_smplx +from fdanyone.errors import ConfigurationError, FourDAnyoneError +from fdanyone.io import AtomicResultDirectory, remove_tree, write_json +from fdanyone.motion.gvhmr import MotionStageHook, validate_gvhmr +from fdanyone.motion.result import MotionResult +from fdanyone.video import ( + CanonicalClip, + decode_canonical_clip, + validate_required_video_codecs, + verify_lossless_video, + write_gvhmr_video, +) +from fdanyone.views import ViewPlan, resolve_view_plan + +LOGGER = logging.getLogger("fdanyone") + +if TYPE_CHECKING: + from fdanyone.model.inference import DenoiseStepHook + from fdanyone.skeleton.pipeline import Conditioning + + +def _data_paths(data_dir: str, video_path: str) -> tuple[Path, Path, Path]: + data_root = Path(data_dir).expanduser().resolve() + run_name = Path(video_path).stem + return ( + data_root, + data_root / "gvhmr" / "results" / run_name, + data_root / "fdanyone" / run_name, + ) + + +def _discard_scratch(path: Path) -> None: + """Best-effort cleanup that can never invalidate a published result. + + Some network filesystems keep an open, hidden tombstone after a file is + unlinked. Such a tombstone may remain ``EBUSY`` until this process exits, + so cleanup must not be part of the atomic publication transaction. + """ + + try: + remove_tree(path) + except OSError as exc: + LOGGER.warning( + "Could not remove temporary files at %s (%s). " + "The result is unaffected; the hidden scratch directory can be removed after this process exits.", + path, + exc, + ) + + +def _worker_environment() -> dict[str, str]: + """Give the short-lived GVHMR workers this checkout and stable CUDA flags.""" + + environment = os.environ.copy() + environment.update( + { + "TORCH_CUDNN_V8_API_DISABLED": "1", + "CUDNN_FRONTEND_DISABLE": "1", + "CUDNN_LOGINFO_DBG": "0", + "CUDNN_LOGDEST_DBG": "stderr", + "CUDA_DEVICE_MAX_CONNECTIONS": "1", + "NVIDIA_TF32_OVERRIDE": "0", + } + ) + environment.pop("PYTHONHOME", None) + environment["PYTHONPATH"] = str(Path(__file__).resolve().parent.parent) + return environment + + +def _cpu_worker_environment() -> dict[str, str]: + """Run pure CPU overlap workers without exposing a CUDA device.""" + + environment: dict[str, str] = _worker_environment() + environment["CUDA_VISIBLE_DEVICES"] = "" + return environment + + +def _run_motion( + *, + working_video: Path, + output_dir: Path, + gvhmr_root: Path, + device: str, + worker_python: str, + clip_metadata: Path, + inline_workers: bool = False, + on_motion_stage: MotionStageHook | None = None, +): + output_dir.mkdir(parents=True, exist_ok=True) + request_path = output_dir / ".motion-worker-request.json" + result_dir = output_dir / "result" + if inline_workers: + from fdanyone.motion.worker import run_motion_worker + + run_motion_worker( + gvhmr_root=gvhmr_root, + working_video=working_video, + clip_metadata=clip_metadata, + output_dir=output_dir / "runtime", + result_dir=result_dir, + device=device, + on_stage=on_motion_stage, + ) + return MotionResult.load(result_dir) + write_json( + request_path, + { + "gvhmr_root": str(gvhmr_root), + "working_video": str(working_video), + "clip_metadata": str(clip_metadata), + "output_dir": str(output_dir / "runtime"), + "result_dir": str(result_dir), + "device": device, + }, + ) + try: + subprocess.run( + [worker_python, "-m", "fdanyone.motion.worker", str(request_path)], + check=True, + env=_worker_environment(), + ) + finally: + request_path.unlink(missing_ok=True) + return MotionResult.load(result_dir) + + +def _build_conditioning( + *, + regressor: Path, + foreground_model: Path, + gvhmr_root: Path, + output_dir: Path, + device: str, + worker_python: str, + working_video: Path, + clip_metadata: Path, + motion_result_dir: Path, + view_plan: ViewPlan, + settings: ModeSettings, + inline_workers: bool = False, +) -> Conditioning: + from fdanyone.skeleton.pipeline import Conditioning + + if inline_workers: + from fdanyone.skeleton.worker import run_skeleton_worker + + run_skeleton_worker( + working_video=working_video, + clip_metadata=clip_metadata, + motion_result_dir=motion_result_dir, + regressor_path=regressor, + foreground_model_path=foreground_model, + gvhmr_root=gvhmr_root, + output_dir=output_dir, + device=device, + view_plan=view_plan, + defer_target_skeletons=settings.overlap_target_skeletons, + ) + else: + request_path = output_dir.parent / ".skeleton-worker-request.json" + write_json( + request_path, + { + "working_video": str(working_video), + "clip_metadata": str(clip_metadata), + "motion_result_dir": str(motion_result_dir), + "regressor_path": str(regressor), + "foreground_model_path": str(foreground_model), + "gvhmr_root": str(gvhmr_root), + "output_dir": str(output_dir), + "device": device, + "view_plan": view_plan.to_dict(), + "defer_target_skeletons": settings.overlap_target_skeletons, + }, + ) + try: + subprocess.run( + [ + worker_python, + "-m", + "fdanyone.skeleton.worker", + str(request_path), + ], + check=True, + env=_worker_environment(), + ) + finally: + request_path.unlink(missing_ok=True) + conditioning: Conditioning = Conditioning.load( + output_dir, + allow_pending_targets=settings.overlap_target_skeletons, + skeleton_video_decoder=settings.skeleton_video_decoder, + ) + render_request: Path = output_dir / "target-render-request.json" + if not render_request.is_file(): + return conditioning + + # The renderer stays a subprocess even under ``inline_workers``: it needs no + # GPU, and both it and its completion evidence fail closed unless + # CUDA_VISIBLE_DEVICES is disabled, which the parent process cannot offer. + render_log: Path = output_dir / "target-render.log" + with render_log.open("wb") as log_handle: + render_process: subprocess.Popen[bytes] = subprocess.Popen( + [worker_python, "-m", "fdanyone.skeleton.render_worker", str(render_request)], + env=_cpu_worker_environment(), + stdout=log_handle, + stderr=subprocess.STDOUT, + ) + completed: bool = False + failure: FourDAnyoneError | None = None + + def _validate_target_render() -> None: + return_code: int = render_process.wait() + if return_code != 0: + details: str = render_log.read_text(errors="replace")[-4000:] + raise FourDAnyoneError( + "CPU target skeleton renderer failed with " + f"exit {return_code}. Log tail:\n{details}" + ) + done_path: Path = output_dir / "target-render.done.json" + if not done_path.is_file(): + raise FourDAnyoneError( + f"CPU target skeleton renderer exited without its completion record: {done_path}." + ) + try: + render_evidence: object = json.loads(done_path.read_text()) + except json.JSONDecodeError as exc: + raise FourDAnyoneError( + f"CPU target skeleton completion record is invalid: {done_path}." + ) from exc + if ( + not isinstance(render_evidence, dict) + or int(render_evidence.get("schema_version", 0)) != 1 + or int(render_evidence.get("rendered_views", -1)) + != conditioning.view_plan.num_target_views + or render_evidence.get("cuda_visible_devices") not in ("", "-1") + ): + raise FourDAnyoneError( + "CPU target skeleton completion evidence failed its authenticity contract." + ) + metadata_path: Path = output_dir / "metadata.json" + conditioning_metadata: object = json.loads(metadata_path.read_text()) + if not isinstance(conditioning_metadata, dict): + raise FourDAnyoneError(f"Conditioning metadata must be an object: {metadata_path}.") + conditioning_metadata["target_render_overlap"] = render_evidence + write_json(metadata_path, conditioning_metadata) + + def wait_for_targets() -> None: + # Re-entry (the pipeline-level cleanup) must repeat the original + # verdict, not mask an informative failure with a generic one. + nonlocal completed, failure + if completed: + if failure is not None: + raise failure + return + completed = True + try: + _validate_target_render() + except FourDAnyoneError as exc: + failure = exc + raise + + return conditioning.with_target_waiter(wait_for_targets) + + +@dataclass(frozen=True) +class PreparedRun: + """Everything the generation phase needs from one prepared source clip.""" + + settings: ModeSettings + """Inference policy shared by both phases.""" + clip: CanonicalClip + """Decoded canonical clip on the frozen frame contract.""" + motion: MotionResult + """Published GVHMR result validated against the clip.""" + conditioning: Conditioning + """Source, proposal, and target conditioning on the resolved camera grid.""" + checkpoint: Path + """Resolved 4DAnyone DiT checkpoint.""" + base_assets: BaseAssets + """Resolved VAE, text-encoder, and tokenizer locations.""" + prompt_embedding_path: Path | None + """Exported prompt context standing in for the T5 encoder, when supplied.""" + model_identity: dict + """Model coordinates recorded in the published metadata.""" + device: str + """Selected CUDA device.""" + scratch: Path + """Hidden working directory that ``release_run`` removes.""" + motion_dir: Path + """Published GVHMR result directory.""" + result_dir: Path + """Destination the generation phase publishes to.""" + pipeline_started: float + """``time.monotonic`` reading taken when preparation began.""" + + +def prepare_run( + *, + settings: ModeSettings, + video_path: str, + data_dir: str, + model_dir: str, + checkpoint_path: str | None, + mhr70_regressor_path: str | None, + gvhmr_root: str, + device: str, + start_time: float, + target_fps: str | float, + views_per_layer: int, + layer_pitches: list[int], + start_yaw: int, + yaw_span: int, + views_per_group: int | str, + enable_rcp: bool, + enable_tcr: bool, + inline_workers: bool = False, + on_motion_stage: MotionStageHook | None = None, + prompt_embedding_path: Path | None = None, +) -> PreparedRun: + """Recover motion and build conditioning for one source video.""" + + pipeline_started = time.monotonic() + logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") + view_plan = resolve_view_plan( + views_per_layer=views_per_layer, + layer_pitches=layer_pitches, + start_yaw=start_yaw, + yaw_span=yaw_span, + views_per_group=views_per_group, + enable_rcp=enable_rcp, + enable_tcr=enable_tcr, + ) + data_root, motion_dir, result_dir = _data_paths(data_dir, video_path) + run_name = Path(video_path).stem + atomic = AtomicResultDirectory(result_dir) + # Fail before asset resolution or video decode; the context manager + # checks again later in case another process creates the path. + if os.path.lexists(atomic.destination): + raise ConfigurationError( + f"4DAnyone result already exists: {atomic.destination}. Choose a new --data_dir or input filename." + ) + validate_required_video_codecs() + device, _ = select_cuda_device(device) + + ensure_example_video(video_path) + # Resolve the licensed body model before starting the much larger public + # model download. Interactive use continues automatically after setup; + # background jobs receive an actionable error instead of hanging. + ensure_smplx(model_dir, gvhmr_root) + ensure_models(model_dir, gvhmr_root) + gvhmr_root, gvhmr_revision = validate_gvhmr(gvhmr_root) + worker_python = os.path.abspath(sys.executable) + + regressor = resolve_regressor(mhr70_regressor_path, model_dir=model_dir) + foreground_model = resolve_foreground_model(model_dir) + canonical_fps = None if str(target_fps).lower() == "auto" else target_fps + clip = decode_canonical_clip( + video_path, + num_frames=INFERENCE.num_frames, + start_time=start_time, + fps=canonical_fps, + ) + data_root.mkdir(parents=True, exist_ok=True) + scratch = Path(tempfile.mkdtemp(prefix=f".{run_name}.scratch-", dir=data_root)) + conditioning: Conditioning | None = None + try: + clip_metadata = scratch / "canonical_clip.json" + clip.write_metadata(clip_metadata) + working_video = write_gvhmr_video(clip, scratch / "canonical_clip.mp4") + + if os.path.lexists(motion_dir): + if motion_dir.is_symlink() or not motion_dir.is_dir(): + raise ConfigurationError(f"GVHMR result path is not a regular directory: {motion_dir}") + motion = MotionResult.load(motion_dir) + if motion.gvhmr_revision != gvhmr_revision: + raise ConfigurationError( + f"Existing GVHMR result at {motion_dir} was produced by " + f"GVHMR@{motion.gvhmr_revision}, not GVHMR@{gvhmr_revision}." + ) + motion.validate_against_clip(clip) + LOGGER.info("Reusing validated GVHMR result at %s", motion_dir) + else: + with AtomicResultDirectory(motion_dir) as motion_work: + motion = _run_motion( + working_video=working_video, + output_dir=scratch / "gvhmr", + gvhmr_root=gvhmr_root, + device=device, + worker_python=worker_python, + clip_metadata=clip_metadata, + inline_workers=inline_workers, + on_motion_stage=on_motion_stage, + ) + motion.validate_against_clip(clip) + motion.save(motion_work) + + checkpoint = resolve_checkpoint(checkpoint_path, model_dir=model_dir) + base_assets = resolve_base_assets( + model_dir, + settings, + have_prompt_embedding=prompt_embedding_path is not None, + ) + # Record the published identity only for the published checkpoint; an + # explicit override must not claim the frozen Hugging Face coordinates. + if checkpoint_path is None: + model_identity = {"checkpoint": CHECKPOINT, "repo_id": HF_REPO_ID, "revision": HF_REVISION} + else: + model_identity = {"checkpoint": checkpoint.name, "source": "local_override"} + + conditioning = _build_conditioning( + regressor=regressor, + foreground_model=foreground_model, + gvhmr_root=gvhmr_root, + output_dir=scratch / "conditioning", + device=device, + worker_python=worker_python, + working_video=working_video, + clip_metadata=clip_metadata, + motion_result_dir=motion_dir, + view_plan=view_plan, + settings=settings, + inline_workers=inline_workers, + ) + if conditioning.num_frames != len(clip.frames) or ( + conditioning.fps_num, + conditioning.fps_den, + ) != ( + clip.fps_num, + clip.fps_den, + ): + raise ConfigurationError("Skeleton conditioning does not match the canonical clip timeline.") + # Re-decode the worker-produced source before it becomes a model tensor. + verify_lossless_video(clip, conditioning.source_video) + except BaseException: + if conditioning is not None and conditioning.target_waiter is not None: + conditioning.wait_for_target_skeletons() + _discard_scratch(scratch) + raise + return PreparedRun( + settings=settings, + clip=clip, + motion=motion, + conditioning=conditioning, + checkpoint=checkpoint, + base_assets=base_assets, + prompt_embedding_path=prompt_embedding_path, + model_identity=model_identity, + device=device, + scratch=scratch, + motion_dir=motion_dir, + result_dir=result_dir, + pipeline_started=pipeline_started, + ) + + +def generate_run( + prepared: PreparedRun, + *, + seed: int, + on_denoise_step: DenoiseStepHook | None = None, +) -> dict: + """Generate and publish every view for one prepared run. + + The prompt embedding is fixed by ``prepare_run``, because whether it exists + decides whether the T5 encoder had to be resolved at all. + """ + + # Heavy rendering and generation are imported only after the motion + # contract has been materialized, keeping CLI/help and CPU tests light. + from fdanyone.model.inference import generate_views + from fdanyone.output import export_result + + with AtomicResultDirectory(prepared.result_dir) as work: + generated = generate_views( + clip=prepared.clip, + conditioning=prepared.conditioning, + checkpoint_path=prepared.checkpoint, + assets=prepared.base_assets, + output_dir=prepared.scratch / "generation", + device=prepared.device, + seed=seed, + settings=prepared.settings, + on_denoise_step=on_denoise_step, + prompt_embedding_path=prepared.prompt_embedding_path, + ) + summary = export_result( + clip=prepared.clip, + conditioning=prepared.conditioning, + generated=generated, + destination=work, + motion=prepared.motion, + model_identity=prepared.model_identity, + pipeline_started=prepared.pipeline_started, + settings=prepared.settings, + ) + summary["result_dir"] = str(prepared.result_dir) + summary["motion_dir"] = str(prepared.motion_dir) + return summary + + +def release_run(prepared: PreparedRun) -> None: + """Settle deferred target rendering and drop the run's scratch directory.""" + + if prepared.conditioning.target_waiter is not None: + prepared.conditioning.wait_for_target_skeletons() + _discard_scratch(prepared.scratch) + + +def run_pipeline( + *, + settings: ModeSettings, + video_path: str, + data_dir: str, + model_dir: str, + checkpoint_path: str | None, + mhr70_regressor_path: str | None, + gvhmr_root: str, + device: str, + start_time: float, + target_fps: str | float, + seed: int, + views_per_layer: int, + layer_pitches: list[int], + start_yaw: int, + yaw_span: int, + views_per_group: int | str, + enable_rcp: bool, + enable_tcr: bool, + inline_workers: bool = False, + on_motion_stage: MotionStageHook | None = None, + on_denoise_step: DenoiseStepHook | None = None, + prompt_embedding_path: Path | None = None, +) -> dict: + """Execute inference and publish reusable GVHMR plus 4DAnyone results.""" + + if seed < 0: + raise ConfigurationError(f"seed must be non-negative, got {seed}.") + prepared = prepare_run( + settings=settings, + video_path=video_path, + data_dir=data_dir, + model_dir=model_dir, + checkpoint_path=checkpoint_path, + mhr70_regressor_path=mhr70_regressor_path, + gvhmr_root=gvhmr_root, + device=device, + start_time=start_time, + target_fps=target_fps, + views_per_layer=views_per_layer, + layer_pitches=layer_pitches, + start_yaw=start_yaw, + yaw_span=yaw_span, + views_per_group=views_per_group, + enable_rcp=enable_rcp, + enable_tcr=enable_tcr, + inline_workers=inline_workers, + on_motion_stage=on_motion_stage, + prompt_embedding_path=prompt_embedding_path, + ) + try: + return generate_run(prepared, seed=seed, on_denoise_step=on_denoise_step) + finally: + release_run(prepared) diff --git a/fdanyone/runs.py b/fdanyone/runs.py new file mode 100644 index 0000000000000000000000000000000000000000..e8df2fe0030ace9eb6fb9946fec7631493fd7a45 --- /dev/null +++ b/fdanyone/runs.py @@ -0,0 +1,161 @@ +"""Read what one finished 4DAnyone inference run left on disk. + +Two readers live here. ``discover_run`` collects every artifact of a clip -- +motion, generated videos, skeletons, and the source clip -- into a +``RunLayout``. The rig readers parse the ``cameras.json`` that +``fdanyone.output`` writes next to the generated videos. +""" + +from __future__ import annotations + +import json +import logging +from dataclasses import dataclass +from pathlib import Path + +from fdanyone.errors import FourDAnyoneError +from fdanyone.io import read_json + +LOGGER = logging.getLogger("fdanyone.runs") + + +# --------------------------------------------------------------------------- +# Run layout +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class RunLayout: + """Files that one inference run left on disk for a single clip.""" + + clip: str + data_dir: Path + result_dir: Path + motion_dir: Path + cameras_json: Path | None + metadata_json: Path | None + source_video: Path | None + dense_videos: tuple[tuple[int, Path], ...] + skeleton_videos: tuple[tuple[int, Path], ...] + sparse_videos: tuple[tuple[int, Path], ...] + + @property + def has_motion(self) -> bool: + return (self.motion_dir / "motion.json").is_file() and (self.motion_dir / "motion.safetensors").is_file() + + +def numbered_videos(directory: Path) -> tuple[tuple[int, Path], ...]: + """Return ``(camera_id, path)`` pairs for ``NN.mp4`` files, ID-ordered.""" + + if not directory.is_dir(): + return () + found: list[tuple[int, Path]] = [] + for path in sorted(directory.glob("*.mp4")): + try: + found.append((int(path.stem), path)) + except ValueError: + LOGGER.warning("Ignoring video with a non-numeric name: %s", path) + return tuple(sorted(found)) + + +def find_source_video(data_dir: Path, clip: str, filename: str | None) -> Path | None: + """Locate the input clip that produced the run.""" + + candidates: list[Path] = [] + if filename: + candidates.append(data_dir / "source" / "pexels" / filename) + for suffix in (".mp4", ".mov", ".mkv", ".webm", ".MP4", ".MOV"): + candidates.append(data_dir / "source" / "pexels" / f"{clip}{suffix}") + for candidate in candidates: + if candidate.is_file(): + return candidate + # Only the input and result trees can hold a clip; the rest of the data + # root is model output that a recursive walk would scan for nothing. + roots = tuple(root for root in (data_dir / "source", data_dir / "fdanyone") if root.is_dir()) + if filename: + matches = sorted(path for root in roots for path in root.rglob(filename) if path.is_file()) + if matches: + return matches[0] + matches = sorted( + path for root in roots for path in root.rglob(f"{clip}.*") if path.is_file() and path.suffix != ".rrd" + ) + return matches[0] if matches else None + + +def discover_run(data_dir: Path, clip: str) -> RunLayout: + """Collect every artifact of ``clip`` and reject an empty run.""" + + result_dir = data_dir / "fdanyone" / clip + motion_dir = data_dir / "gvhmr" / "results" / clip + metadata_json = result_dir / "metadata.json" + cameras_json = result_dir / "cameras.json" + metadata = read_json(metadata_json) + filename = None + if metadata is not None: + filename = str(metadata.get("input", {}).get("filename") or "") or None + + layout = RunLayout( + clip=clip, + data_dir=data_dir, + result_dir=result_dir, + motion_dir=motion_dir, + cameras_json=cameras_json if cameras_json.is_file() else None, + metadata_json=metadata_json if metadata_json.is_file() else None, + source_video=find_source_video(data_dir, clip, filename), + dense_videos=numbered_videos(result_dir / "videos" / "dense"), + skeleton_videos=numbered_videos(result_dir / "skeletons"), + sparse_videos=numbered_videos(result_dir / "videos" / "sparse"), + ) + if not layout.has_motion and not layout.dense_videos and layout.source_video is None: + raise FourDAnyoneError( + f"Nothing to visualize for clip {clip!r}. Expected at least one of:\n" + f" {motion_dir / 'motion.safetensors'} (GVHMR motion)\n" + f" {result_dir / 'videos' / 'dense'} (generated views)\n" + f" a source clip named {clip}.* under {data_dir}\n" + "Run `python inference.py --video_path ` first." + ) + return layout + + +# --------------------------------------------------------------------------- +# Camera rig +# --------------------------------------------------------------------------- + + +def read_cameras(result: Path) -> dict: + """Read the ``cameras.json`` a result directory must carry.""" + + path = result / "cameras.json" + try: + cameras = json.loads(path.read_text()) + except (FileNotFoundError, json.JSONDecodeError) as exc: + raise FourDAnyoneError(f"Cannot read {path}: {exc}") from exc + if not isinstance(cameras, dict): + raise FourDAnyoneError("cameras.json must contain a JSON object.") + return cameras + + +def camera_records(rig: dict) -> list[dict]: + """Validate an OPENCV rig and return its ID-ordered camera records.""" + + if not isinstance(rig, dict) or rig.get("camera_model") != "OPENCV": + raise FourDAnyoneError("cameras.json has no supported OPENCV camera rig.") + cameras = rig.get("cameras") + if not isinstance(cameras, list) or not cameras: + raise FourDAnyoneError("Camera rig must contain at least one camera.") + if [camera.get("camera_id") for camera in cameras if isinstance(camera, dict)] != list(range(len(cameras))): + raise FourDAnyoneError("Camera rig must be ordered by camera ID.") + return cameras + + +def dense_video_paths(result: Path, cameras: list[dict]) -> tuple[Path, ...]: + """Return the generated video of every camera, in rig order.""" + + paths = [] + for camera in cameras: + camera_id = int(camera["camera_id"]) + relative = f"videos/dense/{camera_id:02d}.mp4" + if camera.get("video") != relative: + raise FourDAnyoneError(f"Camera {camera_id:02d} points to the wrong target video.") + paths.append(result / relative) + return tuple(paths) diff --git a/fdanyone/skeleton/__init__.py b/fdanyone/skeleton/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..c0f0eff33ecf0478f8e4fa33755b92565c531ebf --- /dev/null +++ b/fdanyone/skeleton/__init__.py @@ -0,0 +1 @@ +"""MHR70 projection and Goliath40 conditioning renderer.""" diff --git a/fdanyone/skeleton/keypoints.py b/fdanyone/skeleton/keypoints.py new file mode 100644 index 0000000000000000000000000000000000000000..5bae6d900e62a0047e1ab6e272c119638b37f59b --- /dev/null +++ b/fdanyone/skeleton/keypoints.py @@ -0,0 +1,170 @@ +"""Minimal MHR70/Goliath40 conditioning schema. + +The exact keypoint names/order and no-finger link closure were extracted from +``facebookresearch/sapiens2`` revision +``0e51c12d7c7257d88431b2d50e523a7b03004854`` and remain subject to the +Sapiens2 License. A copy is kept at +``third_party/licenses/SAPIENS2_LICENSE.md``. The RGB palette, per-keypoint +color policy, and renderer implementation are original 4DAnyone material +licensed under Apache-2.0. +""" + +from __future__ import annotations + +# The exact names/order below are derived from Sapiens2, Copyright (c) Meta +# Platforms, Inc. and affiliates, under the Sapiens2 License noted above. + +KEYPOINT_NAMES = ( + "nose", + "left-eye", + "right-eye", + "left-ear", + "right-ear", + "left-shoulder", + "right-shoulder", + "left-elbow", + "right-elbow", + "left-hip", + "right-hip", + "left-knee", + "right-knee", + "left-ankle", + "right-ankle", + "left-big-toe-tip", + "left-small-toe-tip", + "left-heel", + "right-big-toe-tip", + "right-small-toe-tip", + "right-heel", + "right-thumb-tip", + "right-thumb-first-joint", + "right-thumb-second-joint", + "right-thumb-third-joint", + "right-index-tip", + "right-index-first-joint", + "right-index-second-joint", + "right-index-third-joint", + "right-middle-tip", + "right-middle-first-joint", + "right-middle-second-joint", + "right-middle-third-joint", + "right-ring-tip", + "right-ring-first-joint", + "right-ring-second-joint", + "right-ring-third-joint", + "right-pinky-tip", + "right-pinky-first-joint", + "right-pinky-second-joint", + "right-pinky-third-joint", + "right-wrist", + "left-thumb-tip", + "left-thumb-first-joint", + "left-thumb-second-joint", + "left-thumb-third-joint", + "left-index-tip", + "left-index-first-joint", + "left-index-second-joint", + "left-index-third-joint", + "left-middle-tip", + "left-middle-first-joint", + "left-middle-second-joint", + "left-middle-third-joint", + "left-ring-tip", + "left-ring-first-joint", + "left-ring-second-joint", + "left-ring-third-joint", + "left-pinky-tip", + "left-pinky-first-joint", + "left-pinky-second-joint", + "left-pinky-third-joint", + "left-wrist", + "left-olecranon", + "right-olecranon", + "left-cubital-fossa", + "right-cubital-fossa", + "left-acromion", + "right-acromion", + "neck", +) + +# First-party 4DAnyone color palette, licensed under Apache-2.0. +BLUE = (116, 192, 252) +GREEN = (130, 186, 129) +ORANGE = (248, 129, 81) +TEAL = (99, 230, 190) +YELLOW = (255, 212, 59) +PINK = (229, 153, 247) +PURPLE = (177, 151, 252) +RED = (255, 135, 135) + +# The schema ids and link endpoints are Sapiens2-derived. The RGB assignments +# are first-party 4DAnyone material. +# (schema id, first keypoint, second keypoint, RGB, major) +LINKS = ( + (0, 13, 11, TEAL, True), + (1, 11, 9, TEAL, True), + (2, 14, 12, YELLOW, True), + (3, 12, 10, YELLOW, True), + (4, 9, 10, BLUE, True), + (5, 5, 9, GREEN, True), + (6, 6, 10, ORANGE, True), + (7, 5, 6, BLUE, True), + (8, 5, 7, TEAL, True), + (9, 6, 8, YELLOW, True), + (10, 7, 62, TEAL, True), + (11, 8, 41, YELLOW, True), + (12, 1, 2, BLUE, True), + (13, 0, 1, GREEN, True), + (14, 0, 2, ORANGE, True), + (15, 1, 3, GREEN, True), + (16, 2, 4, ORANGE, True), + (17, 3, 5, GREEN, True), + (18, 4, 6, ORANGE, True), + (19, 13, 15, TEAL, True), + (20, 13, 16, TEAL, True), + (21, 13, 17, TEAL, True), + (22, 14, 18, YELLOW, True), + (23, 14, 19, YELLOW, True), + (24, 14, 20, YELLOW, True), + (25, 62, 45, YELLOW, False), + (29, 62, 49, PINK, False), + (33, 62, 53, PURPLE, False), + (37, 62, 57, RED, False), + (41, 62, 61, TEAL, False), + (45, 41, 24, YELLOW, False), + (49, 41, 28, PINK, False), + (53, 41, 32, PURPLE, False), + (57, 41, 36, RED, False), + (61, 41, 40, TEAL, False), + (65, 5, 10, PURPLE, True), + (66, 6, 9, PURPLE, True), + (67, 3, 4, PURPLE, True), +) + +EXTRA_KEYPOINT_IDS = frozenset(range(63, 70)) + + +def keypoint_color(keypoint_id: int) -> tuple[int, int, int]: + name = KEYPOINT_NAMES[keypoint_id] + if name == "nose" or name == "neck": + return BLUE + if name in {"left-eye", "left-ear"}: + return GREEN + if name in {"right-eye", "right-ear"}: + return ORANGE + if "thumb" in name: + return YELLOW + if "forefinger" in name or "index" in name: + return PINK + if "middle" in name: + return PURPLE + if "ring" in name: + return RED + if "pinky" in name: + return TEAL + return TEAL if name.startswith("left-") else YELLOW if name.startswith("right-") else BLUE + + +VISIBLE_KEYPOINT_IDS = ( + frozenset({index for _, first, second, _, _ in LINKS for index in (first, second)}) | EXTRA_KEYPOINT_IDS +) diff --git a/fdanyone/skeleton/pipeline.py b/fdanyone/skeleton/pipeline.py new file mode 100644 index 0000000000000000000000000000000000000000..469f0b779bce5d822cbbd1bf85f6f2707a160f47 --- /dev/null +++ b/fdanyone/skeleton/pipeline.py @@ -0,0 +1,839 @@ +"""Build source-aware Goliath40 conditioning for multi-view generation.""" + +from __future__ import annotations + +import json +import logging +from collections.abc import Callable, Iterable +from contextlib import contextmanager +from dataclasses import asdict, dataclass, field, replace +from fractions import Fraction +from pathlib import Path + +import numpy as np + +from fdanyone.assets import BIREFNET_REPO_ID, BIREFNET_REVISION +from fdanyone.config import CAMERA, CROP, FOREGROUND, INFERENCE, SKELETON, CameraConfig +from fdanyone.errors import AssetError, FourDAnyoneError +from fdanyone.foreground import predict_foreground_masks +from fdanyone.geometry.cameras import ( + CAMERA_FRAME, + WORLD_FRAME, + Camera, + camera_grid, + camera_ring, + project_points, + reference_intrinsics, +) +from fdanyone.geometry.crop import ( + Crop, + center_crop, + crop_from_bounds, + mask_bounds, + transform_intrinsics, +) +from fdanyone.geometry.framing import analyze_input_framing, solve_sequence_framing +from fdanyone.io import write_json +from fdanyone.motion.gvhmr import gvhmr_imports, validate_gvhmr +from fdanyone.motion.result import MotionResult +from fdanyone.skeleton.keypoints import KEYPOINT_NAMES +from fdanyone.skeleton.renderer import ( + estimate_body_height, + projected_body_scales, + render_goliath40, +) +from fdanyone.vendor.pytorch3d_compat import ( + install_if_needed as install_pytorch3d_compat, +) +from fdanyone.video import ( + CanonicalClip, + iter_rgb_video, + write_lossless_video, + write_video, +) +from fdanyone.views import ViewPlan + +LOGGER = logging.getLogger("fdanyone") + + +@dataclass(frozen=True) +class SkeletonVideo: + path: Path + crop: Crop + + +@dataclass(frozen=True) +class Conditioning: + root: Path + source_video: Path + source_crop: Crop + target_skeletons: tuple[SkeletonVideo, ...] + rcp_skeletons: tuple[SkeletonVideo, ...] + view_plan: ViewPlan + fps_num: int + fps_den: int + num_frames: int + skeleton_video_decoder: str = "pyav" + """Backend used when loading the rendered skeleton videos.""" + target_waiter: Callable[[], None] | None = field(default=None, repr=False, compare=False) + """Optional completion barrier for deferred target skeleton videos.""" + + def load_source_tensor(self): + return _video_tensor(self.source_video, self.num_frames, crop=self.source_crop) + + def load_skeleton_tensor( + self, + skeletons: Iterable[SkeletonVideo], + *, + device: str = "cuda", + ): + import torch + + videos = [ + _video_tensor( + item.path, + self.num_frames, + crop=item.crop, + decoder=self.skeleton_video_decoder, + device=device, + ) + for item in skeletons + ] + return torch.cat(videos, dim=0) + + def wait_for_target_skeletons(self) -> None: + """Wait for a deferred CPU renderer and validate every target artifact.""" + + if self.target_waiter is not None: + self.target_waiter() + missing: list[str] = [str(item.path) for item in self.target_skeletons if not item.path.is_file()] + if missing: + raise FourDAnyoneError(f"Target skeleton rendering is incomplete: {missing}.") + + def with_target_waiter(self, waiter: Callable[[], None]) -> Conditioning: + """Return this file contract with one process-completion barrier attached.""" + + return replace(self, target_waiter=waiter) + + @classmethod + def load( + cls, + directory: str | Path, + *, + allow_pending_targets: bool = False, + skeleton_video_decoder: str = "pyav", + ) -> Conditioning: + """Load conditioning artifacts, optionally before deferred targets finish.""" + + root = Path(directory).expanduser().resolve() + camera_payload = json.loads((root / "cameras.json").read_text()) + metadata = json.loads((root / "metadata.json").read_text()) + try: + view_plan = ViewPlan.from_dict(metadata["view_plan"]) + except (KeyError, TypeError) as exc: + raise FourDAnyoneError("Conditioning artifacts have no valid view plan.") from exc + records = camera_payload["cameras"] + if [int(record["camera_id"]) for record in records] != list(range(view_plan.num_target_views)): + raise FourDAnyoneError("Target conditioning cameras are not in canonical order.") + for record, view in zip(records, view_plan.target_views, strict=True): + if ( + int(record.get("layer_index", -1)) != view.layer_index + or int(record.get("pitch_degrees", 1000)) != view.pitch + or abs(float(record.get("yaw_degrees", 1000.0)) - view.yaw) > 1e-8 + ): + raise FourDAnyoneError("Target conditioning cameras do not match the resolved view layout.") + if camera_payload.get("front_camera_ids") != list(view_plan.front_camera_ids): + raise FourDAnyoneError("Target conditioning has the wrong frontal-camera IDs.") + rcp_records = camera_payload.get("rcp_cameras", []) + if [int(record["camera_id"]) for record in rcp_records] != list(view_plan.rcp_camera_ids): + raise FourDAnyoneError("RCP conditioning cameras do not match the resolved view plan.") + + def skeletons(camera_records: list[dict]) -> tuple[SkeletonVideo, ...]: + return tuple( + SkeletonVideo(root / record["skeleton_video"], Crop(**record["crop"])) for record in camera_records + ) + + target_skeletons = skeletons(records) + rcp_skeletons = skeletons(rcp_records) + required: tuple[Path, ...] = ( + root / metadata["source_video"], + *(item.path for item in rcp_skeletons), + *((item.path for item in target_skeletons) if not allow_pending_targets else ()), + ) + missing = [str(path) for path in required if not path.is_file()] + if missing: + raise FourDAnyoneError(f"Conditioning artifacts are incomplete: {missing}.") + return cls( + root=root, + source_video=root / metadata["source_video"], + source_crop=Crop(**metadata["source_crop"]), + target_skeletons=target_skeletons, + rcp_skeletons=rcp_skeletons, + view_plan=view_plan, + fps_num=int(metadata["fps_num"]), + fps_den=int(metadata["fps_den"]), + num_frames=int(metadata["num_frames"]), + skeleton_video_decoder=skeleton_video_decoder, + ) + + +@dataclass(frozen=True) +class _BodyGeometry: + vertices_world: np.ndarray + joints_world: np.ndarray + keypoints_world: np.ndarray + keypoints_incam: np.ndarray + motion_world_to_canonical_world: np.ndarray + regressor_metadata: dict[str, int | str] + + +def _assert_nvdec_engaged( + decoder, + *, + decoded_device_type: str, + declared_frames: int, + decoded_frames: int, + expected_frames: int, +) -> None: + """Fail closed unless TorchCodec proves an exact, GPU-native NVDEC decode.""" + + if declared_frames != expected_frames: + raise FourDAnyoneError( + f"NVDEC video declares {declared_frames} frames, expected {expected_frames}." + ) + if decoded_frames != expected_frames: + raise FourDAnyoneError( + f"NVDEC produced {decoded_frames} frames, expected {expected_frames}." + ) + fallback = decoder.cpu_fallback + if not fallback.status_known: + raise FourDAnyoneError(f"NVDEC CPU fallback status is unknown: {fallback}") + if fallback: + raise FourDAnyoneError(f"NVDEC fell back to CPU: {fallback}") + if decoded_device_type != "cuda": + raise FourDAnyoneError( + f"NVDEC returned a {decoded_device_type.upper()} tensor instead of CUDA." + ) + + +def _crop_and_normalize(tensor, crop: Crop | None): + """Apply the scaled crop, restore the output size, and map to [-1, 1].""" + + import torchvision.transforms.functional as transform + from torchvision.transforms import InterpolationMode + + if crop is not None: + height, width = tensor.shape[-2:] + scaled = _scale_crop( + crop, + scale_y=height / crop.original_height, + scale_x=width / crop.original_width, + ) + tensor = transform.crop(tensor, scaled.top, scaled.left, scaled.height, scaled.width) + if tensor.shape[-2:] != (crop.output_height, crop.output_width): + tensor = transform.resize( + tensor, + (crop.output_height, crop.output_width), + interpolation=InterpolationMode.BICUBIC, + antialias=True, + ) + tensor = tensor.clamp_(0.0, 1.0) + return tensor.mul_(2.0).sub_(1.0) + + +def _nvdec_video_tensor( + path: Path, + num_frames: int, + *, + crop: Crop | None, + device: str, +): + """Decode one complete skeleton video with TorchCodec's NVDEC backend.""" + + import av # noqa: F401 + import torch + from torchcodec.decoders import VideoDecoder + + decoder = VideoDecoder( + path, + dimension_order="NCHW", + device=device, + output_dtype=torch.uint8, + seek_mode="exact", + ) + frames = decoder.get_frames_in_range(0, num_frames).data + _assert_nvdec_engaged( + decoder, + decoded_device_type=frames.device.type, + declared_frames=len(decoder), + decoded_frames=int(frames.shape[0]), + expected_frames=num_frames, + ) + tensor = _crop_and_normalize(frames.float().div_(255.0), crop) + return tensor.permute(1, 0, 2, 3).unsqueeze(0).contiguous() + + +def _video_tensor( + path: Path, + num_frames: int, + *, + crop: Crop | None = None, + decoder: str = "pyav", + device: str = "cuda", +): + if decoder == "torchcodec_cuda": + if not device.startswith("cuda"): + raise FourDAnyoneError( + f"TorchCodec skeleton decoding requires a CUDA device, got {device!r}." + ) + return _nvdec_video_tensor( + path, + num_frames, + crop=crop, + device=device, + ) + if decoder != "pyav": + raise FourDAnyoneError(f"Unknown skeleton video decoder {decoder!r}.") + + import torch + + output_frames = [] + for frame in iter_rgb_video(path): + tensor = torch.from_numpy(frame).permute(2, 0, 1).float().div_(255.0) + output_frames.append(_crop_and_normalize(tensor, crop)) + if len(output_frames) != num_frames: + raise FourDAnyoneError(f"Video {path} has {len(output_frames)} decoded frames, expected {num_frames}.") + return torch.stack(output_frames, dim=1).unsqueeze(0).contiguous() + + +def _safe_regressor_metadata(raw_metadata: object, support_shape: tuple[int, ...]) -> dict[str, int | str]: + del raw_metadata + return { + "format": "sparse_vertex_regressor", + "num_keypoints": int(support_shape[0]), + "support_vertices_per_keypoint": int(support_shape[1]), + } + + +def _load_regressor(path: Path, device): + import torch + + data = torch.load(path, map_location="cpu", weights_only=True) + support = data["support_vertex_ids"].detach().long().to(device) + weights = data["weights"].detach().float().to(device) + names = tuple(str(value) for value in data["keypoint_names"]) + if support.shape != weights.shape or support.shape[0] != 70: + raise AssetError(f"Unexpected MHR70 regressor shapes: support={support.shape}, weights={weights.shape}.") + if names != KEYPOINT_NAMES: + raise AssetError("MHR70 regressor keypoint order does not match the frozen Goliath70 schema.") + return support, weights, _safe_regressor_metadata(data.get("metadata"), tuple(support.shape)) + + +@contextmanager +def _gvhmr_geometry_context(gvhmr_root: Path): + install_pytorch3d_compat() + with gvhmr_imports(gvhmr_root): + yield + + +def _body_geometry( + motion: MotionResult, + regressor_path: Path, + gvhmr_root: Path, + device: str, +) -> _BodyGeometry: + import torch + + body_model = gvhmr_root / "inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz" + if not body_model.is_file(): + raise AssetError( + "The licensed SMPL-X body model is missing. Run `python scripts/download_smplx.py`; " + f"expected the GVHMR compatibility link at {body_model}." + ) + utility_root = gvhmr_root / "hmr4d/utils/body_model" + smplx_to_smpl_path = utility_root / "smplx2smpl_sparse.pt" + joint_regressor_path = utility_root / "smpl_neutral_J_regressor.pt" + for path in (smplx_to_smpl_path, joint_regressor_path): + if not path.is_file(): + raise AssetError(f"GVHMR body-model utility is missing: {path}") + + torch_device = torch.device(device) + support, weights, regressor_metadata = _load_regressor(regressor_path, torch_device) + with _gvhmr_geometry_context(gvhmr_root): + from hmr4d.utils.geo_transform import apply_T_on_points, compute_T_ayfz2ay + from hmr4d.utils.smplx_utils import make_smplx + + smplx = make_smplx("supermotion").to(torch_device).eval() + global_parameters = {name: value.to(torch_device) for name, value in motion.smpl_params_global.items()} + incam_parameters = {name: value.to(torch_device) for name, value in motion.smpl_params_incam.items()} + with torch.inference_mode(): + vertices_global = smplx(**global_parameters).vertices.detach() + vertices_incam = smplx(**incam_parameters).vertices.detach() + if tuple(vertices_global.shape[1:]) != (10475, 3) or vertices_incam.shape != vertices_global.shape: + raise FourDAnyoneError( + "Expected matching global/incam SMPL-X vertices [frames,10475,3], got " + f"{tuple(vertices_global.shape)} and {tuple(vertices_incam.shape)}." + ) + keypoints_global = (vertices_global[:, support] * weights[None, :, :, None]).sum(dim=2) + keypoints_incam = (vertices_incam[:, support] * weights[None, :, :, None]).sum(dim=2) + smplx_to_smpl = torch.load(smplx_to_smpl_path, map_location=torch_device, weights_only=True) + joint_regressor = torch.load(joint_regressor_path, map_location=torch_device, weights_only=True) + vertices_smpl = torch.stack([torch.matmul(smplx_to_smpl, frame) for frame in vertices_global]) + offset = torch.einsum("jv,vi->ji", joint_regressor, vertices_smpl[0])[0] + offset = offset.clone() + offset[1] = vertices_smpl[..., 1].min() + vertices_offset = vertices_smpl - offset + first_joints = torch.einsum("jv,lvi->lji", joint_regressor, vertices_offset[[0]]) + transform = compute_T_ayfz2ay(first_joints, inverse=True) + vertices_world = apply_T_on_points(vertices_offset, transform) + keypoints_world = apply_T_on_points(keypoints_global - offset, transform) + joints_world = torch.einsum("jv,lvi->lji", joint_regressor, vertices_world) + + world_transform = transform[0].detach().clone() + world_transform[:3, 3] -= world_transform[:3, :3] @ offset + result = _BodyGeometry( + vertices_world.detach().cpu().numpy().astype(np.float32), + joints_world.detach().cpu().numpy().astype(np.float32), + keypoints_world.detach().cpu().numpy().astype(np.float32), + keypoints_incam.detach().cpu().numpy().astype(np.float32), + world_transform.detach().cpu().numpy().astype(np.float64), + regressor_metadata, + ) + del smplx, vertices_global, vertices_incam, vertices_smpl, vertices_world, keypoints_global, keypoints_world + torch.cuda.empty_cache() + return result + + +def _front_direction(joints: np.ndarray) -> np.ndarray: + first = joints[0] + left = first[1, [0, 2]] - first[2, [0, 2]] + first[16, [0, 2]] - first[17, [0, 2]] + norm = float(np.linalg.norm(left)) + if norm <= 1e-8: + return np.array([0.0, 0.0, -1.0], dtype=np.float64) + left /= norm + return np.array([left[1], 0.0, -left[0]], dtype=np.float64) + + +def _projection_shape(height: int, width: int, max_render_height: int = 1280) -> tuple[int, int]: + divisor = 2 + while height / divisor > max_render_height: + divisor += 1 + return height // divisor, width // divisor + + +def _output_skeleton_shape(height: int, width: int) -> tuple[int, int]: + while max(height, width) > INFERENCE.skeleton_max_dimension: + height //= 2 + width //= 2 + return max(2, height - height % 2), max(2, width - width % 2) + + +def _scale_crop(crop: Crop, *, scale_y: float, scale_x: float) -> Crop: + original_height = max(1, round(crop.original_height * scale_y)) + original_width = max(1, round(crop.original_width * scale_x)) + top = min(max(0, round(crop.top * scale_y)), original_height - 1) + left = min(max(0, round(crop.left * scale_x)), original_width - 1) + height = min(max(1, round(crop.height * scale_y)), original_height - top) + width = min(max(1, round(crop.width * scale_x)), original_width - left) + return Crop( + top, + left, + height, + width, + original_height, + original_width, + crop.output_height, + crop.output_width, + ) + + +def _cropped_camera(camera: Camera, crop: Crop) -> Camera: + intrinsic = transform_intrinsics(np.asarray(camera.K), crop) + return replace( + camera, + K=tuple(tuple(float(value) for value in row) for row in intrinsic), + image_width=crop.output_width, + image_height=crop.output_height, + ) + + +def _resized_camera(camera: Camera, image_height: int, image_width: int) -> Camera: + intrinsic = np.asarray(camera.K, dtype=np.float64).copy() + intrinsic[0] *= image_width / camera.image_width + intrinsic[1] *= image_height / camera.image_height + intrinsic[2, 2] = 1.0 + return replace( + camera, + K=tuple(tuple(float(value) for value in row) for row in intrinsic), + image_width=image_width, + image_height=image_height, + ) + + +def render_skeleton_video( + *, + keypoints_world_fkc: np.ndarray, + camera: Camera, + output_path: str | Path, + num_frames: int, + canvas_height: int, + canvas_width: int, + output_height: int, + output_width: int, + body_height: float, + focal_pixels: float, + fps: Fraction, + crf: int, + preset: str, +) -> Path: + """Project and render one prepared CPU-only skeleton video. + + Args: + keypoints_world_fkc: Float32 world keypoints shaped ``[frames, keypoints, 3]``. + camera: Frozen camera used for every frame. + output_path: MP4 destination. + num_frames: Required frame count. + canvas_height: Projection-space image height. + canvas_width: Projection-space image width. + output_height: Encoded skeleton height. + output_width: Encoded skeleton width. + body_height: Robust 3D body height in metres. + focal_pixels: Geometric-mean focal length in projection pixels. + fps: Exact encoded frame rate. + crf: H.264 constant-rate-factor value. + preset: libx264 speed preset. + + Returns: + The rendered MP4 path. + """ + + if keypoints_world_fkc.shape[0] != num_frames: + raise FourDAnyoneError( + f"Prepared target render has {keypoints_world_fkc.shape[0]} frames, expected {num_frames}." + ) + projected_frames_fkc: list[np.ndarray] = [] + depth_frames_fk: list[np.ndarray] = [] + frame_points_kc: np.ndarray + for frame_points_kc in keypoints_world_fkc: + projected_kc: np.ndarray + depth_k: np.ndarray + projected_kc, depth_k, _ = project_points(frame_points_kc, camera) + projected_frames_fkc.append(projected_kc) + depth_frames_fk.append(depth_k) + keypoints_2d_fkc: np.ndarray = np.stack(projected_frames_fkc) + keypoint_depths_fk: np.ndarray = np.stack(depth_frames_fk) + body_scales_f: np.ndarray = projected_body_scales( + keypoint_depths_fk, + KEYPOINT_NAMES, + body_height, + focal_pixels, + ) + + def rendered_frames() -> Iterable[np.ndarray]: + scores_k: np.ndarray = np.ones(len(KEYPOINT_NAMES), dtype=np.float32) + frame_index: int + for frame_index in range(num_frames): + yield render_goliath40( + keypoints_2d_fkc[frame_index], + keypoint_depths_fk[frame_index], + scores_k, + canvas_height=canvas_height, + canvas_width=canvas_width, + output_height=output_height, + output_width=output_width, + body_scale_px=float(body_scales_f[frame_index]), + ) + + return write_video( + iter(rendered_frames()), + output_path, + fps, + crf=crf, + preset=preset, + ) + + +def build_skeleton_conditioning( + *, + motion: MotionResult, + clip: CanonicalClip, + regressor_path: str | Path, + foreground_model_path: str | Path, + gvhmr_root: str | Path, + output_dir: str | Path, + device: str, + view_plan: ViewPlan, + defer_target_skeletons: bool = False, +) -> Conditioning: + """Build source, RCP, and target conditioning on one camera grid.""" + + gvhmr_root, _ = validate_gvhmr(gvhmr_root) + root = Path(output_dir).expanduser().resolve() + root.mkdir(parents=True, exist_ok=True) + + LOGGER.info("Estimating source foreground masks with BiRefNet") + masks = predict_foreground_masks(clip.rgb_frames, foreground_model_path, device) + geometry = _body_geometry(motion, Path(regressor_path), gvhmr_root, device) + input_framing = analyze_input_framing( + geometry.keypoints_incam, + KEYPOINT_NAMES, + motion.K_fullimg.detach().cpu().numpy(), + motion.observed_keypoints_2d.detach().cpu().numpy(), + masks, + ) + + targets = geometry.vertices_world.mean(axis=1) + targets[:, 1] = 0.0 + center = targets.mean(axis=0) + front_direction = _front_direction(geometry.joints_world) + reference_K = reference_intrinsics(clip.height, clip.width) + framing_pitches = tuple(dict.fromkeys((*view_plan.layer_pitches, int(CAMERA.pitch_degrees)))) + + def camera_factory(radius: float, target_height: float) -> tuple[Camera, ...]: + # A full safety ring at every requested pitch keeps framing independent + # of view density while covering all target and canonical RCP cameras. + candidate_center = center.copy() + candidate_center[1] = target_height + return tuple( + camera + for layer_index, pitch in enumerate(framing_pitches) + for camera in camera_ring( + center=candidate_center, + front_direction=front_direction, + K=reference_K, + image_height=clip.height, + image_width=clip.width, + radius=radius, + target_height=target_height, + spec=CameraConfig(count=CAMERA.count, pitch_degrees=float(pitch)), + layer_index=layer_index, + camera_id_offset=layer_index * CAMERA.count, + ) + ) + + projection_height, projection_width = _projection_shape(clip.height, clip.width) + framing = solve_sequence_framing( + geometry.keypoints_world, + KEYPOINT_NAMES, + input_framing, + camera_factory, + projection_width / projection_height, + ) + LOGGER.info( + "Adaptive framing: input=%s confidence=%.3f radius=%.3f target_height=%.3f f/H=%.3f", + input_framing.label, + input_framing.confidence, + framing.radius, + framing.target_height, + framing.focal_normalized, + ) + center[1] = framing.target_height + raw_intrinsic = reference_intrinsics( + clip.height, + clip.width, + focal_normalized=framing.focal_normalized, + ) + raw_target_cameras = camera_grid( + center=center, + front_direction=front_direction, + K=raw_intrinsic, + image_height=clip.height, + image_width=clip.width, + radius=framing.radius, + target_height=framing.target_height, + views_per_layer=view_plan.views_per_layer, + layer_pitches=view_plan.layer_pitches, + start_yaw=view_plan.start_yaw, + yaw_span=view_plan.yaw_span, + ) + + source_crop = crop_from_bounds( + bounds=mask_bounds(masks, CROP.mask_threshold), + image_height=clip.height, + image_width=clip.width, + output_height=INFERENCE.height, + output_width=INFERENCE.width, + margins=CROP.margins, + allow_upscale=CROP.allow_upscale, + ) + target_crop = center_crop(clip.height, clip.width, INFERENCE.height, INFERENCE.width) + source_video = write_lossless_video(clip, root / "source.mkv") + skeleton_root = root / "goliath40" + skeleton_root.mkdir() + skeleton_height, skeleton_width = _output_skeleton_shape(clip.height, clip.width) + body_height = estimate_body_height(geometry.keypoints_world, KEYPOINT_NAMES) + keypoints_path: Path = root / "keypoints_3d.npy" + np.save(keypoints_path, geometry.keypoints_world) + focal_pixels: float = float(np.sqrt(raw_intrinsic[0, 0] * raw_intrinsic[1, 1])) + + def render_skeleton(camera: Camera, path: Path) -> Path: + return render_skeleton_video( + keypoints_world_fkc=geometry.keypoints_world, + camera=camera, + output_path=path, + num_frames=motion.num_frames, + canvas_height=clip.height, + canvas_width=clip.width, + output_height=skeleton_height, + output_width=skeleton_width, + body_height=body_height, + focal_pixels=focal_pixels, + fps=clip.fps, + crf=INFERENCE.skeleton_h264_crf, + preset=INFERENCE.h264_preset, + ) + + target_paths: tuple[Path, ...] = tuple( + skeleton_root / f"{camera.camera_id:02d}.mp4" for camera in raw_target_cameras + ) + cropped_target_cameras: tuple[Camera, ...] = tuple( + _cropped_camera(camera, target_crop) for camera in raw_target_cameras + ) + + if not view_plan.enable_rcp: + raw_rcp_cameras: tuple[Camera, ...] = () + cropped_rcp_cameras: tuple[Camera, ...] = () + elif view_plan.is_canonical_target_ring: + raw_rcp_cameras = tuple(raw_target_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids) + cropped_rcp_cameras = tuple(cropped_target_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids) + else: + canonical_cameras: tuple[Camera, ...] = camera_ring( + center=center, + front_direction=front_direction, + K=raw_intrinsic, + image_height=clip.height, + image_width=clip.width, + radius=framing.radius, + target_height=framing.target_height, + layer_index=-1, + ) + raw_rcp_cameras = tuple(canonical_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids) + cropped_rcp_cameras = tuple(_cropped_camera(camera, target_crop) for camera in raw_rcp_cameras) + + if defer_target_skeletons: + if raw_rcp_cameras: + rcp_root: Path = root / "rcp_goliath40" + rcp_root.mkdir() + rcp_paths: tuple[Path, ...] = tuple( + render_skeleton(camera, rcp_root / f"{camera.camera_id:02d}.mp4") + for camera in raw_rcp_cameras + ) + else: + rcp_paths = () + else: + target_paths = tuple( + render_skeleton(camera, path) + for camera, path in zip(raw_target_cameras, target_paths, strict=True) + ) + if not raw_rcp_cameras: + rcp_paths = () + elif view_plan.is_canonical_target_ring: + rcp_paths = tuple(target_paths[camera_id] for camera_id in view_plan.rcp_camera_ids) + else: + rcp_root = root / "rcp_goliath40" + rcp_root.mkdir() + rcp_paths = tuple( + render_skeleton(camera, rcp_root / f"{camera.camera_id:02d}.mp4") + for camera in raw_rcp_cameras + ) + + def camera_records( + raw_cameras: tuple[Camera, ...], + cropped_cameras: tuple[Camera, ...], + paths: tuple[Path, ...], + ) -> list[dict]: + return [ + { + **camera.to_dict(), + "crop": asdict(target_crop), + "raw_camera": raw.to_dict(), + "skeleton_camera": _resized_camera(raw, skeleton_height, skeleton_width).to_dict(), + "skeleton_video": path.relative_to(root).as_posix(), + } + for raw, camera, path in zip(raw_cameras, cropped_cameras, paths, strict=True) + ] + + framing_payload = framing.to_dict() + camera_payload = { + "camera_model": "OPENCV", + "world_frame": WORLD_FRAME, + "camera_frame": CAMERA_FRAME, + "front_camera_ids": list(view_plan.front_camera_ids), + "motion_world": motion.motion_world, + "motion_world_to_canonical_world": geometry.motion_world_to_canonical_world.tolist(), + "ring_center": center.tolist(), + "framing": framing_payload, + "cameras": camera_records(raw_target_cameras, cropped_target_cameras, target_paths), + "rcp_cameras": camera_records(raw_rcp_cameras, cropped_rcp_cameras, rcp_paths), + } + write_json(root / "cameras.json", camera_payload) + write_json( + root / "metadata.json", + { + "view_plan": view_plan.to_dict(), + "num_frames": motion.num_frames, + "fps_num": clip.fps_num, + "fps_den": clip.fps_den, + "source_video": source_video.name, + "source_crop": asdict(source_crop), + "source_crop_policy": { + "subject": "fmask", + "threshold": CROP.mask_threshold, + "margins": CROP.margins, + "allow_upscale": CROP.allow_upscale, + }, + "foreground_model": { + "repo_id": BIREFNET_REPO_ID, + "revision": BIREFNET_REVISION, + "image_size": FOREGROUND.image_size, + "batch_size": FOREGROUND.batch_size, + }, + "framing": framing_payload, + "regressor_metadata": geometry.regressor_metadata, + "keypoint_names": KEYPOINT_NAMES, + "visible_keypoint_set": "goliath40", + "skeleton_codec_boundary": f"libx264_crf{INFERENCE.skeleton_h264_crf}", + "skeleton_canvas": {"height": skeleton_height, "width": skeleton_width}, + "skeleton_draw_scale": { + "mode": "kp3d", + "body_height_3d": body_height, + "body_reference_px": SKELETON.draw_body_reference_px, + }, + "target_render_deferred": defer_target_skeletons, + }, + ) + if defer_target_skeletons: + write_json( + root / "target-render-request.json", + { + "schema_version": 1, + "keypoints_3d_path": str(keypoints_path), + "targets": [ + {"camera": camera.to_dict(), "output_path": str(path)} + for camera, path in zip(raw_target_cameras, target_paths, strict=True) + ], + "num_frames": motion.num_frames, + "canvas_height": clip.height, + "canvas_width": clip.width, + "output_height": skeleton_height, + "output_width": skeleton_width, + "body_height": body_height, + "focal_pixels": focal_pixels, + "fps_num": clip.fps_num, + "fps_den": clip.fps_den, + "crf": INFERENCE.skeleton_h264_crf, + "preset": INFERENCE.h264_preset, + "done_path": str(root / "target-render.done.json"), + }, + ) + return Conditioning( + root=root, + source_video=source_video, + source_crop=source_crop, + target_skeletons=tuple(SkeletonVideo(path, target_crop) for path in target_paths), + rcp_skeletons=tuple(SkeletonVideo(path, target_crop) for path in rcp_paths), + view_plan=view_plan, + fps_num=clip.fps_num, + fps_den=clip.fps_den, + num_frames=motion.num_frames, + ) diff --git a/fdanyone/skeleton/render_worker.py b/fdanyone/skeleton/render_worker.py new file mode 100644 index 0000000000000000000000000000000000000000..422f27b04a7e8392ce7dfd5583d17428785a49d1 --- /dev/null +++ b/fdanyone/skeleton/render_worker.py @@ -0,0 +1,100 @@ +"""CPU-only target skeleton renderer for conditioning overlap.""" + +from __future__ import annotations + +import json +import os +import sys +import time +from fractions import Fraction +from pathlib import Path + +import numpy as np + +from fdanyone.errors import FourDAnyoneError +from fdanyone.geometry.cameras import Camera +from fdanyone.io import write_json +from fdanyone.skeleton.pipeline import render_skeleton_video + + +def _camera_from_payload(payload: object) -> Camera: + """Decode a camera record from the conditioning file protocol.""" + + if not isinstance(payload, dict): + raise FourDAnyoneError("Target render camera must be a JSON object.") + return Camera.from_dict(payload) + + +def render_target_skeletons(request_path: str | Path) -> tuple[Path, ...]: + """Render every target in one prepared, CUDA-hidden file request.""" + + started: float = time.monotonic() + request_file: Path = Path(request_path).expanduser().resolve() + request: object = json.loads(request_file.read_text()) + if not isinstance(request, dict) or int(request.get("schema_version", 0)) != 1: + raise FourDAnyoneError("Unsupported target skeleton render request.") + cuda_visible_devices: str | None = os.environ.get("CUDA_VISIBLE_DEVICES") + if cuda_visible_devices not in ("", "-1"): + raise FourDAnyoneError( + "Target skeleton rendering must run with CUDA_VISIBLE_DEVICES disabled." + ) + + keypoints_path: Path = Path(str(request["keypoints_3d_path"])).expanduser().resolve() + keypoints_world_fkc: np.ndarray = np.load(keypoints_path, allow_pickle=False) + if keypoints_world_fkc.dtype != np.float32 or keypoints_world_fkc.ndim != 3: + raise FourDAnyoneError( + "Prepared target keypoints must be float32 [frames,keypoints,3], got " + f"dtype={keypoints_world_fkc.dtype}, shape={keypoints_world_fkc.shape}." + ) + targets: object = request.get("targets") + if not isinstance(targets, list) or not targets: + raise FourDAnyoneError("Target skeleton render request has no targets.") + + rendered_paths: list[Path] = [] + target: object + for target in targets: + if not isinstance(target, dict): + raise FourDAnyoneError("Target skeleton entry must be a JSON object.") + camera: Camera = _camera_from_payload(target["camera"]) + output_path: Path = Path(str(target["output_path"])).expanduser().resolve() + rendered_path: Path = render_skeleton_video( + keypoints_world_fkc=keypoints_world_fkc, + camera=camera, + output_path=output_path, + num_frames=int(request["num_frames"]), + canvas_height=int(request["canvas_height"]), + canvas_width=int(request["canvas_width"]), + output_height=int(request["output_height"]), + output_width=int(request["output_width"]), + body_height=float(request["body_height"]), + focal_pixels=float(request["focal_pixels"]), + fps=Fraction(int(request["fps_num"]), int(request["fps_den"])), + crf=int(request["crf"]), + preset=str(request["preset"]), + ) + rendered_paths.append(rendered_path) + + done_path: Path = Path(str(request["done_path"])).expanduser().resolve() + write_json( + done_path, + { + "schema_version": 1, + "rendered_views": len(rendered_paths), + "elapsed_seconds": time.monotonic() - started, + "cuda_visible_devices": cuda_visible_devices, + "outputs": [str(path) for path in rendered_paths], + }, + ) + return tuple(rendered_paths) + + +def main(request_path: str) -> None: + """Run one target skeleton request.""" + + render_target_skeletons(request_path) + + +if __name__ == "__main__": + if len(sys.argv) != 2: + raise SystemExit("Usage: python -m fdanyone.skeleton.render_worker REQUEST.json") + main(sys.argv[1]) diff --git a/fdanyone/skeleton/renderer.py b/fdanyone/skeleton/renderer.py new file mode 100644 index 0000000000000000000000000000000000000000..09e3d7f4ea15493c494382927026c357df60fb89 --- /dev/null +++ b/fdanyone/skeleton/renderer.py @@ -0,0 +1,230 @@ +"""Dependency-light, depth-aware Goliath40 rasterization.""" + +from __future__ import annotations + +import math + +import cv2 +import numpy as np + +from fdanyone.config import INFERENCE, SKELETON +from fdanyone.skeleton.keypoints import EXTRA_KEYPOINT_IDS, LINKS, VISIBLE_KEYPOINT_IDS, keypoint_color + +_BODY_SCALE_SEGMENTS = ( + (("left-shoulder",), ("right-shoulder",), 0.26), + (("left-hip",), ("right-hip",), 0.18), + (("left-shoulder", "right-shoulder"), ("left-hip", "right-hip"), 0.32), + (("left-shoulder",), ("left-hip",), 0.32), + (("right-shoulder",), ("right-hip",), 0.32), + (("left-shoulder",), ("left-elbow",), 0.19), + (("right-shoulder",), ("right-elbow",), 0.19), + (("left-elbow",), ("left-wrist",), 0.16), + (("right-elbow",), ("right-wrist",), 0.16), + (("left-hip",), ("left-knee",), 0.245), + (("right-hip",), ("right-knee",), 0.245), + (("left-knee",), ("left-ankle",), 0.245), + (("right-knee",), ("right-ankle",), 0.245), +) +_BODY_CENTER_NAMES = ("left-shoulder", "right-shoulder", "left-hip", "right-hip") + + +def _name_index(names) -> dict[str, int]: + normalized = [str(name).strip().lower().replace("_", "-") for name in names] + mapping = dict(zip(normalized, range(len(normalized)), strict=True)) + if len(mapping) != len(normalized): + raise ValueError("Keypoint names must be unique after normalization.") + return mapping + + +def estimate_body_height(keypoints_3d: np.ndarray, names) -> float: + """Robustly infer physical body height from stable anatomical segments.""" + + points = np.asarray(keypoints_3d, dtype=np.float64) + if points.ndim != 3 or points.shape[1:] != (len(names), 3): + raise ValueError(f"Expected keypoints [frames,{len(names)},3], got {points.shape}.") + mapping = _name_index(names) + required = {name for first, second, _ in _BODY_SCALE_SEGMENTS for group in (first, second) for name in group} + missing = sorted(required - mapping.keys()) + if missing: + raise ValueError(f"Missing body-scale keypoints: {missing}.") + + frame_heights = [] + for frame in points: + candidates = [] + for first, second, height_ratio in _BODY_SCALE_SEGMENTS: + point_a = frame[[mapping[name] for name in first]].mean(axis=0) + point_b = frame[[mapping[name] for name in second]].mean(axis=0) + length = float(np.linalg.norm(point_a - point_b)) + if np.isfinite(length) and length > 0: + candidates.append(length / height_ratio) + if candidates: + frame_heights.append(float(np.median(candidates))) + heights = np.asarray(frame_heights, dtype=np.float64) + heights = heights[np.isfinite(heights) & (heights > 0)] + if not heights.size: + raise ValueError("Cannot estimate body height from the supplied keypoints.") + if heights.size >= 5: + lower, upper = np.percentile(heights, [10.0, 90.0]) + trimmed = heights[(heights >= lower) & (heights <= upper)] + if trimmed.size: + heights = trimmed + return float(np.median(heights)) + + +def projected_body_scales( + keypoint_depths: np.ndarray, + names, + body_height_3d: float, + focal_px: float, +) -> np.ndarray: + """Convert physical body height and camera depth to per-frame pixel scale.""" + + depths = np.asarray(keypoint_depths, dtype=np.float64) + if depths.ndim != 2 or depths.shape[1] != len(names): + raise ValueError(f"Expected keypoint depths [frames,{len(names)}], got {depths.shape}.") + if not np.isfinite(body_height_3d) or body_height_3d <= 0: + raise ValueError("body_height_3d must be positive and finite.") + if not np.isfinite(focal_px) or focal_px <= 0: + raise ValueError("focal_px must be positive and finite.") + mapping = _name_index(names) + missing = [name for name in _BODY_CENTER_NAMES if name not in mapping] + if missing: + raise ValueError(f"Missing body-center keypoints: {missing}.") + center_depths = depths[:, [mapping[name] for name in _BODY_CENTER_NAMES]] + valid = np.isfinite(center_depths) & (center_depths > 1e-6) + counts = valid.sum(axis=1) + mean_depths = np.full(depths.shape[0], np.nan, dtype=np.float64) + enough = counts >= 2 + mean_depths[enough] = np.where(valid, center_depths, 0.0).sum(axis=1)[enough] / counts[enough] + scales = np.full(depths.shape[0], np.nan, dtype=np.float32) + scales[enough] = (body_height_3d * focal_px / mean_depths[enough]).astype(np.float32) + return scales + + +def _draw_point(z_buffer, canvas, point, depth, radius, color) -> None: + if not np.isfinite(depth) or depth <= 0: + return + height, width = z_buffer.shape + x, y = point + radius = max(1, int(radius)) + x0, x1 = max(0, x - radius), min(width, x + radius + 1) + y0, y1 = max(0, y - radius), min(height, y + radius + 1) + if x0 >= x1 or y0 >= y1: + return + yy, xx = np.ogrid[y0:y1, x0:x1] + mask = (xx - x) ** 2 + (yy - y) ** 2 <= radius**2 + view = z_buffer[y0:y1, x0:x1] + update = mask & (depth <= view) + view[update] = depth + canvas[y0:y1, x0:x1][update] = color + + +def _draw_line(z_buffer, canvas, p1, p2, d1, d2, thickness, color) -> None: + if not np.isfinite(d1) or not np.isfinite(d2) or max(d1, d2) <= 0: + return + height, width = z_buffer.shape + x1, y1 = p1 + x2, y2 = p2 + dx, dy = float(x2 - x1), float(y2 - y1) + length_sq = dx * dx + dy * dy + if length_sq < 1e-6: + _draw_point(z_buffer, canvas, p1, (d1 + d2) / 2.0, max(1, thickness // 2), color) + return + radius = max(0.5, float(thickness) / 2.0) + pad = int(math.ceil(radius)) + 1 + x0, x3 = max(0, min(x1, x2) - pad), min(width, max(x1, x2) + pad + 1) + y0, y3 = max(0, min(y1, y2) - pad), min(height, max(y1, y2) + pad + 1) + if x0 >= x3 or y0 >= y3: + return + yy, xx = np.ogrid[y0:y3, x0:x3] + t = np.clip(((xx - x1) * dx + (yy - y1) * dy) / length_sq, 0.0, 1.0) + mask = (xx - (x1 + t * dx)) ** 2 + (yy - (y1 + t * dy)) ** 2 <= radius**2 + depth = d1 + t * (d2 - d1) + view = z_buffer[y0:y3, x0:x3] + update = mask & (depth > 0) & (depth <= view) + view[update] = depth[update] + canvas[y0:y3, x0:x3][update] = color + + +def render_goliath40( + keypoints: np.ndarray, + depths: np.ndarray, + scores: np.ndarray, + *, + canvas_height: int, + canvas_width: int, + output_height: int, + output_width: int, + score_threshold: float = 0.3, + body_scale_px: float | None = None, +) -> np.ndarray: + """Render one RGB frame using the reference Sapiens2 sizing rules.""" + + canvas_scale = max(1.0, INFERENCE.skeleton_max_dimension / max(output_height, output_width)) + render_height = int(round(output_height * canvas_scale)) + render_width = int(round(output_width * canvas_scale)) + points = np.asarray(keypoints, dtype=np.float32).copy() + points[:, 0] *= render_width / canvas_width + points[:, 1] *= render_height / canvas_height + depths = np.asarray(depths, dtype=np.float32) + scores = np.asarray(scores, dtype=np.float32) + # Keep the reference renderer's BGR draw -> resize -> RGB conversion order. + # Resizing a channel-permuted uint8 image can differ by one LSB in OpenCV's + # optimized interpolation kernels, so drawing directly in RGB is not quite + # byte-exact even though the colors are semantically identical. + canvas = np.zeros((render_height, render_width, 3), dtype=np.uint8) + z_buffer = np.full((render_height, render_width), np.inf, dtype=np.float32) + line_scale = render_height / 1024.0 + if body_scale_px is not None and np.isfinite(body_scale_px) and body_scale_px > 0: + keypoint_scale = render_height / canvas_height + line_scale = max( + 0.25, + float(body_scale_px) * keypoint_scale / SKELETON.draw_body_reference_px, + ) + base_radius = max(1, int(round(2 * line_scale))) + base_thickness = max(1, int(round(2 * line_scale))) + point_items: dict[int, tuple[tuple[int, int], float, int, tuple[int, int, int]]] = {} + + def remember(index: int, radius: int) -> None: + if scores[index] < score_threshold or not np.isfinite(points[index]).all(): + return + item = ( + (int(round(points[index, 0])), int(round(points[index, 1]))), + float(depths[index]), + radius, + keypoint_color(index)[::-1], + ) + if index not in point_items or radius > point_items[index][2]: + point_items[index] = item + + for _, first, second, color, major in LINKS: + if min(float(scores[first]), float(scores[second])) < score_threshold: + continue + if not np.isfinite(points[[first, second]]).all(): + continue + p1 = (int(round(points[first, 0])), int(round(points[first, 1]))) + p2 = (int(round(points[second, 0])), int(round(points[second, 1]))) + thickness = base_thickness * (2 if major else 1) + _draw_line( + z_buffer, + canvas, + p1, + p2, + float(depths[first]), + float(depths[second]), + thickness, + color[::-1], + ) + radius = max(1, int(round(base_radius * (1.75 if major else 1.0)))) + remember(first, radius) + remember(second, radius) + + for index in EXTRA_KEYPOINT_IDS: + remember(index, max(2, int(round(base_radius * 1.5)))) + for index, (point, depth, radius, color) in point_items.items(): + if index in VISIBLE_KEYPOINT_IDS: + _draw_point(z_buffer, canvas, point, depth, radius, color) + if (render_height, render_width) != (output_height, output_width): + canvas = cv2.resize(canvas, (output_width, output_height), interpolation=cv2.INTER_AREA) + canvas = cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB) + return np.ascontiguousarray(canvas) diff --git a/fdanyone/skeleton/worker.py b/fdanyone/skeleton/worker.py new file mode 100644 index 0000000000000000000000000000000000000000..ebbfc86491b572d37cdcd0ed3e630e2c2a915334 --- /dev/null +++ b/fdanyone/skeleton/worker.py @@ -0,0 +1,66 @@ +"""Private subprocess entry point for licensed body-model conditioning.""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +from fdanyone.device import select_cuda_device +from fdanyone.motion.result import MotionResult +from fdanyone.skeleton.pipeline import build_skeleton_conditioning +from fdanyone.video import load_canonical_working_clip +from fdanyone.views import ViewPlan + + +def run_skeleton_worker( + *, + working_video: str | Path, + clip_metadata: str | Path, + motion_result_dir: str | Path, + regressor_path: str | Path, + foreground_model_path: str | Path, + gvhmr_root: str | Path, + output_dir: str | Path, + device: str, + view_plan: ViewPlan, + defer_target_skeletons: bool = False, +) -> None: + """Build conditioning for one request and publish it under ``output_dir``.""" + + device, _ = select_cuda_device(device) + clip = load_canonical_working_clip(working_video, clip_metadata) + motion = MotionResult.load(motion_result_dir) + build_skeleton_conditioning( + motion=motion, + clip=clip, + regressor_path=regressor_path, + foreground_model_path=foreground_model_path, + gvhmr_root=gvhmr_root, + output_dir=output_dir, + device=device, + view_plan=view_plan, + defer_target_skeletons=defer_target_skeletons, + ) + + +def main(request_path: str) -> None: + request = json.loads(Path(request_path).read_text()) + run_skeleton_worker( + working_video=request["working_video"], + clip_metadata=request["clip_metadata"], + motion_result_dir=request["motion_result_dir"], + regressor_path=request["regressor_path"], + foreground_model_path=request["foreground_model_path"], + gvhmr_root=request["gvhmr_root"], + output_dir=request["output_dir"], + device=request["device"], + view_plan=ViewPlan.from_dict(request["view_plan"]), + defer_target_skeletons=bool(request.get("defer_target_skeletons", False)), + ) + + +if __name__ == "__main__": + if len(sys.argv) != 2: + raise SystemExit("Usage: python -m fdanyone.skeleton.worker REQUEST.json") + main(sys.argv[1]) diff --git a/fdanyone/vendor/__init__.py b/fdanyone/vendor/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..62562d93e5c09e8bcd0ceddb74593037542fbbbb --- /dev/null +++ b/fdanyone/vendor/__init__.py @@ -0,0 +1 @@ +"""Third-party inference code redistributed with its upstream notices.""" diff --git a/fdanyone/vendor/diffsynth/LICENSE b/fdanyone/vendor/diffsynth/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..0e5a49f5f70af9e9d37278f72315d1b1afd34895 --- /dev/null +++ b/fdanyone/vendor/diffsynth/LICENSE @@ -0,0 +1,201 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [2023] [Zhongjie Duan] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/fdanyone/vendor/diffsynth/UPSTREAM.md b/fdanyone/vendor/diffsynth/UPSTREAM.md new file mode 100644 index 0000000000000000000000000000000000000000..1312373566fccba6f541dc1779ce33034df8307b --- /dev/null +++ b/fdanyone/vendor/diffsynth/UPSTREAM.md @@ -0,0 +1,23 @@ +# DiffSynth-Studio provenance + +This directory is a deliberately small inference-only extract from [modelscope/DiffSynth-Studio](https://github.com/modelscope/DiffSynth-Studio), licensed under Apache-2.0. + +- Public base revision: `04e39f7de53df7276a7b40ca1791c2a393e05ff3` +- Research fork revision used by the original experiment: `c00782d90c872c97bda4745a9e6a41a0a4a7c4db` +- `UPSTREAM.patch` SHA-256: `178e6035e451f94a2122fa4d2c876a488546964768b90275549cbf609f97daba` +- Extracted: 2026-07-16 + +The research fork revision was not anonymously reachable when the release contract was audited. `UPSTREAM.patch` therefore records the exact binary-safe diff from the public base to the research revision for the retained Wan/SpaTem/scheduler source files. Unrelated research-fork changes are deliberately excluded. `VENDORED_FILES.txt` is the reviewable extraction manifest. + +4DAnyone subsequently removed registry, downloader, training, image-encoder, camera-controller, unused scheduler, VAE-tiling, and unrelated pipeline surfaces. It also removed the unshipped Group-B Wan branches and retained only the released Group-A multiview, view-pack, and pose path. First-party code constructs the retained models directly and records the selected attention implementation in run metadata. The retained inference math is protected by frozen-tensor and prepared-path parity tests. + +To reproduce the retained research sources before pruning: + +```bash +git clone https://github.com/modelscope/DiffSynth-Studio.git +git -C DiffSynth-Studio checkout 04e39f7de53df7276a7b40ca1791c2a393e05ff3 +git -C DiffSynth-Studio apply --check /path/to/UPSTREAM.patch +git -C DiffSynth-Studio apply /path/to/UPSTREAM.patch +``` + +The patch contains research-code additions and is itself source code. It must remain covered by this directory's Apache-2.0 `LICENSE` and attribution. diff --git a/fdanyone/vendor/diffsynth/UPSTREAM.patch b/fdanyone/vendor/diffsynth/UPSTREAM.patch new file mode 100644 index 0000000000000000000000000000000000000000..c1906b36d1898a08ed8cd4f408238d358644d804 --- /dev/null +++ b/fdanyone/vendor/diffsynth/UPSTREAM.patch @@ -0,0 +1,1334 @@ +diff --git a/diffsynth/models/wan_video_dit.py b/diffsynth/models/wan_video_dit.py +index 1a54728..b663722 100644 +--- a/diffsynth/models/wan_video_dit.py ++++ b/diffsynth/models/wan_video_dit.py +@@ -3,7 +3,7 @@ import torch.nn as nn + import torch.nn.functional as F + import math + from typing import Tuple, Optional +-from einops import rearrange ++from einops import rearrange, repeat + from .utils import hash_state_dict_keys + from .wan_video_camera_controller import SimpleAdapter + try: +@@ -23,7 +23,12 @@ try: + SAGE_ATTN_AVAILABLE = True + except ModuleNotFoundError: + SAGE_ATTN_AVAILABLE = False +- ++ ++try: ++ from spas_sage_attn import spas_sage2_attn_meansim_topk_cuda ++ SPARGE_ATTN_AVAILABLE = True ++except ModuleNotFoundError: ++ SPARGE_ATTN_AVAILABLE = False + + def flash_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, num_heads: int, compatibility_mode=False): + if compatibility_mode: +@@ -46,6 +51,13 @@ def flash_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, num_heads + v = rearrange(v, "b s (n d) -> b s n d", n=num_heads) + x = flash_attn.flash_attn_func(q, k, v) + x = rearrange(x, "b s n d -> b s (n d)", n=num_heads) ++ # TODO: test spas_sage2_attn_meansim_topk_cuda ++ # elif SPARGE_ATTN_AVAILABLE: ++ # q = rearrange(q, "b s (n d) -> b n s d", n=num_heads) ++ # k = rearrange(k, "b s (n d) -> b n s d", n=num_heads) ++ # v = rearrange(v, "b s (n d) -> b n s d", n=num_heads) ++ # x = spas_sage2_attn_meansim_topk_cuda(q, k, v, simthreshd1=-0.1, topk=0.5, pvthreshd=15, is_causal=False) ++ # x = rearrange(x, "b n s d -> b s (n d)", n=num_heads) + elif SAGE_ATTN_AVAILABLE: + q = rearrange(q, "b s (n d) -> b n s d", n=num_heads) + k = rearrange(k, "b s (n d) -> b n s d", n=num_heads) +@@ -186,6 +198,32 @@ class CrossAttention(nn.Module): + return self.o(x) + + ++class CrossAttentionSrcCam(nn.Module): ++ def __init__(self, dim: int, num_heads: int, eps: float = 1e-6): ++ super().__init__() ++ self.dim = dim ++ self.num_heads = num_heads ++ self.head_dim = dim // num_heads ++ ++ self.q = nn.Linear(dim, dim) ++ self.k = nn.Linear(dim, dim) ++ self.v = nn.Linear(dim, dim) ++ self.o = nn.Linear(dim, dim) ++ self.norm_q = RMSNorm(dim, eps=eps) ++ self.norm_k = RMSNorm(dim, eps=eps) ++ ++ self.attn = AttentionModule(self.num_heads) ++ ++ def forward(self, x, x_src, freqs, freqs_src): ++ q = self.norm_q(self.q(x)) ++ k = self.norm_k(self.k(x_src)) ++ v = self.v(x_src) ++ q = rope_apply(q, freqs, self.num_heads) ++ k = rope_apply(k, freqs_src, self.num_heads) ++ x = self.attn(q, k, v) ++ return self.o(x) ++ ++ + class GateModule(nn.Module): + def __init__(self,): + super().__init__() +@@ -200,6 +238,13 @@ class DiTBlock(nn.Module): + self.num_heads = num_heads + self.ffn_dim = ffn_dim + ++ self.disable_video_attn = False ++ self.use_4d_attn = False ++ self.use_mvs_attn = False ++ self.use_src_self_attn = False ++ self.use_src_cross_attn = False ++ self.use_cam_encoder = False ++ + self.self_attn = SelfAttention(dim, num_heads, eps) + self.cross_attn = CrossAttention( + dim, num_heads, eps, has_image_input=has_image_input) +@@ -211,7 +256,10 @@ class DiTBlock(nn.Module): + self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + self.gate = GateModule() + +- def forward(self, x, context, t_mod, freqs): ++ def forward(self, x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, shape): ++ # breakpoint() ++ v, f, h, w = shape ++ + has_seq = len(t_mod.shape) == 4 + chunk_dim = 2 if has_seq else 1 + # msa: multi-head self-attention mlp: multi-layer perceptron +@@ -222,12 +270,80 @@ class DiTBlock(nn.Module): + shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2), + shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2), + ) +- input_x = modulate(self.norm1(x), shift_msa, scale_msa) +- x = self.gate(x, gate_msa, self.self_attn(input_x, freqs)) ++ ++ # video self-attention ++ if not self.disable_video_attn: ++ input_x = modulate(self.norm1(x), shift_msa, scale_msa) ++ if self.use_4d_attn: ++ input_x = rearrange(input_x, "v fhw c -> (v fhw) c").unsqueeze(0) ++ freqs = repeat(freqs, "fhw 1 c -> (v fhw) 1 c", v=v) ++ input_x = self.self_attn(input_x, freqs) ++ if self.use_4d_attn: ++ input_x = rearrange(input_x.squeeze(0), "(v fhw) c -> v fhw c", v=v) ++ x = self.gate(x, gate_msa, input_x) ++ ++ # source-view self-attention ++ if self.use_src_self_attn: ++ x_cat = torch.cat([x, x_src], dim=1) ++ shift_src, scale_src, gate_src = ( ++ self.modulation_src.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1) ++ input_x_cat = modulate(self.norm1_src(x_cat), shift_src, scale_src) ++ ++ freqs_cat = torch.cat([freqs, freqs_src], dim=0) ++ input_x_cat = self.self_attn_src(input_x_cat, freqs_cat) ++ x_cat = self.gate(x_cat, gate_src, input_x_cat) ++ len_src = x_src.shape[1] ++ x, x_src = x_cat[:, :-len_src, ...], x_cat[:, -len_src:, ...] ++ ++ # source-view cross-attention ++ if self.use_src_cross_attn: ++ shift_src, scale_src, gate_src = ( ++ self.modulation_src.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1) ++ input_x = modulate(self.norm1_src(x), shift_src, scale_src) ++ input_x_src = self.norm1_src(x_src) ++ ++ input_x = self.cross_attn_src(input_x, input_x_src, freqs, freqs_src) ++ x = self.gate(x, gate_src, input_x) ++ ++ # multiview self-attention ++ if self.use_mvs_attn: ++ shift_mvs, scale_mvs, gate_mvs = ( ++ self.modulation_mvs.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1) ++ input_x = modulate(self.norm1_mvs(x), shift_mvs, scale_mvs) ++ ++ # add camera embedding before multiview attention ++ if self.use_cam_encoder and cam_emb is not None: ++ cam_proj = self.cam_encoder(cam_emb) # (v, 1, dim) ++ cam_proj = cam_proj.unsqueeze(2).unsqueeze(3).expand(-1, f, h, w, -1) # (v, f, h, w, dim) ++ cam_proj = rearrange(cam_proj, "v f h w d -> v (f h w) d") ++ input_x = input_x + cam_proj ++ ++ input_x = rearrange(input_x, "v (f h w) c -> f (v h w) c", v=v, f=f, h=h, w=w) ++ input_x = self.self_attn_mvs(input_x, freqs_mvs) ++ input_x = rearrange(input_x, "f (v h w) c -> v (f h w) c", v=v, f=f, h=h, w=w) ++ ++ # projector wraps multiview attention output ++ if self.use_cam_encoder: ++ input_x = self.projector(input_x) ++ ++ x = self.gate(x, gate_mvs, input_x) ++ ++ # prompt cross-attention ++ context = repeat(context, "1 l c -> v l c", v=x.shape[0]) + x = x + self.cross_attn(self.norm3(x), context) ++ ++ # feed-forward network + input_x = modulate(self.norm2(x), shift_mlp, scale_mlp) + x = self.gate(x, gate_mlp, self.ffn(input_x)) +- return x ++ ++ if self.use_src_self_attn: ++ # for src self-attention, the x_src is updated as well ++ x_src = x_src + self.cross_attn(self.norm3(x_src), context) ++ ++ input_x_src = modulate(self.norm2(x_src), shift_mlp, scale_mlp) ++ x_src = self.gate(x_src, gate_mlp, self.ffn(input_x_src)) ++ ++ return x, x_src + + + class MLP(torch.nn.Module): +@@ -264,11 +380,57 @@ class Head(nn.Module): + shift, scale = (self.modulation.unsqueeze(0).to(dtype=t_mod.dtype, device=t_mod.device) + t_mod.unsqueeze(2)).chunk(2, dim=2) + x = (self.head(self.norm(x) * (1 + scale.squeeze(2)) + shift.squeeze(2))) + else: ++ if t_mod.shape[0] != 1: ++ t_mod = t_mod[:, None, :] # [b, d] -> [b, 1, d] + shift, scale = (self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(2, dim=1) + x = (self.head(self.norm(x) * (1 + scale) + shift)) + return x + + ++def pad_for_3d_conv(x, kernel_size): ++ """Pad to be divisible by kernel_size. From FramePack.""" ++ _, _, t, h, w = x.shape ++ pt, ph, pw = kernel_size ++ pad_t = (pt - (t % pt)) % pt ++ pad_h = (ph - (h % ph)) % ph ++ pad_w = (pw - (w % pw)) % pw ++ if pad_t == 0 and pad_h == 0 and pad_w == 0: ++ return x ++ return F.pad(x, (0, pad_w, 0, pad_h, 0, pad_t), mode='replicate') ++ ++ ++class ViewPackEmbedding(nn.Module): ++ """Multi-resolution spatial patch embedding for clean source views. ++ Spatial-only downsampling: temporal dim preserved. ++ Ref: HunyuanVideoPatchEmbedForCleanLatents in FramePack. ++ ++ Supported configurations (all exactly fill 4 quadrants): ++ - 4×2x views (each 2x view → 1 quadrant) ++ - 3×2x + 4×4x (3 quadrants from 2x + 1 quadrant from 4×4x tile) ++ """ ++ ++ def __init__(self, in_dim, dim, patch_size): ++ super().__init__() ++ pt, ph, pw = patch_size # (1, 2, 2) for Wan ++ # 2x: spatial 2x downsample relative to 1x -> kernel (1, 4, 4) ++ self.proj_2x = nn.Conv3d(in_dim, dim, kernel_size=(pt, ph*2, pw*2), stride=(pt, ph*2, pw*2)) ++ # 4x: spatial 4x downsample relative to 1x -> kernel (1, 8, 8) ++ self.proj_4x = nn.Conv3d(in_dim, dim, kernel_size=(pt, ph*4, pw*4), stride=(pt, ph*4, pw*4)) ++ ++ @torch.no_grad() ++ def initialize_from_patch_embedding(self, patch_embedding: nn.Conv3d): ++ """FramePack-style init: tile spatial dims and scale by 1/area_ratio.""" ++ weight = patch_embedding.weight.detach().clone() # (dim, in_dim, 1, 2, 2) ++ bias = patch_embedding.bias.detach().clone() ++ sd = { ++ 'proj_2x.weight': repeat(weight, 'b c t h w -> b c t (h 2) (w 2)') / 4.0, ++ 'proj_2x.bias': bias.clone(), ++ 'proj_4x.weight': repeat(weight, 'b c t h w -> b c t (h 4) (w 4)') / 16.0, ++ 'proj_4x.bias': bias.clone(), ++ } ++ self.load_state_dict(sd) ++ ++ + class WanModel(torch.nn.Module): + def __init__( + self, +@@ -355,32 +517,121 @@ class WanModel(torch.nn.Module): + + def forward(self, + x: torch.Tensor, ++ x_src: torch.Tensor, + timestep: torch.Tensor, + context: torch.Tensor, ++ skeletons: Optional[torch.Tensor] = None, ++ cam_emb: Optional[torch.Tensor] = None, ++ drop_viewpack_tokens: bool = False, + clip_feature: Optional[torch.Tensor] = None, + y: Optional[torch.Tensor] = None, + use_gradient_checkpointing: bool = False, + use_gradient_checkpointing_offload: bool = False, + **kwargs, + ): +- t = self.time_embedding( +- sinusoidal_embedding_1d(self.freq_dim, timestep)) +- t_mod = self.time_projection(t).unflatten(1, (6, self.dim)) ++ # breakpoint() + context = self.text_embedding(context) +- ++ + if self.has_image_input: + x = torch.cat([x, y], dim=1) # (b, c_x + c_y, f, h, w) + clip_embdding = self.img_emb(clip_feature) + context = torch.cat([clip_embdding, context], dim=1) +- ++ + x, (f, h, w) = self.patchify(x) +- ++ + freqs = torch.cat([ + self.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), + self.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), + self.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) + ], dim=-1).reshape(f * h * w, 1, -1).to(x.device) +- ++ ++ # build packed views from 1x/2x/4x source views ++ v_src = x_src.shape[0] ++ if v_src == 1: ++ # 1×1x: no 2x/4x sources ++ x_src_2x, x_src_4x = None, None ++ elif v_src == 5: ++ # 1×1x + 4×2x ++ x_src, x_src_2x, x_src_4x = x_src[:1], x_src[1:], None ++ elif v_src == 8: ++ # 1×1x + 3×2x + 4×4x ++ x_src, x_src_2x, x_src_4x = x_src[:1], x_src[1:4], x_src[4:] ++ else: ++ raise ValueError(f"Unsupported number of source views: {v_src}") ++ ++ # 1x source: always packed as one extra view ++ x_src, (f_src, h_src, w_src) = self.patchify(x_src) ++ if not self.use_src_attn: ++ # use viewpack for 1x source views ++ x = torch.cat([x, x_src], dim=0) ++ x_src = None ++ freqs_src = None ++ v_pack = 1 ++ else: ++ # use src self/cross-attention for 1x source views ++ x_src = repeat(x_src, "1 fhw_src c -> v fhw_src c", v=x.shape[0]) ++ freqs_src = torch.cat([ ++ self.freqs[0][self.freqs_src_shift:f_src + self.freqs_src_shift].view(f_src, 1, 1, -1).expand(f_src, h_src, w_src, -1), ++ self.freqs[1][:h_src].view(1, h_src, 1, -1).expand(f_src, h_src, w_src, -1), ++ self.freqs[2][:w_src].view(1, 1, w_src, -1).expand(f_src, h_src, w_src, -1) ++ ], dim=-1).reshape(f_src * h_src * w_src, 1, -1).to(x.device) ++ v_pack = 0 ++ ++ # 2x/4x source views: packed after 1x source views ++ if self.use_viewpack: ++ if x_src_2x is not None: ++ x_src_2x = pad_for_3d_conv(x_src_2x, self.viewpack_embedding.proj_2x.kernel_size) ++ x_src_2x = self.viewpack_embedding.proj_2x(x_src_2x) # (v_2x, dim, f, h//2, w//2) ++ ++ if x_src_4x is not None: ++ # tile 4x source views into 2x source views ++ x_src_4x = pad_for_3d_conv(x_src_4x, self.viewpack_embedding.proj_4x.kernel_size) ++ x_src_4x = self.viewpack_embedding.proj_4x(x_src_4x) # (4, dim, f, h//4, w//4) ++ x_src_4x = rearrange(x_src_4x, '(g1 g2) c f h w -> 1 c f (g1 h) (g2 w)', g1=2, g2=2) # (1, dim, f, h//2, w//2) ++ f_2x, h_2x, w_2x = x_src_2x.shape[2:] # crop padding surplus to match 2x source views (f, h, w) ++ x_src_4x = x_src_4x[:, :, :f_2x, :h_2x, :w_2x] ++ x_src_2x = torch.cat([x_src_2x, x_src_4x], dim=0) # (4, dim, f, h//2, w//2) ++ ++ # tile 2x source views into 1x source views ++ x_pack = rearrange(x_src_2x, '(g1 g2) c f h w -> 1 c f (g1 h) (g2 w)', g1=2, g2=2) ++ x_pack = x_pack[:, :, :f, :h, :w] # crop padding surplus to match 1x source views (f, h, w) ++ x_pack = rearrange(x_pack, '1 c f h w -> 1 (f h w) c') ++ x_pack = x_pack.to(dtype=x.dtype) ++ if drop_viewpack_tokens: ++ # Keep viewpack parameters in the autograd graph across distributed ranks. ++ zero_dependency = x_pack.float().mean().to(dtype=x.dtype) * 0.0 ++ x = x + zero_dependency ++ else: ++ x = torch.cat([x, x_pack], dim=0) ++ v_pack += 1 ++ ++ timestep = torch.cat([timestep, torch.zeros(v_pack, device=timestep.device, dtype=timestep.dtype)]) ++ if skeletons is not None: ++ skeletons = torch.cat([skeletons, -torch.ones_like(skeletons[:1]).expand(v_pack, -1, -1, -1, -1)], dim=0) ++ ++ # Expand cam_emb for viewpack views (zero vectors for packed source views) ++ if cam_emb is not None: ++ cam_emb = torch.cat([cam_emb, torch.zeros(v_pack, cam_emb.shape[-1], ++ device=cam_emb.device, dtype=cam_emb.dtype)], dim=0) ++ cam_emb = cam_emb.unsqueeze(1) # (v, 1, 12) ++ ++ # Compute time embeddings (after concat since timestep may have been extended) ++ t = self.time_embedding( ++ sinusoidal_embedding_1d(self.freq_dim, timestep).to(x.dtype)) ++ t_mod = self.time_projection(t).unflatten(1, (6, self.dim)) ++ ++ if self.use_pose_encoder: ++ skeleton_latents = self.pose_encoder(skeletons) ++ skeleton_tokens = rearrange(skeleton_latents, "v c f h w -> v (f h w) c") ++ x = x + skeleton_tokens ++ ++ v = x.shape[0] ++ freqs_mvs = torch.cat([ ++ self.freqs[0][:v].view(v, 1, 1, -1).expand(v, h, w, -1), ++ self.freqs[1][:h].view(1, h, 1, -1).expand(v, h, w, -1), ++ self.freqs[2][:w].view(1, 1, w, -1).expand(v, h, w, -1) ++ ], dim=-1).reshape(v * h * w, 1, -1).to(x.device) ++ + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) +@@ -390,19 +641,24 @@ class WanModel(torch.nn.Module): + if self.training and use_gradient_checkpointing: + if use_gradient_checkpointing_offload: + with torch.autograd.graph.save_on_cpu(): +- x = torch.utils.checkpoint.checkpoint( ++ x, x_src = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), +- x, context, t_mod, freqs, ++ x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w), + use_reentrant=False, + ) + else: +- x = torch.utils.checkpoint.checkpoint( ++ x, x_src = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), +- x, context, t_mod, freqs, ++ x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w), + use_reentrant=False, + ) + else: +- x = block(x, context, t_mod, freqs) ++ x, x_src = block(x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w)) ++ ++ # Strip added views (viewpack) ++ if v_pack > 0: ++ x = x[:-v_pack, ...] ++ t = t[:-v_pack, ...] + + x = self.head(x, t) + x = self.unpatchify(x, (f, h, w)) +@@ -411,7 +667,7 @@ class WanModel(torch.nn.Module): + @staticmethod + def state_dict_converter(): + return WanModelStateDictConverter() +- ++ + + class WanModelStateDictConverter: + def __init__(self): +diff --git a/diffsynth/models/wan_video_pose_encoder.py b/diffsynth/models/wan_video_pose_encoder.py +new file mode 100644 +index 0000000..5180625 +--- /dev/null ++++ b/diffsynth/models/wan_video_pose_encoder.py +@@ -0,0 +1,81 @@ ++import torch ++import torch.nn as nn ++import numpy as np ++from torch.nn import init ++ ++# PoseEncoder is 3D version of PoseNet in MimicMotion: ++# https://github.com/Tencent/MimicMotion/blob/c053153a1d124abae8c08568925ae88debc63001/mimicmotion/modules/pose_net.py ++ ++ ++class PoseEncoder(nn.Module): ++ def __init__(self, out_dim=5120, in_channels=3): ++ super().__init__() ++ ++ if out_dim in (5120, 1536): ++ # Wan2.1-T2V-14B / 1.3B ++ t_strides = (1, 1, 1, 2, 2) # downsampled by 4 ++ s_strides = (2, 2, 1, 2, 2) # downsampled by 16 ++ kernel_size = (3, 3, 3) ++ elif out_dim == 3072: ++ # Wan2.2-TI2V-5B ++ t_strides = (1, 1, 1, 2, 2) # downsampled by 4 ++ s_strides = (2, 2, 2, 2, 2) # downsampled by 32 ++ kernel_size = (3, 4, 4) ++ else: ++ raise ValueError(f"Invalid out_dim: {out_dim}") ++ ++ strides = [(t, s, s) for t, s in zip(t_strides, s_strides)] ++ ++ self.conv_layers = nn.Sequential( ++ nn.Conv3d(in_channels, in_channels, kernel_size=3, stride=1, padding=1), ++ nn.SiLU(), ++ nn.Conv3d(in_channels, 16, kernel_size=kernel_size, stride=strides[0], padding=(1, 1, 1)), ++ nn.SiLU(), ++ nn.Conv3d(16, 16, kernel_size=3, stride=1, padding=1), ++ nn.SiLU(), ++ nn.Conv3d(16, 32, kernel_size=kernel_size, stride=strides[1], padding=(1, 1, 1)), ++ nn.SiLU(), ++ nn.Conv3d(32, 32, kernel_size=3, stride=1, padding=1), ++ nn.SiLU(), ++ nn.Conv3d(32, 64, kernel_size=kernel_size, stride=strides[2], padding=(1, 1, 1)), ++ nn.SiLU(), ++ nn.Conv3d(64, 64, kernel_size=3, stride=1, padding=1), ++ nn.SiLU(), ++ nn.Conv3d(64, 128, kernel_size=kernel_size, stride=strides[3], padding=(1, 1, 1)), ++ nn.SiLU(), ++ nn.Conv3d(128, 128, kernel_size=3, stride=1, padding=1), ++ nn.SiLU(), ++ nn.Conv3d(128, 256, kernel_size=kernel_size, stride=strides[4], padding=(1, 1, 1)), ++ nn.SiLU(), ++ ) ++ ++ self.final_proj = nn.Conv3d(256, out_dim, kernel_size=1) ++ ++ self.scale = nn.Parameter(torch.ones(1) * 2.0) ++ ++ self._initialize_weights() ++ ++ def _initialize_weights(self): ++ for m in self.modules(): ++ if isinstance(m, nn.Conv3d): ++ # He (Kaiming) initialization in fan‑in mode ++ receptive = np.prod(m.kernel_size) * m.in_channels ++ init.normal_(m.weight, mean=0.0, std=np.sqrt(2.0 / receptive)) ++ if m.bias is not None: ++ init.zeros_(m.bias) ++ # start with zero output so model behaves like unconditional ++ init.zeros_(self.final_proj.weight) ++ if self.final_proj.bias is not None: ++ init.zeros_(self.final_proj.bias) ++ ++ def forward(self, x: torch.Tensor) -> torch.Tensor: ++ """ ++ x: (B, C, F, H, W) -> latent grid matching the DiT patch tokens. ++ Wan2.1 uses F/4, H/16, W/16; Wan2.2-TI2V-5B uses F/4, H/32, W/32. ++ """ ++ # Wan pattern: 1 -> 4 -> 4 -> ... ++ x = torch.cat([x[:, :, :1].repeat(1, 1, 3, 1, 1), x], dim=2) ++ ++ x = self.conv_layers(x) ++ x = self.final_proj(x) ++ return x * self.scale +diff --git a/diffsynth/models/wan_video_vae.py b/diffsynth/models/wan_video_vae.py +index 397a2e7..43057ff 100644 +--- a/diffsynth/models/wan_video_vae.py ++++ b/diffsynth/models/wan_video_vae.py +@@ -1121,7 +1121,7 @@ class WanVideoVAE(nn.Module): + weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device) + values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device) + +- for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"): ++ for h, h_, w, w_ in tasks: + hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device) + hidden_states_batch = self.model.decode(hidden_states_batch, self.scale).to(data_device) + +@@ -1173,7 +1173,7 @@ class WanVideoVAE(nn.Module): + weight = torch.zeros((1, 1, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device) + values = torch.zeros((1, self.z_dim, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device) + +- for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"): ++ for h, h_, w, w_ in tasks: + hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device) + hidden_states_batch = self.model.encode(hidden_states_batch, self.scale).to(data_device) + +@@ -1216,10 +1216,9 @@ class WanVideoVAE(nn.Module): + + + def encode(self, videos, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)): +- + videos = [video.to("cpu") for video in videos] + hidden_states = [] +- for video in videos: ++ for video in tqdm(videos, desc="VAE encoding", disable=not tiled): + video = video.unsqueeze(0) + if tiled: + tile_size = (tile_size[0] * self.upsampling_factor, tile_size[1] * self.upsampling_factor) +@@ -1234,11 +1233,18 @@ class WanVideoVAE(nn.Module): + + + def decode(self, hidden_states, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)): +- if tiled: +- video = self.tiled_decode(hidden_states, device, tile_size, tile_stride) +- else: +- video = self.single_decode(hidden_states, device) +- return video ++ hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states] ++ videos = [] ++ for hidden_state in tqdm(hidden_states, desc="VAE decoding", disable=not tiled): ++ hidden_state = hidden_state.unsqueeze(0) ++ if tiled: ++ video = self.tiled_decode(hidden_state, device, tile_size, tile_stride) ++ else: ++ video = self.single_decode(hidden_state, device) ++ video = video.squeeze(0) ++ videos.append(video) ++ videos = torch.stack(videos) ++ return videos + + + @staticmethod +diff --git a/diffsynth/pipelines/__init__.py b/diffsynth/pipelines/__init__.py +index e2ad551..f878ad8 100644 +--- a/diffsynth/pipelines/__init__.py ++++ b/diffsynth/pipelines/__init__.py +@@ -12,4 +12,5 @@ from .pipeline_runner import SDVideoPipelineRunner + from .hunyuan_video import HunyuanVideoPipeline + from .step_video import StepVideoPipeline + from .wan_video import WanVideoPipeline ++from .wan_video_spatem import WanVideoSpaTemPipeline + KolorsImagePipeline = SDXLImagePipeline +diff --git a/diffsynth/pipelines/wan_video_spatem.py b/diffsynth/pipelines/wan_video_spatem.py +new file mode 100644 +index 0000000..85b4f65 +--- /dev/null ++++ b/diffsynth/pipelines/wan_video_spatem.py +@@ -0,0 +1,659 @@ ++from ..models import ModelManager ++from ..models.wan_video_dit import WanModel ++from ..models.wan_video_pose_encoder import PoseEncoder ++from ..models.wan_video_text_encoder import WanTextEncoder ++from ..models.wan_video_vae import WanVideoVAE ++from ..models.wan_video_image_encoder import WanImageEncoder ++from ..schedulers.flow_match import FlowMatchScheduler ++from ..schedulers.bride_match import BridgeMatchScheduler ++from ..pipelines.base import BasePipeline ++from ..prompters import WanPrompter ++import torch, os ++import torch.nn as nn ++import numpy as np ++import torch.nn.functional as F ++from PIL import Image ++from tqdm import tqdm ++from typing import Optional, Union ++from functools import partial ++from einops import rearrange ++ ++from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear ++from ..models.wan_video_text_encoder import T5RelativeEmbedding, T5LayerNorm ++from ..models.wan_video_dit import RMSNorm, SelfAttention, CrossAttentionSrcCam, ViewPackEmbedding ++from ..models.wan_video_vae import RMS_norm, CausalConv3d, Upsample ++from ..utils import ModelConfig ++ ++ ++class WanVideoSpaTemPipeline(BasePipeline): ++ ++ def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None): ++ super().__init__(device=device, torch_dtype=torch_dtype) ++ self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True) ++ self.prompter = WanPrompter(tokenizer_path=tokenizer_path) ++ self.text_encoder: WanTextEncoder = None ++ self.image_encoder: WanImageEncoder = None ++ self.dit: WanModel = None ++ self.vae: WanVideoVAE = None ++ self.model_names = ["text_encoder", "dit", "vae"] ++ self.height_division_factor = 16 ++ self.width_division_factor = 16 ++ ++ def enable_vram_management(self, num_persistent_param_in_dit=None): ++ dtype = next(iter(self.text_encoder.parameters())).dtype ++ enable_vram_management( ++ self.text_encoder, ++ module_map={ ++ torch.nn.Linear: AutoWrappedLinear, ++ torch.nn.Embedding: AutoWrappedModule, ++ T5RelativeEmbedding: AutoWrappedModule, ++ T5LayerNorm: AutoWrappedModule, ++ }, ++ module_config=dict( ++ offload_dtype=dtype, ++ offload_device="cpu", ++ onload_dtype=dtype, ++ onload_device="cpu", ++ computation_dtype=self.torch_dtype, ++ computation_device=self.device, ++ ), ++ ) ++ dtype = next(iter(self.dit.parameters())).dtype ++ enable_vram_management( ++ self.dit, ++ module_map={ ++ torch.nn.Linear: AutoWrappedLinear, ++ torch.nn.Conv3d: AutoWrappedModule, ++ torch.nn.LayerNorm: AutoWrappedModule, ++ RMSNorm: AutoWrappedModule, ++ }, ++ module_config=dict( ++ offload_dtype=dtype, ++ offload_device="cpu", ++ onload_dtype=dtype, ++ onload_device=self.device, ++ computation_dtype=self.torch_dtype, ++ computation_device=self.device, ++ ), ++ max_num_param=num_persistent_param_in_dit, ++ overflow_module_config=dict( ++ offload_dtype=dtype, ++ offload_device="cpu", ++ onload_dtype=dtype, ++ onload_device="cpu", ++ computation_dtype=self.torch_dtype, ++ computation_device=self.device, ++ ), ++ ) ++ dtype = next(iter(self.vae.parameters())).dtype ++ enable_vram_management( ++ self.vae, ++ module_map={ ++ torch.nn.Linear: AutoWrappedLinear, ++ torch.nn.Conv2d: AutoWrappedModule, ++ RMS_norm: AutoWrappedModule, ++ CausalConv3d: AutoWrappedModule, ++ Upsample: AutoWrappedModule, ++ torch.nn.SiLU: AutoWrappedModule, ++ torch.nn.Dropout: AutoWrappedModule, ++ }, ++ module_config=dict( ++ offload_dtype=dtype, ++ offload_device="cpu", ++ onload_dtype=dtype, ++ onload_device=self.device, ++ computation_dtype=self.torch_dtype, ++ computation_device=self.device, ++ ), ++ ) ++ if self.image_encoder is not None: ++ dtype = next(iter(self.image_encoder.parameters())).dtype ++ enable_vram_management( ++ self.image_encoder, ++ module_map={ ++ torch.nn.Linear: AutoWrappedLinear, ++ torch.nn.Conv2d: AutoWrappedModule, ++ torch.nn.LayerNorm: AutoWrappedModule, ++ }, ++ module_config=dict( ++ offload_dtype=dtype, ++ offload_device="cpu", ++ onload_dtype=dtype, ++ onload_device="cpu", ++ computation_dtype=dtype, ++ computation_device=self.device, ++ ), ++ ) ++ self.enable_cpu_offload() ++ ++ def fetch_models(self, model_manager: ModelManager): ++ text_encoder_model_and_path = model_manager.fetch_model("wan_video_text_encoder", require_model_path=True) ++ if text_encoder_model_and_path is not None: ++ self.text_encoder, tokenizer_path = text_encoder_model_and_path ++ self.prompter.fetch_models(self.text_encoder) ++ self.prompter.fetch_tokenizer(os.path.join(os.path.dirname(tokenizer_path), "google/umt5-xxl")) ++ self.dit = model_manager.fetch_model("wan_video_dit") ++ self.vae = model_manager.fetch_model("wan_video_vae") ++ self.image_encoder = model_manager.fetch_model("wan_video_image_encoder") ++ ++ @staticmethod ++ def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None): ++ if device is None: ++ device = model_manager.device ++ if torch_dtype is None: ++ torch_dtype = model_manager.torch_dtype ++ pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=torch_dtype) ++ pipe.fetch_models(model_manager) ++ return pipe ++ ++ @staticmethod ++ def from_pretrained( ++ torch_dtype: torch.dtype = torch.bfloat16, ++ device: Union[str, torch.device] = "cuda", ++ model_configs: list[ModelConfig] = [], ++ tokenizer_config: ModelConfig = ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="google/*"), ++ redirect_common_files: bool = True, ++ use_usp=False, ++ ): ++ # Redirect model path ++ if redirect_common_files: ++ redirect_dict = { ++ "models_t5_umt5-xxl-enc-bf16.pth": "Wan-AI/Wan2.1-T2V-1.3B", ++ "Wan2.1_VAE.pth": "Wan-AI/Wan2.1-T2V-1.3B", ++ "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth": "Wan-AI/Wan2.1-I2V-14B-480P", ++ } ++ for model_config in model_configs: ++ if model_config.origin_file_pattern is None or model_config.model_id is None: ++ continue ++ if ( ++ model_config.origin_file_pattern in redirect_dict ++ and model_config.model_id != redirect_dict[model_config.origin_file_pattern] ++ ): ++ print( ++ f"To avoid repeatedly downloading model files, ({model_config.model_id}, {model_config.origin_file_pattern}) is redirected to ({redirect_dict[model_config.origin_file_pattern]}, {model_config.origin_file_pattern}). You can use `redirect_common_files=False` to disable file redirection." ++ ) ++ model_config.model_id = redirect_dict[model_config.origin_file_pattern] ++ ++ # Initialize pipeline ++ pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=torch_dtype) ++ if use_usp: ++ pipe.initialize_usp() ++ ++ # Download and load models ++ model_manager = ModelManager() ++ for model_config in model_configs: ++ model_config.download_if_necessary(use_usp=use_usp) ++ model_manager.load_model( ++ model_config.path, ++ device=model_config.offload_device or device, ++ torch_dtype=model_config.offload_dtype or torch_dtype, ++ ) ++ ++ # Load models ++ pipe.text_encoder = model_manager.fetch_model("wan_video_text_encoder") ++ dit = model_manager.fetch_model("wan_video_dit", index=2) ++ if isinstance(dit, list): ++ pipe.dit, pipe.dit2 = dit ++ else: ++ pipe.dit = dit ++ pipe.vae = model_manager.fetch_model("wan_video_vae") ++ pipe.image_encoder = model_manager.fetch_model("wan_video_image_encoder") ++ pipe.motion_controller = model_manager.fetch_model("wan_video_motion_controller") ++ pipe.vace = model_manager.fetch_model("wan_video_vace") ++ ++ # Size division factor ++ if pipe.vae is not None: ++ pipe.height_division_factor = pipe.vae.upsampling_factor * 2 ++ pipe.width_division_factor = pipe.vae.upsampling_factor * 2 ++ ++ # Initialize tokenizer ++ tokenizer_config.local_model_path = model_configs[0].local_model_path ++ tokenizer_config.skip_download = model_configs[0].skip_download ++ tokenizer_config.download_if_necessary(use_usp=use_usp) ++ pipe.prompter.fetch_models(pipe.text_encoder) ++ pipe.prompter.fetch_tokenizer(tokenizer_config.path) ++ ++ # Unified Sequence Parallel ++ if use_usp: ++ pipe.enable_usp() ++ ++ return pipe ++ ++ def init_spatem_modules( ++ self, ++ disable_video_attn: bool = False, ++ use_4d_attn: bool = False, ++ use_mvs_attn: bool = False, ++ use_src_self_attn: bool = False, ++ use_src_cross_attn: bool = False, ++ freqs_src_shift: int = 121, ++ use_viewpack: bool = True, ++ viewpack_dropout_prob: float = 0.0, ++ use_pose_encoder: bool = True, ++ pose_encoder_type: str = "rgb", ++ use_cam_encoder: bool = False, ++ range_4d_attn: tuple[int, int, int] = (0, None, 2), ++ range_mvs_attn: tuple[int, int, int] = (1, None, 2), ++ range_src_self_attn: tuple[int, int, int] = (0, None, 2), ++ range_src_cross_attn: tuple[int, int, int] = (0, None, 2), ++ use_lbm: bool = False, ++ fill_wpmask_with_noise: bool = False, ++ ): ++ # breakpoint() ++ device, dtype = self.dit.patch_embedding.weight.device, self.dit.patch_embedding.weight.dtype ++ ++ if disable_video_attn: ++ # todo: delete self_attn layers from the model ++ if use_4d_attn: ++ raise ValueError("Cannot use 4D attention when video attention is disabled") ++ for block in self.dit.blocks: ++ block.disable_video_attn = True ++ ++ if use_4d_attn: ++ b, e, s = range_4d_attn ++ for block in self.dit.blocks[b:e:s]: ++ block.use_4d_attn = True ++ ++ if use_mvs_attn: ++ b, e, s = range_mvs_attn ++ for block in self.dit.blocks[b:e:s]: ++ block.use_mvs_attn = True ++ ++ dim = block.self_attn.q.weight.shape[0] ++ block.modulation_mvs = nn.Parameter(block.modulation[:, :3, :].detach().clone()) ++ block.norm1_mvs = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to( ++ device=device, dtype=dtype ++ ) ++ block.self_attn_mvs = SelfAttention(dim, block.self_attn.num_heads, block.self_attn.norm_q.eps).to( ++ device=device, dtype=dtype ++ ) ++ block.self_attn_mvs.load_state_dict(block.self_attn.state_dict(), strict=True) ++ ++ if not 0.0 <= viewpack_dropout_prob <= 1.0: ++ raise ValueError("viewpack_dropout_prob should be between 0 and 1") ++ if viewpack_dropout_prob > 0.0 and not use_viewpack: ++ raise ValueError("viewpack_dropout_prob requires use_viewpack=True") ++ ++ if use_viewpack: ++ viewpack_emb = ViewPackEmbedding( ++ in_dim=self.dit.patch_embedding.weight.shape[1], ++ dim=self.dit.patch_embedding.weight.shape[0], ++ patch_size=list(self.dit.patch_embedding.kernel_size), ++ ) ++ viewpack_emb.initialize_from_patch_embedding(self.dit.patch_embedding) ++ self.dit.viewpack_embedding = viewpack_emb.to(device=device, dtype=dtype) ++ elif use_src_self_attn: ++ if use_src_cross_attn: ++ raise ValueError("Cannot use both src self-attention and src cross-attention") ++ ++ b, e, s = range_src_self_attn ++ for block in self.dit.blocks[b:e:s]: ++ block.use_src_self_attn = True ++ ++ dim = block.self_attn.q.weight.shape[0] ++ block.modulation_src = nn.Parameter(block.modulation[:, :3, :].detach().clone()) ++ block.norm1_src = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to( ++ device=device, dtype=dtype ++ ) ++ block.self_attn_src = SelfAttention(dim, block.self_attn.num_heads, block.self_attn.norm_q.eps).to( ++ device=device, dtype=dtype ++ ) ++ block.self_attn_src.load_state_dict(block.self_attn.state_dict(), strict=True) ++ elif use_src_cross_attn: ++ b, e, s = range_src_cross_attn ++ for block in self.dit.blocks[b:e:s]: ++ block.use_src_cross_attn = True ++ ++ dim = block.self_attn.q.weight.shape[0] ++ block.modulation_src = nn.Parameter(block.modulation[:, :3, :].detach().clone()) ++ block.norm1_src = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to( ++ device=device, dtype=dtype ++ ) ++ block.cross_attn_src = CrossAttentionSrcCam( ++ dim, block.self_attn.num_heads, block.self_attn.norm_q.eps ++ ).to(device=device, dtype=dtype) ++ block.cross_attn_src.load_state_dict(block.self_attn.state_dict(), strict=True) ++ ++ if use_pose_encoder: ++ if pose_encoder_type == "rgb": ++ in_channels = 3 ++ elif pose_encoder_type == "rgbd": ++ in_channels = 4 ++ else: ++ raise ValueError(f"Invalid pose_encoder_type: {pose_encoder_type}") ++ pose_encoder = PoseEncoder(out_dim=self.dit.patch_embedding.out_channels, in_channels=in_channels) ++ self.dit.pose_encoder = pose_encoder.to(device=device, dtype=dtype) ++ ++ if use_cam_encoder: ++ dim = self.dit.blocks[0].self_attn.q.weight.shape[0] ++ for block in self.dit.blocks: ++ block.use_cam_encoder = True ++ block.cam_encoder = nn.Linear(12, dim).to(device=device, dtype=dtype) ++ block.projector = nn.Linear(dim, dim).to(device=device, dtype=dtype) ++ block.cam_encoder.weight.data.zero_() ++ block.cam_encoder.bias.data.zero_() ++ block.projector.weight = nn.Parameter(torch.eye(dim, device=device, dtype=dtype)) ++ block.projector.bias = nn.Parameter(torch.zeros(dim, device=device, dtype=dtype)) ++ ++ if use_lbm: ++ # TODO: hard-code for now ++ self.scheduler = BridgeMatchScheduler() ++ self.dit.fill_wpmask_with_noise = fill_wpmask_with_noise ++ ++ self.dit.use_pose_encoder = use_pose_encoder ++ self.dit.use_cam_encoder = use_cam_encoder ++ self.dit.use_viewpack = use_viewpack ++ self.dit.viewpack_dropout_prob = viewpack_dropout_prob ++ self.dit.use_src_attn = use_src_self_attn or use_src_cross_attn ++ self.dit.use_lbm = use_lbm ++ self.dit.freqs_src_shift = freqs_src_shift ++ ++ def denoising_model(self): ++ return self.dit ++ ++ def encode_prompt(self, prompt, positive=True): ++ prompt_emb = self.prompter.encode_prompt(prompt, positive=positive) ++ return {"context": prompt_emb} ++ ++ def encode_image(self, image, num_frames, height, width): ++ image = self.preprocess_image(image.resize((width, height))).to(self.device) ++ clip_context = self.image_encoder.encode_image([image]) ++ msk = torch.ones(1, num_frames, height // 8, width // 8, device=self.device) ++ msk[:, 1:] = 0 ++ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) ++ msk = msk.view(1, msk.shape[1] // 4, 4, height // 8, width // 8) ++ msk = msk.transpose(1, 2)[0] ++ ++ vae_input = torch.concat( ++ [image.transpose(0, 1), torch.zeros(3, num_frames - 1, height, width).to(image.device)], dim=1 ++ ) ++ y = self.vae.encode([vae_input.to(dtype=self.torch_dtype, device=self.device)], device=self.device)[0] ++ y = torch.concat([msk, y]) ++ y = y.unsqueeze(0) ++ clip_context = clip_context.to(dtype=self.torch_dtype, device=self.device) ++ y = y.to(dtype=self.torch_dtype, device=self.device) ++ return {"clip_feature": clip_context, "y": y} ++ ++ def tensor2video(self, frames): ++ frames = rearrange(frames, "c f h w -> f h w c") ++ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) ++ frames = [Image.fromarray(frame) for frame in frames] ++ return frames ++ ++ def prepare_extra_input(self, latents=None): ++ return {} ++ ++ def encode_video(self, input_video, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)): ++ latents = self.vae.encode( ++ input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride ++ ) ++ return latents ++ ++ def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)): ++ frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride) ++ return frames ++ ++ def encode_fmask(self, mask, size): ++ f_, h_, w_ = size ++ g_ = (mask.shape[2] - 1) // (f_ - 1) ++ ++ # union of the frames in each latent (the first frame is encoded independently) ++ mask = torch.cat([mask[:, :, :1].repeat(1, 1, g_ - 1, 1, 1), mask], dim=2) ++ mask = rearrange(mask, "v c (f g) h w -> v c f g h w", f=f_, g=g_) ++ mask = mask.max(dim=3).values ++ ++ # interpolate along the spatial dimensions ++ mask = rearrange(mask, "v c f h w -> (v f) c h w") ++ mask = F.interpolate(mask, size=(h_, w_), mode="area") ++ mask = rearrange(mask, "(v f) c h w -> v c f h w", f=f_) ++ return mask ++ ++ def encode_wpmask(self, mask, size): ++ f_, h_, w_ = size ++ g_ = (mask.shape[2] - 1) // (f_ - 1) ++ ++ # intersection of the frames in each latent (the first frame is encoded independently) ++ mask = torch.cat([mask[:, :, :1].repeat(1, 1, g_ - 1, 1, 1), mask], dim=2) ++ mask = rearrange(mask, "v c (f g) h w -> v c f g h w", f=f_, g=g_) ++ mask = mask.min(dim=3).values ++ ++ # interpolate along the spatial dimensions ++ mask = rearrange(mask, "v c f h w -> (v f) c h w") ++ mask = F.interpolate(mask, size=(h_, w_), mode="area") ++ mask = rearrange(mask, "(v f) c h w -> v c f h w", f=f_) ++ return mask ++ ++ @torch.no_grad() ++ def __call__( ++ self, ++ prompt, ++ negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", ++ src_videos: torch.Tensor = None, ++ skeletons: torch.Tensor = None, ++ wpvideos: torch.Tensor = None, ++ wpmasks: torch.Tensor = None, ++ cam_emb: torch.Tensor = None, ++ input_image: Image.Image = None, ++ input_video: torch.Tensor = None, ++ denoising_strength: float = 1.0, ++ seed: int = None, ++ rand_device: str = "cpu", ++ height: int = 832, ++ width: int = 480, ++ num_frames: int = None, ++ cfg_scale: float = 5.0, ++ num_inference_steps: int = 50, ++ sigma_shift: float = 5.0, ++ tiled: bool = True, ++ tile_size: tuple[int, int] = (52, 30), ++ tile_stride: tuple[int, int] = (26, 15), ++ tea_cache_l1_thresh: float = None, ++ tea_cache_model_id: str = "", ++ progress_bar_cmd=partial(tqdm, desc="Denoising"), ++ progress_bar_st=None, ++ return_tensor=False, ++ ): ++ # breakpoint() ++ assert num_frames is None, "num_frames is not supported for WanVideoSpaTemPipeline" ++ assert input_image is None, "input_image is not supported for WanVideoSpaTemPipeline" ++ assert input_video is None, "input_video is not supported for WanVideoSpaTemPipeline" ++ assert tea_cache_l1_thresh is None, "tea_cache_l1_thresh is not supported for WanVideoSpaTemPipeline" ++ ++ # Parameter check ++ height, width = self.check_resize_height_width(height, width) ++ ++ # Tiler parameters ++ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride} ++ ++ # Scheduler ++ if self.dit.use_lbm: ++ # bridge matching scheduler ++ self.scheduler.set_timesteps(num_inference_steps) ++ else: ++ # flow matching scheduler ++ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift) ++ ++ src_videos = src_videos.to(dtype=self.torch_dtype, device=self.device) ++ if skeletons is not None: ++ skeletons = skeletons.to(dtype=self.torch_dtype, device=self.device) ++ if wpvideos is not None: ++ wpvideos = wpvideos.to(dtype=self.torch_dtype, device=self.device) ++ if wpmasks is not None: ++ wpmasks = wpmasks.to(dtype=self.torch_dtype, device=self.device) ++ if cam_emb is not None: ++ cam_emb = cam_emb.to(dtype=self.torch_dtype, device=self.device) ++ ++ if skeletons is not None: ++ num_cameras = skeletons.shape[0] ++ num_frames = skeletons.shape[2] ++ elif wpvideos is not None: ++ num_cameras = wpvideos.shape[0] ++ num_frames = wpvideos.shape[2] ++ else: ++ raise ValueError("Either skeletons or wpvideos must be provided") ++ ++ # Initialize noise ++ noise_shape = ( ++ num_cameras, ++ self.vae.model.z_dim, ++ (num_frames - 1) // 4 + 1, ++ height // self.vae.upsampling_factor, ++ width // self.vae.upsampling_factor, ++ ) ++ noise = self.generate_noise(noise_shape, seed=seed, device=rand_device, dtype=torch.float32) ++ noise = noise.to(dtype=self.torch_dtype, device=self.device) ++ ++ if input_video is not None: ++ self.load_models_to_device(["vae"]) ++ input_video = self.preprocess_images(input_video) ++ input_video = torch.stack(input_video, dim=2).to(dtype=self.torch_dtype, device=self.device) ++ latents = self.encode_video(input_video, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device) ++ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0]) ++ else: ++ latents = noise ++ ++ # Encode source video ++ self.load_models_to_device(["vae"]) ++ src_latents = self.encode_video(src_videos, **tiler_kwargs) ++ src_latents = src_latents.to(dtype=self.torch_dtype, device=self.device) ++ src_latents_nega = torch.zeros_like(src_latents) ++ ++ # Latent bridge matching ++ if self.dit.use_lbm: ++ if skeletons is not None: ++ # skeleton-based: use primary src_latents as bridge source ++ lbm_src_latents = src_latents[:1].expand_as(latents) ++ elif wpvideos is not None: ++ lbm_src_latents = self.encode_video(wpvideos, **tiler_kwargs).to( ++ dtype=self.torch_dtype, device=self.device ++ ) ++ if self.dit.fill_wpmask_with_noise: ++ wpmask_latents = self.encode_wpmask(wpmasks, size=lbm_src_latents.shape[-3:]) ++ lbm_src_latents = lbm_src_latents * wpmask_latents + noise * (1 - wpmask_latents) ++ ++ latents = lbm_src_latents ++ ++ # Encode prompts ++ self.load_models_to_device(["text_encoder"]) ++ prompt_emb_posi = self.encode_prompt(prompt, positive=True) ++ if cfg_scale != 1.0: ++ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False) ++ ++ # Encode image ++ if input_image is not None and self.image_encoder is not None: ++ self.load_models_to_device(["image_encoder", "vae"]) ++ image_emb = self.encode_image(input_image, num_frames, height, width) ++ else: ++ image_emb = {} ++ ++ # Extra input ++ extra_input = self.prepare_extra_input(latents) ++ ++ # Denoise ++ self.load_models_to_device(["dit"]) ++ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)): ++ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device) ++ timestep = torch.cat([timestep] * num_cameras, dim=0) ++ ++ # Inference ++ noise_pred_posi = self.denoising_model()( ++ x=latents, ++ x_src=src_latents, ++ timestep=timestep, ++ skeletons=skeletons, ++ cam_emb=cam_emb, ++ **prompt_emb_posi, ++ **image_emb, ++ **extra_input, ++ ) ++ if cfg_scale != 1.0: ++ noise_pred_nega = self.denoising_model()( ++ x=latents, ++ x_src=src_latents_nega, ++ timestep=timestep, ++ skeletons=skeletons, ++ cam_emb=cam_emb, ++ **prompt_emb_nega, ++ **image_emb, ++ **extra_input, ++ ) ++ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega) ++ else: ++ noise_pred = noise_pred_posi ++ ++ # Scheduler ++ latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents) ++ ++ # Decode ++ self.load_models_to_device(["vae"]) ++ pred_videos = self.decode_video(latents, **tiler_kwargs) ++ ++ if return_tensor: ++ return pred_videos ++ ++ self.load_models_to_device([]) ++ pred_video_list = [] ++ for pred_video in pred_videos: ++ pred_video_list.append(self.tensor2video(pred_video)) ++ return pred_video_list ++ ++ ++class TeaCache: ++ def __init__(self, num_inference_steps, rel_l1_thresh, model_id): ++ self.num_inference_steps = num_inference_steps ++ self.step = 0 ++ self.accumulated_rel_l1_distance = 0 ++ self.previous_modulated_input = None ++ self.rel_l1_thresh = rel_l1_thresh ++ self.previous_residual = None ++ self.previous_hidden_states = None ++ ++ self.coefficients_dict = { ++ "Wan2.1-T2V-1.3B": [-5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02], ++ "Wan2.1-T2V-14B": [-3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01], ++ "Wan2.1-I2V-14B-480P": [2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01], ++ "Wan2.1-I2V-14B-720P": [8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02], ++ } ++ if model_id not in self.coefficients_dict: ++ supported_model_ids = ", ".join([i for i in self.coefficients_dict]) ++ raise ValueError( ++ f"{model_id} is not a supported TeaCache model id. Please choose a valid model id in ({supported_model_ids})." ++ ) ++ self.coefficients = self.coefficients_dict[model_id] ++ ++ def check(self, dit: WanModel, x, t_mod): ++ modulated_inp = t_mod.clone() ++ if self.step == 0 or self.step == self.num_inference_steps - 1: ++ should_calc = True ++ self.accumulated_rel_l1_distance = 0 ++ else: ++ coefficients = self.coefficients ++ rescale_func = np.poly1d(coefficients) ++ self.accumulated_rel_l1_distance += rescale_func( ++ ( ++ (modulated_inp - self.previous_modulated_input).abs().mean() ++ / self.previous_modulated_input.abs().mean() ++ ) ++ .cpu() ++ .item() ++ ) ++ if self.accumulated_rel_l1_distance < self.rel_l1_thresh: ++ should_calc = False ++ else: ++ should_calc = True ++ self.accumulated_rel_l1_distance = 0 ++ self.previous_modulated_input = modulated_inp ++ self.step += 1 ++ if self.step == self.num_inference_steps: ++ self.step = 0 ++ if should_calc: ++ self.previous_hidden_states = x.clone() ++ return not should_calc ++ ++ def store(self, hidden_states): ++ self.previous_residual = hidden_states - self.previous_hidden_states ++ self.previous_hidden_states = None ++ ++ def update(self, hidden_states): ++ hidden_states = hidden_states + self.previous_residual ++ return hidden_states +diff --git a/diffsynth/schedulers/__init__.py b/diffsynth/schedulers/__init__.py +index 0ec4325..03cae00 100644 +--- a/diffsynth/schedulers/__init__.py ++++ b/diffsynth/schedulers/__init__.py +@@ -1,3 +1,4 @@ + from .ddim import EnhancedDDIMScheduler + from .continuous_ode import ContinuousODEScheduler + from .flow_match import FlowMatchScheduler ++from .bride_match import BridgeMatchScheduler +diff --git a/diffsynth/schedulers/bride_match.py b/diffsynth/schedulers/bride_match.py +new file mode 100644 +index 0000000..c157702 +--- /dev/null ++++ b/diffsynth/schedulers/bride_match.py +@@ -0,0 +1,71 @@ ++import torch, math ++ ++ ++class BridgeMatchScheduler: ++ ++ def __init__( ++ self, ++ num_train_timesteps=1000, ++ num_inference_steps=8, ++ sigma_max=1.0, ++ bridge_noise_sigma=0.005, ++ ): ++ self.sigma_max = sigma_max ++ self.sigma_min = sigma_max / num_train_timesteps ++ self.num_train_timesteps = num_train_timesteps # train timesteps for base model ++ self.num_inference_steps = num_inference_steps # inference steps for bridge matching ++ self.bridge_noise_sigma = bridge_noise_sigma ++ ++ self.set_timesteps(self.num_inference_steps) ++ ++ def set_timesteps(self, num_inference_steps=8, training=False): ++ sigma_start = self.sigma_max ++ sigma_end = self.sigma_max / num_inference_steps ++ self.sigmas = torch.linspace(sigma_start, sigma_end, num_inference_steps) ++ self.timesteps = self.sigmas * self.num_train_timesteps ++ ++ self.training = training ++ ++ def retrieve_sigma(self, timestep): ++ if isinstance(timestep, torch.Tensor): ++ timestep = timestep.cpu() ++ timestep_id = torch.argmin((self.timesteps - timestep).abs()) ++ sigma = self.sigmas[timestep_id] ++ return sigma ++ ++ def get_noise_term(self, sigma, sample): ++ # bridge noise term == 0 when sigma == 1 or 0 ++ return self.bridge_noise_sigma * (sigma * (1.0 - sigma)) ** 0.5 * torch.randn_like(sample) ++ ++ def step(self, model_output, timestep, sample, to_final=False, **kwargs): ++ if isinstance(timestep, torch.Tensor): ++ timestep = timestep.cpu() ++ timestep_id = torch.argmin((self.timesteps - timestep).abs()) ++ sigma = self.sigmas[timestep_id] ++ if to_final or timestep_id + 1 >= len(self.timesteps): ++ sigma_ = 0 ++ else: ++ sigma_ = self.sigmas[timestep_id + 1] ++ ++ prev_sample = sample + model_output * (sigma_ - sigma) + self.get_noise_term(sigma_, sample) ++ return prev_sample ++ ++ def add_noise(self, tgt_sample, src_sample, timestep): ++ sigma = self.retrieve_sigma(timestep) ++ noisy_sample = sigma * src_sample + (1 - sigma) * tgt_sample + self.get_noise_term(sigma, tgt_sample) ++ return noisy_sample ++ ++ def training_target(self, tgt_sample, noisy_sample, timestep): ++ sigma = self.retrieve_sigma(timestep) ++ ++ target = (noisy_sample - tgt_sample) / sigma ++ return target ++ ++ def denoised_sample(self, prediction, noisy_sample, timestep): ++ sigma = self.retrieve_sigma(timestep) ++ ++ sample = noisy_sample - prediction * sigma ++ return sample ++ ++ def training_weight(self, timestep): ++ return 1.0 +diff --git a/diffsynth/schedulers/flow_match.py b/diffsynth/schedulers/flow_match.py +index 6a8e235..0cd3e08 100644 +--- a/diffsynth/schedulers/flow_match.py ++++ b/diffsynth/schedulers/flow_match.py +@@ -98,7 +98,12 @@ class FlowMatchScheduler(): + def training_target(self, sample, noise, timestep): + target = noise - sample + return target +- ++ ++ ++ def denoised_sample(self, prediction, noise, timestep): ++ sample = noise - prediction ++ return sample ++ + + def training_weight(self, timestep): + timestep_id = torch.argmin((self.timesteps - timestep.to(self.timesteps.device)).abs()) diff --git a/fdanyone/vendor/diffsynth/VENDORED_FILES.txt b/fdanyone/vendor/diffsynth/VENDORED_FILES.txt new file mode 100644 index 0000000000000000000000000000000000000000..45540202934e0e22f49219459bd86ab746a2a5be --- /dev/null +++ b/fdanyone/vendor/diffsynth/VENDORED_FILES.txt @@ -0,0 +1,22 @@ +# Paths are relative to fdanyone/vendor/diffsynth. +LICENSE +UPSTREAM.md +UPSTREAM.patch +VENDORED_FILES.txt +__init__.py +models/__init__.py +models/utils.py +models/wan_video_dit.py +models/wan_video_pose_encoder.py +models/wan_video_text_encoder.py +models/wan_video_vae.py +pipelines/__init__.py +pipelines/base.py +pipelines/wan_video_spatem.py +prompters/__init__.py +prompters/base_prompter.py +prompters/wan_prompter.py +schedulers/__init__.py +schedulers/flow_match.py +vram_management/__init__.py +vram_management/layers.py diff --git a/fdanyone/vendor/diffsynth/__init__.py b/fdanyone/vendor/diffsynth/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..3d29a80af230aad176139fcf012f4505a8da9f5f --- /dev/null +++ b/fdanyone/vendor/diffsynth/__init__.py @@ -0,0 +1 @@ +"""Minimal DiffSynth-Studio inference closure used by 4DAnyone.""" diff --git a/fdanyone/vendor/diffsynth/pipelines/__init__.py b/fdanyone/vendor/diffsynth/pipelines/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..3117b0d6a7cdd6bb9eaa4c0088ff0ac415b2f025 --- /dev/null +++ b/fdanyone/vendor/diffsynth/pipelines/__init__.py @@ -0,0 +1,5 @@ +"""Vendored inference pipelines.""" + +from .wan_video_spatem import WanVideoSpaTemPipeline + +__all__ = ["WanVideoSpaTemPipeline"] diff --git a/fdanyone/vendor/diffsynth/pipelines/base.py b/fdanyone/vendor/diffsynth/pipelines/base.py new file mode 100644 index 0000000000000000000000000000000000000000..2a4f01cff55dc0fcca02dc5234227bd65efc7434 --- /dev/null +++ b/fdanyone/vendor/diffsynth/pipelines/base.py @@ -0,0 +1,127 @@ +import torch +import numpy as np +from PIL import Image +from torchvision.transforms import GaussianBlur + + + +class BasePipeline(torch.nn.Module): + + def __init__(self, device="cuda", torch_dtype=torch.float16, height_division_factor=64, width_division_factor=64): + super().__init__() + self.device = device + self.torch_dtype = torch_dtype + self.height_division_factor = height_division_factor + self.width_division_factor = width_division_factor + self.cpu_offload = False + self.model_names = [] + + + def check_resize_height_width(self, height, width): + if height % self.height_division_factor != 0: + height = (height + self.height_division_factor - 1) // self.height_division_factor * self.height_division_factor + print(f"The height cannot be evenly divided by {self.height_division_factor}. We round it up to {height}.") + if width % self.width_division_factor != 0: + width = (width + self.width_division_factor - 1) // self.width_division_factor * self.width_division_factor + print(f"The width cannot be evenly divided by {self.width_division_factor}. We round it up to {width}.") + return height, width + + + def preprocess_image(self, image): + image = torch.Tensor(np.array(image, dtype=np.float32) * (2 / 255) - 1).permute(2, 0, 1).unsqueeze(0) + return image + + + def preprocess_images(self, images): + return [self.preprocess_image(image) for image in images] + + + def vae_output_to_image(self, vae_output): + image = vae_output[0].cpu().float().permute(1, 2, 0).numpy() + image = Image.fromarray(((image / 2 + 0.5).clip(0, 1) * 255).astype("uint8")) + return image + + + def vae_output_to_video(self, vae_output): + video = vae_output.cpu().permute(1, 2, 0).numpy() + video = [Image.fromarray(((image / 2 + 0.5).clip(0, 1) * 255).astype("uint8")) for image in video] + return video + + + def merge_latents(self, value, latents, masks, scales, blur_kernel_size=33, blur_sigma=10.0): + if len(latents) > 0: + blur = GaussianBlur(kernel_size=blur_kernel_size, sigma=blur_sigma) + height, width = value.shape[-2:] + weight = torch.ones_like(value) + for latent, mask, scale in zip(latents, masks, scales): + mask = self.preprocess_image(mask.resize((width, height))).mean(dim=1, keepdim=True) > 0 + mask = mask.repeat(1, latent.shape[1], 1, 1).to(dtype=latent.dtype, device=latent.device) + mask = blur(mask) + value += latent * mask * scale + weight += mask * scale + value /= weight + return value + + + def control_noise_via_local_prompts(self, prompt_emb_global, prompt_emb_locals, masks, mask_scales, inference_callback, special_kwargs=None, special_local_kwargs_list=None): + if special_kwargs is None: + noise_pred_global = inference_callback(prompt_emb_global) + else: + noise_pred_global = inference_callback(prompt_emb_global, special_kwargs) + if special_local_kwargs_list is None: + noise_pred_locals = [inference_callback(prompt_emb_local) for prompt_emb_local in prompt_emb_locals] + else: + noise_pred_locals = [inference_callback(prompt_emb_local, special_kwargs) for prompt_emb_local, special_kwargs in zip(prompt_emb_locals, special_local_kwargs_list)] + noise_pred = self.merge_latents(noise_pred_global, noise_pred_locals, masks, mask_scales) + return noise_pred + + + def extend_prompt(self, prompt, local_prompts, masks, mask_scales): + local_prompts = local_prompts or [] + masks = masks or [] + mask_scales = mask_scales or [] + extended_prompt_dict = self.prompter.extend_prompt(prompt) + prompt = extended_prompt_dict.get("prompt", prompt) + local_prompts += extended_prompt_dict.get("prompts", []) + masks += extended_prompt_dict.get("masks", []) + mask_scales += [100.0] * len(extended_prompt_dict.get("masks", [])) + return prompt, local_prompts, masks, mask_scales + + + def enable_cpu_offload(self): + self.cpu_offload = True + + + def load_models_to_device(self, loadmodel_names=[]): + # only load models to device if cpu_offload is enabled + if not self.cpu_offload: + return + # offload the unneeded models to cpu + for model_name in self.model_names: + if model_name not in loadmodel_names: + model = getattr(self, model_name) + if model is not None: + if hasattr(model, "vram_management_enabled") and model.vram_management_enabled: + for module in model.modules(): + if hasattr(module, "offload"): + module.offload() + else: + model.cpu() + # load the needed models to device + for model_name in loadmodel_names: + model = getattr(self, model_name) + if model is not None: + if hasattr(model, "vram_management_enabled") and model.vram_management_enabled: + for module in model.modules(): + if hasattr(module, "onload"): + module.onload() + else: + model.to(self.device) + # fresh the cuda cache + torch.cuda.empty_cache() + + + def generate_noise(self, shape, seed=None, device="cpu", dtype=torch.float16): + generator = None if seed is None else torch.Generator(device).manual_seed(seed) + noise = torch.randn(shape, generator=generator, device=device, dtype=dtype) + return noise diff --git a/fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py b/fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py new file mode 100644 index 0000000000000000000000000000000000000000..861f0d8ee8c7b2708c5850412c425536895bd2ac --- /dev/null +++ b/fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py @@ -0,0 +1,93 @@ +import torch +import torch.nn as nn + +from ..models.wan_video_dit import SelfAttention, ViewPackEmbedding, WanModel +from ..models.wan_video_pose_encoder import PoseEncoder +from ..models.wan_video_text_encoder import WanTextEncoder +from ..models.wan_video_vae import WanVideoVAE +from ..pipelines.base import BasePipeline +from ..prompters.wan_prompter import WanPrompter +from ..schedulers.flow_match import FlowMatchScheduler + + +class WanVideoSpaTemPipeline(BasePipeline): + def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None): + super().__init__(device=device, torch_dtype=torch_dtype) + self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True) + self.prompter = WanPrompter(tokenizer_path=tokenizer_path) + self.text_encoder: WanTextEncoder = None + self.dit: WanModel = None + self.vae: WanVideoVAE = None + self.model_names = ["text_encoder", "dit", "vae"] + self.height_division_factor = 16 + self.width_division_factor = 16 + + def init_spatem_modules( + self, + use_mvs_attn: bool = False, + use_viewpack: bool = True, + viewpack_dropout_prob: float = 0.0, + use_pose_encoder: bool = True, + pose_encoder_type: str = "rgb", + range_mvs_attn: tuple[int, int, int] = (1, None, 2), + ): + device = self.dit.patch_embedding.weight.device + dtype = self.dit.patch_embedding.weight.dtype + + if use_mvs_attn: + begin, end, stride = range_mvs_attn + for block in self.dit.blocks[begin:end:stride]: + block.use_mvs_attn = True + dim = block.self_attn.q.weight.shape[0] + block.modulation_mvs = nn.Parameter( + block.modulation[:, :3, :].detach().clone() + ) + block.norm1_mvs = nn.LayerNorm( + dim, + eps=block.norm1.eps, + elementwise_affine=False, + ).to(device=device, dtype=dtype) + block.self_attn_mvs = SelfAttention( + dim, + block.self_attn.num_heads, + block.self_attn.norm_q.eps, + ).to(device=device, dtype=dtype) + block.self_attn_mvs.load_state_dict( + block.self_attn.state_dict(), + strict=True, + ) + + if not 0.0 <= viewpack_dropout_prob <= 1.0: + raise ValueError("viewpack_dropout_prob should be between 0 and 1") + if viewpack_dropout_prob > 0.0 and not use_viewpack: + raise ValueError("viewpack_dropout_prob requires use_viewpack=True") + if use_viewpack: + viewpack_emb = ViewPackEmbedding( + in_dim=self.dit.patch_embedding.weight.shape[1], + dim=self.dit.patch_embedding.weight.shape[0], + patch_size=list(self.dit.patch_embedding.kernel_size), + ) + viewpack_emb.initialize_from_patch_embedding(self.dit.patch_embedding) + self.dit.viewpack_embedding = viewpack_emb.to( + device=device, + dtype=dtype, + ) + + if use_pose_encoder: + if pose_encoder_type != "rgb": + raise ValueError(f"Invalid pose_encoder_type: {pose_encoder_type}") + pose_encoder = PoseEncoder( + out_dim=self.dit.patch_embedding.out_channels, + in_channels=3, + ) + self.dit.pose_encoder = pose_encoder.to(device=device, dtype=dtype) + + self.dit.use_pose_encoder = use_pose_encoder + self.dit.use_viewpack = use_viewpack + self.dit.viewpack_dropout_prob = viewpack_dropout_prob + + def encode_video(self, input_video): + return self.vae.encode(input_video, device=self.device) + + def decode_video(self, latents): + return self.vae.decode(latents, device=self.device) diff --git a/fdanyone/vendor/diffsynth/prompters/__init__.py b/fdanyone/vendor/diffsynth/prompters/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d98693193926df2e1ad6fdf9282f17b7068dfad1 --- /dev/null +++ b/fdanyone/vendor/diffsynth/prompters/__init__.py @@ -0,0 +1,5 @@ +"""Vendored Wan prompt helpers.""" + +from .wan_prompter import WanPrompter + +__all__ = ["WanPrompter"] diff --git a/fdanyone/vendor/diffsynth/prompters/base_prompter.py b/fdanyone/vendor/diffsynth/prompters/base_prompter.py new file mode 100644 index 0000000000000000000000000000000000000000..749fc3fc67b626b67986df1e7e6d91527e679eaa --- /dev/null +++ b/fdanyone/vendor/diffsynth/prompters/base_prompter.py @@ -0,0 +1,69 @@ +import torch + + + +def tokenize_long_prompt(tokenizer, prompt, max_length=None): + # Get model_max_length from self.tokenizer + length = tokenizer.model_max_length if max_length is None else max_length + + # To avoid the warning. set self.tokenizer.model_max_length to +oo. + tokenizer.model_max_length = 99999999 + + # Tokenize it! + input_ids = tokenizer(prompt, return_tensors="pt").input_ids + + # Determine the real length. + max_length = (input_ids.shape[1] + length - 1) // length * length + + # Restore tokenizer.model_max_length + tokenizer.model_max_length = length + + # Tokenize it again with fixed length. + input_ids = tokenizer( + prompt, + return_tensors="pt", + padding="max_length", + max_length=max_length, + truncation=True + ).input_ids + + # Reshape input_ids to fit the text encoder. + num_sentence = input_ids.shape[1] // length + input_ids = input_ids.reshape((num_sentence, length)) + + return input_ids + + + +class BasePrompter: + def __init__(self): + self.refiners = [] + self.extenders = [] + + + def load_prompt_refiners(self, model_manager, refiner_classes=[]): + for refiner_class in refiner_classes: + refiner = refiner_class.from_model_manager(model_manager) + self.refiners.append(refiner) + + def load_prompt_extenders(self, model_manager, extender_classes=[]): + for extender_class in extender_classes: + extender = extender_class.from_model_manager(model_manager) + self.extenders.append(extender) + + + @torch.no_grad() + def process_prompt(self, prompt, positive=True): + if isinstance(prompt, list): + prompt = [self.process_prompt(prompt_, positive=positive) for prompt_ in prompt] + else: + for refiner in self.refiners: + prompt = refiner(prompt, positive=positive) + return prompt + + @torch.no_grad() + def extend_prompt(self, prompt:str, positive=True): + extended_prompt = dict(prompt=prompt) + for extender in self.extenders: + extended_prompt = extender(extended_prompt) + return extended_prompt diff --git a/fdanyone/vendor/diffsynth/prompters/wan_prompter.py b/fdanyone/vendor/diffsynth/prompters/wan_prompter.py new file mode 100644 index 0000000000000000000000000000000000000000..01a765d3cb3bf2ee4d06553fd061ed7dd75443b2 --- /dev/null +++ b/fdanyone/vendor/diffsynth/prompters/wan_prompter.py @@ -0,0 +1,109 @@ +from .base_prompter import BasePrompter +from ..models.wan_video_text_encoder import WanTextEncoder +from transformers import AutoTokenizer +import os, torch +import ftfy +import html +import string +import regex as re + + +def basic_clean(text): + text = ftfy.fix_text(text) + text = html.unescape(html.unescape(text)) + return text.strip() + + +def whitespace_clean(text): + text = re.sub(r'\s+', ' ', text) + text = text.strip() + return text + + +def canonicalize(text, keep_punctuation_exact_string=None): + text = text.replace('_', ' ') + if keep_punctuation_exact_string: + text = keep_punctuation_exact_string.join( + part.translate(str.maketrans('', '', string.punctuation)) + for part in text.split(keep_punctuation_exact_string)) + else: + text = text.translate(str.maketrans('', '', string.punctuation)) + text = text.lower() + text = re.sub(r'\s+', ' ', text) + return text.strip() + + +class HuggingfaceTokenizer: + + def __init__(self, name, seq_len=None, clean=None, **kwargs): + assert clean in (None, 'whitespace', 'lower', 'canonicalize') + self.name = name + self.seq_len = seq_len + self.clean = clean + + # init tokenizer + self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs) + self.vocab_size = self.tokenizer.vocab_size + + def __call__(self, sequence, **kwargs): + return_mask = kwargs.pop('return_mask', False) + + # arguments + _kwargs = {'return_tensors': 'pt'} + if self.seq_len is not None: + _kwargs.update({ + 'padding': 'max_length', + 'truncation': True, + 'max_length': self.seq_len + }) + _kwargs.update(**kwargs) + + # tokenization + if isinstance(sequence, str): + sequence = [sequence] + if self.clean: + sequence = [self._clean(u) for u in sequence] + ids = self.tokenizer(sequence, **_kwargs) + + # output + if return_mask: + return ids.input_ids, ids.attention_mask + else: + return ids.input_ids + + def _clean(self, text): + if self.clean == 'whitespace': + text = whitespace_clean(basic_clean(text)) + elif self.clean == 'lower': + text = whitespace_clean(basic_clean(text)).lower() + elif self.clean == 'canonicalize': + text = canonicalize(basic_clean(text)) + return text + + +class WanPrompter(BasePrompter): + + def __init__(self, tokenizer_path=None, text_len=512): + super().__init__() + self.text_len = text_len + self.text_encoder = None + self.fetch_tokenizer(tokenizer_path) + + def fetch_tokenizer(self, tokenizer_path=None): + if tokenizer_path is not None: + self.tokenizer = HuggingfaceTokenizer(name=tokenizer_path, seq_len=self.text_len, clean='whitespace') + + def fetch_models(self, text_encoder: WanTextEncoder = None): + self.text_encoder = text_encoder + + def encode_prompt(self, prompt, positive=True, device="cuda"): + prompt = self.process_prompt(prompt, positive=positive) + + ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True) + ids = ids.to(device) + mask = mask.to(device) + seq_lens = mask.gt(0).sum(dim=1).long() + prompt_emb = self.text_encoder(ids, mask) + for i, v in enumerate(seq_lens): + prompt_emb[:, v:] = 0 + return prompt_emb diff --git a/fdanyone/vendor/diffsynth/schedulers/__init__.py b/fdanyone/vendor/diffsynth/schedulers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..79d355241c1e83a8614048f62b95478c4cecd210 --- /dev/null +++ b/fdanyone/vendor/diffsynth/schedulers/__init__.py @@ -0,0 +1,5 @@ +"""Vendored diffusion schedulers.""" + +from .flow_match import FlowMatchScheduler + +__all__ = ["FlowMatchScheduler"] diff --git a/fdanyone/vendor/diffsynth/schedulers/flow_match.py b/fdanyone/vendor/diffsynth/schedulers/flow_match.py new file mode 100644 index 0000000000000000000000000000000000000000..0cd3e08a9811d7761b11add30a54dfc8c51c5ac9 --- /dev/null +++ b/fdanyone/vendor/diffsynth/schedulers/flow_match.py @@ -0,0 +1,125 @@ +import torch, math + + + +class FlowMatchScheduler(): + + def __init__( + self, + num_inference_steps=100, + num_train_timesteps=1000, + shift=3.0, + sigma_max=1.0, + sigma_min=0.003/1.002, + inverse_timesteps=False, + extra_one_step=False, + reverse_sigmas=False, + exponential_shift=False, + exponential_shift_mu=None, + shift_terminal=None, + ): + self.num_train_timesteps = num_train_timesteps + self.shift = shift + self.sigma_max = sigma_max + self.sigma_min = sigma_min + self.inverse_timesteps = inverse_timesteps + self.extra_one_step = extra_one_step + self.reverse_sigmas = reverse_sigmas + self.exponential_shift = exponential_shift + self.exponential_shift_mu = exponential_shift_mu + self.shift_terminal = shift_terminal + self.set_timesteps(num_inference_steps) + + + def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, dynamic_shift_len=None): + if shift is not None: + self.shift = shift + sigma_start = self.sigma_min + (self.sigma_max - self.sigma_min) * denoising_strength + if self.extra_one_step: + self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1] + else: + self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps) + if self.inverse_timesteps: + self.sigmas = torch.flip(self.sigmas, dims=[0]) + if self.exponential_shift: + mu = self.calculate_shift(dynamic_shift_len) if dynamic_shift_len is not None else self.exponential_shift_mu + self.sigmas = math.exp(mu) / (math.exp(mu) + (1 / self.sigmas - 1)) + else: + self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas) + if self.shift_terminal is not None: + one_minus_z = 1 - self.sigmas + scale_factor = one_minus_z[-1] / (1 - self.shift_terminal) + self.sigmas = 1 - (one_minus_z / scale_factor) + if self.reverse_sigmas: + self.sigmas = 1 - self.sigmas + self.timesteps = self.sigmas * self.num_train_timesteps + if training: + x = self.timesteps + y = torch.exp(-2 * ((x - num_inference_steps / 2) / num_inference_steps) ** 2) + y_shifted = y - y.min() + bsmntw_weighing = y_shifted * (num_inference_steps / y_shifted.sum()) + self.linear_timesteps_weights = bsmntw_weighing + self.training = True + else: + self.training = False + + + def step(self, model_output, timestep, sample, to_final=False, **kwargs): + if isinstance(timestep, torch.Tensor): + timestep = timestep.cpu() + timestep_id = torch.argmin((self.timesteps - timestep).abs()) + sigma = self.sigmas[timestep_id] + if to_final or timestep_id + 1 >= len(self.timesteps): + sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0 + else: + sigma_ = self.sigmas[timestep_id + 1] + prev_sample = sample + model_output * (sigma_ - sigma) + return prev_sample + + + def return_to_timestep(self, timestep, sample, sample_stablized): + if isinstance(timestep, torch.Tensor): + timestep = timestep.cpu() + timestep_id = torch.argmin((self.timesteps - timestep).abs()) + sigma = self.sigmas[timestep_id] + model_output = (sample - sample_stablized) / sigma + return model_output + + + def add_noise(self, original_samples, noise, timestep): + if isinstance(timestep, torch.Tensor): + timestep = timestep.cpu() + timestep_id = torch.argmin((self.timesteps - timestep).abs()) + sigma = self.sigmas[timestep_id] + sample = (1 - sigma) * original_samples + sigma * noise + return sample + + + def training_target(self, sample, noise, timestep): + target = noise - sample + return target + + + def denoised_sample(self, prediction, noise, timestep): + sample = noise - prediction + return sample + + + def training_weight(self, timestep): + timestep_id = torch.argmin((self.timesteps - timestep.to(self.timesteps.device)).abs()) + weights = self.linear_timesteps_weights[timestep_id] + return weights + + + def calculate_shift( + self, + image_seq_len, + base_seq_len: int = 256, + max_seq_len: int = 8192, + base_shift: float = 0.5, + max_shift: float = 0.9, + ): + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return mu diff --git a/fdanyone/vendor/diffsynth/vram_management/__init__.py b/fdanyone/vendor/diffsynth/vram_management/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..69a388db1dea2d5699b716260dfa0902c27c0ab5 --- /dev/null +++ b/fdanyone/vendor/diffsynth/vram_management/__init__.py @@ -0,0 +1 @@ +from .layers import * diff --git a/fdanyone/vendor/diffsynth/vram_management/layers.py b/fdanyone/vendor/diffsynth/vram_management/layers.py new file mode 100644 index 0000000000000000000000000000000000000000..58eb26dfcb46524b02ec84f44619ace773582b81 --- /dev/null +++ b/fdanyone/vendor/diffsynth/vram_management/layers.py @@ -0,0 +1,213 @@ +import torch, copy + +from ..models.utils import init_weights_on_device + + +def cast_to(weight, dtype, device): + r = torch.empty_like(weight, dtype=dtype, device=device) + r.copy_(weight) + return r + + +class AutoTorchModule(torch.nn.Module): + def __init__(self): + super().__init__() + + def check_free_vram(self): + gpu_mem_state = torch.cuda.mem_get_info(self.computation_device) + used_memory = (gpu_mem_state[1] - gpu_mem_state[0]) / (1024 ** 3) + return used_memory < self.vram_limit + + def offload(self): + if self.state != 0: + self.to(dtype=self.offload_dtype, device=self.offload_device) + self.state = 0 + + def onload(self): + if self.state != 1: + self.to(dtype=self.onload_dtype, device=self.onload_device) + self.state = 1 + + def keep(self): + if self.state != 2: + self.to(dtype=self.computation_dtype, device=self.computation_device) + self.state = 2 + + +class AutoWrappedModule(AutoTorchModule): + def __init__(self, module: torch.nn.Module, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, **kwargs): + super().__init__() + self.module = module.to(dtype=offload_dtype, device=offload_device) + self.offload_dtype = offload_dtype + self.offload_device = offload_device + self.onload_dtype = onload_dtype + self.onload_device = onload_device + self.computation_dtype = computation_dtype + self.computation_device = computation_device + self.vram_limit = vram_limit + self.state = 0 + + def forward(self, *args, **kwargs): + if self.state == 2: + module = self.module + else: + if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device: + module = self.module + elif self.vram_limit is not None and self.check_free_vram(): + self.keep() + module = self.module + else: + module = copy.deepcopy(self.module).to(dtype=self.computation_dtype, device=self.computation_device) + return module(*args, **kwargs) + + +class WanAutoCastLayerNorm(torch.nn.LayerNorm, AutoTorchModule): + def __init__(self, module: torch.nn.LayerNorm, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, **kwargs): + with init_weights_on_device(device=torch.device("meta")): + super().__init__(module.normalized_shape, eps=module.eps, elementwise_affine=module.elementwise_affine, bias=module.bias is not None, dtype=offload_dtype, device=offload_device) + self.weight = module.weight + self.bias = module.bias + self.offload_dtype = offload_dtype + self.offload_device = offload_device + self.onload_dtype = onload_dtype + self.onload_device = onload_device + self.computation_dtype = computation_dtype + self.computation_device = computation_device + self.vram_limit = vram_limit + self.state = 0 + + def forward(self, x, *args, **kwargs): + if self.state == 2: + weight, bias = self.weight, self.bias + else: + if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device: + weight, bias = self.weight, self.bias + elif self.vram_limit is not None and self.check_free_vram(): + self.keep() + weight, bias = self.weight, self.bias + else: + weight = None if self.weight is None else cast_to(self.weight, self.computation_dtype, self.computation_device) + bias = None if self.bias is None else cast_to(self.bias, self.computation_dtype, self.computation_device) + with torch.amp.autocast(device_type=x.device.type): + x = torch.nn.functional.layer_norm(x.float(), self.normalized_shape, weight, bias, self.eps).type_as(x) + return x + + +class AutoWrappedLinear(torch.nn.Linear, AutoTorchModule): + def __init__(self, module: torch.nn.Linear, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device, vram_limit, name="", **kwargs): + with init_weights_on_device(device=torch.device("meta")): + super().__init__(in_features=module.in_features, out_features=module.out_features, bias=module.bias is not None, dtype=offload_dtype, device=offload_device) + self.weight = module.weight + self.bias = module.bias + self.offload_dtype = offload_dtype + self.offload_device = offload_device + self.onload_dtype = onload_dtype + self.onload_device = onload_device + self.computation_dtype = computation_dtype + self.computation_device = computation_device + self.vram_limit = vram_limit + self.state = 0 + self.name = name + self.lora_A_weights = [] + self.lora_B_weights = [] + self.lora_merger = None + self.enable_fp8 = computation_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fnuz] + + def fp8_linear( + self, + input: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + device = input.device + origin_dtype = input.dtype + origin_shape = input.shape + input = input.reshape(-1, origin_shape[-1]) + + x_max = torch.max(torch.abs(input), dim=-1, keepdim=True).values + fp8_max = 448.0 + # For float8_e4m3fnuz, the maximum representable value is half of that of e4m3fn. + # To avoid overflow and ensure numerical compatibility during FP8 computation, + # we scale down the input by 2.0 in advance. + # This scaling will be compensated later during the final result scaling. + if self.computation_dtype == torch.float8_e4m3fnuz: + fp8_max = fp8_max / 2.0 + scale_a = torch.clamp(x_max / fp8_max, min=1.0).float().to(device=device) + scale_b = torch.ones((weight.shape[0], 1)).to(device=device) + input = input / (scale_a + 1e-8) + input = input.to(self.computation_dtype) + weight = weight.to(self.computation_dtype) + bias = bias.to(torch.bfloat16) + + result = torch._scaled_mm( + input, + weight.T, + scale_a=scale_a, + scale_b=scale_b.T, + bias=bias, + out_dtype=origin_dtype, + ) + new_shape = origin_shape[:-1] + result.shape[-1:] + result = result.reshape(new_shape) + return result + + def forward(self, x, *args, **kwargs): + # VRAM management + if self.state == 2: + weight, bias = self.weight, self.bias + else: + if self.onload_dtype == self.computation_dtype and self.onload_device == self.computation_device: + weight, bias = self.weight, self.bias + elif self.vram_limit is not None and self.check_free_vram(): + self.keep() + weight, bias = self.weight, self.bias + else: + weight = cast_to(self.weight, self.computation_dtype, self.computation_device) + bias = None if self.bias is None else cast_to(self.bias, self.computation_dtype, self.computation_device) + + # Linear forward + if self.enable_fp8: + out = self.fp8_linear(x, weight, bias) + else: + out = torch.nn.functional.linear(x, weight, bias) + + # LoRA + if len(self.lora_A_weights) == 0: + # No LoRA + return out + elif self.lora_merger is None: + # Native LoRA inference + for lora_A, lora_B in zip(self.lora_A_weights, self.lora_B_weights): + out = out + x @ lora_A.T @ lora_B.T + else: + # LoRA fusion + lora_output = [] + for lora_A, lora_B in zip(self.lora_A_weights, self.lora_B_weights): + lora_output.append(x @ lora_A.T @ lora_B.T) + lora_output = torch.stack(lora_output) + out = self.lora_merger(out, lora_output) + return out + + +def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict, module_config: dict, max_num_param=None, overflow_module_config: dict = None, total_num_param=0, vram_limit=None, name_prefix=""): + for name, module in model.named_children(): + layer_name = name if name_prefix == "" else name_prefix + "." + name + for source_module, target_module in module_map.items(): + if isinstance(module, source_module): + num_param = sum(p.numel() for p in module.parameters()) + if max_num_param is not None and total_num_param + num_param > max_num_param: + module_config_ = overflow_module_config + else: + module_config_ = module_config + module_ = target_module(module, **module_config_, vram_limit=vram_limit, name=layer_name) + setattr(model, name, module_) + total_num_param += num_param + break + else: + total_num_param = enable_vram_management_recursively(module, module_map, module_config, max_num_param, overflow_module_config, total_num_param, vram_limit=vram_limit, name_prefix=layer_name) + return total_num_param + + +def enable_vram_management(model: torch.nn.Module, module_map: dict, module_config: dict, max_num_param=None, overflow_module_config: dict = None, vram_limit=None): + enable_vram_management_recursively(model, module_map, module_config, max_num_param, overflow_module_config, total_num_param=0, vram_limit=vram_limit) + model.vram_management_enabled = True diff --git a/fdanyone/vendor/pytorch3d_compat/LICENSE b/fdanyone/vendor/pytorch3d_compat/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..c55382ff0992d90ae5ecb2cd9ac624ccd20bda4d --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/LICENSE @@ -0,0 +1,30 @@ +BSD License + +For PyTorch3D software + +Copyright (c) Meta Platforms, Inc. and affiliates. All rights reserved. + +Redistribution and use in source and binary forms, with or without modification, +are permitted provided that the following conditions are met: + + * Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. + + * Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. + + * Neither the name Meta nor the names of its contributors may be used to + endorse or promote products derived from this software without specific + prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND +ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR +ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES +(INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; +LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON +ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS +SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. diff --git a/fdanyone/vendor/pytorch3d_compat/UPSTREAM.md b/fdanyone/vendor/pytorch3d_compat/UPSTREAM.md new file mode 100644 index 0000000000000000000000000000000000000000..094596423df1315cbef38cea9acf873f6a2557b9 --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/UPSTREAM.md @@ -0,0 +1,8 @@ +# PyTorch3D compatibility subset + +The two retained source files are copied byte-for-byte from [`facebookresearch/pytorch3d@f34104c`](https://github.com/facebookresearch/pytorch3d/tree/f34104cf6ebefacd7b7e07955ee7aaa823e616ac) (release `v0.7.6`): + +- `common/datatypes.py` +- `transforms/rotation_conversions.py` + +PyTorch3D is BSD-licensed; its license is preserved in `LICENSE`. The smaller `__init__.py` files and native-PyTorch KNN fallback are 4DAnyone adapter code. They expose only the rotation/geometry surface imported by classic GVHMR inference and intentionally omit PyTorch3D's renderer and compiled operators. diff --git a/fdanyone/vendor/pytorch3d_compat/VENDORED_FILES.txt b/fdanyone/vendor/pytorch3d_compat/VENDORED_FILES.txt new file mode 100644 index 0000000000000000000000000000000000000000..dc0a7abebbf08445d28c5db4cb9d3aedee52cc19 --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/VENDORED_FILES.txt @@ -0,0 +1,11 @@ +# Complete file list for this compatibility subset. +LICENSE +UPSTREAM.md +VENDORED_FILES.txt +__init__.py +common/__init__.py +common/datatypes.py +ops/__init__.py +ops/knn.py +transforms/__init__.py +transforms/rotation_conversions.py diff --git a/fdanyone/vendor/pytorch3d_compat/__init__.py b/fdanyone/vendor/pytorch3d_compat/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..ab7a327612845499a9ccca9b2a476aac08e75907 --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/__init__.py @@ -0,0 +1,38 @@ +"""Minimal PyTorch3D 0.7.6 compatibility surface for classic GVHMR. + +The official GVHMR inference path needs rotation conversions but imports the +compiled PyTorch3D package through broader training/demo modules. 4DAnyone +registers this BSD-licensed, inference-only subset when full PyTorch3D is not +installed, allowing GVHMR and generation to share the same modern PyTorch +environment. +""" + +from __future__ import annotations + +import importlib +import importlib.util +import sys + + +def install_if_needed() -> bool: + """Expose the compatibility modules as ``pytorch3d`` when it is absent.""" + + if importlib.util.find_spec("pytorch3d") is not None: + return False + + root = sys.modules[__name__] + transforms = importlib.import_module(f"{__name__}.transforms") + ops = importlib.import_module(f"{__name__}.ops") + knn = importlib.import_module(f"{__name__}.ops.knn") + sys.modules.update( + { + "pytorch3d": root, + "pytorch3d.transforms": transforms, + "pytorch3d.ops": ops, + "pytorch3d.ops.knn": knn, + } + ) + return True + + +__all__ = ["install_if_needed"] diff --git a/fdanyone/vendor/pytorch3d_compat/common/__init__.py b/fdanyone/vendor/pytorch3d_compat/common/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..818c53ce146ec336abff251a0363fd71db381c72 --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/common/__init__.py @@ -0,0 +1,5 @@ +"""Small type helpers retained from PyTorch3D 0.7.6.""" + +from .datatypes import Device, get_device, make_device + +__all__ = ["Device", "get_device", "make_device"] diff --git a/fdanyone/vendor/pytorch3d_compat/common/datatypes.py b/fdanyone/vendor/pytorch3d_compat/common/datatypes.py new file mode 100644 index 0000000000000000000000000000000000000000..03fe3efc54dd81044ee579ee0aba8641eaa6b834 --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/common/datatypes.py @@ -0,0 +1,58 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from typing import Optional, Union + +import torch + + +Device = Union[str, torch.device] + + +def make_device(device: Device) -> torch.device: + """ + Makes an actual torch.device object from the device specified as + either a string or torch.device object. If the device is `cuda` without + a specific index, the index of the current device is assigned. + + Args: + device: Device (as str or torch.device) + + Returns: + A matching torch.device object + """ + device = torch.device(device) if isinstance(device, str) else device + if device.type == "cuda" and device.index is None: + # If cuda but with no index, then the current cuda device is indicated. + # In that case, we fix to that device + device = torch.device(f"cuda:{torch.cuda.current_device()}") + return device + + +def get_device(x, device: Optional[Device] = None) -> torch.device: + """ + Gets the device of the specified variable x if it is a tensor, or + falls back to a default CPU device otherwise. Allows overriding by + providing an explicit device. + + Args: + x: a torch.Tensor to get the device from or another type + device: Device (as str or torch.device) to fall back to + + Returns: + A matching torch.device object + """ + + # User overrides device + if device is not None: + return make_device(device) + + # Set device based on input tensor + if torch.is_tensor(x): + return x.device + + # Default device is cpu + return torch.device("cpu") diff --git a/fdanyone/vendor/pytorch3d_compat/ops/__init__.py b/fdanyone/vendor/pytorch3d_compat/ops/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a6d8c6f9000a1d0b3e3fbb3e5d43e7740c2318db --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/ops/__init__.py @@ -0,0 +1,5 @@ +"""Native-PyTorch operations needed to import classic GVHMR.""" + +from . import knn + +__all__ = ["knn"] diff --git a/fdanyone/vendor/pytorch3d_compat/ops/knn.py b/fdanyone/vendor/pytorch3d_compat/ops/knn.py new file mode 100644 index 0000000000000000000000000000000000000000..8663614d5999ab3daafbe5c9d85ca4ca5dc7a07f --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/ops/knn.py @@ -0,0 +1,22 @@ +"""Inference-only fallback for the one optional GVHMR KNN helper.""" + +from __future__ import annotations + + +def knn_points(p1, p2, *, K: int = 1, return_nn: bool = False): + """Return the PyTorch3D-compatible tuple using native PyTorch operations.""" + + import torch + + if p1.ndim != 3 or p2.ndim != 3 or p1.shape[0] != p2.shape[0]: + raise ValueError("p1 and p2 must have shapes (N, P, D) and (N, Q, D).") + if K < 1 or p2.shape[1] < K: + raise ValueError(f"K must be in [1, {p2.shape[1]}], got {K}.") + distances = torch.cdist(p1, p2).square() + squared_distances, indices = distances.topk(K, dim=-1, largest=False, sorted=True) + neighbors = None + if return_nn: + expanded = p2[:, None].expand(-1, p1.shape[1], -1, -1) + gather_index = indices[..., None].expand(-1, -1, -1, p2.shape[-1]) + neighbors = torch.gather(expanded, 2, gather_index) + return squared_distances, indices, neighbors diff --git a/fdanyone/vendor/pytorch3d_compat/transforms/__init__.py b/fdanyone/vendor/pytorch3d_compat/transforms/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..32dd82aae06a9a849e4e2ef46c63e659131c15fd --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/transforms/__init__.py @@ -0,0 +1,44 @@ +"""PyTorch3D 0.7.6 rotation conversions used by classic GVHMR inference.""" + +from .rotation_conversions import ( + axis_angle_to_matrix, + euler_angles_to_matrix, + matrix_to_axis_angle, + matrix_to_quaternion, + matrix_to_rotation_6d, + quaternion_to_axis_angle, + quaternion_to_matrix, + rotation_6d_to_matrix, +) + + +def so3_exp_map(log_rot, eps: float = 0.0001): + """Match the PyTorch3D 0.7.6 SO(3) exponential map used by GVHMR.""" + + del eps + if log_rot.ndim != 2 or log_rot.shape[1] != 3: + raise ValueError("Input tensor shape has to be Nx3.") + return axis_angle_to_matrix(log_rot) + + +def so3_log_map(rotation, eps: float = 0.0001, cos_bound: float = 1e-4): + """Match the PyTorch3D 0.7.6 SO(3) logarithm used by GVHMR.""" + + del eps, cos_bound + if rotation.ndim != 3 or rotation.shape[1:] != (3, 3): + raise ValueError("Input has to be a batch of 3x3 Tensors.") + return matrix_to_axis_angle(rotation) + + +__all__ = [ + "axis_angle_to_matrix", + "euler_angles_to_matrix", + "matrix_to_axis_angle", + "matrix_to_quaternion", + "matrix_to_rotation_6d", + "quaternion_to_axis_angle", + "quaternion_to_matrix", + "rotation_6d_to_matrix", + "so3_exp_map", + "so3_log_map", +] diff --git a/fdanyone/vendor/pytorch3d_compat/transforms/rotation_conversions.py b/fdanyone/vendor/pytorch3d_compat/transforms/rotation_conversions.py new file mode 100644 index 0000000000000000000000000000000000000000..459441ca184ff484e252b2b4e4fc86b9b24d4c0e --- /dev/null +++ b/fdanyone/vendor/pytorch3d_compat/transforms/rotation_conversions.py @@ -0,0 +1,596 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from typing import Optional + +import torch +import torch.nn.functional as F + +from ..common.datatypes import Device + + +""" +The transformation matrices returned from the functions in this file assume +the points on which the transformation will be applied are column vectors. +i.e. the R matrix is structured as + + R = [ + [Rxx, Rxy, Rxz], + [Ryx, Ryy, Ryz], + [Rzx, Rzy, Rzz], + ] # (3, 3) + +This matrix can be applied to column vectors by post multiplication +by the points e.g. + + points = [[0], [1], [2]] # (3 x 1) xyz coordinates of a point + transformed_points = R * points + +To apply the same matrix to points which are row vectors, the R matrix +can be transposed and pre multiplied by the points: + +e.g. + points = [[0, 1, 2]] # (1 x 3) xyz coordinates of a point + transformed_points = points * R.transpose(1, 0) +""" + + +def quaternion_to_matrix(quaternions: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as quaternions to rotation matrices. + + Args: + quaternions: quaternions with real part first, + as tensor of shape (..., 4). + + Returns: + Rotation matrices as tensor of shape (..., 3, 3). + """ + r, i, j, k = torch.unbind(quaternions, -1) + # pyre-fixme[58]: `/` is not supported for operand types `float` and `Tensor`. + two_s = 2.0 / (quaternions * quaternions).sum(-1) + + o = torch.stack( + ( + 1 - two_s * (j * j + k * k), + two_s * (i * j - k * r), + two_s * (i * k + j * r), + two_s * (i * j + k * r), + 1 - two_s * (i * i + k * k), + two_s * (j * k - i * r), + two_s * (i * k - j * r), + two_s * (j * k + i * r), + 1 - two_s * (i * i + j * j), + ), + -1, + ) + return o.reshape(quaternions.shape[:-1] + (3, 3)) + + +def _copysign(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """ + Return a tensor where each element has the absolute value taken from the, + corresponding element of a, with sign taken from the corresponding + element of b. This is like the standard copysign floating-point operation, + but is not careful about negative 0 and NaN. + + Args: + a: source tensor. + b: tensor whose signs will be used, of the same shape as a. + + Returns: + Tensor of the same shape as a with the signs of b. + """ + signs_differ = (a < 0) != (b < 0) + return torch.where(signs_differ, -a, a) + + +def _sqrt_positive_part(x: torch.Tensor) -> torch.Tensor: + """ + Returns torch.sqrt(torch.max(0, x)) + but with a zero subgradient where x is 0. + """ + ret = torch.zeros_like(x) + positive_mask = x > 0 + ret[positive_mask] = torch.sqrt(x[positive_mask]) + return ret + + +def matrix_to_quaternion(matrix: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as rotation matrices to quaternions. + + Args: + matrix: Rotation matrices as tensor of shape (..., 3, 3). + + Returns: + quaternions with real part first, as tensor of shape (..., 4). + """ + if matrix.size(-1) != 3 or matrix.size(-2) != 3: + raise ValueError(f"Invalid rotation matrix shape {matrix.shape}.") + + batch_dim = matrix.shape[:-2] + m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind( + matrix.reshape(batch_dim + (9,)), dim=-1 + ) + + q_abs = _sqrt_positive_part( + torch.stack( + [ + 1.0 + m00 + m11 + m22, + 1.0 + m00 - m11 - m22, + 1.0 - m00 + m11 - m22, + 1.0 - m00 - m11 + m22, + ], + dim=-1, + ) + ) + + # we produce the desired quaternion multiplied by each of r, i, j, k + quat_by_rijk = torch.stack( + [ + # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and + # `int`. + torch.stack([q_abs[..., 0] ** 2, m21 - m12, m02 - m20, m10 - m01], dim=-1), + # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and + # `int`. + torch.stack([m21 - m12, q_abs[..., 1] ** 2, m10 + m01, m02 + m20], dim=-1), + # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and + # `int`. + torch.stack([m02 - m20, m10 + m01, q_abs[..., 2] ** 2, m12 + m21], dim=-1), + # pyre-fixme[58]: `**` is not supported for operand types `Tensor` and + # `int`. + torch.stack([m10 - m01, m20 + m02, m21 + m12, q_abs[..., 3] ** 2], dim=-1), + ], + dim=-2, + ) + + # We floor here at 0.1 but the exact level is not important; if q_abs is small, + # the candidate won't be picked. + flr = torch.tensor(0.1).to(dtype=q_abs.dtype, device=q_abs.device) + quat_candidates = quat_by_rijk / (2.0 * q_abs[..., None].max(flr)) + + # if not for numerical problems, quat_candidates[i] should be same (up to a sign), + # forall i; we pick the best-conditioned one (with the largest denominator) + out = quat_candidates[ + F.one_hot(q_abs.argmax(dim=-1), num_classes=4) > 0.5, : + ].reshape(batch_dim + (4,)) + return standardize_quaternion(out) + + +def _axis_angle_rotation(axis: str, angle: torch.Tensor) -> torch.Tensor: + """ + Return the rotation matrices for one of the rotations about an axis + of which Euler angles describe, for each value of the angle given. + + Args: + axis: Axis label "X" or "Y or "Z". + angle: any shape tensor of Euler angles in radians + + Returns: + Rotation matrices as tensor of shape (..., 3, 3). + """ + + cos = torch.cos(angle) + sin = torch.sin(angle) + one = torch.ones_like(angle) + zero = torch.zeros_like(angle) + + if axis == "X": + R_flat = (one, zero, zero, zero, cos, -sin, zero, sin, cos) + elif axis == "Y": + R_flat = (cos, zero, sin, zero, one, zero, -sin, zero, cos) + elif axis == "Z": + R_flat = (cos, -sin, zero, sin, cos, zero, zero, zero, one) + else: + raise ValueError("letter must be either X, Y or Z.") + + return torch.stack(R_flat, -1).reshape(angle.shape + (3, 3)) + + +def euler_angles_to_matrix(euler_angles: torch.Tensor, convention: str) -> torch.Tensor: + """ + Convert rotations given as Euler angles in radians to rotation matrices. + + Args: + euler_angles: Euler angles in radians as tensor of shape (..., 3). + convention: Convention string of three uppercase letters from + {"X", "Y", and "Z"}. + + Returns: + Rotation matrices as tensor of shape (..., 3, 3). + """ + if euler_angles.dim() == 0 or euler_angles.shape[-1] != 3: + raise ValueError("Invalid input euler angles.") + if len(convention) != 3: + raise ValueError("Convention must have 3 letters.") + if convention[1] in (convention[0], convention[2]): + raise ValueError(f"Invalid convention {convention}.") + for letter in convention: + if letter not in ("X", "Y", "Z"): + raise ValueError(f"Invalid letter {letter} in convention string.") + matrices = [ + _axis_angle_rotation(c, e) + for c, e in zip(convention, torch.unbind(euler_angles, -1)) + ] + # return functools.reduce(torch.matmul, matrices) + return torch.matmul(torch.matmul(matrices[0], matrices[1]), matrices[2]) + + +def _angle_from_tan( + axis: str, other_axis: str, data, horizontal: bool, tait_bryan: bool +) -> torch.Tensor: + """ + Extract the first or third Euler angle from the two members of + the matrix which are positive constant times its sine and cosine. + + Args: + axis: Axis label "X" or "Y or "Z" for the angle we are finding. + other_axis: Axis label "X" or "Y or "Z" for the middle axis in the + convention. + data: Rotation matrices as tensor of shape (..., 3, 3). + horizontal: Whether we are looking for the angle for the third axis, + which means the relevant entries are in the same row of the + rotation matrix. If not, they are in the same column. + tait_bryan: Whether the first and third axes in the convention differ. + + Returns: + Euler Angles in radians for each matrix in data as a tensor + of shape (...). + """ + + i1, i2 = {"X": (2, 1), "Y": (0, 2), "Z": (1, 0)}[axis] + if horizontal: + i2, i1 = i1, i2 + even = (axis + other_axis) in ["XY", "YZ", "ZX"] + if horizontal == even: + return torch.atan2(data[..., i1], data[..., i2]) + if tait_bryan: + return torch.atan2(-data[..., i2], data[..., i1]) + return torch.atan2(data[..., i2], -data[..., i1]) + + +def _index_from_letter(letter: str) -> int: + if letter == "X": + return 0 + if letter == "Y": + return 1 + if letter == "Z": + return 2 + raise ValueError("letter must be either X, Y or Z.") + + +def matrix_to_euler_angles(matrix: torch.Tensor, convention: str) -> torch.Tensor: + """ + Convert rotations given as rotation matrices to Euler angles in radians. + + Args: + matrix: Rotation matrices as tensor of shape (..., 3, 3). + convention: Convention string of three uppercase letters. + + Returns: + Euler angles in radians as tensor of shape (..., 3). + """ + if len(convention) != 3: + raise ValueError("Convention must have 3 letters.") + if convention[1] in (convention[0], convention[2]): + raise ValueError(f"Invalid convention {convention}.") + for letter in convention: + if letter not in ("X", "Y", "Z"): + raise ValueError(f"Invalid letter {letter} in convention string.") + if matrix.size(-1) != 3 or matrix.size(-2) != 3: + raise ValueError(f"Invalid rotation matrix shape {matrix.shape}.") + i0 = _index_from_letter(convention[0]) + i2 = _index_from_letter(convention[2]) + tait_bryan = i0 != i2 + if tait_bryan: + central_angle = torch.asin( + matrix[..., i0, i2] * (-1.0 if i0 - i2 in [-1, 2] else 1.0) + ) + else: + central_angle = torch.acos(matrix[..., i0, i0]) + + o = ( + _angle_from_tan( + convention[0], convention[1], matrix[..., i2], False, tait_bryan + ), + central_angle, + _angle_from_tan( + convention[2], convention[1], matrix[..., i0, :], True, tait_bryan + ), + ) + return torch.stack(o, -1) + + +def random_quaternions( + n: int, dtype: Optional[torch.dtype] = None, device: Optional[Device] = None +) -> torch.Tensor: + """ + Generate random quaternions representing rotations, + i.e. versors with nonnegative real part. + + Args: + n: Number of quaternions in a batch to return. + dtype: Type to return. + device: Desired device of returned tensor. Default: + uses the current device for the default tensor type. + + Returns: + Quaternions as tensor of shape (N, 4). + """ + if isinstance(device, str): + device = torch.device(device) + o = torch.randn((n, 4), dtype=dtype, device=device) + s = (o * o).sum(1) + o = o / _copysign(torch.sqrt(s), o[:, 0])[:, None] + return o + + +def random_rotations( + n: int, dtype: Optional[torch.dtype] = None, device: Optional[Device] = None +) -> torch.Tensor: + """ + Generate random rotations as 3x3 rotation matrices. + + Args: + n: Number of rotation matrices in a batch to return. + dtype: Type to return. + device: Device of returned tensor. Default: if None, + uses the current device for the default tensor type. + + Returns: + Rotation matrices as tensor of shape (n, 3, 3). + """ + quaternions = random_quaternions(n, dtype=dtype, device=device) + return quaternion_to_matrix(quaternions) + + +def random_rotation( + dtype: Optional[torch.dtype] = None, device: Optional[Device] = None +) -> torch.Tensor: + """ + Generate a single random 3x3 rotation matrix. + + Args: + dtype: Type to return + device: Device of returned tensor. Default: if None, + uses the current device for the default tensor type + + Returns: + Rotation matrix as tensor of shape (3, 3). + """ + return random_rotations(1, dtype, device)[0] + + +def standardize_quaternion(quaternions: torch.Tensor) -> torch.Tensor: + """ + Convert a unit quaternion to a standard form: one in which the real + part is non negative. + + Args: + quaternions: Quaternions with real part first, + as tensor of shape (..., 4). + + Returns: + Standardized quaternions as tensor of shape (..., 4). + """ + return torch.where(quaternions[..., 0:1] < 0, -quaternions, quaternions) + + +def quaternion_raw_multiply(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """ + Multiply two quaternions. + Usual torch rules for broadcasting apply. + + Args: + a: Quaternions as tensor of shape (..., 4), real part first. + b: Quaternions as tensor of shape (..., 4), real part first. + + Returns: + The product of a and b, a tensor of quaternions shape (..., 4). + """ + aw, ax, ay, az = torch.unbind(a, -1) + bw, bx, by, bz = torch.unbind(b, -1) + ow = aw * bw - ax * bx - ay * by - az * bz + ox = aw * bx + ax * bw + ay * bz - az * by + oy = aw * by - ax * bz + ay * bw + az * bx + oz = aw * bz + ax * by - ay * bx + az * bw + return torch.stack((ow, ox, oy, oz), -1) + + +def quaternion_multiply(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """ + Multiply two quaternions representing rotations, returning the quaternion + representing their composition, i.e. the versor with nonnegative real part. + Usual torch rules for broadcasting apply. + + Args: + a: Quaternions as tensor of shape (..., 4), real part first. + b: Quaternions as tensor of shape (..., 4), real part first. + + Returns: + The product of a and b, a tensor of quaternions of shape (..., 4). + """ + ab = quaternion_raw_multiply(a, b) + return standardize_quaternion(ab) + + +def quaternion_invert(quaternion: torch.Tensor) -> torch.Tensor: + """ + Given a quaternion representing rotation, get the quaternion representing + its inverse. + + Args: + quaternion: Quaternions as tensor of shape (..., 4), with real part + first, which must be versors (unit quaternions). + + Returns: + The inverse, a tensor of quaternions of shape (..., 4). + """ + + scaling = torch.tensor([1, -1, -1, -1], device=quaternion.device) + return quaternion * scaling + + +def quaternion_apply(quaternion: torch.Tensor, point: torch.Tensor) -> torch.Tensor: + """ + Apply the rotation given by a quaternion to a 3D point. + Usual torch rules for broadcasting apply. + + Args: + quaternion: Tensor of quaternions, real part first, of shape (..., 4). + point: Tensor of 3D points of shape (..., 3). + + Returns: + Tensor of rotated points of shape (..., 3). + """ + if point.size(-1) != 3: + raise ValueError(f"Points are not in 3D, {point.shape}.") + real_parts = point.new_zeros(point.shape[:-1] + (1,)) + point_as_quaternion = torch.cat((real_parts, point), -1) + out = quaternion_raw_multiply( + quaternion_raw_multiply(quaternion, point_as_quaternion), + quaternion_invert(quaternion), + ) + return out[..., 1:] + + +def axis_angle_to_matrix(axis_angle: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as axis/angle to rotation matrices. + + Args: + axis_angle: Rotations given as a vector in axis angle form, + as a tensor of shape (..., 3), where the magnitude is + the angle turned anticlockwise in radians around the + vector's direction. + + Returns: + Rotation matrices as tensor of shape (..., 3, 3). + """ + return quaternion_to_matrix(axis_angle_to_quaternion(axis_angle)) + + +def matrix_to_axis_angle(matrix: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as rotation matrices to axis/angle. + + Args: + matrix: Rotation matrices as tensor of shape (..., 3, 3). + + Returns: + Rotations given as a vector in axis angle form, as a tensor + of shape (..., 3), where the magnitude is the angle + turned anticlockwise in radians around the vector's + direction. + """ + return quaternion_to_axis_angle(matrix_to_quaternion(matrix)) + + +def axis_angle_to_quaternion(axis_angle: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as axis/angle to quaternions. + + Args: + axis_angle: Rotations given as a vector in axis angle form, + as a tensor of shape (..., 3), where the magnitude is + the angle turned anticlockwise in radians around the + vector's direction. + + Returns: + quaternions with real part first, as tensor of shape (..., 4). + """ + angles = torch.norm(axis_angle, p=2, dim=-1, keepdim=True) + half_angles = angles * 0.5 + eps = 1e-6 + small_angles = angles.abs() < eps + sin_half_angles_over_angles = torch.empty_like(angles) + sin_half_angles_over_angles[~small_angles] = ( + torch.sin(half_angles[~small_angles]) / angles[~small_angles] + ) + # for x small, sin(x/2) is about x/2 - (x/2)^3/6 + # so sin(x/2)/x is about 1/2 - (x*x)/48 + sin_half_angles_over_angles[small_angles] = ( + 0.5 - (angles[small_angles] * angles[small_angles]) / 48 + ) + quaternions = torch.cat( + [torch.cos(half_angles), axis_angle * sin_half_angles_over_angles], dim=-1 + ) + return quaternions + + +def quaternion_to_axis_angle(quaternions: torch.Tensor) -> torch.Tensor: + """ + Convert rotations given as quaternions to axis/angle. + + Args: + quaternions: quaternions with real part first, + as tensor of shape (..., 4). + + Returns: + Rotations given as a vector in axis angle form, as a tensor + of shape (..., 3), where the magnitude is the angle + turned anticlockwise in radians around the vector's + direction. + """ + norms = torch.norm(quaternions[..., 1:], p=2, dim=-1, keepdim=True) + half_angles = torch.atan2(norms, quaternions[..., :1]) + angles = 2 * half_angles + eps = 1e-6 + small_angles = angles.abs() < eps + sin_half_angles_over_angles = torch.empty_like(angles) + sin_half_angles_over_angles[~small_angles] = ( + torch.sin(half_angles[~small_angles]) / angles[~small_angles] + ) + # for x small, sin(x/2) is about x/2 - (x/2)^3/6 + # so sin(x/2)/x is about 1/2 - (x*x)/48 + sin_half_angles_over_angles[small_angles] = ( + 0.5 - (angles[small_angles] * angles[small_angles]) / 48 + ) + return quaternions[..., 1:] / sin_half_angles_over_angles + + +def rotation_6d_to_matrix(d6: torch.Tensor) -> torch.Tensor: + """ + Converts 6D rotation representation by Zhou et al. [1] to rotation matrix + using Gram--Schmidt orthogonalization per Section B of [1]. + Args: + d6: 6D rotation representation, of size (*, 6) + + Returns: + batch of rotation matrices of size (*, 3, 3) + + [1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H. + On the Continuity of Rotation Representations in Neural Networks. + IEEE Conference on Computer Vision and Pattern Recognition, 2019. + Retrieved from http://arxiv.org/abs/1812.07035 + """ + + a1, a2 = d6[..., :3], d6[..., 3:] + b1 = F.normalize(a1, dim=-1) + b2 = a2 - (b1 * a2).sum(-1, keepdim=True) * b1 + b2 = F.normalize(b2, dim=-1) + b3 = torch.cross(b1, b2, dim=-1) + return torch.stack((b1, b2, b3), dim=-2) + + +def matrix_to_rotation_6d(matrix: torch.Tensor) -> torch.Tensor: + """ + Converts rotation matrices to 6D rotation representation by Zhou et al. [1] + by dropping the last row. Note that 6D representation is not unique. + Args: + matrix: batch of rotation matrices of size (*, 3, 3) + + Returns: + 6D rotation representation, of size (*, 6) + + [1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H. + On the Continuity of Rotation Representations in Neural Networks. + IEEE Conference on Computer Vision and Pattern Recognition, 2019. + Retrieved from http://arxiv.org/abs/1812.07035 + """ + batch_dim = matrix.size()[:-2] + return matrix[..., :2, :].clone().reshape(batch_dim + (6,)) diff --git a/fdanyone/video.py b/fdanyone/video.py new file mode 100644 index 0000000000000000000000000000000000000000..7867e48fb7961c36e1fbfc7c592a6fa414b67c22 --- /dev/null +++ b/fdanyone/video.py @@ -0,0 +1,540 @@ +"""PTS-aware canonical video decoding and encoding.""" + +from __future__ import annotations + +import itertools +import json +import math +from collections.abc import Iterator +from concurrent.futures import Executor, Future +from dataclasses import dataclass +from fractions import Fraction +from pathlib import Path + +import av +import numpy as np + +from fdanyone.config import INFERENCE +from fdanyone.errors import VideoContractError + +AUTO_DOWNSAMPLE_FPS = tuple( + Fraction(numerator, denominator) for numerator, denominator in INFERENCE.auto_downsample_fps +) +SAMPLING_POLICY = INFERENCE.temporal_sampling_policy +INTEGER_RATIO_TOLERANCE = 1e-3 + + +def validate_required_video_codecs() -> None: + missing = [] + for codec in ("libx264", "libx264rgb"): + try: + av.codec.Codec(codec, "w") + except (ValueError, av.FFmpegError): + missing.append(codec) + if missing: + raise VideoContractError( + f"The installed FFmpeg/PyAV build lacks required video encoders: {missing}. " + "Install a PyAV build with libx264 and libx264rgb support." + ) + + +def choose_canonical_fps(input_rate: Fraction) -> Fraction: + """Keep the source rate unless it supports clean integer downsampling. + + High-frame-rate sources may be reduced to an exact 24/25/30-family + divisor (48 -> 24, 50 -> 25, 60 -> 30, 120 -> 30). Rates without such a + divisor, including 40 FPS, are preserved exactly. + """ + + if input_rate <= 0: + raise VideoContractError(f"Input frame rate must be positive, got {input_rate}.") + + divisible: list[tuple[Fraction, int]] = [] + for candidate in AUTO_DOWNSAMPLE_FPS: + ratio = float(input_rate / candidate) + multiple = round(ratio) + if multiple >= 2 and abs(ratio - multiple) <= INTEGER_RATIO_TOLERANCE: + divisible.append((candidate, multiple)) + if divisible: + _, multiple = max(divisible, key=lambda match: float(match[0])) + return input_rate / multiple + + return input_rate + + +def _stream_rate(stream: av.video.stream.VideoStream) -> Fraction: + for value in (stream.average_rate, stream.guessed_rate, stream.base_rate): + if value is not None and value > 0: + return Fraction(value.numerator, value.denominator) + raise VideoContractError("The input video does not declare a usable frame rate.") + + +def _normalize_rotation_degrees(raw: object) -> int: + """Normalize a display-matrix rotation to a supported CCW quarter turn.""" + + try: + rotation = int(round(float(raw))) % 360 + except (TypeError, ValueError): + rotation = 0 + if rotation not in (0, 90, 180, 270): + raise VideoContractError(f"Unsupported video rotation metadata: {raw!r} degrees.") + return rotation + + +def _rotation_degrees(stream: av.video.stream.VideoStream) -> int: + return _normalize_rotation_degrees(stream.metadata.get("rotate", "0")) + + +def _frame_rotation_degrees(frame: av.VideoFrame, metadata_rotation: int) -> int: + """Prefer FFmpeg display-matrix side data over the legacy rotate tag.""" + + frame_rotation = getattr(frame, "rotation", 0) + if frame_rotation is None or float(frame_rotation) == 0.0: + return metadata_rotation + return _normalize_rotation_degrees(frame_rotation) + + +def _frame_timestamp(frame: av.VideoFrame, stream: av.video.stream.VideoStream, index: int) -> Fraction: + if frame.pts is not None and frame.time_base is not None: + return Fraction(frame.pts * frame.time_base) + rate = _stream_rate(stream) + return Fraction(index, 1) / rate + + +def _frame_to_rgb(frame: av.VideoFrame, rotation_degrees: int) -> np.ndarray: + rgb = frame.to_ndarray(format="rgb24") + if rotation_degrees == 90: + rgb = np.rot90(rgb, k=1) + elif rotation_degrees == 180: + rgb = np.rot90(rgb, k=2) + elif rotation_degrees == 270: + rgb = np.rot90(rgb, k=3) + return np.ascontiguousarray(rgb) + + +@dataclass(frozen=True) +class CanonicalFrame: + rgb: np.ndarray + source_index: int + source_pts: int | None + source_timestamp: Fraction + canonical_timestamp: Fraction + + +@dataclass(frozen=True) +class CanonicalClip: + source_path: Path + source_size_bytes: int + source_mtime_ns: int + fps: Fraction + input_rate: Fraction + source_time_base: Fraction + start_time: Fraction + frames: tuple[CanonicalFrame, ...] + rotation_degrees: int + + @property + def height(self) -> int: + return int(self.frames[0].rgb.shape[0]) + + @property + def width(self) -> int: + return int(self.frames[0].rgb.shape[1]) + + @property + def fps_num(self) -> int: + return self.fps.numerator + + @property + def fps_den(self) -> int: + return self.fps.denominator + + @property + def rgb_frames(self) -> tuple[np.ndarray, ...]: + return tuple(frame.rgb for frame in self.frames) + + def metadata(self) -> dict: + return { + # The subprocess boundary only needs source identity, not a + # machine-specific absolute path. Keep the file-protocol payload + # safe to share by storing the basename. + "source_path": self.source_path.name, + "source_size_bytes": self.source_size_bytes, + "source_mtime_ns": self.source_mtime_ns, + "fps_num": self.fps_num, + "fps_den": self.fps_den, + "input_rate_num": self.input_rate.numerator, + "input_rate_den": self.input_rate.denominator, + "source_time_base_num": self.source_time_base.numerator, + "source_time_base_den": self.source_time_base.denominator, + "start_time_num": self.start_time.numerator, + "start_time_den": self.start_time.denominator, + "sampling_policy": SAMPLING_POLICY, + "num_frames": len(self.frames), + "height": self.height, + "width": self.width, + "rotation_degrees_applied": self.rotation_degrees, + "frames": [ + { + "canonical_index": index, + "canonical_timestamp_sec": float(frame.canonical_timestamp), + "source_index": frame.source_index, + "source_pts": frame.source_pts, + "source_timestamp_sec": float(frame.source_timestamp), + } + for index, frame in enumerate(self.frames) + ], + } + + def write_metadata(self, path: str | Path) -> None: + Path(path).write_text(json.dumps(self.metadata(), indent=2, sort_keys=True) + "\n") + + +@dataclass(frozen=True) +class _DecodedFrame: + rgb: np.ndarray + index: int + pts: int | None + timestamp: Fraction + time_base: Fraction + rotation_degrees: int + + +def _decode_frames( + container: av.container.InputContainer, + stream: av.video.stream.VideoStream, + rotation: int, +) -> Iterator[_DecodedFrame]: + observed_rotation = None + for index, frame in enumerate(container.decode(stream)): + frame_rotation = _frame_rotation_degrees(frame, rotation) + if observed_rotation is None: + observed_rotation = frame_rotation + elif frame_rotation != observed_rotation: + raise VideoContractError( + f"Video display rotation changes at frame {index}: {observed_rotation} -> {frame_rotation}." + ) + yield _DecodedFrame( + rgb=_frame_to_rgb(frame, frame_rotation), + index=index, + pts=frame.pts, + timestamp=_frame_timestamp(frame, stream, index), + time_base=( + Fraction(frame.time_base) + if frame.time_base is not None + else Fraction(stream.time_base) + if stream.time_base is not None + else Fraction(1, 1) / _stream_rate(stream) + ), + rotation_degrees=frame_rotation, + ) + + +def decode_canonical_clip( + video_path: str | Path, + *, + num_frames: int = 121, + start_time: float = 0.0, + fps: str | int | float | Fraction | None = None, +) -> CanonicalClip: + """Decode one canonical clip, selecting frames by source presentation time.""" + + path = Path(video_path).expanduser().resolve() + if not path.is_file(): + raise VideoContractError(f"Input video does not exist: {path}") + if num_frames <= 0: + raise VideoContractError(f"num_frames must be positive, got {num_frames}.") + if not math.isfinite(start_time) or start_time < 0: + raise VideoContractError(f"start_time must be non-negative, got {start_time}.") + + source_stat = path.stat() + + with av.open(str(path), mode="r") as container: + if not container.streams.video: + raise VideoContractError(f"Input has no video stream: {path}") + stream = container.streams.video[0] + input_rate = _stream_rate(stream) + try: + output_rate = choose_canonical_fps(input_rate) if fps is None else Fraction(str(fps)) + except (ValueError, ZeroDivisionError) as exc: + raise VideoContractError(f"Cannot parse target fps {fps!r} as a rational frame rate.") from exc + if output_rate <= 0: + raise VideoContractError(f"Target fps must be positive, got {output_rate}.") + metadata_rotation = _rotation_degrees(stream) + decoded = _decode_frames(container, stream, metadata_rotation) + try: + previous = next(decoded) + except StopIteration as exc: + raise VideoContractError(f"Input video contains no decodable frames: {path}") from exc + + origin = previous.timestamp + start_offset = Fraction(str(start_time)) + start = origin + start_offset + targets = [start + Fraction(index, 1) / output_rate for index in range(num_frames)] + selected: list[_DecodedFrame] = [] + target_index = 0 + current = previous + + for current in decoded: + if current.timestamp < previous.timestamp: + raise VideoContractError( + f"Input presentation timestamps are not monotonic at source frame {current.index}: " + f"{float(current.timestamp):.6f}s < {float(previous.timestamp):.6f}s." + ) + while target_index < num_frames and targets[target_index] <= current.timestamp: + target = targets[target_index] + candidate = previous if abs(previous.timestamp - target) <= abs(current.timestamp - target) else current + if selected and candidate.index < selected[-1].index: + raise VideoContractError("Temporal sampling produced non-monotonic source-frame order.") + selected.append(candidate) + target_index += 1 + if target_index >= num_frames: + break + previous = current + + max_error = Fraction(3, 4) / output_rate + if target_index < num_frames: + # A faster canonical clock can legitimately select the final source + # frame more than once, just as it may reuse frames in the middle of + # the clip. Stop once the nearest-frame error exceeds the same gap + # bound enforced below; that remains a short-input failure rather + # than temporal padding. + while target_index < num_frames and abs(current.timestamp - targets[target_index]) <= max_error: + if selected and current.index < selected[-1].index: + raise VideoContractError("Temporal sampling produced non-monotonic source-frame order.") + selected.append(current) + target_index += 1 + + if len(selected) != num_frames: + duration = float(current.timestamp - origin) + required = float(Fraction(num_frames - 1, 1) / output_rate + start_offset) + raise VideoContractError( + f"Input is too short for {num_frames} frames at {float(output_rate):.6f} FPS from " + f"start_time={start_time}: decoded duration={duration:.3f}s, required={required:.3f}s." + ) + + errors = [abs(frame.timestamp - target) for frame, target in zip(selected, targets, strict=True)] + if max(errors) > max_error: + raise VideoContractError( + "Input timestamps contain a gap too large for stable sampling: " + f"max error={float(max(errors)):.6f}s, limit={float(max_error):.6f}s." + ) + expected_shape = selected[0].rgb.shape + if any(frame.rgb.shape != expected_shape for frame in selected[1:]): + raise VideoContractError("Input frame dimensions change inside the selected canonical clip.") + source_time_base = selected[0].time_base + if any(frame.time_base != source_time_base for frame in selected[1:]): + raise VideoContractError("Input frame time base changes inside the selected canonical clip.") + + canonical_frames = tuple( + CanonicalFrame( + rgb=frame.rgb, + source_index=frame.index, + source_pts=frame.pts, + source_timestamp=frame.timestamp, + canonical_timestamp=Fraction(index, 1) / output_rate, + ) + for index, frame in enumerate(selected) + ) + + final_stat = path.stat() + if (final_stat.st_size, final_stat.st_mtime_ns) != (source_stat.st_size, source_stat.st_mtime_ns): + raise VideoContractError(f"Input video changed while it was being decoded: {path}") + return CanonicalClip( + source_path=path, + source_size_bytes=source_stat.st_size, + source_mtime_ns=source_stat.st_mtime_ns, + fps=output_rate, + input_rate=input_rate, + source_time_base=source_time_base, + start_time=start_offset, + frames=canonical_frames, + rotation_degrees=selected[0].rotation_degrees, + ) + + +def write_lossless_video(clip: CanonicalClip, path: str | Path) -> Path: + """Write an FFV1 working video. A subsequent decode must be RGB-identical.""" + + output_path = Path(path).expanduser().resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + with av.open(str(output_path), mode="w", format="matroska") as container: + stream = container.add_stream("ffv1", rate=clip.fps) + stream.width = clip.width + stream.height = clip.height + # FFV1 does not expose 8-bit GBR planar. BGR0 is an exact 8-bit RGB + # representation (the fourth byte is padding) and round-trips to rgb24. + stream.pix_fmt = "bgr0" + for index, canonical_frame in enumerate(clip.frames): + frame = av.VideoFrame.from_ndarray(canonical_frame.rgb, format="rgb24") + frame.pts = index + frame.time_base = Fraction(1, 1) / clip.fps + for packet in stream.encode(frame): + container.mux(packet) + for packet in stream.encode(): + container.mux(packet) + verify_lossless_video(clip, output_path) + return output_path + + +def write_gvhmr_video(clip: CanonicalClip, path: str | Path) -> Path: + """Write the frame-counted, RGB-lossless MP4 consumed by GVHMR. + + GVHMR's imageio metadata probe cannot determine the frame count of an + FFV1 Matroska stream. ``libx264rgb`` in lossless mode preserves every RGB + byte while placing an exact frame count in the MP4 index. + """ + + output_path = Path(path).expanduser().resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + with av.open(str(output_path), mode="w") as container: + stream = container.add_stream("libx264rgb", rate=clip.fps) + stream.width = clip.width + stream.height = clip.height + stream.pix_fmt = "rgb24" + stream.options = {"crf": "0", "preset": "medium"} + for index, canonical_frame in enumerate(clip.frames): + frame = av.VideoFrame.from_ndarray(canonical_frame.rgb, format="rgb24") + frame.pts = index + frame.time_base = Fraction(1, 1) / clip.fps + for packet in stream.encode(frame): + container.mux(packet) + for packet in stream.encode(): + container.mux(packet) + verify_lossless_video(clip, output_path) + with av.open(str(output_path), mode="r") as container: + declared_frames = int(container.streams.video[0].frames) + if declared_frames != len(clip.frames): + raise VideoContractError( + f"Backend MP4 declares {declared_frames} frames, expected {len(clip.frames)}: {output_path}." + ) + return output_path + + +def verify_lossless_video(clip: CanonicalClip, path: str | Path) -> None: + decoded = iter_rgb_video(path) + sentinel = object() + for index, (actual, expected) in enumerate(itertools.zip_longest(decoded, clip.rgb_frames, fillvalue=sentinel)): + if actual is sentinel or expected is sentinel or not np.array_equal(actual, expected): + raise VideoContractError(f"Lossless working-video verification failed at frame {index}.") + + +def write_video( + frames: Iterator[np.ndarray] | tuple[np.ndarray, ...], + path: str | Path, + fps: Fraction, + *, + crf: int = 18, + preset: str = "medium", +) -> Path: + """Encode RGB frames as a broadly playable H.264 MP4.""" + + output_path = Path(path).expanduser().resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + iterator = iter(frames) + try: + first = next(iterator) + except StopIteration as exc: + raise VideoContractError("Cannot encode an empty video.") from exc + height, width = first.shape[:2] + with av.open(str(output_path), mode="w") as container: + stream = container.add_stream("libx264", rate=fps) + stream.width = width + stream.height = height + stream.pix_fmt = "yuv420p" + stream.options = {"crf": str(crf), "preset": preset} + for index, rgb in enumerate(itertools.chain((first,), iterator)): + if rgb.shape != first.shape: + raise VideoContractError(f"Frame {index} has shape {rgb.shape}, expected {first.shape}.") + frame = av.VideoFrame.from_ndarray(np.ascontiguousarray(rgb), format="rgb24") + frame.pts = index + frame.time_base = Fraction(1, 1) / fps + for packet in stream.encode(frame): + container.mux(packet) + for packet in stream.encode(): + container.mux(packet) + return output_path + + +def write_video_async( + executor: Executor, + frames: Iterator[np.ndarray] | tuple[np.ndarray, ...], + path: str | Path, + fps: Fraction, + *, + crf: int = 18, + preset: str = "medium", +) -> Future[Path]: + """Materialize frames on the caller, then submit the CPU-only H.264 tail.""" + + materialized: tuple[np.ndarray, ...] = tuple(frames) + return executor.submit( + write_video, + materialized, + path, + fps, + crf=crf, + preset=preset, + ) + + +def iter_rgb_video(path: str | Path) -> Iterator[np.ndarray]: + """Stream RGB frames without interpreting or resampling timestamps.""" + + video_path = Path(path).expanduser().resolve() + with av.open(str(video_path), mode="r") as container: + if not container.streams.video: + raise VideoContractError(f"Video has no stream: {video_path}") + for frame in container.decode(container.streams.video[0]): + yield np.ascontiguousarray(frame.to_ndarray(format="rgb24")) + + +def read_rgb_video(path: str | Path, *, expected_frames: int | None = None) -> tuple[np.ndarray, ...]: + """Decode an RGB video without interpreting or resampling its timestamps.""" + + frames = tuple(iter_rgb_video(path)) + if expected_frames is not None and len(frames) != expected_frames: + raise VideoContractError( + f"{Path(path).expanduser().resolve()} has {len(frames)} frames, expected {expected_frames}." + ) + return frames + + +def load_canonical_working_clip(video_path: str | Path, metadata_path: str | Path) -> CanonicalClip: + """Rehydrate the canonical clip inside a short-lived worker.""" + + video_path = Path(video_path).expanduser().resolve() + metadata = json.loads(Path(metadata_path).read_text()) + if metadata.get("sampling_policy") != SAMPLING_POLICY: + raise VideoContractError(f"Unsupported canonical sampling policy: {metadata.get('sampling_policy')!r}.") + records = metadata["frames"] + decoded = read_rgb_video(video_path, expected_frames=int(metadata["num_frames"])) + if len(records) != len(decoded): + raise VideoContractError( + f"Canonical metadata has {len(records)} frame records, but {video_path} has {len(decoded)} frames." + ) + frames = [] + for canonical_index, (rgb, record) in enumerate(zip(decoded, records, strict=True)): + if int(record["canonical_index"]) != canonical_index: + raise VideoContractError("Canonical subprocess metadata is not in frame-index order.") + frames.append( + CanonicalFrame( + rgb=rgb, + source_index=int(record["source_index"]), + source_pts=None if record["source_pts"] is None else int(record["source_pts"]), + source_timestamp=Fraction(str(record["source_timestamp_sec"])), + canonical_timestamp=Fraction(canonical_index, 1) + / Fraction(int(metadata["fps_num"]), int(metadata["fps_den"])), + ) + ) + return CanonicalClip( + source_path=Path(metadata["source_path"]), + source_size_bytes=int(metadata["source_size_bytes"]), + source_mtime_ns=int(metadata["source_mtime_ns"]), + fps=Fraction(int(metadata["fps_num"]), int(metadata["fps_den"])), + input_rate=Fraction(int(metadata["input_rate_num"]), int(metadata["input_rate_den"])), + source_time_base=Fraction(int(metadata["source_time_base_num"]), int(metadata["source_time_base_den"])), + start_time=Fraction(int(metadata["start_time_num"]), int(metadata["start_time_den"])), + frames=tuple(frames), + rotation_degrees=int(metadata["rotation_degrees_applied"]), + ) diff --git a/fdanyone/views.py b/fdanyone/views.py new file mode 100644 index 0000000000000000000000000000000000000000..bc260d4c67c2bbdb14566874169888e0a32d0826 --- /dev/null +++ b/fdanyone/views.py @@ -0,0 +1,204 @@ +"""Resolve the reader-facing target-view layout and inference grouping.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass + +from fdanyone.config import CAMERA +from fdanyone.errors import ConfigurationError + +VALID_VIEWS_PER_GROUP = (4, 6) +MIN_PITCH = -15 +MAX_PITCH = 45 + +# RCP uses the canonical proposal cameras seen during training. The resolved +# group size selects a prefix; at most the first four become target references. +RCP_CAMERA_ORDER = (4, 9, 14, 19, 0, 12) + + +@dataclass(frozen=True) +class TargetView: + """One requested camera in the layer-major public view order.""" + + camera_id: int + layer_index: int + pitch: int + yaw: float + + +@dataclass(frozen=True) +class ViewPlan: + """Validated target cameras and method components for one run.""" + + views_per_layer: int + layer_pitches: tuple[int, ...] + start_yaw: int + yaw_span: int + views_per_group: int + enable_rcp: bool + enable_tcr: bool + + @property + def num_layers(self) -> int: + return len(self.layer_pitches) + + @property + def num_target_views(self) -> int: + return self.views_per_layer * self.num_layers + + @property + def groups_per_layer(self) -> int: + return self.views_per_layer // self.views_per_group + + @property + def num_groups(self) -> int: + return self.groups_per_layer * self.num_layers + + @property + def closed_yaw(self) -> bool: + return self.yaw_span == 360 + + @property + def tcr_active(self) -> bool: + return self.enable_tcr + + @property + def target_views(self) -> tuple[TargetView, ...]: + step = self.yaw_span / self.views_per_layer + return tuple( + TargetView( + camera_id=layer_index * self.views_per_layer + view_index, + layer_index=layer_index, + pitch=pitch, + yaw=self.start_yaw + view_index * step, + ) + for layer_index, pitch in enumerate(self.layer_pitches) + for view_index in range(self.views_per_layer) + ) + + @property + def front_camera_ids(self) -> tuple[int, ...]: + return tuple(view.camera_id for view in self.target_views if abs(view.yaw % 360.0) < 1e-8) + + @property + def rcp_camera_ids(self) -> tuple[int, ...]: + return RCP_CAMERA_ORDER[: self.views_per_group] if self.enable_rcp else () + + @property + def is_canonical_target_ring(self) -> bool: + return ( + self.views_per_layer == CAMERA.count + and self.layer_pitches == (int(CAMERA.pitch_degrees),) + and self.start_yaw == 0 + and self.yaw_span == 360 + ) + + def to_dict(self) -> dict[str, object]: + return { + "views_per_layer": self.views_per_layer, + "layer_pitches": list(self.layer_pitches), + "start_yaw": self.start_yaw, + "yaw_span": self.yaw_span, + "views_per_group": self.views_per_group, + "enable_rcp": self.enable_rcp, + "enable_tcr": self.enable_tcr, + } + + @classmethod + def from_dict(cls, value: object) -> ViewPlan: + if not isinstance(value, dict): + raise ConfigurationError("View plan must be a JSON object.") + try: + return resolve_view_plan( + views_per_layer=value["views_per_layer"], + layer_pitches=value["layer_pitches"], + start_yaw=value["start_yaw"], + yaw_span=value["yaw_span"], + views_per_group=value["views_per_group"], + enable_rcp=value["enable_rcp"], + enable_tcr=value["enable_tcr"], + ) + except KeyError as exc: + raise ConfigurationError(f"View plan is missing {exc.args[0]!r}.") from None + + +def _integer(name: str, value: object) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise ConfigurationError(f"{name} must be an integer, got {value!r}.") + return value + + +def _layer_pitches(value: object) -> tuple[int, ...]: + if isinstance(value, (str, bytes)) or not isinstance(value, Sequence) or not value: + raise ConfigurationError("layer_pitches must be a non-empty list of integer degrees.") + pitches = tuple(_integer("Each layer pitch", pitch) for pitch in value) + if len(set(pitches)) != len(pitches): + raise ConfigurationError(f"layer_pitches must not contain duplicates, got {list(pitches)}.") + invalid = [pitch for pitch in pitches if not MIN_PITCH <= pitch <= MAX_PITCH] + if invalid: + raise ConfigurationError( + f"Each layer pitch must be between {MIN_PITCH} and {MAX_PITCH} degrees, got {invalid}." + ) + return pitches + + +def _group_size(value: int | str, views_per_layer: int) -> int: + if isinstance(value, str): + if value.lower() == "auto": + divisors = tuple(size for size in VALID_VIEWS_PER_GROUP if views_per_layer % size == 0) + if not divisors: + raise ConfigurationError(f"views_per_layer ({views_per_layer}) must be divisible by 4 or 6.") + return max(divisors) + try: + value = int(value) + except ValueError: + raise ConfigurationError( + f"views_per_group must be 'auto' or one of {VALID_VIEWS_PER_GROUP}, got {value!r}." + ) from None + value = _integer("views_per_group", value) + if value not in VALID_VIEWS_PER_GROUP: + raise ConfigurationError(f"views_per_group must be one of {VALID_VIEWS_PER_GROUP}, got {value!r}.") + if views_per_layer % value: + raise ConfigurationError(f"views_per_layer ({views_per_layer}) must be divisible by views_per_group ({value}).") + return value + + +def resolve_view_plan( + *, + views_per_layer: int = 24, + layer_pitches: Sequence[int] = (15,), + start_yaw: int = 0, + yaw_span: int = 360, + views_per_group: int | str = "auto", + enable_rcp: bool = True, + enable_tcr: bool = True, +) -> ViewPlan: + """Validate the compact CLI settings before expensive work starts.""" + + views_per_layer = _integer("views_per_layer", views_per_layer) + if views_per_layer <= 0: + raise ConfigurationError(f"views_per_layer must be positive, got {views_per_layer}.") + pitches = _layer_pitches(layer_pitches) + start_yaw = _integer("start_yaw", start_yaw) + start_yaw = (start_yaw + 180) % 360 - 180 + yaw_span = _integer("yaw_span", yaw_span) + if not 0 < yaw_span <= 360: + raise ConfigurationError(f"yaw_span must be between 1 and 360 degrees, got {yaw_span}.") + resolved_group_size = _group_size(views_per_group, views_per_layer) + if not isinstance(enable_rcp, bool): + raise ConfigurationError(f"enable_rcp must be True or False, got {enable_rcp!r}.") + if not isinstance(enable_tcr, bool): + raise ConfigurationError(f"enable_tcr must be True or False, got {enable_tcr!r}.") + + # Up to six requested targets are cheaper and clearer to generate directly. + rcp_active = enable_rcp and views_per_layer * len(pitches) > 6 + return ViewPlan( + views_per_layer=views_per_layer, + layer_pitches=pitches, + start_yaw=start_yaw, + yaw_span=yaw_span, + views_per_group=resolved_group_size, + enable_rcp=rcp_active, + enable_tcr=enable_tcr, + ) diff --git a/fdanyone/viz.py b/fdanyone/viz.py new file mode 100644 index 0000000000000000000000000000000000000000..19ef209ff0d6b87549f85cfdbb903f6405f93e65 --- /dev/null +++ b/fdanyone/viz.py @@ -0,0 +1,22 @@ +"""Rerun names shared by every 4DAnyone recording. + +The ``rerun`` import stays function-local so that importing this module never +drags the Viewer SDK into an environment that only needs the names. +""" + +from __future__ import annotations + +from fractions import Fraction + +APPLICATION_ID = "4danyone" +FRAME_TIMELINE = "frame" +TIME_TIMELINE = "time" + + +def set_frame(index: int, fps: Fraction) -> None: + """Put the next log calls on the shared frame and seconds timelines.""" + + import rerun as rr + + rr.set_time(FRAME_TIMELINE, sequence=index) + rr.set_time(TIME_TIMELINE, duration=float(Fraction(index, 1) / fps)) diff --git a/sync_vendor.sh b/sync_vendor.sh new file mode 100755 index 0000000000000000000000000000000000000000..c87a999f22da0c1d36cc737a8f1bc2d2929471c6 --- /dev/null +++ b/sync_vendor.sh @@ -0,0 +1,112 @@ +#!/usr/bin/env bash +# Copy the inference-path subset of the fdanyone package into this Space. +# +# Reconstruction (nerfstudio, FreeTimeGS) and the Fire-based script shim are not +# part of the Space, so they stay out of the tree and out of pixi.toml. +# +# Usage: ./sync_vendor.sh [source-checkout] +set -euo pipefail + +SOURCE="${1:-$HOME/0Dev/personal/4DAnyone-5090}" +DEST="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +if [[ ! -d "$SOURCE/fdanyone" ]]; then + echo "error: no fdanyone package under $SOURCE" >&2 + exit 1 +fi + +REVISION="$(git -C "$SOURCE" rev-parse HEAD)" +BRANCH="$(git -C "$SOURCE" rev-parse --abbrev-ref HEAD)" +GVHMR_REVISION="$(git -C "$SOURCE" submodule status third_party/GVHMR | sed 's/^[-+ ]//;s/ .*//')" + +if [[ -n "$(git -C "$SOURCE" status --porcelain -- fdanyone)" ]]; then + echo "error: $SOURCE/fdanyone has uncommitted changes; the recorded SHA would be a lie" >&2 + exit 1 +fi + +rsync -a --delete \ + --exclude '__pycache__/' \ + --exclude 'freetimegs/' \ + --exclude 'nerfstudio/' \ + --exclude 'vendor/freetimegs/' \ + --exclude 'cli.py' \ + "$SOURCE/fdanyone/" "$DEST/fdanyone/" + +# The Space ships the exported prompt embedding instead of the 11 GB UMT5-XXL +# encoder, and it never runs reconstruction. prepare_run() calls ensure_models(), +# which downloads every entry of MODEL_FILES inside the GPU allocation, so the +# encoder, its tokenizer, and the perceptual VGG-19 must leave that set. +python3 - "$DEST/fdanyone/assets.py" <<'PATCH' +import sys +from pathlib import Path + +path = Path(sys.argv[1]) +before = """MODEL_FILES = ( + CHECKPOINT, + MHR70_REGRESSOR, + WAN_VAE, + TEXT_ENCODER, + *TOKENIZER_FILES, + GVHMR_CHECKPOINT, + HMR2_CHECKPOINT, + VITPOSE_CHECKPOINT, + YOLO_CHECKPOINT, + PERCEPTUAL_VGG19, +)""" +after = """# Space patch (sync_vendor.sh): the exported prompt embedding replaces the +# UMT5-XXL encoder and its tokenizer, and reconstruction never runs here. +MODEL_FILES = ( + CHECKPOINT, + MHR70_REGRESSOR, + WAN_VAE, + GVHMR_CHECKPOINT, + HMR2_CHECKPOINT, + VITPOSE_CHECKPOINT, + YOLO_CHECKPOINT, +)""" +source = path.read_text() +if before not in source: + raise SystemExit(f"error: MODEL_FILES in {path} no longer matches the patched form") +path.write_text(source.replace(before, after)) +PATCH + +cat > "$DEST/PROVENANCE.md" < | +| Branch | \`$BRANCH\` | +| Commit | \`$REVISION\` | +| Synced | $(date -u +%Y-%m-%dT%H:%M:%SZ) | + +## Excluded from the copy + +- \`fdanyone/nerfstudio/\`, \`fdanyone/freetimegs/\`, \`fdanyone/vendor/freetimegs/\` — 3DGS + and 4DGS reconstruction, which this Space does not run. +- \`fdanyone/cli.py\` — the Fire shim for \`scripts/\`, which is not copied either. +- \`__pycache__/\`. + +## Patched in the copy + +- \`fdanyone/assets.py\`: \`MODEL_FILES\` drops \`models_t5_umt5-xxl-enc-bf16.pth\`, + \`4danyone/umt5-xxl/\`, and the perceptual VGG-19. \`prepare_run\` calls + \`ensure_models\`, which downloads every missing entry — inside the ZeroGPU + allocation. The Space passes \`prompt_embedding_path\`, so the 11 GB encoder is + never loaded, and it must never be fetched either. + +## GVHMR + +GVHMR is a git submodule of the source repository and is deliberately absent +here. \`download_assets.py\` clones it into the ephemeral disk at boot, at the +pinned revision below, and \`fdanyone_app.py\` passes that path as \`gvhmr_root\`. + +| Item | Value | +| --- | --- | +| Repository | | +| Revision | \`$GVHMR_REVISION\` | +EOF + +echo "synced fdanyone@$REVISION ($BRANCH); GVHMR pinned at $GVHMR_REVISION"