Spaces:
Running on Zero
Vendor the inference-path subset of the fdanyone package
Browse filesThe Space serves one pipeline, so it copies only what that pipeline touches:
reconstruction (nerfstudio, FreeTimeGS) and the Fire shim for scripts/ stay
out of the tree and out of pixi.toml. sync_vendor.sh regenerates the copy and
records the source SHA, so fdanyone/ is never edited by hand.
One patch rides along. prepare_run calls ensure_models, which downloads every
missing entry of MODEL_FILES -- inside the ZeroGPU allocation. The Space ships
the exported prompt embedding and never runs reconstruction, so the UMT5-XXL
encoder, its tokenizer, and the perceptual VGG-19 leave that set. Skipping the
load is not enough; the 11 GB file must never be fetched either.
GVHMR is a submodule upstream and is deliberately absent here: download_assets.py
clones it into the ephemeral disk at the pinned revision.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
- .gitattributes +1 -0
- .gitignore +10 -0
- PROVENANCE.md +36 -0
- fdanyone/__init__.py +1 -0
- fdanyone/assets.py +191 -0
- fdanyone/config.py +200 -0
- fdanyone/device.py +25 -0
- fdanyone/download.py +427 -0
- fdanyone/errors.py +17 -0
- fdanyone/foreground.py +65 -0
- fdanyone/geometry/__init__.py +1 -0
- fdanyone/geometry/cameras.py +250 -0
- fdanyone/geometry/crop.py +144 -0
- fdanyone/geometry/framing.py +579 -0
- fdanyone/io.py +115 -0
- fdanyone/model/__init__.py +1 -0
- fdanyone/model/inference.py +689 -0
- fdanyone/model/loader.py +426 -0
- fdanyone/model/prepared.py +247 -0
- fdanyone/model/profiling.py +147 -0
- fdanyone/model/quantization.py +87 -0
- fdanyone/model/routing.py +69 -0
- fdanyone/model/tiny_decoder.py +60 -0
- fdanyone/model/turbo_lora.py +125 -0
- fdanyone/motion/__init__.py +5 -0
- fdanyone/motion/body.py +143 -0
- fdanyone/motion/gvhmr.py +332 -0
- fdanyone/motion/result.py +188 -0
- fdanyone/motion/worker.py +53 -0
- fdanyone/output.py +218 -0
- fdanyone/pipeline.py +591 -0
- fdanyone/runs.py +161 -0
- fdanyone/skeleton/__init__.py +1 -0
- fdanyone/skeleton/keypoints.py +170 -0
- fdanyone/skeleton/pipeline.py +839 -0
- fdanyone/skeleton/render_worker.py +100 -0
- fdanyone/skeleton/renderer.py +230 -0
- fdanyone/skeleton/worker.py +66 -0
- fdanyone/vendor/__init__.py +1 -0
- fdanyone/vendor/diffsynth/LICENSE +201 -0
- fdanyone/vendor/diffsynth/UPSTREAM.md +23 -0
- fdanyone/vendor/diffsynth/UPSTREAM.patch +1334 -0
- fdanyone/vendor/diffsynth/VENDORED_FILES.txt +22 -0
- fdanyone/vendor/diffsynth/__init__.py +1 -0
- fdanyone/vendor/diffsynth/pipelines/__init__.py +5 -0
- fdanyone/vendor/diffsynth/pipelines/base.py +127 -0
- fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py +93 -0
- fdanyone/vendor/diffsynth/prompters/__init__.py +5 -0
- fdanyone/vendor/diffsynth/prompters/base_prompter.py +69 -0
- fdanyone/vendor/diffsynth/prompters/wan_prompter.py +109 -0
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
*.mp4 filter=lfs diff=lfs merge=lfs -text
|
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Pixi's materialized environments; pixi.lock is the source of truth.
|
| 2 |
+
.pixi/
|
| 3 |
+
|
| 4 |
+
# Everything download_assets.py fetches at boot, plus per-run scratch.
|
| 5 |
+
models/
|
| 6 |
+
data/
|
| 7 |
+
|
| 8 |
+
__pycache__/
|
| 9 |
+
*.py[cod]
|
| 10 |
+
.pytest_cache/
|
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Provenance
|
| 2 |
+
|
| 3 |
+
`fdanyone/` is a copy, not a fork. Regenerate it with `./sync_vendor.sh`.
|
| 4 |
+
|
| 5 |
+
| Item | Value |
|
| 6 |
+
| --- | --- |
|
| 7 |
+
| Source repository | <https://github.com/pablovela5620/4DAnyone-5090> |
|
| 8 |
+
| Branch | `space-streaming` |
|
| 9 |
+
| Commit | `0cc334c2b260ff19b23d423ac8123d318ed602ff` |
|
| 10 |
+
| Synced | 2026-08-27T05:15:32Z |
|
| 11 |
+
|
| 12 |
+
## Excluded from the copy
|
| 13 |
+
|
| 14 |
+
- `fdanyone/nerfstudio/`, `fdanyone/freetimegs/`, `fdanyone/vendor/freetimegs/` — 3DGS
|
| 15 |
+
and 4DGS reconstruction, which this Space does not run.
|
| 16 |
+
- `fdanyone/cli.py` — the Fire shim for `scripts/`, which is not copied either.
|
| 17 |
+
- `__pycache__/`.
|
| 18 |
+
|
| 19 |
+
## Patched in the copy
|
| 20 |
+
|
| 21 |
+
- `fdanyone/assets.py`: `MODEL_FILES` drops `models_t5_umt5-xxl-enc-bf16.pth`,
|
| 22 |
+
`4danyone/umt5-xxl/`, and the perceptual VGG-19. `prepare_run` calls
|
| 23 |
+
`ensure_models`, which downloads every missing entry — inside the ZeroGPU
|
| 24 |
+
allocation. The Space passes `prompt_embedding_path`, so the 11 GB encoder is
|
| 25 |
+
never loaded, and it must never be fetched either.
|
| 26 |
+
|
| 27 |
+
## GVHMR
|
| 28 |
+
|
| 29 |
+
GVHMR is a git submodule of the source repository and is deliberately absent
|
| 30 |
+
here. `download_assets.py` clones it into the ephemeral disk at boot, at the
|
| 31 |
+
pinned revision below, and `fdanyone_app.py` passes that path as `gvhmr_root`.
|
| 32 |
+
|
| 33 |
+
| Item | Value |
|
| 34 |
+
| --- | --- |
|
| 35 |
+
| Repository | <https://github.com/zju3dv/GVHMR> |
|
| 36 |
+
| Revision | `6ec3ca39336c50492c0fae65fba2fb831fc7d866` |
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""4DAnyone inference."""
|
|
@@ -0,0 +1,191 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Locate the model files used by 4DAnyone.
|
| 2 |
+
|
| 3 |
+
Every published file is anchored by one immutable Hugging Face revision and
|
| 4 |
+
downloaded on demand. ``fdanyone.download`` fetches missing files;
|
| 5 |
+
the resolvers here only locate them.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
from dataclasses import dataclass
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
from typing import TYPE_CHECKING
|
| 13 |
+
|
| 14 |
+
from fdanyone.errors import AssetError
|
| 15 |
+
|
| 16 |
+
if TYPE_CHECKING:
|
| 17 |
+
from fdanyone.config import ModeSettings
|
| 18 |
+
|
| 19 |
+
HF_REPO_ID = "AntResearch/4DAnyone"
|
| 20 |
+
HF_REVISION = "7850985888b56aabf09e69480b73248f1a76bcbe"
|
| 21 |
+
|
| 22 |
+
BIREFNET_REPO_ID = "ZhengPeng7/BiRefNet"
|
| 23 |
+
BIREFNET_REVISION = "e2bf8e4460fc8fa32bba5ea4d94b3233d367b0e4"
|
| 24 |
+
BIREFNET_DIR = "birefnet"
|
| 25 |
+
BIREFNET_FILES = (
|
| 26 |
+
"BiRefNet_config.py",
|
| 27 |
+
"birefnet.py",
|
| 28 |
+
"config.json",
|
| 29 |
+
"model.safetensors",
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
CHECKPOINT = "4danyone/model.safetensors"
|
| 33 |
+
MHR70_REGRESSOR = "4danyone/smplx_to_goliath70.pt"
|
| 34 |
+
WAN_VAE = "4danyone/Wan2.2_VAE.pth"
|
| 35 |
+
TEXT_ENCODER = "4danyone/models_t5_umt5-xxl-enc-bf16.pth"
|
| 36 |
+
TOKENIZER_DIR = "4danyone/umt5-xxl"
|
| 37 |
+
TOKENIZER_FILES = tuple(
|
| 38 |
+
f"{TOKENIZER_DIR}/{name}"
|
| 39 |
+
for name in ("special_tokens_map.json", "spiece.model", "tokenizer.json", "tokenizer_config.json")
|
| 40 |
+
)
|
| 41 |
+
|
| 42 |
+
GVHMR_CHECKPOINT = "gvhmr/gvhmr_siga24_release.ckpt"
|
| 43 |
+
HMR2_CHECKPOINT = "gvhmr/epoch=10-step=25000.ckpt"
|
| 44 |
+
VITPOSE_CHECKPOINT = "gvhmr/vitpose-h-multi-coco.pth"
|
| 45 |
+
YOLO_CHECKPOINT = "gvhmr/yolov8x.pt"
|
| 46 |
+
PERCEPTUAL_VGG19 = "perceptual/imagenet-vgg-verydeep-19-conv.safetensors"
|
| 47 |
+
TAEW2_2 = "fps-assets/taehv/taew2_2.pth"
|
| 48 |
+
TAEW2_2_URL = (
|
| 49 |
+
"https://raw.githubusercontent.com/madebyollin/taehv/"
|
| 50 |
+
"e743234f3217ab3d1570f65642ab06596d1bd7c5/taew2_2.pth"
|
| 51 |
+
)
|
| 52 |
+
TAEW2_2_SHA256 = "d053e216ca50e2bb837bbcd79b85f0366bea00e5938025572382a773b74c559a"
|
| 53 |
+
TURBO_LORA = (
|
| 54 |
+
"fps-assets/turbo-lora/LoRAs/Wan22-Turbo/"
|
| 55 |
+
"Wan22_TI2V_5B_Turbo_lora_rank_64_fp16.safetensors"
|
| 56 |
+
)
|
| 57 |
+
TURBO_LORA_REPO_ID = "Kijai/WanVideo_comfy"
|
| 58 |
+
TURBO_LORA_REVISION = "86c2b0442e01eeee630b48fd7efc0cd37af03252"
|
| 59 |
+
TURBO_LORA_REPO_FILE = (
|
| 60 |
+
"LoRAs/Wan22-Turbo/Wan22_TI2V_5B_Turbo_lora_rank_64_fp16.safetensors"
|
| 61 |
+
)
|
| 62 |
+
TURBO_LORA_SHA256 = "0ace5244e3d1256f884662c261b017249796cf5b95f05d5ed93cc02a478967b8"
|
| 63 |
+
|
| 64 |
+
SMPLX_MODEL = "body_models/smplx/SMPLX_NEUTRAL.npz"
|
| 65 |
+
|
| 66 |
+
# Space patch (sync_vendor.sh): the exported prompt embedding replaces the
|
| 67 |
+
# UMT5-XXL encoder and its tokenizer, and reconstruction never runs here.
|
| 68 |
+
MODEL_FILES = (
|
| 69 |
+
CHECKPOINT,
|
| 70 |
+
MHR70_REGRESSOR,
|
| 71 |
+
WAN_VAE,
|
| 72 |
+
GVHMR_CHECKPOINT,
|
| 73 |
+
HMR2_CHECKPOINT,
|
| 74 |
+
VITPOSE_CHECKPOINT,
|
| 75 |
+
YOLO_CHECKPOINT,
|
| 76 |
+
)
|
| 77 |
+
|
| 78 |
+
EXAMPLE_FILES = (
|
| 79 |
+
"data/source/pexels/10331522-uhd_2160_4096_25fps.mp4",
|
| 80 |
+
"data/source/pexels/2785536-uhd_2160_3840_25fps.mp4",
|
| 81 |
+
"data/source/pexels/5435720-uhd_2160_4096_25fps.mp4",
|
| 82 |
+
"data/source/pexels/5885633-hd_1080_1920_25fps.mp4",
|
| 83 |
+
"data/source/pexels/5999210-uhd_2160_4096_25fps.mp4",
|
| 84 |
+
"data/source/pexels/6980035-uhd_2160_4096_30fps.mp4",
|
| 85 |
+
"data/source/pexels/7080903-hd_1080_1920_30fps.mp4",
|
| 86 |
+
"data/source/pexels/7480858-uhd_2160_3840_25fps.mp4",
|
| 87 |
+
)
|
| 88 |
+
|
| 89 |
+
# Upstream GVHMR resolves its model files relative to its own checkout, so the
|
| 90 |
+
# install commands link each downloaded file to the location GVHMR expects.
|
| 91 |
+
GVHMR_LINKS = (
|
| 92 |
+
(GVHMR_CHECKPOINT, "inputs/checkpoints/gvhmr/gvhmr_siga24_release.ckpt"),
|
| 93 |
+
(HMR2_CHECKPOINT, "inputs/checkpoints/hmr2/epoch=10-step=25000.ckpt"),
|
| 94 |
+
(VITPOSE_CHECKPOINT, "inputs/checkpoints/vitpose/vitpose-h-multi-coco.pth"),
|
| 95 |
+
(YOLO_CHECKPOINT, "inputs/checkpoints/yolo/yolov8x.pt"),
|
| 96 |
+
(SMPLX_MODEL, "inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz"),
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
@dataclass(frozen=True, slots=True)
|
| 101 |
+
class BaseAssets:
|
| 102 |
+
vae: Path
|
| 103 |
+
"""Wan VAE checkpoint."""
|
| 104 |
+
text_encoder: Path | None
|
| 105 |
+
"""UMT5 text-encoder checkpoint; ``None`` when a prompt embedding replaces it."""
|
| 106 |
+
tokenizer: Path | None
|
| 107 |
+
"""UMT5 tokenizer directory; ``None`` when a prompt embedding replaces it."""
|
| 108 |
+
tiny_decoder: Path | None
|
| 109 |
+
"""Pinned TAEW2.2 checkpoint when turbo mode needs it."""
|
| 110 |
+
turbo_lora: Path | None
|
| 111 |
+
"""Pinned Wan2.2 Turbo-LoRA when turbo mode needs it."""
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def _require_file(path: Path, label: str, command: str) -> Path:
|
| 115 |
+
resolved = path.expanduser().resolve()
|
| 116 |
+
if not resolved.is_file():
|
| 117 |
+
raise AssetError(f"{label} does not exist: {resolved}. Run `python {command}` to install it.")
|
| 118 |
+
return resolved
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def resolve_checkpoint(path: str | Path | None = None, model_dir: str | Path = "models") -> Path:
|
| 122 |
+
if path is not None:
|
| 123 |
+
resolved = Path(path).expanduser().resolve()
|
| 124 |
+
if not resolved.is_file():
|
| 125 |
+
raise AssetError(f"Checkpoint override does not exist: {resolved}")
|
| 126 |
+
return resolved
|
| 127 |
+
return _require_file(Path(model_dir) / CHECKPOINT, "Checkpoint", "scripts/download_model.py")
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def resolve_regressor(path: str | Path | None = None, model_dir: str | Path = "models") -> Path:
|
| 131 |
+
if path is not None:
|
| 132 |
+
resolved = Path(path).expanduser().resolve()
|
| 133 |
+
if not resolved.is_file():
|
| 134 |
+
raise AssetError(f"MHR70 regressor override does not exist: {resolved}")
|
| 135 |
+
return resolved
|
| 136 |
+
return _require_file(Path(model_dir) / MHR70_REGRESSOR, "MHR70 regressor", "scripts/download_model.py")
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def resolve_foreground_model(model_dir: str | Path = "models") -> Path:
|
| 140 |
+
root = Path(model_dir).expanduser() / BIREFNET_DIR
|
| 141 |
+
for relative in BIREFNET_FILES:
|
| 142 |
+
_require_file(root / relative, "BiRefNet file", "scripts/download_model.py")
|
| 143 |
+
return root.resolve()
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def resolve_perceptual_vgg19(model_dir: str | Path = "models") -> Path:
|
| 147 |
+
"""Resolve the converted VGG-19 weights used by perceptual reconstruction."""
|
| 148 |
+
|
| 149 |
+
return _require_file(
|
| 150 |
+
Path(model_dir) / PERCEPTUAL_VGG19,
|
| 151 |
+
"Perceptual VGG-19 weights",
|
| 152 |
+
"scripts/download_model.py",
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
def resolve_base_assets(
|
| 157 |
+
model_dir: str | Path,
|
| 158 |
+
settings: ModeSettings,
|
| 159 |
+
*,
|
| 160 |
+
have_prompt_embedding: bool = False,
|
| 161 |
+
) -> BaseAssets:
|
| 162 |
+
"""Resolve the local VAE, T5 encoder, and tokenizer.
|
| 163 |
+
|
| 164 |
+
The encoder and its tokenizer serve only the single fixed prompt, so a
|
| 165 |
+
caller that already holds the exported embedding needs neither and must
|
| 166 |
+
not be asked for the 11 GB checkpoint.
|
| 167 |
+
"""
|
| 168 |
+
|
| 169 |
+
root = Path(model_dir).expanduser()
|
| 170 |
+
text_encoder: Path | None = None
|
| 171 |
+
tokenizer: Path | None = None
|
| 172 |
+
if not have_prompt_embedding:
|
| 173 |
+
for relative in TOKENIZER_FILES:
|
| 174 |
+
_require_file(root / relative, "Tokenizer file", "scripts/download_model.py")
|
| 175 |
+
text_encoder = _require_file(root / TEXT_ENCODER, "Text encoder", "scripts/download_model.py")
|
| 176 |
+
tokenizer = (root / TOKENIZER_DIR).expanduser().resolve()
|
| 177 |
+
return BaseAssets(
|
| 178 |
+
vae=_require_file(root / WAN_VAE, "VAE", "scripts/download_model.py"),
|
| 179 |
+
text_encoder=text_encoder,
|
| 180 |
+
tokenizer=tokenizer,
|
| 181 |
+
tiny_decoder=(
|
| 182 |
+
_require_file(root / TAEW2_2, "TAEW2.2 checkpoint", "scripts/download_model.py")
|
| 183 |
+
if settings.tiny_decoders
|
| 184 |
+
else None
|
| 185 |
+
),
|
| 186 |
+
turbo_lora=(
|
| 187 |
+
_require_file(root / TURBO_LORA, "Turbo-LoRA", "scripts/download_model.py")
|
| 188 |
+
if settings.turbo_lora
|
| 189 |
+
else None
|
| 190 |
+
),
|
| 191 |
+
)
|
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Fixed model and preprocessing settings used by the released method.
|
| 2 |
+
|
| 3 |
+
Only reader-useful choices live in the CLI. These values describe the trained
|
| 4 |
+
model and therefore stay together here instead of being exposed as knobs.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
from collections.abc import Mapping
|
| 10 |
+
from dataclasses import dataclass
|
| 11 |
+
from types import MappingProxyType
|
| 12 |
+
from typing import Literal
|
| 13 |
+
|
| 14 |
+
ModeName = Literal["turbo", "reference"]
|
| 15 |
+
|
| 16 |
+
MVS_ATTENTION_RANGE: tuple[int | None, int | None, int | None] = (0, None, 1)
|
| 17 |
+
USE_VIEWPACK: bool = True
|
| 18 |
+
USE_POSE_ENCODER: bool = True
|
| 19 |
+
POSE_ENCODER_TYPE: str = "rgb"
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
@dataclass(frozen=True, slots=True)
|
| 23 |
+
class InferenceConstants:
|
| 24 |
+
"""Architecture and media constants shared by both inference modes."""
|
| 25 |
+
|
| 26 |
+
num_frames: int = 121
|
| 27 |
+
"""Frames generated for every view."""
|
| 28 |
+
height: int = 1280
|
| 29 |
+
"""Output height in pixels."""
|
| 30 |
+
width: int = 704
|
| 31 |
+
"""Output width in pixels."""
|
| 32 |
+
prompt: str = "视频中的人在做动作"
|
| 33 |
+
"""Fixed positive prompt used by the released model."""
|
| 34 |
+
auto_downsample_fps: tuple[tuple[int, int], ...] = (
|
| 35 |
+
(24, 1),
|
| 36 |
+
(24000, 1001),
|
| 37 |
+
(25, 1),
|
| 38 |
+
(30, 1),
|
| 39 |
+
(30000, 1001),
|
| 40 |
+
)
|
| 41 |
+
"""Input rates that may be reduced to their exact supported divisor."""
|
| 42 |
+
temporal_sampling_policy: str = "nearest_source_pts_on_zero_based_cfr_clock"
|
| 43 |
+
"""Canonical clip sampling policy."""
|
| 44 |
+
rcp_jpeg_quality: int = 85
|
| 45 |
+
"""JPEG quality at the proposal-to-target boundary."""
|
| 46 |
+
skeleton_h264_crf: int = 17
|
| 47 |
+
"""CRF used for skeleton conditioning videos."""
|
| 48 |
+
target_h264_crf: int = 18
|
| 49 |
+
"""CRF used for generated target videos."""
|
| 50 |
+
h264_preset: str = "medium"
|
| 51 |
+
"""H.264 encoder preset."""
|
| 52 |
+
skeleton_max_dimension: int = 2048
|
| 53 |
+
"""Largest skeleton-render canvas dimension."""
|
| 54 |
+
denoising_strength: float = 1.0
|
| 55 |
+
"""Flow-match denoising strength."""
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
@dataclass(frozen=True, slots=True)
|
| 59 |
+
class ModeSettings:
|
| 60 |
+
"""One complete, immutable inference policy."""
|
| 61 |
+
|
| 62 |
+
mode: ModeName
|
| 63 |
+
"""Public CLI and metadata name."""
|
| 64 |
+
num_inference_steps: int
|
| 65 |
+
"""Number of flow-matching denoising steps."""
|
| 66 |
+
scheduler_shift: float
|
| 67 |
+
"""Flow-match sigma shift."""
|
| 68 |
+
stream_dit_weights: bool
|
| 69 |
+
"""Whether all wrapped DiT weights stream from host memory."""
|
| 70 |
+
fp8_w8a8: bool
|
| 71 |
+
"""Whether safe interior projections use dynamic per-tensor FP8."""
|
| 72 |
+
regional_compile: bool
|
| 73 |
+
"""Whether repeated DiT blocks use the fixed turbo compile policy."""
|
| 74 |
+
overlap_target_skeletons: bool = False
|
| 75 |
+
"""Whether target skeleton rendering overlaps the proposal stage."""
|
| 76 |
+
bf16_block_glue: bool = False
|
| 77 |
+
"""Run transformer normalization, modulation, and gates in BF16."""
|
| 78 |
+
direct_rcp_latent_handoff: bool = False
|
| 79 |
+
"""Feed generated RCP latents to target conditioning without re-encoding."""
|
| 80 |
+
async_video_encode: bool = False
|
| 81 |
+
"""Overlap CPU x264 encoding with the next per-view GPU VAE decode."""
|
| 82 |
+
tiny_decoders: bool = False
|
| 83 |
+
"""Whether proposal and target views use the pinned TAEW2.2 decoder."""
|
| 84 |
+
nvdec_skeletons: bool = False
|
| 85 |
+
"""Whether skeleton videos use fail-closed CUDA decoding."""
|
| 86 |
+
turbo_lora: bool = False
|
| 87 |
+
"""Whether to merge the pinned Wan2.2 Turbo-LoRA before quantization."""
|
| 88 |
+
|
| 89 |
+
@property
|
| 90 |
+
def dit_pose_batch_size(self) -> int | None:
|
| 91 |
+
"""Return the validated FP8 pose-activation batch."""
|
| 92 |
+
|
| 93 |
+
return 4 if self.fp8_w8a8 else None
|
| 94 |
+
|
| 95 |
+
@property
|
| 96 |
+
def exact_attention(self) -> bool:
|
| 97 |
+
"""Return whether this mode requires exact SDPA attention."""
|
| 98 |
+
|
| 99 |
+
return self.mode == "reference"
|
| 100 |
+
|
| 101 |
+
@property
|
| 102 |
+
def skeleton_video_decoder(self) -> str:
|
| 103 |
+
"""Return the concrete skeleton-video backend."""
|
| 104 |
+
|
| 105 |
+
return "torchcodec_cuda" if self.nvdec_skeletons else "pyav"
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
MODES: Mapping[ModeName, ModeSettings] = MappingProxyType(
|
| 109 |
+
{
|
| 110 |
+
"turbo": ModeSettings(
|
| 111 |
+
mode="turbo",
|
| 112 |
+
num_inference_steps=4,
|
| 113 |
+
scheduler_shift=17.0,
|
| 114 |
+
stream_dit_weights=False,
|
| 115 |
+
fp8_w8a8=True,
|
| 116 |
+
regional_compile=True,
|
| 117 |
+
overlap_target_skeletons=True,
|
| 118 |
+
bf16_block_glue=True,
|
| 119 |
+
direct_rcp_latent_handoff=True,
|
| 120 |
+
async_video_encode=True,
|
| 121 |
+
tiny_decoders=True,
|
| 122 |
+
nvdec_skeletons=True,
|
| 123 |
+
turbo_lora=True,
|
| 124 |
+
),
|
| 125 |
+
"reference": ModeSettings(
|
| 126 |
+
mode="reference",
|
| 127 |
+
num_inference_steps=24,
|
| 128 |
+
scheduler_shift=5.0,
|
| 129 |
+
stream_dit_weights=True,
|
| 130 |
+
fp8_w8a8=False,
|
| 131 |
+
regional_compile=False,
|
| 132 |
+
),
|
| 133 |
+
}
|
| 134 |
+
)
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
@dataclass(frozen=True)
|
| 138 |
+
class CameraConfig:
|
| 139 |
+
count: int = 24
|
| 140 |
+
pitch_degrees: float = 15.0
|
| 141 |
+
|
| 142 |
+
def __post_init__(self) -> None:
|
| 143 |
+
if self.count <= 0:
|
| 144 |
+
raise ValueError("Camera count must be positive.")
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
@dataclass(frozen=True)
|
| 148 |
+
class ForegroundConfig:
|
| 149 |
+
"""Pinned standard BiRefNet inference contract."""
|
| 150 |
+
|
| 151 |
+
image_size: tuple[int, int] = (1024, 1024)
|
| 152 |
+
batch_size: int = 4
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
@dataclass(frozen=True)
|
| 156 |
+
class FramingConfig:
|
| 157 |
+
"""Sequence-level camera solve matching the current GVHMR demo."""
|
| 158 |
+
|
| 159 |
+
reference_radius: float = 3.0
|
| 160 |
+
reference_target_height: float = 1.0
|
| 161 |
+
reference_focal_normalized: float = 1664.0 / 1280.0
|
| 162 |
+
height_target_ratio: float = 0.80
|
| 163 |
+
height_percentile: float = 95.0
|
| 164 |
+
width_target_ratio: float = 0.90
|
| 165 |
+
width_percentile: float = 80.0
|
| 166 |
+
min_radius: float = 1.5
|
| 167 |
+
max_radius: float = 8.0
|
| 168 |
+
input_min_confidence: float = 0.55
|
| 169 |
+
max_focal_normalized: float = 4.0
|
| 170 |
+
cutoff_target_ratio: float = 0.99
|
| 171 |
+
cutoff_percentile: float = 80.0
|
| 172 |
+
|
| 173 |
+
|
| 174 |
+
@dataclass(frozen=True)
|
| 175 |
+
class CropConfig:
|
| 176 |
+
"""Source-mask crop; generated cameras use a plain center aspect crop."""
|
| 177 |
+
|
| 178 |
+
margin_top: float = 0.04
|
| 179 |
+
margin_right: float = 0.04
|
| 180 |
+
margin_bottom: float = 0.04
|
| 181 |
+
margin_left: float = 0.04
|
| 182 |
+
allow_upscale: bool = True
|
| 183 |
+
mask_threshold: float = 0.05
|
| 184 |
+
|
| 185 |
+
@property
|
| 186 |
+
def margins(self) -> tuple[float, float, float, float]:
|
| 187 |
+
return (self.margin_top, self.margin_right, self.margin_bottom, self.margin_left)
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
@dataclass(frozen=True)
|
| 191 |
+
class SkeletonConfig:
|
| 192 |
+
draw_body_reference_px: float = 640.0
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
INFERENCE = InferenceConstants()
|
| 196 |
+
CAMERA = CameraConfig()
|
| 197 |
+
FOREGROUND = ForegroundConfig()
|
| 198 |
+
FRAMING = FramingConfig()
|
| 199 |
+
CROP = CropConfig()
|
| 200 |
+
SKELETON = SkeletonConfig()
|
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""CUDA device selection shared by pipeline and isolated workers."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from fdanyone.errors import ConfigurationError
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def select_cuda_device(device: str) -> tuple[str, int]:
|
| 9 |
+
"""Validate, select, and normalize one CUDA device."""
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
|
| 13 |
+
try:
|
| 14 |
+
requested = torch.device(device)
|
| 15 |
+
except (RuntimeError, TypeError, ValueError) as exc:
|
| 16 |
+
raise ConfigurationError(f"Invalid CUDA device {device!r}.") from exc
|
| 17 |
+
if requested.type != "cuda" or not torch.cuda.is_available():
|
| 18 |
+
raise ConfigurationError(f"4DAnyone requires an available CUDA device, got {device!r}.")
|
| 19 |
+
index = torch.cuda.current_device() if requested.index is None else requested.index
|
| 20 |
+
if index < 0 or index >= torch.cuda.device_count():
|
| 21 |
+
raise ConfigurationError(
|
| 22 |
+
f"CUDA device index {index} is unavailable; visible device count is {torch.cuda.device_count()}."
|
| 23 |
+
)
|
| 24 |
+
torch.cuda.set_device(index)
|
| 25 |
+
return f"cuda:{index}", index
|
|
@@ -0,0 +1,427 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Download the published 4DAnyone assets from Hugging Face.
|
| 2 |
+
|
| 3 |
+
Missing model checkpoints and bundled example clips are fetched automatically
|
| 4 |
+
when inference needs them; the scripts under ``scripts/`` pre-fetch the same
|
| 5 |
+
files. SMPL-X is licensed separately, so first-run inference starts its
|
| 6 |
+
interactive installer only when a terminal is available.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from __future__ import annotations
|
| 10 |
+
|
| 11 |
+
import getpass
|
| 12 |
+
import logging
|
| 13 |
+
import os
|
| 14 |
+
import shlex
|
| 15 |
+
import shutil
|
| 16 |
+
import sys
|
| 17 |
+
import tempfile
|
| 18 |
+
import urllib.error
|
| 19 |
+
import urllib.parse
|
| 20 |
+
import urllib.request
|
| 21 |
+
import zipfile
|
| 22 |
+
from pathlib import Path, PurePosixPath
|
| 23 |
+
|
| 24 |
+
from fdanyone.assets import (
|
| 25 |
+
BIREFNET_DIR,
|
| 26 |
+
BIREFNET_FILES,
|
| 27 |
+
BIREFNET_REPO_ID,
|
| 28 |
+
BIREFNET_REVISION,
|
| 29 |
+
EXAMPLE_FILES,
|
| 30 |
+
GVHMR_LINKS,
|
| 31 |
+
HF_REPO_ID,
|
| 32 |
+
HF_REVISION,
|
| 33 |
+
MODEL_FILES,
|
| 34 |
+
PERCEPTUAL_VGG19,
|
| 35 |
+
SMPLX_MODEL,
|
| 36 |
+
TAEW2_2,
|
| 37 |
+
TAEW2_2_SHA256,
|
| 38 |
+
TAEW2_2_URL,
|
| 39 |
+
TURBO_LORA,
|
| 40 |
+
TURBO_LORA_REPO_FILE,
|
| 41 |
+
TURBO_LORA_REPO_ID,
|
| 42 |
+
TURBO_LORA_REVISION,
|
| 43 |
+
TURBO_LORA_SHA256,
|
| 44 |
+
resolve_perceptual_vgg19,
|
| 45 |
+
)
|
| 46 |
+
from fdanyone.errors import AssetError
|
| 47 |
+
|
| 48 |
+
LOGGER = logging.getLogger("fdanyone")
|
| 49 |
+
|
| 50 |
+
SMPLX_HOME = "https://smpl-x.is.tue.mpg.de/"
|
| 51 |
+
SMPLX_DOWNLOAD_URL = "https://download.is.tue.mpg.de/download.php?domain=smplx&sfile=models_smplx_v1_1.zip"
|
| 52 |
+
SMPLX_ARCHIVE_MEMBER = ("models", "smplx", "SMPLX_NEUTRAL.npz")
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _snapshot(
|
| 56 |
+
allow_patterns: list[str],
|
| 57 |
+
local_dir: Path,
|
| 58 |
+
*,
|
| 59 |
+
repo_id: str = HF_REPO_ID,
|
| 60 |
+
revision: str = HF_REVISION,
|
| 61 |
+
) -> None:
|
| 62 |
+
try:
|
| 63 |
+
from huggingface_hub import snapshot_download
|
| 64 |
+
except ImportError as exc:
|
| 65 |
+
raise AssetError("Install requirements.txt before downloading assets.") from exc
|
| 66 |
+
|
| 67 |
+
try:
|
| 68 |
+
snapshot_download(
|
| 69 |
+
repo_id=repo_id,
|
| 70 |
+
revision=revision,
|
| 71 |
+
allow_patterns=allow_patterns,
|
| 72 |
+
local_dir=local_dir,
|
| 73 |
+
)
|
| 74 |
+
except Exception as exc:
|
| 75 |
+
raise AssetError(
|
| 76 |
+
f"Could not download {repo_id}@{revision}. Check the network connection and Hugging Face access."
|
| 77 |
+
) from exc
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def ensure_foreground_model(model_dir: str | Path = "models") -> Path:
|
| 81 |
+
root = Path(model_dir).expanduser().resolve() / BIREFNET_DIR
|
| 82 |
+
missing = [relative for relative in BIREFNET_FILES if not (root / relative).is_file()]
|
| 83 |
+
if missing:
|
| 84 |
+
LOGGER.info("Downloading BiRefNet foreground model (first run only)")
|
| 85 |
+
_snapshot(
|
| 86 |
+
missing,
|
| 87 |
+
root,
|
| 88 |
+
repo_id=BIREFNET_REPO_ID,
|
| 89 |
+
revision=BIREFNET_REVISION,
|
| 90 |
+
)
|
| 91 |
+
return root
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def require_gvhmr_checkout(gvhmr_root: str | Path) -> Path:
|
| 95 |
+
root = Path(gvhmr_root).expanduser().resolve()
|
| 96 |
+
if not (root / "hmr4d/__init__.py").is_file():
|
| 97 |
+
raise AssetError(
|
| 98 |
+
f"GVHMR is not initialized at {root}. Run `git submodule update --init third_party/GVHMR` first."
|
| 99 |
+
)
|
| 100 |
+
return root
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def _ensure_link(source: Path, destination: Path) -> None:
|
| 104 |
+
source = source.expanduser().resolve()
|
| 105 |
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
| 106 |
+
if destination.is_symlink():
|
| 107 |
+
try:
|
| 108 |
+
if destination.resolve(strict=True).samefile(source):
|
| 109 |
+
return
|
| 110 |
+
except FileNotFoundError:
|
| 111 |
+
pass
|
| 112 |
+
destination.unlink()
|
| 113 |
+
elif destination.exists():
|
| 114 |
+
if destination.samefile(source):
|
| 115 |
+
return
|
| 116 |
+
raise AssetError(
|
| 117 |
+
f"GVHMR asset location is occupied by an unrelated file: {destination}. "
|
| 118 |
+
f"Move it away so the downloaded {source.name} can be linked."
|
| 119 |
+
)
|
| 120 |
+
relative = os.path.relpath(source, start=destination.parent)
|
| 121 |
+
destination.symlink_to(relative)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def create_classic_gvhmr_links(
|
| 125 |
+
model_dir: str | Path = "models",
|
| 126 |
+
gvhmr_root: str | Path = "third_party/GVHMR",
|
| 127 |
+
*,
|
| 128 |
+
require_models: bool = True,
|
| 129 |
+
require_smplx: bool = True,
|
| 130 |
+
) -> Path:
|
| 131 |
+
"""Create the ignored compatibility links expected by upstream GVHMR."""
|
| 132 |
+
|
| 133 |
+
root = require_gvhmr_checkout(gvhmr_root)
|
| 134 |
+
models = Path(model_dir).expanduser().resolve()
|
| 135 |
+
for relative, target in GVHMR_LINKS:
|
| 136 |
+
source = models / relative
|
| 137 |
+
if not source.is_file():
|
| 138 |
+
required = require_smplx if relative == SMPLX_MODEL else require_models
|
| 139 |
+
if required:
|
| 140 |
+
command = "scripts/download_smplx.py" if relative == SMPLX_MODEL else "scripts/download_model.py"
|
| 141 |
+
raise AssetError(f"Model file is missing: {source}. Run `python {command}` first.")
|
| 142 |
+
continue
|
| 143 |
+
_ensure_link(source, root / target)
|
| 144 |
+
return root
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def ensure_models(
|
| 148 |
+
model_dir: str | Path = "models",
|
| 149 |
+
gvhmr_root: str | Path = "third_party/GVHMR",
|
| 150 |
+
) -> Path:
|
| 151 |
+
"""Download any missing published model file and refresh the GVHMR links."""
|
| 152 |
+
|
| 153 |
+
require_gvhmr_checkout(gvhmr_root)
|
| 154 |
+
models = Path(model_dir).expanduser().resolve()
|
| 155 |
+
missing = [relative for relative in MODEL_FILES if not (models / relative).is_file()]
|
| 156 |
+
if missing:
|
| 157 |
+
LOGGER.info("Downloading %d model files from %s (first run only)", len(missing), HF_REPO_ID)
|
| 158 |
+
# Repository paths match the local layout, so download straight into
|
| 159 |
+
# place; huggingface_hub stages and resumes partial files itself.
|
| 160 |
+
_snapshot(missing, models)
|
| 161 |
+
ensure_foreground_model(models)
|
| 162 |
+
create_classic_gvhmr_links(models, gvhmr_root, require_smplx=False)
|
| 163 |
+
return models
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def _verify_sha256(path: Path, expected: str, label: str) -> None:
|
| 167 |
+
import hashlib
|
| 168 |
+
|
| 169 |
+
digest = hashlib.sha256()
|
| 170 |
+
with path.open("rb") as handle:
|
| 171 |
+
for chunk in iter(lambda: handle.read(1 << 20), b""):
|
| 172 |
+
digest.update(chunk)
|
| 173 |
+
if digest.hexdigest() != expected:
|
| 174 |
+
path.unlink(missing_ok=True)
|
| 175 |
+
raise AssetError(f"{label} failed its SHA-256 check; the download was removed. Re-run it.")
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
def ensure_turbo_assets(model_dir: str | Path = "models") -> Path:
|
| 179 |
+
"""Download the pinned turbo-mode checkpoints (TAEW2.2 decoder, Turbo-LoRA)."""
|
| 180 |
+
|
| 181 |
+
models = Path(model_dir).expanduser().resolve()
|
| 182 |
+
taew = models / TAEW2_2
|
| 183 |
+
if not taew.is_file():
|
| 184 |
+
LOGGER.info("Downloading the TAEW2.2 tiny decoder (first run only)")
|
| 185 |
+
taew.parent.mkdir(parents=True, exist_ok=True)
|
| 186 |
+
staged = taew.with_suffix(".part")
|
| 187 |
+
try:
|
| 188 |
+
urllib.request.urlretrieve(TAEW2_2_URL, staged)
|
| 189 |
+
except urllib.error.URLError as exc:
|
| 190 |
+
raise AssetError(f"Could not download the TAEW2.2 decoder from {TAEW2_2_URL}.") from exc
|
| 191 |
+
_verify_sha256(staged, TAEW2_2_SHA256, "TAEW2.2 decoder")
|
| 192 |
+
staged.replace(taew)
|
| 193 |
+
lora = models / TURBO_LORA
|
| 194 |
+
if not lora.is_file():
|
| 195 |
+
LOGGER.info(
|
| 196 |
+
"Downloading the Wan2.2 Turbo-LoRA from %s (first run only)", TURBO_LORA_REPO_ID
|
| 197 |
+
)
|
| 198 |
+
_snapshot(
|
| 199 |
+
[TURBO_LORA_REPO_FILE],
|
| 200 |
+
models / "fps-assets" / "turbo-lora",
|
| 201 |
+
repo_id=TURBO_LORA_REPO_ID,
|
| 202 |
+
revision=TURBO_LORA_REVISION,
|
| 203 |
+
)
|
| 204 |
+
_verify_sha256(lora, TURBO_LORA_SHA256, "Turbo-LoRA")
|
| 205 |
+
return models
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def ensure_perceptual_vgg19(model_dir: str | Path = "models") -> Path:
|
| 209 |
+
"""Download only the optional VGG-19 reconstruction asset when missing."""
|
| 210 |
+
|
| 211 |
+
models = Path(model_dir).expanduser().resolve()
|
| 212 |
+
destination = models / PERCEPTUAL_VGG19
|
| 213 |
+
if not destination.is_file():
|
| 214 |
+
LOGGER.info("Downloading the perceptual VGG-19 model (first use only)")
|
| 215 |
+
_snapshot([PERCEPTUAL_VGG19], models)
|
| 216 |
+
return resolve_perceptual_vgg19(models)
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
def download_model(
|
| 220 |
+
model_dir: str = "models",
|
| 221 |
+
gvhmr_root: str = "third_party/GVHMR",
|
| 222 |
+
) -> dict[str, str]:
|
| 223 |
+
"""Download the published model checkpoints."""
|
| 224 |
+
|
| 225 |
+
models = ensure_models(model_dir, gvhmr_root)
|
| 226 |
+
ensure_turbo_assets(models)
|
| 227 |
+
return {
|
| 228 |
+
"models": str(models),
|
| 229 |
+
"revision": HF_REVISION,
|
| 230 |
+
"foreground_revision": BIREFNET_REVISION,
|
| 231 |
+
"turbo_lora_revision": TURBO_LORA_REVISION,
|
| 232 |
+
}
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
def download_example(data_dir: str = "data") -> dict[str, str]:
|
| 236 |
+
"""Download the bundled example clips."""
|
| 237 |
+
|
| 238 |
+
data = Path(data_dir).expanduser().resolve()
|
| 239 |
+
destinations = {relative: data / Path(relative).relative_to("data") for relative in EXAMPLE_FILES}
|
| 240 |
+
missing = [relative for relative, destination in destinations.items() if not destination.is_file()]
|
| 241 |
+
if missing:
|
| 242 |
+
# Repository paths carry a leading ``data/`` prefix while --data_dir is
|
| 243 |
+
# the local root itself, so stage the snapshot and move each file.
|
| 244 |
+
staging = data / ".download"
|
| 245 |
+
_snapshot(missing, staging)
|
| 246 |
+
for relative in missing:
|
| 247 |
+
destination = destinations[relative]
|
| 248 |
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
| 249 |
+
(staging / relative).replace(destination)
|
| 250 |
+
shutil.rmtree(staging)
|
| 251 |
+
return {"examples": str(data / "source/pexels"), "revision": HF_REVISION}
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
def ensure_example_video(video_path: str | Path) -> Path:
|
| 255 |
+
"""Fetch a bundled example clip when its expected file is missing."""
|
| 256 |
+
|
| 257 |
+
path = Path(video_path).expanduser()
|
| 258 |
+
if path.is_file():
|
| 259 |
+
return path
|
| 260 |
+
matches = [relative for relative in EXAMPLE_FILES if PurePosixPath(relative).name == path.name]
|
| 261 |
+
if not matches:
|
| 262 |
+
raise AssetError(f"Input video does not exist: {path.resolve()}")
|
| 263 |
+
LOGGER.info("Downloading the bundled example clip %s", path.name)
|
| 264 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 265 |
+
# Stage beside the destination so the final rename stays on one filesystem.
|
| 266 |
+
staging = path.parent / ".download"
|
| 267 |
+
_snapshot(matches[:1], staging)
|
| 268 |
+
(staging / matches[0]).replace(path)
|
| 269 |
+
shutil.rmtree(staging)
|
| 270 |
+
return path
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
def _parse_interactive_path(value: str) -> Path:
|
| 274 |
+
try:
|
| 275 |
+
parts = shlex.split(value.strip())
|
| 276 |
+
except ValueError as exc:
|
| 277 |
+
raise AssetError(f"Could not parse the archive path: {exc}") from None
|
| 278 |
+
if len(parts) != 1:
|
| 279 |
+
raise AssetError("Enter one ZIP or SMPLX_NEUTRAL.npz path.")
|
| 280 |
+
return Path(parts[0]).expanduser()
|
| 281 |
+
|
| 282 |
+
|
| 283 |
+
def _copy_model_from_source(source: Path, destination: Path) -> None:
|
| 284 |
+
if source.name == "SMPLX_NEUTRAL.npz":
|
| 285 |
+
shutil.copyfile(source, destination)
|
| 286 |
+
return
|
| 287 |
+
if zipfile.is_zipfile(source):
|
| 288 |
+
try:
|
| 289 |
+
with zipfile.ZipFile(source) as archive:
|
| 290 |
+
candidates = [
|
| 291 |
+
info
|
| 292 |
+
for info in archive.infolist()
|
| 293 |
+
if not info.is_dir() and PurePosixPath(info.filename).parts[-3:] == SMPLX_ARCHIVE_MEMBER
|
| 294 |
+
]
|
| 295 |
+
if len(candidates) != 1:
|
| 296 |
+
raise AssetError(
|
| 297 |
+
"The archive must contain exactly one models/smplx/SMPLX_NEUTRAL.npz file. "
|
| 298 |
+
"Download models_smplx_v1_1.zip from the official SMPL-X website."
|
| 299 |
+
)
|
| 300 |
+
with archive.open(candidates[0]) as model, destination.open("wb") as output:
|
| 301 |
+
shutil.copyfileobj(model, output, length=8 * 1024 * 1024)
|
| 302 |
+
except zipfile.BadZipFile:
|
| 303 |
+
raise AssetError(f"SMPL-X archive is invalid: {source}") from None
|
| 304 |
+
return
|
| 305 |
+
raise AssetError("Select models_smplx_v1_1.zip or SMPLX_NEUTRAL.npz.")
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
def install_smplx(
|
| 309 |
+
source_path: str | Path,
|
| 310 |
+
model_dir: str | Path = "models",
|
| 311 |
+
gvhmr_root: str | Path = "third_party/GVHMR",
|
| 312 |
+
) -> Path:
|
| 313 |
+
"""Install a user-provided official ZIP or neutral NPZ."""
|
| 314 |
+
|
| 315 |
+
source = Path(source_path).expanduser().resolve()
|
| 316 |
+
if not source.is_file():
|
| 317 |
+
raise AssetError(f"SMPL-X source does not exist: {source}")
|
| 318 |
+
|
| 319 |
+
target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL
|
| 320 |
+
target.parent.mkdir(parents=True, exist_ok=True)
|
| 321 |
+
temporary = target.parent / f".{target.name}.download-{os.getpid()}"
|
| 322 |
+
try:
|
| 323 |
+
_copy_model_from_source(source, temporary)
|
| 324 |
+
temporary.replace(target)
|
| 325 |
+
finally:
|
| 326 |
+
temporary.unlink(missing_ok=True)
|
| 327 |
+
create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True)
|
| 328 |
+
return target
|
| 329 |
+
|
| 330 |
+
|
| 331 |
+
def _download_official(username: str, password: str, destination: Path) -> None:
|
| 332 |
+
payload = urllib.parse.urlencode({"username": username, "password": password}).encode()
|
| 333 |
+
request = urllib.request.Request(
|
| 334 |
+
SMPLX_DOWNLOAD_URL,
|
| 335 |
+
data=payload,
|
| 336 |
+
headers={"User-Agent": "4DAnyone SMPL-X installer"},
|
| 337 |
+
method="POST",
|
| 338 |
+
)
|
| 339 |
+
try:
|
| 340 |
+
with urllib.request.urlopen(request, timeout=120) as response, destination.open("wb") as output:
|
| 341 |
+
shutil.copyfileobj(response, output, length=8 * 1024 * 1024)
|
| 342 |
+
except (OSError, urllib.error.URLError) as exc:
|
| 343 |
+
raise AssetError(f"Official SMPL-X download failed: {exc}") from None
|
| 344 |
+
if not zipfile.is_zipfile(destination):
|
| 345 |
+
raise AssetError(
|
| 346 |
+
"The SMPL-X website did not return a ZIP archive. Check the account, license acceptance, or website."
|
| 347 |
+
)
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
def _prompt_for_archive(model_dir: str, gvhmr_root: str) -> dict[str, str] | None:
|
| 351 |
+
print(f"Download models_smplx_v1_1.zip from:\n {SMPLX_DOWNLOAD_URL}")
|
| 352 |
+
while True:
|
| 353 |
+
try:
|
| 354 |
+
value = input("Archive path (drag the downloaded ZIP here): ").strip()
|
| 355 |
+
except EOFError:
|
| 356 |
+
value = ""
|
| 357 |
+
if not value:
|
| 358 |
+
print("SMPL-X setup cancelled; the downloaded ZIP was not modified.")
|
| 359 |
+
return None
|
| 360 |
+
try:
|
| 361 |
+
installed = install_smplx(_parse_interactive_path(value), model_dir, gvhmr_root)
|
| 362 |
+
except AssetError as exc:
|
| 363 |
+
print(f"error: {exc}")
|
| 364 |
+
continue
|
| 365 |
+
return {"installed": str(installed)}
|
| 366 |
+
|
| 367 |
+
|
| 368 |
+
def download_smplx(
|
| 369 |
+
archive_path: str | None = None,
|
| 370 |
+
model_dir: str = "models",
|
| 371 |
+
gvhmr_root: str = "third_party/GVHMR",
|
| 372 |
+
) -> dict[str, str] | None:
|
| 373 |
+
"""Install the separately licensed SMPL-X neutral body model."""
|
| 374 |
+
|
| 375 |
+
target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL
|
| 376 |
+
if target.is_file():
|
| 377 |
+
create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True)
|
| 378 |
+
return {"installed": str(target)}
|
| 379 |
+
|
| 380 |
+
if archive_path is not None:
|
| 381 |
+
return {"installed": str(install_smplx(archive_path, model_dir, gvhmr_root))}
|
| 382 |
+
|
| 383 |
+
print(f"SMPL-X requires a free account and license acceptance at {SMPLX_HOME}")
|
| 384 |
+
try:
|
| 385 |
+
accepted = input("Have you registered and accepted the SMPL-X license? [y/N]: ").strip().lower()
|
| 386 |
+
except EOFError:
|
| 387 |
+
accepted = ""
|
| 388 |
+
if accepted in {"y", "yes"}:
|
| 389 |
+
username = input("SMPL-X username or email: ").strip()
|
| 390 |
+
password = getpass.getpass("SMPL-X password: ")
|
| 391 |
+
if username and password:
|
| 392 |
+
with tempfile.TemporaryDirectory(prefix="fdanyone-smplx-") as temporary_dir:
|
| 393 |
+
archive = Path(temporary_dir) / "models_smplx_v1_1.zip"
|
| 394 |
+
try:
|
| 395 |
+
_download_official(username, password, archive)
|
| 396 |
+
installed = install_smplx(archive, model_dir, gvhmr_root)
|
| 397 |
+
except AssetError as exc:
|
| 398 |
+
print(f"Automatic download was unavailable: {exc}")
|
| 399 |
+
else:
|
| 400 |
+
return {"installed": str(installed)}
|
| 401 |
+
return _prompt_for_archive(model_dir, gvhmr_root)
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
def ensure_smplx(
|
| 405 |
+
model_dir: str | Path = "models",
|
| 406 |
+
gvhmr_root: str | Path = "third_party/GVHMR",
|
| 407 |
+
) -> Path:
|
| 408 |
+
"""Install SMPL-X interactively on first use, without blocking jobs."""
|
| 409 |
+
|
| 410 |
+
require_gvhmr_checkout(gvhmr_root)
|
| 411 |
+
target = Path(model_dir).expanduser().resolve() / SMPLX_MODEL
|
| 412 |
+
if target.is_file():
|
| 413 |
+
create_classic_gvhmr_links(model_dir, gvhmr_root, require_models=False, require_smplx=True)
|
| 414 |
+
return target
|
| 415 |
+
|
| 416 |
+
if not getattr(sys.stdin, "isatty", lambda: False)():
|
| 417 |
+
raise AssetError(
|
| 418 |
+
f"SMPL-X is not installed at {target}, and inference has no interactive terminal. "
|
| 419 |
+
"Run `python scripts/download_smplx.py` before starting this job."
|
| 420 |
+
)
|
| 421 |
+
|
| 422 |
+
print("SMPL-X is required and has not been installed; starting its licensed setup.")
|
| 423 |
+
result = download_smplx(model_dir=str(model_dir), gvhmr_root=str(gvhmr_root))
|
| 424 |
+
if result is None or not target.is_file():
|
| 425 |
+
raise AssetError("SMPL-X setup was cancelled; inference cannot continue.")
|
| 426 |
+
LOGGER.info("SMPL-X installed; continuing inference")
|
| 427 |
+
return target
|
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Project-specific errors with actionable user-facing messages."""
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
class FourDAnyoneError(RuntimeError):
|
| 5 |
+
"""Base class for expected pipeline failures."""
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class ConfigurationError(FourDAnyoneError):
|
| 9 |
+
"""Raised when a frozen inference contract is violated."""
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class AssetError(FourDAnyoneError):
|
| 13 |
+
"""Raised when a model or gated asset is missing or invalid."""
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class VideoContractError(FourDAnyoneError):
|
| 17 |
+
"""Raised when an input cannot provide the canonical 121-frame clip."""
|
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Pinned BiRefNet inference over the canonical source clip."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import gc
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
from PIL import Image
|
| 10 |
+
|
| 11 |
+
from fdanyone.config import FOREGROUND
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def predict_foreground_masks(
|
| 15 |
+
frames: tuple[np.ndarray, ...],
|
| 16 |
+
model_path: str | Path,
|
| 17 |
+
device: str,
|
| 18 |
+
*,
|
| 19 |
+
batch_size: int = FOREGROUND.batch_size,
|
| 20 |
+
) -> np.ndarray:
|
| 21 |
+
"""Return full-raster 8-bit foreground masks for the canonical clip."""
|
| 22 |
+
|
| 23 |
+
import torch
|
| 24 |
+
from torchvision import transforms
|
| 25 |
+
from torchvision.transforms.functional import to_pil_image
|
| 26 |
+
from transformers import AutoModelForImageSegmentation
|
| 27 |
+
|
| 28 |
+
if not frames:
|
| 29 |
+
raise ValueError("Foreground inference requires at least one frame.")
|
| 30 |
+
if batch_size <= 0:
|
| 31 |
+
raise ValueError("batch_size must be positive.")
|
| 32 |
+
shape = frames[0].shape
|
| 33 |
+
if any(frame.dtype != np.uint8 or frame.shape != shape for frame in frames):
|
| 34 |
+
raise ValueError("Foreground frames must share one RGB uint8 raster.")
|
| 35 |
+
|
| 36 |
+
model = AutoModelForImageSegmentation.from_pretrained(
|
| 37 |
+
str(Path(model_path).expanduser().resolve()),
|
| 38 |
+
local_files_only=True,
|
| 39 |
+
trust_remote_code=True,
|
| 40 |
+
)
|
| 41 |
+
model = model.eval().half().to(device)
|
| 42 |
+
transform = transforms.Compose(
|
| 43 |
+
[
|
| 44 |
+
transforms.Resize(FOREGROUND.image_size),
|
| 45 |
+
transforms.ToTensor(),
|
| 46 |
+
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
| 47 |
+
]
|
| 48 |
+
)
|
| 49 |
+
output: list[np.ndarray] = []
|
| 50 |
+
try:
|
| 51 |
+
for start in range(0, len(frames), batch_size):
|
| 52 |
+
images = [Image.fromarray(frame, mode="RGB") for frame in frames[start : start + batch_size]]
|
| 53 |
+
inputs = torch.stack([transform(image) for image in images]).to(device=device, dtype=torch.float16)
|
| 54 |
+
with torch.inference_mode():
|
| 55 |
+
predictions = model(inputs)[-1].sigmoid().cpu()
|
| 56 |
+
for image, prediction in zip(images, predictions, strict=True):
|
| 57 |
+
mask = to_pil_image(prediction).resize(image.size).convert("L")
|
| 58 |
+
output.append(np.asarray(mask, dtype=np.uint8).copy())
|
| 59 |
+
del inputs, predictions
|
| 60 |
+
finally:
|
| 61 |
+
del model
|
| 62 |
+
gc.collect()
|
| 63 |
+
if torch.cuda.is_available():
|
| 64 |
+
torch.cuda.empty_cache()
|
| 65 |
+
return np.stack(output)
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Camera and crop geometry."""
|
|
@@ -0,0 +1,250 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Canonical uniform camera rings and camera serialization."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from dataclasses import asdict, dataclass
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
|
| 9 |
+
from fdanyone.config import CAMERA, FRAMING, CameraConfig
|
| 10 |
+
|
| 11 |
+
WORLD_FRAME = {
|
| 12 |
+
"name": "canonical_human_world",
|
| 13 |
+
"handedness": "right",
|
| 14 |
+
"axes": {
|
| 15 |
+
"x": "right in the source-facing ring camera (the subject's anatomical left)",
|
| 16 |
+
"y": "up",
|
| 17 |
+
"z": "front; the subject initially faces +z and the source-facing camera lies on the +z side",
|
| 18 |
+
},
|
| 19 |
+
"origin": "initial root projected to the ground plane",
|
| 20 |
+
"units": "meters",
|
| 21 |
+
}
|
| 22 |
+
|
| 23 |
+
CAMERA_FRAME = {
|
| 24 |
+
"name": "opencv_camera",
|
| 25 |
+
"handedness": "right",
|
| 26 |
+
"axes": {"x": "image right", "y": "image down", "z": "forward from camera into the scene"},
|
| 27 |
+
"matrix_convention": "column vectors: x_camera = world_to_camera @ x_world_homogeneous",
|
| 28 |
+
"intrinsics_convention": "pixels with origin at the top-left",
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _normalize(vector: np.ndarray) -> np.ndarray:
|
| 33 |
+
norm = float(np.linalg.norm(vector))
|
| 34 |
+
if norm <= 1e-12:
|
| 35 |
+
raise ValueError("Cannot normalize a zero-length direction vector.")
|
| 36 |
+
return vector / norm
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def look_at_pytorch3d(position: np.ndarray, target: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
| 40 |
+
"""Return the column-vector w2c used by PyTorch3D's look_at_rotation."""
|
| 41 |
+
|
| 42 |
+
position = np.asarray(position, dtype=np.float64)
|
| 43 |
+
target = np.asarray(target, dtype=np.float64)
|
| 44 |
+
up = np.array([0.0, 1.0, 0.0], dtype=np.float64)
|
| 45 |
+
z_axis = _normalize(target - position)
|
| 46 |
+
x_axis = _normalize(np.cross(up, z_axis))
|
| 47 |
+
y_axis = _normalize(np.cross(z_axis, x_axis))
|
| 48 |
+
rotation = np.stack([x_axis, y_axis, z_axis], axis=0)
|
| 49 |
+
translation = -(rotation @ position)
|
| 50 |
+
return rotation, translation
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def pytorch3d_to_opencv(rotation: np.ndarray, translation: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
| 54 |
+
rotation_cv = np.asarray(rotation, dtype=np.float64).copy()
|
| 55 |
+
translation_cv = np.asarray(translation, dtype=np.float64).copy()
|
| 56 |
+
rotation_cv[:2] *= -1.0
|
| 57 |
+
translation_cv[:2] *= -1.0
|
| 58 |
+
return rotation_cv, translation_cv
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def homogeneous_w2c(rotation: np.ndarray, translation: np.ndarray) -> np.ndarray:
|
| 62 |
+
matrix = np.eye(4, dtype=np.float64)
|
| 63 |
+
matrix[:3, :3] = rotation
|
| 64 |
+
matrix[:3, 3] = translation
|
| 65 |
+
return matrix
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
@dataclass(frozen=True)
|
| 69 |
+
class Camera:
|
| 70 |
+
camera_id: int
|
| 71 |
+
layer_index: int
|
| 72 |
+
yaw_degrees: float
|
| 73 |
+
azimuth_degrees: float
|
| 74 |
+
pitch_degrees: float
|
| 75 |
+
position: tuple[float, float, float]
|
| 76 |
+
K: tuple[tuple[float, float, float], ...]
|
| 77 |
+
world_to_camera: tuple[tuple[float, float, float, float], ...]
|
| 78 |
+
camera_to_world: tuple[tuple[float, float, float, float], ...]
|
| 79 |
+
image_width: int
|
| 80 |
+
image_height: int
|
| 81 |
+
|
| 82 |
+
def to_dict(self) -> dict:
|
| 83 |
+
return asdict(self)
|
| 84 |
+
|
| 85 |
+
@classmethod
|
| 86 |
+
def from_dict(cls, payload: dict) -> Camera:
|
| 87 |
+
"""Invert ``to_dict``, coercing JSON lists back to the tuple fields."""
|
| 88 |
+
|
| 89 |
+
return cls(
|
| 90 |
+
camera_id=int(payload["camera_id"]),
|
| 91 |
+
layer_index=int(payload["layer_index"]),
|
| 92 |
+
yaw_degrees=float(payload["yaw_degrees"]),
|
| 93 |
+
azimuth_degrees=float(payload["azimuth_degrees"]),
|
| 94 |
+
pitch_degrees=float(payload["pitch_degrees"]),
|
| 95 |
+
position=tuple(float(value) for value in payload["position"]),
|
| 96 |
+
K=tuple(tuple(float(value) for value in row) for row in payload["K"]),
|
| 97 |
+
world_to_camera=tuple(
|
| 98 |
+
tuple(float(value) for value in row) for row in payload["world_to_camera"]
|
| 99 |
+
),
|
| 100 |
+
camera_to_world=tuple(
|
| 101 |
+
tuple(float(value) for value in row) for row in payload["camera_to_world"]
|
| 102 |
+
),
|
| 103 |
+
image_width=int(payload["image_width"]),
|
| 104 |
+
image_height=int(payload["image_height"]),
|
| 105 |
+
)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def reference_intrinsics(
|
| 109 |
+
image_height: int,
|
| 110 |
+
image_width: int,
|
| 111 |
+
max_render_height: int = 1280,
|
| 112 |
+
*,
|
| 113 |
+
focal_normalized: float = FRAMING.reference_focal_normalized,
|
| 114 |
+
) -> np.ndarray:
|
| 115 |
+
"""Reproduce the renderer's downscale-then-rescale intrinsic construction."""
|
| 116 |
+
|
| 117 |
+
divisor = 2
|
| 118 |
+
while image_height / divisor > max_render_height:
|
| 119 |
+
divisor += 1
|
| 120 |
+
render_height = image_height // divisor
|
| 121 |
+
render_width = image_width // divisor
|
| 122 |
+
scale = image_height / render_height
|
| 123 |
+
if focal_normalized <= 0:
|
| 124 |
+
raise ValueError("focal_normalized must be positive.")
|
| 125 |
+
focal = focal_normalized * render_height
|
| 126 |
+
intrinsic = np.array(
|
| 127 |
+
[[focal, 0.0, render_width / 2.0], [0.0, focal, render_height / 2.0], [0.0, 0.0, 1.0]],
|
| 128 |
+
dtype=np.float64,
|
| 129 |
+
)
|
| 130 |
+
intrinsic *= scale
|
| 131 |
+
intrinsic[2, 2] = 1.0
|
| 132 |
+
return intrinsic
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def camera_ring(
|
| 136 |
+
*,
|
| 137 |
+
center: np.ndarray,
|
| 138 |
+
front_direction: np.ndarray,
|
| 139 |
+
K: np.ndarray,
|
| 140 |
+
image_height: int,
|
| 141 |
+
image_width: int,
|
| 142 |
+
radius: float = FRAMING.reference_radius,
|
| 143 |
+
target_height: float = FRAMING.reference_target_height,
|
| 144 |
+
spec: CameraConfig = CAMERA,
|
| 145 |
+
start_yaw_degrees: float = 0.0,
|
| 146 |
+
yaw_span_degrees: float = 360.0,
|
| 147 |
+
layer_index: int = 0,
|
| 148 |
+
camera_id_offset: int = 0,
|
| 149 |
+
) -> tuple[Camera, ...]:
|
| 150 |
+
center = np.asarray(center, dtype=np.float64).copy()
|
| 151 |
+
front_direction = np.asarray(front_direction, dtype=np.float64)
|
| 152 |
+
front_xz = front_direction[[0, 2]]
|
| 153 |
+
front_azimuth = np.arctan2(front_xz[1], front_xz[0])
|
| 154 |
+
# Relative yaw zero faces the person's front; start_yaw chooses where the
|
| 155 |
+
# first camera lies and IDs then advance uniformly through the span.
|
| 156 |
+
azimuth_start = front_azimuth + np.pi + np.deg2rad(start_yaw_degrees)
|
| 157 |
+
target = center.copy()
|
| 158 |
+
if radius <= 0:
|
| 159 |
+
raise ValueError("radius must be positive.")
|
| 160 |
+
target[1] = target_height
|
| 161 |
+
camera_height = target_height + radius * np.tan(np.deg2rad(spec.pitch_degrees))
|
| 162 |
+
|
| 163 |
+
cameras: list[Camera] = []
|
| 164 |
+
yaw_span_radians = np.deg2rad(yaw_span_degrees)
|
| 165 |
+
for view_index in range(spec.count):
|
| 166 |
+
camera_id = camera_id_offset + view_index
|
| 167 |
+
yaw_degrees = start_yaw_degrees + view_index / spec.count * yaw_span_degrees
|
| 168 |
+
azimuth = azimuth_start + view_index / spec.count * yaw_span_radians
|
| 169 |
+
position = np.array(
|
| 170 |
+
[
|
| 171 |
+
center[0] + radius * np.cos(azimuth),
|
| 172 |
+
camera_height,
|
| 173 |
+
center[2] + radius * np.sin(azimuth),
|
| 174 |
+
],
|
| 175 |
+
dtype=np.float64,
|
| 176 |
+
)
|
| 177 |
+
rotation_p3d, translation_p3d = look_at_pytorch3d(position, target)
|
| 178 |
+
rotation_cv, translation_cv = pytorch3d_to_opencv(rotation_p3d, translation_p3d)
|
| 179 |
+
w2c = homogeneous_w2c(rotation_cv, translation_cv)
|
| 180 |
+
c2w = np.linalg.inv(w2c)
|
| 181 |
+
cameras.append(
|
| 182 |
+
Camera(
|
| 183 |
+
camera_id=camera_id,
|
| 184 |
+
layer_index=layer_index,
|
| 185 |
+
yaw_degrees=float(yaw_degrees),
|
| 186 |
+
azimuth_degrees=float(np.rad2deg(azimuth) % 360.0),
|
| 187 |
+
pitch_degrees=spec.pitch_degrees,
|
| 188 |
+
position=tuple(float(value) for value in position),
|
| 189 |
+
K=tuple(tuple(float(value) for value in row) for row in K),
|
| 190 |
+
world_to_camera=tuple(tuple(float(value) for value in row) for row in w2c),
|
| 191 |
+
camera_to_world=tuple(tuple(float(value) for value in row) for row in c2w),
|
| 192 |
+
image_width=image_width,
|
| 193 |
+
image_height=image_height,
|
| 194 |
+
)
|
| 195 |
+
)
|
| 196 |
+
return tuple(cameras)
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
def camera_grid(
|
| 200 |
+
*,
|
| 201 |
+
center: np.ndarray,
|
| 202 |
+
front_direction: np.ndarray,
|
| 203 |
+
K: np.ndarray,
|
| 204 |
+
image_height: int,
|
| 205 |
+
image_width: int,
|
| 206 |
+
views_per_layer: int,
|
| 207 |
+
layer_pitches: tuple[int, ...],
|
| 208 |
+
start_yaw: int,
|
| 209 |
+
yaw_span: int,
|
| 210 |
+
radius: float = FRAMING.reference_radius,
|
| 211 |
+
target_height: float = FRAMING.reference_target_height,
|
| 212 |
+
) -> tuple[Camera, ...]:
|
| 213 |
+
"""Create a layer-major target camera grid."""
|
| 214 |
+
|
| 215 |
+
return tuple(
|
| 216 |
+
camera
|
| 217 |
+
for layer_index, pitch in enumerate(layer_pitches)
|
| 218 |
+
for camera in camera_ring(
|
| 219 |
+
center=center,
|
| 220 |
+
front_direction=front_direction,
|
| 221 |
+
K=K,
|
| 222 |
+
image_height=image_height,
|
| 223 |
+
image_width=image_width,
|
| 224 |
+
radius=radius,
|
| 225 |
+
target_height=target_height,
|
| 226 |
+
spec=CameraConfig(count=views_per_layer, pitch_degrees=float(pitch)),
|
| 227 |
+
start_yaw_degrees=float(start_yaw),
|
| 228 |
+
yaw_span_degrees=float(yaw_span),
|
| 229 |
+
layer_index=layer_index,
|
| 230 |
+
camera_id_offset=layer_index * views_per_layer,
|
| 231 |
+
)
|
| 232 |
+
)
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
def project_points(points_world: np.ndarray, camera: Camera) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
| 236 |
+
points = np.asarray(points_world, dtype=np.float64)
|
| 237 |
+
w2c = np.asarray(camera.world_to_camera)
|
| 238 |
+
K = np.asarray(camera.K)
|
| 239 |
+
points_h = np.concatenate([points, np.ones((points.shape[0], 1), dtype=points.dtype)], axis=1)
|
| 240 |
+
points_camera = points_h @ w2c[:3].T
|
| 241 |
+
homogeneous = points_camera @ K.T
|
| 242 |
+
xy = homogeneous[:, :2] / np.maximum(homogeneous[:, 2:3], 1e-8)
|
| 243 |
+
valid = (
|
| 244 |
+
(points_camera[:, 2] > 0.0)
|
| 245 |
+
& (xy[:, 0] >= 0.0)
|
| 246 |
+
& (xy[:, 0] < camera.image_width)
|
| 247 |
+
& (xy[:, 1] >= 0.0)
|
| 248 |
+
& (xy[:, 1] < camera.image_height)
|
| 249 |
+
)
|
| 250 |
+
return xy.astype(np.float32), points_camera[:, 2].astype(np.float32), valid
|
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Deterministic source-mask and center-aspect crops."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import math
|
| 6 |
+
from dataclasses import dataclass
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
@dataclass(frozen=True)
|
| 12 |
+
class Crop:
|
| 13 |
+
top: int
|
| 14 |
+
left: int
|
| 15 |
+
height: int
|
| 16 |
+
width: int
|
| 17 |
+
original_height: int
|
| 18 |
+
original_width: int
|
| 19 |
+
output_height: int
|
| 20 |
+
output_width: int
|
| 21 |
+
|
| 22 |
+
@property
|
| 23 |
+
def scale_y(self) -> float:
|
| 24 |
+
return self.output_height / self.height
|
| 25 |
+
|
| 26 |
+
@property
|
| 27 |
+
def scale_x(self) -> float:
|
| 28 |
+
return self.output_width / self.width
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def center_crop(image_height: int, image_width: int, output_height: int, output_width: int) -> Crop:
|
| 32 |
+
aspect_ratio = output_height / output_width
|
| 33 |
+
if image_height / image_width >= aspect_ratio:
|
| 34 |
+
crop_width = image_width
|
| 35 |
+
crop_height = min(image_height, max(1, int(round(crop_width * aspect_ratio))))
|
| 36 |
+
else:
|
| 37 |
+
crop_height = image_height
|
| 38 |
+
crop_width = min(image_width, max(1, int(round(crop_height / aspect_ratio))))
|
| 39 |
+
return Crop(
|
| 40 |
+
top=(image_height - crop_height) // 2,
|
| 41 |
+
left=(image_width - crop_width) // 2,
|
| 42 |
+
height=crop_height,
|
| 43 |
+
width=crop_width,
|
| 44 |
+
original_height=image_height,
|
| 45 |
+
original_width=image_width,
|
| 46 |
+
output_height=output_height,
|
| 47 |
+
output_width=output_width,
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def mask_bounds(masks: np.ndarray, threshold: float) -> tuple[int, int, int, int] | None:
|
| 52 |
+
"""Return union bounds as left/top inclusive and right/bottom exclusive."""
|
| 53 |
+
|
| 54 |
+
values = np.asarray(masks)
|
| 55 |
+
if values.ndim == 4 and values.shape[1] == 1:
|
| 56 |
+
values = values[:, 0]
|
| 57 |
+
if values.ndim not in (2, 3):
|
| 58 |
+
raise ValueError(f"Expected masks [H,W] or [F,H,W], got {values.shape}.")
|
| 59 |
+
if not 0 <= threshold <= 1:
|
| 60 |
+
raise ValueError("threshold must be in [0, 1].")
|
| 61 |
+
cutoff = threshold * 255.0 if np.issubdtype(values.dtype, np.integer) else threshold
|
| 62 |
+
union = values.max(axis=0) if values.ndim == 3 else values
|
| 63 |
+
foreground = np.asarray(union > cutoff)
|
| 64 |
+
rows = np.flatnonzero(foreground.any(axis=1))
|
| 65 |
+
columns = np.flatnonzero(foreground.any(axis=0))
|
| 66 |
+
if rows.size == 0 or columns.size == 0:
|
| 67 |
+
return None
|
| 68 |
+
return int(columns[0]), int(rows[0]), int(columns[-1] + 1), int(rows[-1] + 1)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def expand_bounds(bounds, margins, image_height: int, image_width: int):
|
| 72 |
+
if bounds is None:
|
| 73 |
+
return None
|
| 74 |
+
xmin, ymin, xmax, ymax = (float(value) for value in bounds)
|
| 75 |
+
top, right, bottom, left = (float(value) for value in margins)
|
| 76 |
+
box_width, box_height = xmax - xmin, ymax - ymin
|
| 77 |
+
return (
|
| 78 |
+
max(0, int(math.floor(xmin - left * box_width))),
|
| 79 |
+
max(0, int(math.floor(ymin - top * box_height))),
|
| 80 |
+
min(image_width, int(math.ceil(xmax + right * box_width))),
|
| 81 |
+
min(image_height, int(math.ceil(ymax + bottom * box_height))),
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def crop_from_bounds(
|
| 86 |
+
*,
|
| 87 |
+
bounds: tuple[int, int, int, int] | None,
|
| 88 |
+
image_height: int,
|
| 89 |
+
image_width: int,
|
| 90 |
+
output_height: int,
|
| 91 |
+
output_width: int,
|
| 92 |
+
margins: tuple[float, float, float, float],
|
| 93 |
+
allow_upscale: bool = True,
|
| 94 |
+
) -> Crop:
|
| 95 |
+
bounds = expand_bounds(bounds, margins, image_height, image_width)
|
| 96 |
+
if bounds is None:
|
| 97 |
+
return center_crop(image_height, image_width, output_height, output_width)
|
| 98 |
+
|
| 99 |
+
xmin, ymin, xmax, ymax = bounds
|
| 100 |
+
required_height = float(ymax - ymin)
|
| 101 |
+
required_width = float(xmax - xmin)
|
| 102 |
+
if not allow_upscale:
|
| 103 |
+
required_height = max(required_height, min(output_height, image_height))
|
| 104 |
+
required_width = max(required_width, min(output_width, image_width))
|
| 105 |
+
|
| 106 |
+
aspect_ratio = output_height / output_width
|
| 107 |
+
crop_height = max(required_height, required_width * aspect_ratio)
|
| 108 |
+
crop_width = crop_height / aspect_ratio
|
| 109 |
+
max_height = min(float(image_height), float(image_width) * aspect_ratio)
|
| 110 |
+
crop_height = max(1, min(image_height, int(math.ceil(min(crop_height, max_height)))))
|
| 111 |
+
crop_width = max(1, min(image_width, int(math.ceil(crop_height / aspect_ratio))))
|
| 112 |
+
|
| 113 |
+
left_low = max(0.0, xmax - crop_width)
|
| 114 |
+
left_high = min(float(xmin), image_width - crop_width)
|
| 115 |
+
top_low = max(0.0, ymax - crop_height)
|
| 116 |
+
top_high = min(float(ymin), image_height - crop_height)
|
| 117 |
+
if left_low <= left_high:
|
| 118 |
+
left = int(math.floor((left_low + left_high) / 2.0))
|
| 119 |
+
else:
|
| 120 |
+
left = max(0, min(int(math.floor((xmin + xmax - crop_width) / 2.0)), image_width - crop_width))
|
| 121 |
+
if top_low <= top_high:
|
| 122 |
+
top = int(math.floor((top_low + top_high) / 2.0))
|
| 123 |
+
else:
|
| 124 |
+
top = max(0, min(int(math.floor((ymin + ymax - crop_height) / 2.0)), image_height - crop_height))
|
| 125 |
+
return Crop(
|
| 126 |
+
top=top,
|
| 127 |
+
left=left,
|
| 128 |
+
height=crop_height,
|
| 129 |
+
width=crop_width,
|
| 130 |
+
original_height=image_height,
|
| 131 |
+
original_width=image_width,
|
| 132 |
+
output_height=output_height,
|
| 133 |
+
output_width=output_width,
|
| 134 |
+
)
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def transform_intrinsics(K: np.ndarray, crop: Crop) -> np.ndarray:
|
| 138 |
+
transformed = np.asarray(K, dtype=np.float64).copy()
|
| 139 |
+
transformed[0, 2] -= crop.left
|
| 140 |
+
transformed[1, 2] -= crop.top
|
| 141 |
+
transformed[0] *= crop.scale_x
|
| 142 |
+
transformed[1] *= crop.scale_y
|
| 143 |
+
transformed[2, 2] = 1.0
|
| 144 |
+
return transformed
|
|
@@ -0,0 +1,579 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Source-aware static-camera framing for the canonical camera ring."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from collections.abc import Callable, Sequence
|
| 6 |
+
from dataclasses import asdict, dataclass
|
| 7 |
+
|
| 8 |
+
import cv2
|
| 9 |
+
import numpy as np
|
| 10 |
+
|
| 11 |
+
from fdanyone.config import FRAMING, FramingConfig
|
| 12 |
+
from fdanyone.geometry.cameras import Camera
|
| 13 |
+
|
| 14 |
+
FULL_BODY_BOTTOM = 0.82
|
| 15 |
+
CLOSE_UP_BOTTOM = 0.35
|
| 16 |
+
HALF_BODY_BOTTOM = 0.42
|
| 17 |
+
_ANATOMY_COORDS = (0.0, 0.18, 0.45, 0.72, 0.94, 1.0)
|
| 18 |
+
_FINGER_TOKENS = ("thumb", "index", "middle", "ring", "pinky")
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
@dataclass(frozen=True)
|
| 22 |
+
class InputFraming:
|
| 23 |
+
label: str
|
| 24 |
+
visible_body_bottom: float
|
| 25 |
+
visible_body_bottom_p50: float
|
| 26 |
+
visible_body_bottom_p80: float
|
| 27 |
+
visible_height_ratio: float
|
| 28 |
+
wrist_out_ratio_x: float
|
| 29 |
+
wrist_out_ratio_y: float
|
| 30 |
+
valid_frame_ratio: float
|
| 31 |
+
torso_valid_ratio: float
|
| 32 |
+
projection_alignment_error_ratio: float | None
|
| 33 |
+
confidence: float
|
| 34 |
+
|
| 35 |
+
def to_dict(self) -> dict[str, object]:
|
| 36 |
+
return {**asdict(self), "fmask_used": True}
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
@dataclass(frozen=True)
|
| 40 |
+
class AdaptiveThresholds:
|
| 41 |
+
closeup_strength: float
|
| 42 |
+
height_target_ratio: float
|
| 43 |
+
height_percentile: float
|
| 44 |
+
width_target_ratio: float
|
| 45 |
+
width_percentile: float
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
@dataclass(frozen=True)
|
| 49 |
+
class RadiusSolve:
|
| 50 |
+
radius: float
|
| 51 |
+
height_ratio: float
|
| 52 |
+
width_ratio: float
|
| 53 |
+
bound: str | None
|
| 54 |
+
limiting_constraint: str
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
@dataclass(frozen=True)
|
| 58 |
+
class FocalSolve:
|
| 59 |
+
focal_normalized: float
|
| 60 |
+
height_ratio: float
|
| 61 |
+
width_ratio: float
|
| 62 |
+
bound: str | None
|
| 63 |
+
limiting_constraint: str
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
@dataclass(frozen=True)
|
| 67 |
+
class SequenceFraming:
|
| 68 |
+
radius: float
|
| 69 |
+
target_height: float
|
| 70 |
+
focal_normalized: float
|
| 71 |
+
input: InputFraming
|
| 72 |
+
input_applied: bool
|
| 73 |
+
radius_solve: RadiusSolve
|
| 74 |
+
adaptive_thresholds: AdaptiveThresholds | None = None
|
| 75 |
+
focal_solve: FocalSolve | None = None
|
| 76 |
+
cutoff_ratio: float | None = None
|
| 77 |
+
target_bound: str | None = None
|
| 78 |
+
|
| 79 |
+
def to_dict(self) -> dict[str, object]:
|
| 80 |
+
return {
|
| 81 |
+
"method": "sequence_input_profile_static_radius_target_focal",
|
| 82 |
+
"radius": self.radius,
|
| 83 |
+
"target_height": self.target_height,
|
| 84 |
+
"focal_normalized": self.focal_normalized,
|
| 85 |
+
"input_framing_applied": self.input_applied,
|
| 86 |
+
"input_framing": self.input.to_dict(),
|
| 87 |
+
"adaptive_thresholds": (None if self.adaptive_thresholds is None else asdict(self.adaptive_thresholds)),
|
| 88 |
+
"radius_solver": asdict(self.radius_solve),
|
| 89 |
+
"focal_solver": None if self.focal_solve is None else asdict(self.focal_solve),
|
| 90 |
+
"cutoff_ratio": self.cutoff_ratio,
|
| 91 |
+
"target_bound": self.target_bound,
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def _name_index(names: Sequence[str]) -> dict[str, int]:
|
| 96 |
+
normalized = [str(name).strip().lower().replace("_", "-") for name in names]
|
| 97 |
+
mapping = dict(zip(normalized, range(len(normalized)), strict=True))
|
| 98 |
+
if len(mapping) != len(normalized):
|
| 99 |
+
raise ValueError("Keypoint names must be unique after normalization.")
|
| 100 |
+
return mapping
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def _required_indices(names: Sequence[str]) -> dict[str, int]:
|
| 104 |
+
required = ["nose", "left-eye", "right-eye", "left-ear", "right-ear", "neck"]
|
| 105 |
+
for side in ("left", "right"):
|
| 106 |
+
required.extend(
|
| 107 |
+
f"{side}-{part}"
|
| 108 |
+
for part in (
|
| 109 |
+
"shoulder",
|
| 110 |
+
"hip",
|
| 111 |
+
"knee",
|
| 112 |
+
"ankle",
|
| 113 |
+
"big-toe-tip",
|
| 114 |
+
"small-toe-tip",
|
| 115 |
+
"heel",
|
| 116 |
+
)
|
| 117 |
+
)
|
| 118 |
+
mapping = _name_index(names)
|
| 119 |
+
missing = [name for name in required if name not in mapping]
|
| 120 |
+
if missing:
|
| 121 |
+
raise ValueError(f"Missing framing keypoints: {missing}.")
|
| 122 |
+
return {name: mapping[name] for name in required}
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def _sample_path(anchors: Sequence[np.ndarray], samples_per_segment: int) -> tuple[np.ndarray, np.ndarray]:
|
| 126 |
+
point_chunks: list[np.ndarray] = []
|
| 127 |
+
coordinate_chunks: list[np.ndarray] = []
|
| 128 |
+
for index, (start, end) in enumerate(zip(anchors, anchors[1:], strict=False)):
|
| 129 |
+
endpoint = index == len(anchors) - 2
|
| 130 |
+
weights = np.linspace(
|
| 131 |
+
0.0,
|
| 132 |
+
1.0,
|
| 133 |
+
samples_per_segment + int(endpoint),
|
| 134 |
+
endpoint=endpoint,
|
| 135 |
+
dtype=np.float64,
|
| 136 |
+
)
|
| 137 |
+
point_chunks.append(start[:, None] * (1.0 - weights[None, :, None]) + end[:, None] * weights[None, :, None])
|
| 138 |
+
coordinate_chunks.append(_ANATOMY_COORDS[index] * (1.0 - weights) + _ANATOMY_COORDS[index + 1] * weights)
|
| 139 |
+
return np.concatenate(point_chunks, axis=1), np.concatenate(coordinate_chunks)
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def anatomy_samples(
|
| 143 |
+
keypoints: np.ndarray,
|
| 144 |
+
names: Sequence[str],
|
| 145 |
+
samples_per_segment: int = 8,
|
| 146 |
+
) -> tuple[np.ndarray, np.ndarray]:
|
| 147 |
+
points = np.asarray(keypoints, dtype=np.float64)
|
| 148 |
+
if points.ndim != 3 or points.shape[1:] != (len(names), 3) or not np.isfinite(points).all():
|
| 149 |
+
raise ValueError(f"Expected finite keypoints [frames,{len(names)},3], got {points.shape}.")
|
| 150 |
+
if samples_per_segment < 2:
|
| 151 |
+
raise ValueError("samples_per_segment must be at least two.")
|
| 152 |
+
ids = _required_indices(names)
|
| 153 |
+
face = np.mean(
|
| 154 |
+
points[:, [ids[name] for name in ("nose", "left-eye", "right-eye", "left-ear", "right-ear")]], axis=1
|
| 155 |
+
)
|
| 156 |
+
head_top = face + 0.65 * (face - points[:, ids["neck"]])
|
| 157 |
+
paths: list[np.ndarray] = []
|
| 158 |
+
coordinates: list[np.ndarray] = []
|
| 159 |
+
for side in ("left", "right"):
|
| 160 |
+
foot = np.mean(
|
| 161 |
+
points[:, [ids[f"{side}-{part}"] for part in ("big-toe-tip", "small-toe-tip", "heel")]],
|
| 162 |
+
axis=1,
|
| 163 |
+
)
|
| 164 |
+
path, path_coordinates = _sample_path(
|
| 165 |
+
[
|
| 166 |
+
head_top,
|
| 167 |
+
points[:, ids[f"{side}-shoulder"]],
|
| 168 |
+
points[:, ids[f"{side}-hip"]],
|
| 169 |
+
points[:, ids[f"{side}-knee"]],
|
| 170 |
+
points[:, ids[f"{side}-ankle"]],
|
| 171 |
+
foot,
|
| 172 |
+
],
|
| 173 |
+
samples_per_segment,
|
| 174 |
+
)
|
| 175 |
+
paths.append(path)
|
| 176 |
+
coordinates.append(path_coordinates)
|
| 177 |
+
return np.concatenate(paths, axis=1), np.concatenate(coordinates)
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def project_incam(points: np.ndarray, intrinsics: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
| 181 |
+
points = np.asarray(points, dtype=np.float64)
|
| 182 |
+
cameras = np.asarray(intrinsics, dtype=np.float64)
|
| 183 |
+
if points.ndim != 3 or points.shape[-1] != 3:
|
| 184 |
+
raise ValueError(f"Expected points [frames,keypoints,3], got {points.shape}.")
|
| 185 |
+
if cameras.shape == (3, 3):
|
| 186 |
+
cameras = np.broadcast_to(cameras, (points.shape[0], 3, 3))
|
| 187 |
+
if cameras.shape == (1, 3, 3):
|
| 188 |
+
cameras = np.broadcast_to(cameras, (points.shape[0], 3, 3))
|
| 189 |
+
if cameras.shape != (points.shape[0], 3, 3):
|
| 190 |
+
raise ValueError(f"Expected intrinsics [{points.shape[0]},3,3], got {cameras.shape}.")
|
| 191 |
+
if not np.isfinite(points).all() or not np.isfinite(cameras).all():
|
| 192 |
+
raise ValueError("Projection inputs must be finite.")
|
| 193 |
+
homogeneous = np.einsum("fij,fkj->fki", cameras, points)
|
| 194 |
+
depths = points[..., 2]
|
| 195 |
+
xy = homogeneous[..., :2] / np.maximum(homogeneous[..., 2:3], 1e-8)
|
| 196 |
+
return xy, depths
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
def _inside(xy: np.ndarray, depths: np.ndarray, width: int, height: int) -> np.ndarray:
|
| 200 |
+
return (depths > 1e-6) & (xy[..., 0] >= 0) & (xy[..., 0] < width) & (xy[..., 1] >= 0) & (xy[..., 1] < height)
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def _mask_support(xy: np.ndarray, inside: np.ndarray, masks: np.ndarray) -> np.ndarray:
|
| 204 |
+
masks = np.asarray(masks)
|
| 205 |
+
if masks.ndim != 3 or masks.shape[0] != xy.shape[0]:
|
| 206 |
+
raise ValueError(f"Expected masks [frames,height,width], got {masks.shape}.")
|
| 207 |
+
height, width = masks.shape[1:]
|
| 208 |
+
patch_radius = max(3, int(round(min(width, height) * 0.005)))
|
| 209 |
+
kernel = np.ones((2 * patch_radius + 1,) * 2, dtype=np.uint8)
|
| 210 |
+
support = np.zeros_like(inside)
|
| 211 |
+
for frame_index, mask in enumerate(masks):
|
| 212 |
+
dilated = cv2.dilate((mask >= 64).astype(np.uint8), kernel)
|
| 213 |
+
candidates = np.flatnonzero(inside[frame_index])
|
| 214 |
+
if candidates.size:
|
| 215 |
+
pixels = np.rint(xy[frame_index, candidates]).astype(np.int64)
|
| 216 |
+
pixels[:, 0] = np.clip(pixels[:, 0], 0, width - 1)
|
| 217 |
+
pixels[:, 1] = np.clip(pixels[:, 1], 0, height - 1)
|
| 218 |
+
support[frame_index, candidates] = dilated[pixels[:, 1], pixels[:, 0]] > 0
|
| 219 |
+
return support
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
def _torso_valid(vitpose: np.ndarray, width: int, height: int) -> np.ndarray:
|
| 223 |
+
detector = np.asarray(vitpose, dtype=np.float64)
|
| 224 |
+
if detector.ndim != 3 or detector.shape[1:] != (17, 3):
|
| 225 |
+
raise ValueError(f"Expected VitPose [frames,17,3], got {detector.shape}.")
|
| 226 |
+
valid = (
|
| 227 |
+
(detector[..., 2] >= 0.3)
|
| 228 |
+
& (detector[..., 0] >= 0)
|
| 229 |
+
& (detector[..., 0] < width)
|
| 230 |
+
& (detector[..., 1] >= 0)
|
| 231 |
+
& (detector[..., 1] < height)
|
| 232 |
+
)
|
| 233 |
+
shoulders = valid[:, 5] & valid[:, 6]
|
| 234 |
+
torso = shoulders & (valid[:, 11] | valid[:, 12])
|
| 235 |
+
return shoulders if float(torso.mean()) < 0.5 else torso
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
def _alignment_error(
|
| 239 |
+
projected: np.ndarray,
|
| 240 |
+
names: Sequence[str],
|
| 241 |
+
vitpose: np.ndarray,
|
| 242 |
+
width: int,
|
| 243 |
+
height: int,
|
| 244 |
+
) -> float | None:
|
| 245 |
+
coco_names = (
|
| 246 |
+
"nose",
|
| 247 |
+
"left-eye",
|
| 248 |
+
"right-eye",
|
| 249 |
+
"left-ear",
|
| 250 |
+
"right-ear",
|
| 251 |
+
"left-shoulder",
|
| 252 |
+
"right-shoulder",
|
| 253 |
+
"left-elbow",
|
| 254 |
+
"right-elbow",
|
| 255 |
+
"left-wrist",
|
| 256 |
+
"right-wrist",
|
| 257 |
+
"left-hip",
|
| 258 |
+
"right-hip",
|
| 259 |
+
"left-knee",
|
| 260 |
+
"right-knee",
|
| 261 |
+
"left-ankle",
|
| 262 |
+
"right-ankle",
|
| 263 |
+
)
|
| 264 |
+
mapping = _name_index(names)
|
| 265 |
+
if any(name not in mapping for name in coco_names):
|
| 266 |
+
return None
|
| 267 |
+
detector = np.asarray(vitpose, dtype=np.float64)
|
| 268 |
+
valid = (
|
| 269 |
+
(detector[..., 2] >= 0.3)
|
| 270 |
+
& (detector[..., 0] >= 0)
|
| 271 |
+
& (detector[..., 0] < width)
|
| 272 |
+
& (detector[..., 1] >= 0)
|
| 273 |
+
& (detector[..., 1] < height)
|
| 274 |
+
)
|
| 275 |
+
if not valid.any():
|
| 276 |
+
return None
|
| 277 |
+
errors = np.linalg.norm(projected[:, [mapping[name] for name in coco_names]] - detector[..., :2], axis=-1)
|
| 278 |
+
return float(np.median(errors[valid]) / height)
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
def analyze_input_framing(
|
| 282 |
+
incam_keypoints: np.ndarray,
|
| 283 |
+
names: Sequence[str],
|
| 284 |
+
intrinsics: np.ndarray,
|
| 285 |
+
vitpose: np.ndarray,
|
| 286 |
+
masks: np.ndarray,
|
| 287 |
+
) -> InputFraming:
|
| 288 |
+
masks = np.asarray(masks)
|
| 289 |
+
if masks.ndim != 3:
|
| 290 |
+
raise ValueError(f"Expected masks [frames,height,width], got {masks.shape}.")
|
| 291 |
+
frame_count, height, width = masks.shape
|
| 292 |
+
if np.asarray(incam_keypoints).shape[0] != frame_count or np.asarray(vitpose).shape[0] != frame_count:
|
| 293 |
+
raise ValueError("Input-framing arrays do not share one frame count.")
|
| 294 |
+
|
| 295 |
+
anatomy, coordinates = anatomy_samples(incam_keypoints, names)
|
| 296 |
+
anatomy_xy, anatomy_depth = project_incam(anatomy, intrinsics)
|
| 297 |
+
support = _mask_support(anatomy_xy, _inside(anatomy_xy, anatomy_depth, width, height), masks)
|
| 298 |
+
torso_valid = _torso_valid(vitpose, width, height)
|
| 299 |
+
bottoms = np.full(frame_count, np.nan)
|
| 300 |
+
visible_heights = np.full(frame_count, np.nan)
|
| 301 |
+
for frame_index in range(frame_count):
|
| 302 |
+
valid = np.flatnonzero(support[frame_index])
|
| 303 |
+
if valid.size and torso_valid[frame_index]:
|
| 304 |
+
bottoms[frame_index] = float(coordinates[valid].max())
|
| 305 |
+
y = anatomy_xy[frame_index, valid, 1]
|
| 306 |
+
visible_heights[frame_index] = float(np.clip((y.max() - y.min()) / height, 0.0, 1.0))
|
| 307 |
+
valid_frames = np.isfinite(bottoms)
|
| 308 |
+
if not valid_frames.any():
|
| 309 |
+
raise ValueError("No valid frames for input framing analysis.")
|
| 310 |
+
|
| 311 |
+
projected, depths = project_incam(incam_keypoints, intrinsics)
|
| 312 |
+
mapping = _name_index(names)
|
| 313 |
+
wrists = [mapping[name] for name in ("left-wrist", "right-wrist")]
|
| 314 |
+
wrist_xy = projected[:, wrists]
|
| 315 |
+
wrist_depth = depths[:, wrists]
|
| 316 |
+
wrist_out_x = (wrist_depth <= 1e-6) | (wrist_xy[..., 0] < 0) | (wrist_xy[..., 0] >= width)
|
| 317 |
+
wrist_out_y = (wrist_depth <= 1e-6) | (wrist_xy[..., 1] < 0) | (wrist_xy[..., 1] >= height)
|
| 318 |
+
alignment = _alignment_error(projected, names, vitpose, width, height)
|
| 319 |
+
valid_ratio = float(valid_frames.mean())
|
| 320 |
+
torso_ratio = float(torso_valid.mean())
|
| 321 |
+
alignment_score = 0.6 if alignment is None else float(np.clip(1.0 - alignment / 0.15, 0.0, 1.0))
|
| 322 |
+
confidence = float(np.clip(0.45 * valid_ratio + 0.35 * torso_ratio + 0.20 * alignment_score, 0.0, 1.0))
|
| 323 |
+
bottom = float(np.percentile(bottoms[valid_frames], 20.0))
|
| 324 |
+
if bottom >= FULL_BODY_BOTTOM:
|
| 325 |
+
label = "full_body"
|
| 326 |
+
elif bottom >= HALF_BODY_BOTTOM:
|
| 327 |
+
label = "half_body"
|
| 328 |
+
else:
|
| 329 |
+
label = "close_up"
|
| 330 |
+
finite_heights = visible_heights[np.isfinite(visible_heights)]
|
| 331 |
+
return InputFraming(
|
| 332 |
+
label=label,
|
| 333 |
+
visible_body_bottom=round(bottom, 6),
|
| 334 |
+
visible_body_bottom_p50=round(float(np.percentile(bottoms[valid_frames], 50.0)), 6),
|
| 335 |
+
visible_body_bottom_p80=round(float(np.percentile(bottoms[valid_frames], 80.0)), 6),
|
| 336 |
+
visible_height_ratio=round(float(np.percentile(finite_heights, 50.0)) if finite_heights.size else 0.0, 6),
|
| 337 |
+
wrist_out_ratio_x=round(float(np.any(wrist_out_x, axis=1).mean()), 6),
|
| 338 |
+
wrist_out_ratio_y=round(float(np.any(wrist_out_y, axis=1).mean()), 6),
|
| 339 |
+
valid_frame_ratio=round(valid_ratio, 6),
|
| 340 |
+
torso_valid_ratio=round(torso_ratio, 6),
|
| 341 |
+
projection_alignment_error_ratio=None if alignment is None else round(alignment, 6),
|
| 342 |
+
confidence=round(confidence, 6),
|
| 343 |
+
)
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
def _camera_coordinates(points: np.ndarray, cameras: Sequence[Camera]) -> np.ndarray:
|
| 347 |
+
world = np.asarray(points, dtype=np.float64)
|
| 348 |
+
points_h = np.concatenate([world, np.ones((*world.shape[:-1], 1), dtype=np.float64)], axis=-1)
|
| 349 |
+
w2c = np.asarray([camera.world_to_camera for camera in cameras], dtype=np.float64)
|
| 350 |
+
return np.einsum("cij,fkj->cfki", w2c[:, :3], points_h)
|
| 351 |
+
|
| 352 |
+
|
| 353 |
+
def projected_axis_ratios(
|
| 354 |
+
points: np.ndarray,
|
| 355 |
+
cameras: Sequence[Camera],
|
| 356 |
+
focal_normalized: float,
|
| 357 |
+
*,
|
| 358 |
+
axis: int,
|
| 359 |
+
axis_scale: float = 1.0,
|
| 360 |
+
) -> np.ndarray:
|
| 361 |
+
camera_points = _camera_coordinates(points, cameras)
|
| 362 |
+
depths = camera_points[..., 2]
|
| 363 |
+
coordinates = focal_normalized * camera_points[..., axis] / np.maximum(depths, 1e-6) / axis_scale
|
| 364 |
+
ratios = coordinates.max(axis=2) - coordinates.min(axis=2)
|
| 365 |
+
ratios[np.any(depths <= 1e-6, axis=2)] = np.inf
|
| 366 |
+
return ratios.max(axis=0)
|
| 367 |
+
|
| 368 |
+
|
| 369 |
+
def _selected(points: np.ndarray, names: Sequence[str], *, exclude_hands: bool, exclude_fingers: bool) -> np.ndarray:
|
| 370 |
+
excluded = _FINGER_TOKENS
|
| 371 |
+
if exclude_hands:
|
| 372 |
+
excluded = ("wrist", *_FINGER_TOKENS)
|
| 373 |
+
elif not exclude_fingers:
|
| 374 |
+
excluded = ()
|
| 375 |
+
ids = [index for index, name in enumerate(names) if not any(token in name.lower() for token in excluded)]
|
| 376 |
+
return np.asarray(points)[:, ids]
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
def solve_radius(
|
| 380 |
+
points: np.ndarray,
|
| 381 |
+
names: Sequence[str],
|
| 382 |
+
camera_factory: Callable[[float, float], Sequence[Camera]],
|
| 383 |
+
aspect_ratio: float,
|
| 384 |
+
spec: FramingConfig = FRAMING,
|
| 385 |
+
) -> RadiusSolve:
|
| 386 |
+
height_points = _selected(points, names, exclude_hands=True, exclude_fingers=True)
|
| 387 |
+
width_points = _selected(points, names, exclude_hands=False, exclude_fingers=True)
|
| 388 |
+
|
| 389 |
+
def evaluate(radius: float) -> tuple[float, float, float]:
|
| 390 |
+
cameras = camera_factory(radius, spec.reference_target_height)
|
| 391 |
+
heights = projected_axis_ratios(height_points, cameras, spec.reference_focal_normalized, axis=1)
|
| 392 |
+
widths = projected_axis_ratios(
|
| 393 |
+
width_points, cameras, spec.reference_focal_normalized, axis=0, axis_scale=aspect_ratio
|
| 394 |
+
)
|
| 395 |
+
height = float(np.percentile(heights, spec.height_percentile))
|
| 396 |
+
width = float(np.percentile(widths, spec.width_percentile))
|
| 397 |
+
return max(height / spec.height_target_ratio, width / spec.width_target_ratio), height, width
|
| 398 |
+
|
| 399 |
+
def result(radius: float, values: tuple[float, float, float], bound: str | None) -> RadiusSolve:
|
| 400 |
+
_, height, width = values
|
| 401 |
+
limiting = "width" if width / spec.width_target_ratio > height / spec.height_target_ratio else "height"
|
| 402 |
+
return RadiusSolve(radius, height, width, bound, limiting)
|
| 403 |
+
|
| 404 |
+
lower_values = evaluate(spec.min_radius)
|
| 405 |
+
if lower_values[0] <= 1.0:
|
| 406 |
+
return result(spec.min_radius, lower_values, "min")
|
| 407 |
+
upper_values = evaluate(spec.max_radius)
|
| 408 |
+
if upper_values[0] > 1.0:
|
| 409 |
+
return result(spec.max_radius, upper_values, "max")
|
| 410 |
+
lower, upper, best = spec.min_radius, spec.max_radius, upper_values
|
| 411 |
+
for _ in range(24):
|
| 412 |
+
midpoint = (lower + upper) / 2
|
| 413 |
+
values = evaluate(midpoint)
|
| 414 |
+
if values[0] > 1.0:
|
| 415 |
+
lower = midpoint
|
| 416 |
+
else:
|
| 417 |
+
upper, best = midpoint, values
|
| 418 |
+
if values[0] <= 1.0 and abs(values[0] - 1.0) <= 1e-4:
|
| 419 |
+
break
|
| 420 |
+
return result(upper, best, None)
|
| 421 |
+
|
| 422 |
+
|
| 423 |
+
def adaptive_thresholds(profile: InputFraming) -> AdaptiveThresholds:
|
| 424 |
+
strength = float(
|
| 425 |
+
np.clip((FULL_BODY_BOTTOM - profile.visible_body_bottom) / (FULL_BODY_BOTTOM - CLOSE_UP_BOTTOM), 0.0, 1.0)
|
| 426 |
+
)
|
| 427 |
+
return AdaptiveThresholds(strength, 0.80 + 0.12 * strength, 95.0, 0.90 + 0.20 * strength, 80.0 - 30.0 * strength)
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
def _cutoff_points(points: np.ndarray, coordinates: np.ndarray, cutoff: float) -> np.ndarray:
|
| 431 |
+
path_length = points.shape[1] // 2
|
| 432 |
+
output = []
|
| 433 |
+
for offset in (0, path_length):
|
| 434 |
+
local_coordinates = coordinates[offset : offset + path_length]
|
| 435 |
+
local_points = points[:, offset : offset + path_length]
|
| 436 |
+
after = int(np.searchsorted(local_coordinates, cutoff, side="right"))
|
| 437 |
+
if after == 0:
|
| 438 |
+
output.append(local_points[:, 0])
|
| 439 |
+
elif after >= len(local_coordinates):
|
| 440 |
+
output.append(local_points[:, -1])
|
| 441 |
+
else:
|
| 442 |
+
before = after - 1
|
| 443 |
+
weight = (cutoff - local_coordinates[before]) / (local_coordinates[after] - local_coordinates[before])
|
| 444 |
+
output.append(local_points[:, before] * (1.0 - weight) + local_points[:, after] * weight)
|
| 445 |
+
return np.stack(output, axis=1)
|
| 446 |
+
|
| 447 |
+
|
| 448 |
+
def _visible_anatomy(
|
| 449 |
+
points: np.ndarray,
|
| 450 |
+
names: Sequence[str],
|
| 451 |
+
bottom: float,
|
| 452 |
+
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
| 453 |
+
samples, coordinates = anatomy_samples(points, names)
|
| 454 |
+
core = samples[:, coordinates <= bottom + 1e-8]
|
| 455 |
+
cutoff = _cutoff_points(samples, coordinates, bottom)
|
| 456 |
+
if not np.any(np.isclose(coordinates, bottom, atol=1e-8)):
|
| 457 |
+
core = np.concatenate([core, cutoff], axis=1)
|
| 458 |
+
mapping = _name_index(names)
|
| 459 |
+
arm_tokens = ("shoulder", "acromion", "elbow", "olecranon", "cubital-fossa", "wrist")
|
| 460 |
+
arms = np.asarray(points)[
|
| 461 |
+
:, [index for name, index in mapping.items() if any(token in name for token in arm_tokens)]
|
| 462 |
+
]
|
| 463 |
+
return (
|
| 464 |
+
np.asarray(core, dtype=np.float32),
|
| 465 |
+
np.asarray(np.concatenate([core, arms], axis=1), dtype=np.float32),
|
| 466 |
+
np.asarray(cutoff, dtype=np.float32),
|
| 467 |
+
)
|
| 468 |
+
|
| 469 |
+
|
| 470 |
+
def _solve_focal(
|
| 471 |
+
core: np.ndarray,
|
| 472 |
+
width_points: np.ndarray,
|
| 473 |
+
cameras: Sequence[Camera],
|
| 474 |
+
thresholds: AdaptiveThresholds,
|
| 475 |
+
aspect_ratio: float,
|
| 476 |
+
spec: FramingConfig,
|
| 477 |
+
) -> tuple[FocalSolve, np.ndarray]:
|
| 478 |
+
unit_heights = projected_axis_ratios(core, cameras, 1.0, axis=1)
|
| 479 |
+
unit_widths = projected_axis_ratios(width_points, cameras, 1.0, axis=0, axis_scale=aspect_ratio)
|
| 480 |
+
unit_height = float(np.percentile(unit_heights, thresholds.height_percentile))
|
| 481 |
+
unit_width = float(np.percentile(unit_widths, thresholds.width_percentile))
|
| 482 |
+
height_focal = thresholds.height_target_ratio / unit_height
|
| 483 |
+
width_focal = thresholds.width_target_ratio / unit_width if unit_width > 0 else np.inf
|
| 484 |
+
unconstrained = min(height_focal, width_focal)
|
| 485 |
+
focal = float(np.clip(unconstrained, spec.reference_focal_normalized, spec.max_focal_normalized))
|
| 486 |
+
bound = (
|
| 487 |
+
"min"
|
| 488 |
+
if unconstrained < spec.reference_focal_normalized
|
| 489 |
+
else "max"
|
| 490 |
+
if unconstrained > spec.max_focal_normalized
|
| 491 |
+
else None
|
| 492 |
+
)
|
| 493 |
+
return (
|
| 494 |
+
FocalSolve(
|
| 495 |
+
focal,
|
| 496 |
+
float(np.percentile(unit_heights * focal, thresholds.height_percentile)),
|
| 497 |
+
float(np.percentile(unit_widths * focal, thresholds.width_percentile)),
|
| 498 |
+
bound,
|
| 499 |
+
"width" if width_focal < height_focal else "height",
|
| 500 |
+
),
|
| 501 |
+
unit_heights,
|
| 502 |
+
)
|
| 503 |
+
|
| 504 |
+
|
| 505 |
+
def _cutoff_ratio(points: np.ndarray, cameras: Sequence[Camera], focal: float, percentile: float) -> float:
|
| 506 |
+
camera_points = _camera_coordinates(points, cameras)
|
| 507 |
+
positions = 0.5 + focal * camera_points[..., 1] / np.maximum(camera_points[..., 2], 1e-6)
|
| 508 |
+
positions[camera_points[..., 2] <= 1e-6] = np.nan
|
| 509 |
+
return float(np.percentile(positions[np.isfinite(positions)], percentile))
|
| 510 |
+
|
| 511 |
+
|
| 512 |
+
def solve_sequence_framing(
|
| 513 |
+
keypoints: np.ndarray,
|
| 514 |
+
names: Sequence[str],
|
| 515 |
+
profile: InputFraming,
|
| 516 |
+
camera_factory: Callable[[float, float], Sequence[Camera]],
|
| 517 |
+
aspect_ratio: float,
|
| 518 |
+
spec: FramingConfig = FRAMING,
|
| 519 |
+
) -> SequenceFraming:
|
| 520 |
+
radius_result = solve_radius(keypoints, names, camera_factory, aspect_ratio, spec)
|
| 521 |
+
applied = profile.confidence >= spec.input_min_confidence
|
| 522 |
+
thresholds = adaptive_thresholds(profile) if applied else None
|
| 523 |
+
if not applied or thresholds is None or thresholds.closeup_strength <= 0:
|
| 524 |
+
return SequenceFraming(
|
| 525 |
+
radius_result.radius,
|
| 526 |
+
spec.reference_target_height,
|
| 527 |
+
spec.reference_focal_normalized,
|
| 528 |
+
profile,
|
| 529 |
+
applied,
|
| 530 |
+
radius_result,
|
| 531 |
+
thresholds,
|
| 532 |
+
)
|
| 533 |
+
|
| 534 |
+
core, width_points, cutoff = _visible_anatomy(keypoints, names, profile.visible_body_bottom)
|
| 535 |
+
centers = (core[..., 1].min(axis=1) + core[..., 1].max(axis=1)) / 2
|
| 536 |
+
anatomical_target = float(np.median(centers))
|
| 537 |
+
alignment = min(1.0, 2.0 * thresholds.closeup_strength)
|
| 538 |
+
initial_target = spec.reference_target_height * (1.0 - alignment) + anatomical_target * alignment
|
| 539 |
+
|
| 540 |
+
def evaluate(target_height: float) -> tuple[FocalSolve, float]:
|
| 541 |
+
cameras = camera_factory(radius_result.radius, target_height)
|
| 542 |
+
focal, _ = _solve_focal(core, width_points, cameras, thresholds, aspect_ratio, spec)
|
| 543 |
+
return focal, _cutoff_ratio(cutoff, cameras, focal.focal_normalized, spec.cutoff_percentile)
|
| 544 |
+
|
| 545 |
+
lower, upper = initial_target - 0.5, initial_target + 0.5
|
| 546 |
+
lower_value, upper_value = evaluate(lower), evaluate(upper)
|
| 547 |
+
increasing = upper_value[1] > lower_value[1]
|
| 548 |
+
if not min(lower_value[1], upper_value[1]) <= spec.cutoff_target_ratio <= max(lower_value[1], upper_value[1]):
|
| 549 |
+
if abs(lower_value[1] - spec.cutoff_target_ratio) <= abs(upper_value[1] - spec.cutoff_target_ratio):
|
| 550 |
+
target, value, target_bound = lower, lower_value, "min"
|
| 551 |
+
else:
|
| 552 |
+
target, value, target_bound = upper, upper_value, "max"
|
| 553 |
+
else:
|
| 554 |
+
target, value, target_bound = lower, lower_value, None
|
| 555 |
+
for _ in range(20):
|
| 556 |
+
midpoint = (lower + upper) / 2
|
| 557 |
+
candidate = evaluate(midpoint)
|
| 558 |
+
error = candidate[1] - spec.cutoff_target_ratio
|
| 559 |
+
if abs(error) < abs(value[1] - spec.cutoff_target_ratio):
|
| 560 |
+
target, value = midpoint, candidate
|
| 561 |
+
if abs(error) <= 1e-4:
|
| 562 |
+
break
|
| 563 |
+
if (error < 0) == increasing:
|
| 564 |
+
lower = midpoint
|
| 565 |
+
else:
|
| 566 |
+
upper = midpoint
|
| 567 |
+
|
| 568 |
+
return SequenceFraming(
|
| 569 |
+
radius_result.radius,
|
| 570 |
+
target,
|
| 571 |
+
value[0].focal_normalized,
|
| 572 |
+
profile,
|
| 573 |
+
True,
|
| 574 |
+
radius_result,
|
| 575 |
+
thresholds,
|
| 576 |
+
value[0],
|
| 577 |
+
value[1],
|
| 578 |
+
target_bound,
|
| 579 |
+
)
|
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Filesystem helpers for crash-safe result publication."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import errno
|
| 6 |
+
import json
|
| 7 |
+
import os
|
| 8 |
+
import shutil
|
| 9 |
+
import time
|
| 10 |
+
import uuid
|
| 11 |
+
from contextlib import AbstractContextManager
|
| 12 |
+
from pathlib import Path
|
| 13 |
+
from typing import Any
|
| 14 |
+
|
| 15 |
+
from fdanyone.errors import FourDAnyoneError
|
| 16 |
+
|
| 17 |
+
_RETRYABLE_TREE_ERRORS = {errno.EBUSY, errno.ENOTEMPTY, errno.ESTALE}
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def write_json(path: str | Path, value: object, *, sort_keys: bool = True) -> None:
|
| 21 |
+
target = Path(path)
|
| 22 |
+
target.parent.mkdir(parents=True, exist_ok=True)
|
| 23 |
+
temporary = target.with_name(f".{target.name}.{uuid.uuid4().hex}.tmp")
|
| 24 |
+
temporary.write_text(json.dumps(value, indent=2, sort_keys=sort_keys) + "\n")
|
| 25 |
+
os.replace(temporary, target)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def read_json(path: str | Path | None) -> dict[str, Any] | None:
|
| 29 |
+
"""Read a JSON object, returning None when there is no such file."""
|
| 30 |
+
|
| 31 |
+
if path is None:
|
| 32 |
+
return None
|
| 33 |
+
target = Path(path)
|
| 34 |
+
if not target.is_file():
|
| 35 |
+
return None
|
| 36 |
+
try:
|
| 37 |
+
payload = json.loads(target.read_text())
|
| 38 |
+
except json.JSONDecodeError as exc:
|
| 39 |
+
raise FourDAnyoneError(f"Cannot parse {target}: {exc}") from exc
|
| 40 |
+
if not isinstance(payload, dict):
|
| 41 |
+
raise FourDAnyoneError(f"{target} must contain a JSON object.")
|
| 42 |
+
return payload
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def remove_tree(
|
| 46 |
+
path: str | Path,
|
| 47 |
+
*,
|
| 48 |
+
attempts: int = 8,
|
| 49 |
+
initial_delay_seconds: float = 0.1,
|
| 50 |
+
ignore_errors: bool = False,
|
| 51 |
+
) -> None:
|
| 52 |
+
"""Remove a tree, tolerating short directory-entry lag on network filesystems."""
|
| 53 |
+
|
| 54 |
+
target = Path(path)
|
| 55 |
+
if attempts <= 0:
|
| 56 |
+
raise ValueError(f"attempts must be positive, got {attempts}.")
|
| 57 |
+
for attempt in range(attempts):
|
| 58 |
+
try:
|
| 59 |
+
shutil.rmtree(target)
|
| 60 |
+
return
|
| 61 |
+
except FileNotFoundError:
|
| 62 |
+
return
|
| 63 |
+
except OSError as exc:
|
| 64 |
+
retryable = exc.errno in _RETRYABLE_TREE_ERRORS and attempt + 1 < attempts
|
| 65 |
+
if not retryable:
|
| 66 |
+
if ignore_errors:
|
| 67 |
+
return
|
| 68 |
+
raise
|
| 69 |
+
time.sleep(initial_delay_seconds * (2**attempt))
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
class AtomicResultDirectory(AbstractContextManager[Path]):
|
| 73 |
+
"""Build beside the destination and rename only after all validation passes."""
|
| 74 |
+
|
| 75 |
+
def __init__(self, destination: str | Path):
|
| 76 |
+
expanded = Path(destination).expanduser()
|
| 77 |
+
# Resolve the parent for a stable absolute location, but preserve the
|
| 78 |
+
# leaf itself so a dangling output symlink cannot be followed and
|
| 79 |
+
# mistaken for a nonexistent destination.
|
| 80 |
+
self.destination = expanded.parent.resolve() / expanded.name
|
| 81 |
+
self.working = self.destination.with_name(f".{self.destination.name}.work-{uuid.uuid4().hex[:10]}")
|
| 82 |
+
self._committed = False
|
| 83 |
+
|
| 84 |
+
def _destination_exists(self) -> bool:
|
| 85 |
+
return os.path.lexists(self.destination)
|
| 86 |
+
|
| 87 |
+
def __enter__(self) -> Path:
|
| 88 |
+
if self._destination_exists():
|
| 89 |
+
raise FourDAnyoneError(
|
| 90 |
+
f"Output directory already exists: {self.destination}. Choose a new directory to avoid mixed runs."
|
| 91 |
+
)
|
| 92 |
+
self.working.mkdir(parents=True)
|
| 93 |
+
return self.working
|
| 94 |
+
|
| 95 |
+
def commit(self) -> Path:
|
| 96 |
+
if self._committed:
|
| 97 |
+
return self.destination
|
| 98 |
+
if self._destination_exists():
|
| 99 |
+
raise FourDAnyoneError(
|
| 100 |
+
f"Output directory appeared during inference: {self.destination}. Refusing to overwrite it."
|
| 101 |
+
)
|
| 102 |
+
os.replace(self.working, self.destination)
|
| 103 |
+
self._committed = True
|
| 104 |
+
return self.destination
|
| 105 |
+
|
| 106 |
+
def __exit__(self, exc_type, exc_value, traceback) -> bool:
|
| 107 |
+
if exc_type is None:
|
| 108 |
+
try:
|
| 109 |
+
self.commit()
|
| 110 |
+
except BaseException:
|
| 111 |
+
remove_tree(self.working, ignore_errors=True)
|
| 112 |
+
raise
|
| 113 |
+
else:
|
| 114 |
+
remove_tree(self.working, ignore_errors=True)
|
| 115 |
+
return False
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""4DAnyone model loading and stage inference."""
|
|
@@ -0,0 +1,689 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Hydra/Lightning-free multi-view generation.
|
| 2 |
+
|
| 3 |
+
RCP and final target generation share one source encoding and prompt embedding.
|
| 4 |
+
Target groups execute sequentially on one GPU; TCR optionally shifts their
|
| 5 |
+
membership between denoising steps.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import gc
|
| 11 |
+
import logging
|
| 12 |
+
from collections.abc import Callable, Iterable
|
| 13 |
+
from concurrent.futures import Future, ThreadPoolExecutor
|
| 14 |
+
from dataclasses import dataclass
|
| 15 |
+
from pathlib import Path
|
| 16 |
+
from typing import TYPE_CHECKING, TypeAlias
|
| 17 |
+
|
| 18 |
+
import numpy as np
|
| 19 |
+
from PIL import Image
|
| 20 |
+
|
| 21 |
+
from fdanyone.assets import BaseAssets
|
| 22 |
+
from fdanyone.config import INFERENCE, ModeSettings
|
| 23 |
+
from fdanyone.errors import FourDAnyoneError
|
| 24 |
+
from fdanyone.model.loader import load_pipeline, offload_all, onload_all
|
| 25 |
+
from fdanyone.model.prepared import forward_dynamic, prepare_static_conditioning
|
| 26 |
+
from fdanyone.model.profiling import CudaStageTimer, StageClock, profile_dit_step
|
| 27 |
+
from fdanyone.model.routing import routing_steps
|
| 28 |
+
from fdanyone.model.tiny_decoder import decode_tiny_target_video
|
| 29 |
+
from fdanyone.skeleton.pipeline import Conditioning, SkeletonVideo
|
| 30 |
+
from fdanyone.video import CanonicalClip, write_video, write_video_async
|
| 31 |
+
from fdanyone.views import ViewPlan
|
| 32 |
+
|
| 33 |
+
LOGGER = logging.getLogger("fdanyone")
|
| 34 |
+
|
| 35 |
+
if TYPE_CHECKING:
|
| 36 |
+
import torch
|
| 37 |
+
|
| 38 |
+
DenoiseStepHook: TypeAlias = Callable[[int, tuple[int, ...], "torch.Tensor"], None]
|
| 39 |
+
"""Receives ``(step_index, view_indices, x0_hat)`` for one denoised group."""
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
@dataclass(frozen=True)
|
| 43 |
+
class GeneratedViews:
|
| 44 |
+
"""Paths and measurements produced by one resolved view plan."""
|
| 45 |
+
|
| 46 |
+
rcp_videos: tuple[Path, ...]
|
| 47 |
+
target_videos: tuple[Path, ...]
|
| 48 |
+
view_plan: ViewPlan
|
| 49 |
+
seed: int
|
| 50 |
+
device: str
|
| 51 |
+
elapsed_seconds: dict[str, float]
|
| 52 |
+
peak_vram_allocated_bytes: int
|
| 53 |
+
peak_vram_reserved_bytes: int
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def _empty_cuda_cache() -> None:
|
| 57 |
+
import torch
|
| 58 |
+
|
| 59 |
+
gc.collect()
|
| 60 |
+
if torch.cuda.is_available():
|
| 61 |
+
torch.cuda.empty_cache()
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def _move_model(model, device: str) -> None:
|
| 65 |
+
if getattr(model, "vram_management_enabled", False):
|
| 66 |
+
onload_all(model)
|
| 67 |
+
else:
|
| 68 |
+
model.to(device)
|
| 69 |
+
_empty_cuda_cache()
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def _offload_model(model) -> None:
|
| 73 |
+
if getattr(model, "vram_management_enabled", False):
|
| 74 |
+
offload_all(model)
|
| 75 |
+
else:
|
| 76 |
+
model.to("cpu")
|
| 77 |
+
_empty_cuda_cache()
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def _bf16_autocast():
|
| 81 |
+
"""Match Lightning's ``bf16-mixed`` inference context without Lightning."""
|
| 82 |
+
|
| 83 |
+
import torch
|
| 84 |
+
|
| 85 |
+
return torch.autocast(device_type="cuda", dtype=torch.bfloat16)
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def _channels_last_source_layout(video):
|
| 89 |
+
"""Preserve the frozen source tensor's VFHWC-backed VCFHW layout."""
|
| 90 |
+
|
| 91 |
+
import torch
|
| 92 |
+
|
| 93 |
+
if video.ndim != 5:
|
| 94 |
+
raise FourDAnyoneError(f"Expected a 5D video tensor, got shape {tuple(video.shape)}.")
|
| 95 |
+
return video.contiguous(memory_format=torch.channels_last_3d)
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def _encode_prompt(pipe, device: str) -> dict[str, object]:
|
| 99 |
+
import torch
|
| 100 |
+
|
| 101 |
+
_move_model(pipe.text_encoder, device)
|
| 102 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 103 |
+
context = pipe.prompter.encode_prompt(INFERENCE.prompt, positive=True, device=device)
|
| 104 |
+
prompt = {"context": context.detach().to("cpu")}
|
| 105 |
+
_offload_model(pipe.text_encoder)
|
| 106 |
+
return prompt
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def _load_prompt_embedding(path: Path) -> dict[str, object]:
|
| 110 |
+
"""Read an exported prompt context, refusing one built for another prompt."""
|
| 111 |
+
|
| 112 |
+
from safetensors import safe_open
|
| 113 |
+
|
| 114 |
+
with safe_open(str(path), framework="pt", device="cpu") as handle:
|
| 115 |
+
metadata: dict[str, str] = handle.metadata() or {}
|
| 116 |
+
stored_prompt: str | None = metadata.get("prompt")
|
| 117 |
+
if stored_prompt != INFERENCE.prompt:
|
| 118 |
+
raise FourDAnyoneError(
|
| 119 |
+
f"Prompt embedding {path} was exported for {stored_prompt!r}, "
|
| 120 |
+
f"not the configured prompt {INFERENCE.prompt!r}."
|
| 121 |
+
)
|
| 122 |
+
return {"context": handle.get_tensor("context")}
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def _encode_videos(pipe, videos, device: str):
|
| 126 |
+
import torch
|
| 127 |
+
|
| 128 |
+
videos = videos.to(dtype=pipe.torch_dtype, device=device)
|
| 129 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 130 |
+
latents = pipe.encode_video(videos)
|
| 131 |
+
return latents.detach().to("cpu")
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def _noise(pipe, num_views: int, num_frames: int, seed: int, device: str):
|
| 135 |
+
import torch
|
| 136 |
+
|
| 137 |
+
shape = (
|
| 138 |
+
num_views,
|
| 139 |
+
pipe.vae.model.z_dim,
|
| 140 |
+
(num_frames - 1) // 4 + 1,
|
| 141 |
+
INFERENCE.height // pipe.vae.upsampling_factor,
|
| 142 |
+
INFERENCE.width // pipe.vae.upsampling_factor,
|
| 143 |
+
)
|
| 144 |
+
generator = torch.Generator("cpu").manual_seed(seed)
|
| 145 |
+
return torch.randn(shape, generator=generator, device="cpu", dtype=torch.float32).to(
|
| 146 |
+
dtype=pipe.torch_dtype, device=device
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def _load_skeleton_cache(
|
| 151 |
+
conditioning: Conditioning,
|
| 152 |
+
skeletons: Iterable[SkeletonVideo],
|
| 153 |
+
cache: dict[SkeletonVideo, object],
|
| 154 |
+
device: str,
|
| 155 |
+
) -> None:
|
| 156 |
+
import torch
|
| 157 |
+
|
| 158 |
+
for skeleton in skeletons:
|
| 159 |
+
if skeleton in cache:
|
| 160 |
+
continue
|
| 161 |
+
LOGGER.info("Loading skeleton conditioning from %s", skeleton.path.name)
|
| 162 |
+
tensor = conditioning.load_skeleton_tensor(
|
| 163 |
+
[skeleton],
|
| 164 |
+
device=device,
|
| 165 |
+
).to(dtype=torch.bfloat16, device="cpu")
|
| 166 |
+
cache[skeleton] = tensor.contiguous()
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
def _skeleton_group(
|
| 170 |
+
cache: dict[SkeletonVideo, object],
|
| 171 |
+
skeletons: Iterable[SkeletonVideo],
|
| 172 |
+
device: str,
|
| 173 |
+
*,
|
| 174 |
+
channels_last: bool = False,
|
| 175 |
+
):
|
| 176 |
+
import torch
|
| 177 |
+
|
| 178 |
+
items = tuple(skeletons)
|
| 179 |
+
first = cache[items[0]]
|
| 180 |
+
memory_format = torch.channels_last_3d if channels_last else torch.contiguous_format
|
| 181 |
+
group = torch.empty(
|
| 182 |
+
(len(items), *first.shape[1:]),
|
| 183 |
+
dtype=first.dtype,
|
| 184 |
+
device=device,
|
| 185 |
+
memory_format=memory_format,
|
| 186 |
+
)
|
| 187 |
+
for output_index, skeleton in enumerate(items):
|
| 188 |
+
group[output_index].copy_(cache[skeleton][0])
|
| 189 |
+
return group
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
def _denoise(
|
| 193 |
+
pipe,
|
| 194 |
+
src_latents,
|
| 195 |
+
prompt,
|
| 196 |
+
skeletons: tuple[SkeletonVideo, ...],
|
| 197 |
+
skeleton_cache: dict[SkeletonVideo, object],
|
| 198 |
+
view_plan: ViewPlan | None,
|
| 199 |
+
device: str,
|
| 200 |
+
seed: int,
|
| 201 |
+
stage_timer: CudaStageTimer,
|
| 202 |
+
settings: ModeSettings,
|
| 203 |
+
*,
|
| 204 |
+
on_step: DenoiseStepHook | None = None,
|
| 205 |
+
):
|
| 206 |
+
"""Denoise either one full proposal group or routed target groups."""
|
| 207 |
+
|
| 208 |
+
import torch
|
| 209 |
+
from tqdm.auto import tqdm
|
| 210 |
+
|
| 211 |
+
num_views: int = len(skeletons)
|
| 212 |
+
if view_plan is not None and num_views != view_plan.num_target_views:
|
| 213 |
+
raise FourDAnyoneError(
|
| 214 |
+
f"Target generation requires {view_plan.num_target_views} skeleton views, got {num_views}."
|
| 215 |
+
)
|
| 216 |
+
pipe.scheduler.set_timesteps(
|
| 217 |
+
settings.num_inference_steps,
|
| 218 |
+
denoising_strength=INFERENCE.denoising_strength,
|
| 219 |
+
shift=settings.scheduler_shift,
|
| 220 |
+
)
|
| 221 |
+
latents = _noise(pipe, num_views, INFERENCE.num_frames, seed, device)
|
| 222 |
+
source = src_latents.to(dtype=pipe.torch_dtype, device=device)
|
| 223 |
+
context = {name: value.to(dtype=pipe.torch_dtype, device=device) for name, value in prompt.items()}
|
| 224 |
+
if view_plan is None:
|
| 225 |
+
group_size: int = num_views
|
| 226 |
+
routes = routing_steps(
|
| 227 |
+
views_per_layer=num_views,
|
| 228 |
+
num_layers=1,
|
| 229 |
+
group_size=num_views,
|
| 230 |
+
num_steps=settings.num_inference_steps,
|
| 231 |
+
enable_tcr=False,
|
| 232 |
+
circular=False,
|
| 233 |
+
)
|
| 234 |
+
# Preserve this RCP channels_last_3d payload as data. Its layout selects
|
| 235 |
+
# the banked CUDA kernel and must not be normalized by the merged path.
|
| 236 |
+
prepared_skeletons = _skeleton_group(
|
| 237 |
+
skeleton_cache,
|
| 238 |
+
skeletons,
|
| 239 |
+
device,
|
| 240 |
+
channels_last=True,
|
| 241 |
+
)
|
| 242 |
+
pose_batch_size: int | None = None
|
| 243 |
+
description: str = f"RCP 1-to-{num_views}"
|
| 244 |
+
else:
|
| 245 |
+
group_size = view_plan.views_per_group
|
| 246 |
+
routes = routing_steps(
|
| 247 |
+
views_per_layer=view_plan.views_per_layer,
|
| 248 |
+
num_layers=view_plan.num_layers,
|
| 249 |
+
group_size=group_size,
|
| 250 |
+
num_steps=settings.num_inference_steps,
|
| 251 |
+
enable_tcr=view_plan.tcr_active,
|
| 252 |
+
circular=view_plan.closed_yaw,
|
| 253 |
+
)
|
| 254 |
+
prepared_skeletons = tuple(skeleton_cache[skeleton] for skeleton in skeletons)
|
| 255 |
+
pose_batch_size = settings.dit_pose_batch_size
|
| 256 |
+
description = f"Generate {num_views} target views"
|
| 257 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 258 |
+
prepared = prepare_static_conditioning(
|
| 259 |
+
pipe.dit,
|
| 260 |
+
x_src=source,
|
| 261 |
+
context=context["context"],
|
| 262 |
+
skeletons=prepared_skeletons,
|
| 263 |
+
timesteps=pipe.scheduler.timesteps,
|
| 264 |
+
group_size=group_size,
|
| 265 |
+
pose_batch_size=pose_batch_size,
|
| 266 |
+
stage_timer=stage_timer,
|
| 267 |
+
)
|
| 268 |
+
del prepared_skeletons
|
| 269 |
+
|
| 270 |
+
profile_next: bool = view_plan is not None
|
| 271 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 272 |
+
for step_index, groups in enumerate(tqdm(routes, desc=description)):
|
| 273 |
+
timestep = pipe.scheduler.timesteps[step_index]
|
| 274 |
+
for view_indices in groups:
|
| 275 |
+
with torch.profiler.record_function("dit.route_copies"):
|
| 276 |
+
index = torch.tensor(view_indices, dtype=torch.long, device=device)
|
| 277 |
+
local_latents = torch.index_select(latents, 0, index)
|
| 278 |
+
with profile_dit_step(enabled=profile_next):
|
| 279 |
+
prediction = forward_dynamic(
|
| 280 |
+
pipe.dit,
|
| 281 |
+
x=local_latents,
|
| 282 |
+
prepared=prepared,
|
| 283 |
+
view_indices=index,
|
| 284 |
+
step_index=step_index,
|
| 285 |
+
)
|
| 286 |
+
profile_next = False
|
| 287 |
+
if on_step is not None:
|
| 288 |
+
# Flow matching predicts the noise-to-data velocity, so the
|
| 289 |
+
# clean estimate is one full sigma step along it.
|
| 290 |
+
sigma = pipe.scheduler.sigmas[step_index].to(local_latents.device)
|
| 291 |
+
on_step(step_index, view_indices, local_latents - sigma * prediction)
|
| 292 |
+
local_latents = pipe.scheduler.step(prediction, timestep, local_latents)
|
| 293 |
+
with torch.profiler.record_function("dit.route_copies"):
|
| 294 |
+
latents.index_copy_(0, index, local_latents)
|
| 295 |
+
del local_latents, prediction
|
| 296 |
+
return latents.detach().to("cpu")
|
| 297 |
+
|
| 298 |
+
|
| 299 |
+
def _tensor_frames(video) -> Iterable[np.ndarray]:
|
| 300 |
+
"""Match DiffSynth's float-to-uint8 truncation exactly."""
|
| 301 |
+
|
| 302 |
+
frames = video.detach().float().add_(1.0).mul_(127.5).clamp_(0.0, 255.0).to("cpu")
|
| 303 |
+
for frame_index in range(frames.shape[1]):
|
| 304 |
+
yield frames[:, frame_index].permute(1, 2, 0).numpy().astype(np.uint8)
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
def _save_rcp_jpegs(video, camera_id: int, root: Path) -> Path:
|
| 308 |
+
import torchvision.transforms.functional as transform
|
| 309 |
+
|
| 310 |
+
frame_dir = root / f"{camera_id:06d}"
|
| 311 |
+
frame_dir.mkdir(parents=True, exist_ok=False)
|
| 312 |
+
normalized = video.detach().float().mul(0.5).add_(0.5).clamp_(0.0, 1.0).to("cpu")
|
| 313 |
+
for frame_index in range(normalized.shape[1]):
|
| 314 |
+
image = transform.to_pil_image(normalized[:, frame_index])
|
| 315 |
+
image.save(frame_dir / f"{frame_index:06d}.jpg", quality=INFERENCE.rcp_jpeg_quality)
|
| 316 |
+
return frame_dir
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
def _rcp_reference_video_layout(frame_first_video):
|
| 320 |
+
"""Match ``prepare_batch`` for JPEG-backed ``[V,F,C,H,W]`` data."""
|
| 321 |
+
|
| 322 |
+
if frame_first_video.ndim != 5:
|
| 323 |
+
raise FourDAnyoneError(f"Expected a 5D frame-first video tensor, got shape {tuple(frame_first_video.shape)}.")
|
| 324 |
+
return frame_first_video.permute(0, 2, 1, 3, 4)
|
| 325 |
+
|
| 326 |
+
|
| 327 |
+
def _load_rcp_reference_videos(frame_dirs: Iterable[Path], num_frames: int):
|
| 328 |
+
"""Decode RCP references exactly like the frozen JPEG-backed input path."""
|
| 329 |
+
|
| 330 |
+
import torch
|
| 331 |
+
import torchvision.transforms.functional as transform
|
| 332 |
+
|
| 333 |
+
videos = []
|
| 334 |
+
for frame_dir in frame_dirs:
|
| 335 |
+
frames = []
|
| 336 |
+
for frame_index in range(num_frames):
|
| 337 |
+
path = frame_dir / f"{frame_index:06d}.jpg"
|
| 338 |
+
with Image.open(path) as image:
|
| 339 |
+
frames.append(transform.to_tensor(image.convert("RGB")))
|
| 340 |
+
videos.append(torch.stack(frames, dim=0))
|
| 341 |
+
frame_first_video = torch.stack(videos, dim=0).mul_(2.0).sub_(1.0)
|
| 342 |
+
return _rcp_reference_video_layout(frame_first_video)
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
def _decode_view(pipe, latents, device: str, *, decoder=None):
|
| 346 |
+
"""Decode one full- or tiny-decoder latent view."""
|
| 347 |
+
|
| 348 |
+
import torch
|
| 349 |
+
|
| 350 |
+
if decoder is None:
|
| 351 |
+
return pipe.decode_video(latents.to(dtype=pipe.torch_dtype, device=device))[0]
|
| 352 |
+
return decode_tiny_target_video(
|
| 353 |
+
decoder,
|
| 354 |
+
latents.to(dtype=torch.float16, device=device),
|
| 355 |
+
)[0]
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
def _emit_decoded_videos(
|
| 359 |
+
decoded_views: Iterable[tuple[object, Path]],
|
| 360 |
+
clip: CanonicalClip,
|
| 361 |
+
settings: ModeSettings,
|
| 362 |
+
) -> tuple[Path, ...]:
|
| 363 |
+
"""Write decoded views serially or through the bounded encoder pool."""
|
| 364 |
+
|
| 365 |
+
outputs: list[Path] = []
|
| 366 |
+
encode_futures: list[Future[Path]] = []
|
| 367 |
+
executor: ThreadPoolExecutor | None = (
|
| 368 |
+
ThreadPoolExecutor(max_workers=2) if settings.async_video_encode else None
|
| 369 |
+
)
|
| 370 |
+
try:
|
| 371 |
+
for video, path in decoded_views:
|
| 372 |
+
if executor is None:
|
| 373 |
+
outputs.append(
|
| 374 |
+
write_video(
|
| 375 |
+
_tensor_frames(video),
|
| 376 |
+
path,
|
| 377 |
+
clip.fps,
|
| 378 |
+
crf=INFERENCE.target_h264_crf,
|
| 379 |
+
preset=INFERENCE.h264_preset,
|
| 380 |
+
)
|
| 381 |
+
)
|
| 382 |
+
else:
|
| 383 |
+
outputs.append(path)
|
| 384 |
+
encode_futures.append(
|
| 385 |
+
write_video_async(
|
| 386 |
+
executor,
|
| 387 |
+
_tensor_frames(video),
|
| 388 |
+
path,
|
| 389 |
+
clip.fps,
|
| 390 |
+
crf=INFERENCE.target_h264_crf,
|
| 391 |
+
preset=INFERENCE.h264_preset,
|
| 392 |
+
)
|
| 393 |
+
)
|
| 394 |
+
for future in encode_futures:
|
| 395 |
+
future.result()
|
| 396 |
+
finally:
|
| 397 |
+
if executor is not None:
|
| 398 |
+
executor.shutdown(wait=True, cancel_futures=True)
|
| 399 |
+
return tuple(outputs)
|
| 400 |
+
|
| 401 |
+
|
| 402 |
+
def _decode_rcp(
|
| 403 |
+
pipe,
|
| 404 |
+
latents,
|
| 405 |
+
camera_ids: tuple[int, ...],
|
| 406 |
+
output_dir: Path,
|
| 407 |
+
clip: CanonicalClip,
|
| 408 |
+
device: str,
|
| 409 |
+
settings: ModeSettings,
|
| 410 |
+
*,
|
| 411 |
+
rcp_decoder=None,
|
| 412 |
+
) -> tuple[tuple[Path, ...], tuple[Path, ...]]:
|
| 413 |
+
import torch
|
| 414 |
+
|
| 415 |
+
if latents.shape[0] != len(camera_ids):
|
| 416 |
+
raise FourDAnyoneError(f"RCP decode expected {len(camera_ids)} latent views, got {latents.shape[0]}.")
|
| 417 |
+
frame_root = output_dir / "frames"
|
| 418 |
+
video_root = output_dir / "videos"
|
| 419 |
+
frame_root.mkdir(parents=True, exist_ok=False)
|
| 420 |
+
video_root.mkdir(parents=True, exist_ok=False)
|
| 421 |
+
frame_outputs: list[Path] = []
|
| 422 |
+
|
| 423 |
+
def decoded_views() -> Iterable[tuple[object, Path]]:
|
| 424 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 425 |
+
for latent_index, camera_id in enumerate(camera_ids):
|
| 426 |
+
LOGGER.info("Decoding RCP camera %02d", camera_id)
|
| 427 |
+
camera_latents = latents[latent_index : latent_index + 1]
|
| 428 |
+
video = _decode_view(
|
| 429 |
+
pipe,
|
| 430 |
+
camera_latents,
|
| 431 |
+
device,
|
| 432 |
+
decoder=rcp_decoder,
|
| 433 |
+
)
|
| 434 |
+
if not settings.direct_rcp_latent_handoff:
|
| 435 |
+
video_for_output = video.detach().to("cpu")
|
| 436 |
+
frame_outputs.append(
|
| 437 |
+
_save_rcp_jpegs(video_for_output, camera_id, frame_root)
|
| 438 |
+
)
|
| 439 |
+
else:
|
| 440 |
+
video_for_output = video
|
| 441 |
+
video_path: Path = video_root / f"{camera_id:02d}.mp4"
|
| 442 |
+
yield video_for_output, video_path
|
| 443 |
+
del video
|
| 444 |
+
torch.cuda.empty_cache()
|
| 445 |
+
|
| 446 |
+
video_outputs: tuple[Path, ...] = _emit_decoded_videos(
|
| 447 |
+
decoded_views(),
|
| 448 |
+
clip,
|
| 449 |
+
settings,
|
| 450 |
+
)
|
| 451 |
+
return tuple(frame_outputs), video_outputs
|
| 452 |
+
|
| 453 |
+
|
| 454 |
+
def _decode_targets(
|
| 455 |
+
pipe,
|
| 456 |
+
latents,
|
| 457 |
+
output_dir: Path,
|
| 458 |
+
clip: CanonicalClip,
|
| 459 |
+
device: str,
|
| 460 |
+
settings: ModeSettings,
|
| 461 |
+
*,
|
| 462 |
+
target_decoder=None,
|
| 463 |
+
) -> tuple[Path, ...]:
|
| 464 |
+
import torch
|
| 465 |
+
|
| 466 |
+
video_root = output_dir / "videos"
|
| 467 |
+
video_root.mkdir(parents=True, exist_ok=False)
|
| 468 |
+
|
| 469 |
+
def decoded_views() -> Iterable[tuple[object, Path]]:
|
| 470 |
+
with torch.inference_mode(), _bf16_autocast():
|
| 471 |
+
for camera_id in range(latents.shape[0]):
|
| 472 |
+
LOGGER.info("Decoding target camera %02d", camera_id)
|
| 473 |
+
camera_latents = latents[camera_id : camera_id + 1]
|
| 474 |
+
video = _decode_view(
|
| 475 |
+
pipe,
|
| 476 |
+
camera_latents,
|
| 477 |
+
device,
|
| 478 |
+
decoder=target_decoder,
|
| 479 |
+
)
|
| 480 |
+
path: Path = video_root / f"{camera_id:02d}.mp4"
|
| 481 |
+
yield video, path
|
| 482 |
+
del video
|
| 483 |
+
torch.cuda.empty_cache()
|
| 484 |
+
|
| 485 |
+
return _emit_decoded_videos(decoded_views(), clip, settings)
|
| 486 |
+
|
| 487 |
+
|
| 488 |
+
def generate_views(
|
| 489 |
+
*,
|
| 490 |
+
clip: CanonicalClip,
|
| 491 |
+
conditioning: Conditioning,
|
| 492 |
+
checkpoint_path: str | Path,
|
| 493 |
+
assets: BaseAssets,
|
| 494 |
+
output_dir: str | Path,
|
| 495 |
+
device: str,
|
| 496 |
+
seed: int,
|
| 497 |
+
settings: ModeSettings,
|
| 498 |
+
on_denoise_step: DenoiseStepHook | None = None,
|
| 499 |
+
prompt_embedding_path: Path | None = None,
|
| 500 |
+
) -> GeneratedViews:
|
| 501 |
+
"""Generate the proposal (when enabled) and the requested target views."""
|
| 502 |
+
|
| 503 |
+
import torch
|
| 504 |
+
|
| 505 |
+
if conditioning.num_frames != INFERENCE.num_frames or len(clip.frames) != INFERENCE.num_frames:
|
| 506 |
+
raise FourDAnyoneError("Generation requires the frozen 121-frame contract.")
|
| 507 |
+
if seed < 0:
|
| 508 |
+
raise FourDAnyoneError(f"seed must be non-negative, got {seed}.")
|
| 509 |
+
device_index = int(device.removeprefix("cuda:"))
|
| 510 |
+
LOGGER.info("Using %s (%s)", device, torch.cuda.get_device_name(device_index))
|
| 511 |
+
|
| 512 |
+
root = Path(output_dir).expanduser().resolve()
|
| 513 |
+
root.mkdir(parents=True, exist_ok=False)
|
| 514 |
+
view_plan = conditioning.view_plan
|
| 515 |
+
if len(conditioning.target_skeletons) != view_plan.num_target_views:
|
| 516 |
+
raise FourDAnyoneError("Target skeleton count does not match the resolved view plan.")
|
| 517 |
+
if len(conditioning.rcp_skeletons) != len(view_plan.rcp_camera_ids):
|
| 518 |
+
raise FourDAnyoneError("RCP skeleton count does not match the resolved view plan.")
|
| 519 |
+
loaded = load_pipeline(
|
| 520 |
+
checkpoint_path=checkpoint_path,
|
| 521 |
+
assets=assets,
|
| 522 |
+
device=device,
|
| 523 |
+
settings=settings,
|
| 524 |
+
load_text_encoder=prompt_embedding_path is None,
|
| 525 |
+
)
|
| 526 |
+
pipe = loaded.pipe
|
| 527 |
+
clock = StageClock(cuda=CudaStageTimer(device=device))
|
| 528 |
+
skeleton_cache: dict[SkeletonVideo, object] = {}
|
| 529 |
+
torch.cuda.reset_peak_memory_stats(device_index)
|
| 530 |
+
|
| 531 |
+
with clock.stage("prompt_t5"):
|
| 532 |
+
if prompt_embedding_path is None:
|
| 533 |
+
prompt = _encode_prompt(pipe, device)
|
| 534 |
+
else:
|
| 535 |
+
prompt = _load_prompt_embedding(prompt_embedding_path)
|
| 536 |
+
loaded.release_text_encoder()
|
| 537 |
+
|
| 538 |
+
with clock.stage("source_vae_encode"):
|
| 539 |
+
_move_model(pipe.vae, device)
|
| 540 |
+
source_video = _channels_last_source_layout(conditioning.load_source_tensor())
|
| 541 |
+
source_latents = _encode_videos(pipe, source_video, device)
|
| 542 |
+
del source_video
|
| 543 |
+
_offload_model(pipe.vae)
|
| 544 |
+
|
| 545 |
+
rcp_videos: tuple[Path, ...] = ()
|
| 546 |
+
target_sources = source_latents
|
| 547 |
+
if view_plan.enable_rcp:
|
| 548 |
+
with clock.stage("rcp_prepare"):
|
| 549 |
+
_load_skeleton_cache(
|
| 550 |
+
conditioning,
|
| 551 |
+
conditioning.rcp_skeletons,
|
| 552 |
+
skeleton_cache,
|
| 553 |
+
device,
|
| 554 |
+
)
|
| 555 |
+
_move_model(pipe.dit, device)
|
| 556 |
+
with clock.stage("rcp_dit"):
|
| 557 |
+
rcp_latents = _denoise(
|
| 558 |
+
pipe,
|
| 559 |
+
source_latents,
|
| 560 |
+
prompt,
|
| 561 |
+
conditioning.rcp_skeletons,
|
| 562 |
+
skeleton_cache,
|
| 563 |
+
None,
|
| 564 |
+
device,
|
| 565 |
+
seed,
|
| 566 |
+
clock.cuda,
|
| 567 |
+
settings,
|
| 568 |
+
)
|
| 569 |
+
clock.elapsed["rcp_pose_encoder"] = clock.cuda.elapsed_seconds("pose_encoder")
|
| 570 |
+
with clock.stage("rcp_offload"):
|
| 571 |
+
_offload_model(pipe.dit)
|
| 572 |
+
|
| 573 |
+
with clock.stage("rcp_decode_jpeg_reencode"):
|
| 574 |
+
rcp_root = root / "rcp"
|
| 575 |
+
rcp_root.mkdir()
|
| 576 |
+
rcp_decoder = (
|
| 577 |
+
loaded.target_decoder if settings.tiny_decoders else None
|
| 578 |
+
)
|
| 579 |
+
decode_model = pipe.vae if rcp_decoder is None else rcp_decoder
|
| 580 |
+
_move_model(decode_model, device)
|
| 581 |
+
frame_dirs, rcp_videos = _decode_rcp(
|
| 582 |
+
pipe,
|
| 583 |
+
rcp_latents,
|
| 584 |
+
view_plan.rcp_camera_ids,
|
| 585 |
+
rcp_root,
|
| 586 |
+
clip,
|
| 587 |
+
device,
|
| 588 |
+
settings,
|
| 589 |
+
rcp_decoder=rcp_decoder,
|
| 590 |
+
)
|
| 591 |
+
if rcp_decoder is not None:
|
| 592 |
+
_offload_model(rcp_decoder)
|
| 593 |
+
if settings.direct_rcp_latent_handoff:
|
| 594 |
+
selected_rcp_latents = rcp_latents
|
| 595 |
+
target_sources = torch.cat(
|
| 596 |
+
[source_latents, selected_rcp_latents],
|
| 597 |
+
dim=0,
|
| 598 |
+
)
|
| 599 |
+
else:
|
| 600 |
+
if rcp_decoder is not None:
|
| 601 |
+
_move_model(pipe.vae, device)
|
| 602 |
+
rcp_reference_videos = _load_rcp_reference_videos(
|
| 603 |
+
frame_dirs[:4],
|
| 604 |
+
INFERENCE.num_frames,
|
| 605 |
+
)
|
| 606 |
+
rcp_reference_latents = _encode_videos(
|
| 607 |
+
pipe,
|
| 608 |
+
rcp_reference_videos,
|
| 609 |
+
device,
|
| 610 |
+
)
|
| 611 |
+
selected_rcp_latents = rcp_reference_latents
|
| 612 |
+
# VAE38 encodes batch elements independently, so the existing
|
| 613 |
+
# source encoding matches re-encoding source plus references.
|
| 614 |
+
target_sources = torch.cat(
|
| 615 |
+
[source_latents, selected_rcp_latents],
|
| 616 |
+
dim=0,
|
| 617 |
+
)
|
| 618 |
+
del rcp_reference_videos, rcp_reference_latents
|
| 619 |
+
del rcp_latents
|
| 620 |
+
_offload_model(pipe.vae)
|
| 621 |
+
|
| 622 |
+
conditioning.wait_for_target_skeletons()
|
| 623 |
+
with clock.stage("target_prepare"):
|
| 624 |
+
_load_skeleton_cache(
|
| 625 |
+
conditioning,
|
| 626 |
+
conditioning.target_skeletons,
|
| 627 |
+
skeleton_cache,
|
| 628 |
+
device,
|
| 629 |
+
)
|
| 630 |
+
_move_model(pipe.dit, device)
|
| 631 |
+
with clock.stage("target_dit"):
|
| 632 |
+
target_latents = _denoise(
|
| 633 |
+
pipe,
|
| 634 |
+
target_sources,
|
| 635 |
+
prompt,
|
| 636 |
+
conditioning.target_skeletons,
|
| 637 |
+
skeleton_cache,
|
| 638 |
+
view_plan,
|
| 639 |
+
device,
|
| 640 |
+
seed,
|
| 641 |
+
clock.cuda,
|
| 642 |
+
settings,
|
| 643 |
+
on_step=on_denoise_step,
|
| 644 |
+
)
|
| 645 |
+
# One timer spans both stages, so the target share is what RCP did not spend.
|
| 646 |
+
pose_encoder_seconds = clock.cuda.elapsed_seconds("pose_encoder")
|
| 647 |
+
clock.elapsed["target_pose_encoder"] = pose_encoder_seconds - clock.elapsed.get("rcp_pose_encoder", 0.0)
|
| 648 |
+
clock.elapsed["pose_encoder"] = pose_encoder_seconds
|
| 649 |
+
with clock.stage("target_offload"):
|
| 650 |
+
_offload_model(pipe.dit)
|
| 651 |
+
# The decoded skeleton tensors (~650 MB per view) have no reader past the
|
| 652 |
+
# target DiT stage; release them before the decode/encode phase allocates.
|
| 653 |
+
skeleton_cache.clear()
|
| 654 |
+
|
| 655 |
+
with clock.stage("target_vae_decode"):
|
| 656 |
+
target_root = root / "target"
|
| 657 |
+
target_root.mkdir()
|
| 658 |
+
tiny_target_decoder = (
|
| 659 |
+
loaded.target_decoder if settings.tiny_decoders else None
|
| 660 |
+
)
|
| 661 |
+
target_decoder = (
|
| 662 |
+
pipe.vae if tiny_target_decoder is None else tiny_target_decoder
|
| 663 |
+
)
|
| 664 |
+
_move_model(target_decoder, device)
|
| 665 |
+
target_videos = _decode_targets(
|
| 666 |
+
pipe,
|
| 667 |
+
target_latents,
|
| 668 |
+
target_root,
|
| 669 |
+
clip,
|
| 670 |
+
device,
|
| 671 |
+
settings,
|
| 672 |
+
target_decoder=tiny_target_decoder,
|
| 673 |
+
)
|
| 674 |
+
_offload_model(target_decoder)
|
| 675 |
+
peak_vram_allocated = int(torch.cuda.max_memory_allocated(device_index))
|
| 676 |
+
peak_vram_reserved = int(torch.cuda.max_memory_reserved(device_index))
|
| 677 |
+
|
| 678 |
+
del target_latents, target_sources, source_latents, prompt, skeleton_cache, loaded, pipe
|
| 679 |
+
_empty_cuda_cache()
|
| 680 |
+
return GeneratedViews(
|
| 681 |
+
rcp_videos=rcp_videos,
|
| 682 |
+
target_videos=target_videos,
|
| 683 |
+
view_plan=view_plan,
|
| 684 |
+
seed=seed,
|
| 685 |
+
device=device,
|
| 686 |
+
elapsed_seconds=clock.elapsed,
|
| 687 |
+
peak_vram_allocated_bytes=peak_vram_allocated,
|
| 688 |
+
peak_vram_reserved_bytes=peak_vram_reserved,
|
| 689 |
+
)
|
|
@@ -0,0 +1,426 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Direct, registry-free loading of the frozen Wan/SpaTem inference stack."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import gc
|
| 6 |
+
import logging
|
| 7 |
+
import os
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
from typing import Any
|
| 11 |
+
|
| 12 |
+
from fdanyone.assets import BaseAssets
|
| 13 |
+
from fdanyone.config import (
|
| 14 |
+
MVS_ATTENTION_RANGE,
|
| 15 |
+
POSE_ENCODER_TYPE,
|
| 16 |
+
USE_POSE_ENCODER,
|
| 17 |
+
USE_VIEWPACK,
|
| 18 |
+
ModeSettings,
|
| 19 |
+
)
|
| 20 |
+
from fdanyone.errors import AssetError, ConfigurationError
|
| 21 |
+
|
| 22 |
+
LOGGER = logging.getLogger("fdanyone")
|
| 23 |
+
|
| 24 |
+
WAN22_TI2V_5B_CONFIG = {
|
| 25 |
+
"has_image_input": False,
|
| 26 |
+
"patch_size": (1, 2, 2),
|
| 27 |
+
"in_dim": 48,
|
| 28 |
+
"dim": 3072,
|
| 29 |
+
"ffn_dim": 14336,
|
| 30 |
+
"freq_dim": 256,
|
| 31 |
+
"text_dim": 4096,
|
| 32 |
+
"out_dim": 48,
|
| 33 |
+
"num_heads": 24,
|
| 34 |
+
"num_layers": 30,
|
| 35 |
+
"eps": 1e-6,
|
| 36 |
+
"seperated_timestep": True,
|
| 37 |
+
"require_clip_embedding": False,
|
| 38 |
+
"require_vae_embedding": False,
|
| 39 |
+
"fuse_vae_embedding_in_latents": True,
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
@dataclass
|
| 44 |
+
class LoadedPipeline:
|
| 45 |
+
pipe: object
|
| 46 |
+
target_decoder: object | None = None
|
| 47 |
+
|
| 48 |
+
def release_text_encoder(self) -> None:
|
| 49 |
+
"""Release T5 after the single fixed prompt has been encoded."""
|
| 50 |
+
|
| 51 |
+
import torch
|
| 52 |
+
|
| 53 |
+
self.pipe.text_encoder = None
|
| 54 |
+
self.pipe.prompter.text_encoder = None
|
| 55 |
+
gc.collect()
|
| 56 |
+
if torch.cuda.is_available():
|
| 57 |
+
torch.cuda.empty_cache()
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def _load_checkpoint(path: Path):
|
| 61 |
+
try:
|
| 62 |
+
from safetensors.torch import load_file
|
| 63 |
+
except ImportError as exc:
|
| 64 |
+
raise AssetError("safetensors is required to load the 4DAnyone checkpoint.") from exc
|
| 65 |
+
return load_file(str(path), device="cpu")
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def _strict_assign(module, state_dict: dict, label: str) -> None:
|
| 69 |
+
"""Load into a meta-initialized module without a second parameter copy."""
|
| 70 |
+
|
| 71 |
+
try:
|
| 72 |
+
incompatible = module.load_state_dict(state_dict, strict=True, assign=True)
|
| 73 |
+
except TypeError as exc:
|
| 74 |
+
raise ConfigurationError("4DAnyone requires PyTorch >=2.8 for assign-based model loading.") from exc
|
| 75 |
+
except RuntimeError as exc:
|
| 76 |
+
raise AssetError(f"{label} is incompatible with the released architecture: {exc}") from exc
|
| 77 |
+
if incompatible.missing_keys or incompatible.unexpected_keys:
|
| 78 |
+
raise AssetError(
|
| 79 |
+
f"{label} strict load failed; missing={incompatible.missing_keys}, "
|
| 80 |
+
f"unexpected={incompatible.unexpected_keys}"
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def _load_dit(
|
| 85 |
+
checkpoint_path: Path,
|
| 86 |
+
dtype,
|
| 87 |
+
*,
|
| 88 |
+
settings: ModeSettings,
|
| 89 |
+
turbo_lora_path: Path | None,
|
| 90 |
+
):
|
| 91 |
+
import torch
|
| 92 |
+
|
| 93 |
+
from fdanyone.vendor.diffsynth.models.wan_video_dit import (
|
| 94 |
+
AttentionModule,
|
| 95 |
+
DiTBlock,
|
| 96 |
+
WanModel,
|
| 97 |
+
precompute_freqs_cis_3d,
|
| 98 |
+
)
|
| 99 |
+
from fdanyone.vendor.diffsynth.pipelines.wan_video_spatem import (
|
| 100 |
+
WanVideoSpaTemPipeline,
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
with torch.device("meta"):
|
| 104 |
+
dit = WanModel(**WAN22_TI2V_5B_CONFIG)
|
| 105 |
+
shell = WanVideoSpaTemPipeline(device="cpu", torch_dtype=dtype)
|
| 106 |
+
shell.dit = dit
|
| 107 |
+
shell.init_spatem_modules(
|
| 108 |
+
use_mvs_attn=True,
|
| 109 |
+
range_mvs_attn=MVS_ATTENTION_RANGE,
|
| 110 |
+
# ``ViewPack`` is the upstream module name for RCP references.
|
| 111 |
+
use_viewpack=USE_VIEWPACK,
|
| 112 |
+
use_pose_encoder=USE_POSE_ENCODER,
|
| 113 |
+
pose_encoder_type=POSE_ENCODER_TYPE,
|
| 114 |
+
)
|
| 115 |
+
state_dict = _load_checkpoint(checkpoint_path)
|
| 116 |
+
_strict_assign(dit, state_dict, "4DAnyone DiT checkpoint")
|
| 117 |
+
del state_dict
|
| 118 |
+
for module in dit.modules():
|
| 119 |
+
if isinstance(module, DiTBlock):
|
| 120 |
+
module.use_bf16_block_glue = settings.bf16_block_glue
|
| 121 |
+
if isinstance(module, AttentionModule):
|
| 122 |
+
module.exact = settings.exact_attention
|
| 123 |
+
# ``freqs`` is a derived, non-persistent tensor and therefore is not in the
|
| 124 |
+
# state dict populated above.
|
| 125 |
+
dit.freqs = precompute_freqs_cis_3d(WAN22_TI2V_5B_CONFIG["dim"] // WAN22_TI2V_5B_CONFIG["num_heads"])
|
| 126 |
+
legacy_backend: str | None = os.environ.get("FDANYONE_ATTENTION_BACKEND")
|
| 127 |
+
if legacy_backend:
|
| 128 |
+
assert legacy_backend.lower() in {"sdpa", "sageattention"}
|
| 129 |
+
LOGGER.warning("FDANYONE_ATTENTION_BACKEND is obsolete and ignored; --mode owns attention.")
|
| 130 |
+
if turbo_lora_path is not None:
|
| 131 |
+
from fdanyone.model.turbo_lora import merge_wan_turbo_lora
|
| 132 |
+
|
| 133 |
+
report = merge_wan_turbo_lora(
|
| 134 |
+
dit,
|
| 135 |
+
turbo_lora_path,
|
| 136 |
+
)
|
| 137 |
+
LOGGER.info(
|
| 138 |
+
"Turbo-LoRA merged: modules=%d direct_parameters=%d max_rank=%d",
|
| 139 |
+
report.lora_modules,
|
| 140 |
+
report.direct_parameters,
|
| 141 |
+
report.max_rank,
|
| 142 |
+
)
|
| 143 |
+
return dit.eval().requires_grad_(False)
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def _load_vae(path: Path, dtype):
|
| 147 |
+
import torch
|
| 148 |
+
|
| 149 |
+
from fdanyone.vendor.diffsynth.models.wan_video_vae import WanVideoVAE38
|
| 150 |
+
|
| 151 |
+
state_dict = torch.load(path, map_location="cpu", weights_only=True)
|
| 152 |
+
state_dict = WanVideoVAE38.state_dict_converter().from_civitai(state_dict)
|
| 153 |
+
with torch.device("meta"):
|
| 154 |
+
vae = WanVideoVAE38()
|
| 155 |
+
_strict_assign(vae, state_dict, "Wan2.2 VAE")
|
| 156 |
+
del state_dict
|
| 157 |
+
# Wan's latent normalization tensors are plain attributes rather than
|
| 158 |
+
# registered buffers, so materialize them after meta initialization.
|
| 159 |
+
mean = (
|
| 160 |
+
-0.2289,
|
| 161 |
+
-0.0052,
|
| 162 |
+
-0.1323,
|
| 163 |
+
-0.2339,
|
| 164 |
+
-0.2799,
|
| 165 |
+
0.0174,
|
| 166 |
+
0.1838,
|
| 167 |
+
0.1557,
|
| 168 |
+
-0.1382,
|
| 169 |
+
0.0542,
|
| 170 |
+
0.2813,
|
| 171 |
+
0.0891,
|
| 172 |
+
0.1570,
|
| 173 |
+
-0.0098,
|
| 174 |
+
0.0375,
|
| 175 |
+
-0.1825,
|
| 176 |
+
-0.2246,
|
| 177 |
+
-0.1207,
|
| 178 |
+
-0.0698,
|
| 179 |
+
0.5109,
|
| 180 |
+
0.2665,
|
| 181 |
+
-0.2108,
|
| 182 |
+
-0.2158,
|
| 183 |
+
0.2502,
|
| 184 |
+
-0.2055,
|
| 185 |
+
-0.0322,
|
| 186 |
+
0.1109,
|
| 187 |
+
0.1567,
|
| 188 |
+
-0.0729,
|
| 189 |
+
0.0899,
|
| 190 |
+
-0.2799,
|
| 191 |
+
-0.1230,
|
| 192 |
+
-0.0313,
|
| 193 |
+
-0.1649,
|
| 194 |
+
0.0117,
|
| 195 |
+
0.0723,
|
| 196 |
+
-0.2839,
|
| 197 |
+
-0.2083,
|
| 198 |
+
-0.0520,
|
| 199 |
+
0.3748,
|
| 200 |
+
0.0152,
|
| 201 |
+
0.1957,
|
| 202 |
+
0.1433,
|
| 203 |
+
-0.2944,
|
| 204 |
+
0.3573,
|
| 205 |
+
-0.0548,
|
| 206 |
+
-0.1681,
|
| 207 |
+
-0.0667,
|
| 208 |
+
)
|
| 209 |
+
std = (
|
| 210 |
+
0.4765,
|
| 211 |
+
1.0364,
|
| 212 |
+
0.4514,
|
| 213 |
+
1.1677,
|
| 214 |
+
0.5313,
|
| 215 |
+
0.4990,
|
| 216 |
+
0.4818,
|
| 217 |
+
0.5013,
|
| 218 |
+
0.8158,
|
| 219 |
+
1.0344,
|
| 220 |
+
0.5894,
|
| 221 |
+
1.0901,
|
| 222 |
+
0.6885,
|
| 223 |
+
0.6165,
|
| 224 |
+
0.8454,
|
| 225 |
+
0.4978,
|
| 226 |
+
0.5759,
|
| 227 |
+
0.3523,
|
| 228 |
+
0.7135,
|
| 229 |
+
0.6804,
|
| 230 |
+
0.5833,
|
| 231 |
+
1.4146,
|
| 232 |
+
0.8986,
|
| 233 |
+
0.5659,
|
| 234 |
+
0.7069,
|
| 235 |
+
0.5338,
|
| 236 |
+
0.4889,
|
| 237 |
+
0.4917,
|
| 238 |
+
0.4069,
|
| 239 |
+
0.4999,
|
| 240 |
+
0.6866,
|
| 241 |
+
0.4093,
|
| 242 |
+
0.5709,
|
| 243 |
+
0.6065,
|
| 244 |
+
0.6415,
|
| 245 |
+
0.4944,
|
| 246 |
+
0.5726,
|
| 247 |
+
1.2042,
|
| 248 |
+
0.5458,
|
| 249 |
+
1.6887,
|
| 250 |
+
0.3971,
|
| 251 |
+
1.0600,
|
| 252 |
+
0.3943,
|
| 253 |
+
0.5537,
|
| 254 |
+
0.5444,
|
| 255 |
+
0.4089,
|
| 256 |
+
0.7468,
|
| 257 |
+
0.7744,
|
| 258 |
+
)
|
| 259 |
+
vae.mean = torch.tensor(mean)
|
| 260 |
+
vae.std = torch.tensor(std)
|
| 261 |
+
vae.scale = [vae.mean, 1.0 / vae.std]
|
| 262 |
+
return vae.to(dtype=dtype).eval().requires_grad_(False)
|
| 263 |
+
|
| 264 |
+
|
| 265 |
+
def _load_text_encoder(path: Path, dtype):
|
| 266 |
+
import torch
|
| 267 |
+
|
| 268 |
+
from fdanyone.vendor.diffsynth.models.wan_video_text_encoder import WanTextEncoder
|
| 269 |
+
|
| 270 |
+
state_dict = torch.load(path, map_location="cpu", weights_only=True)
|
| 271 |
+
state_dict = WanTextEncoder.state_dict_converter().from_civitai(state_dict)
|
| 272 |
+
with torch.device("meta"):
|
| 273 |
+
text_encoder = WanTextEncoder()
|
| 274 |
+
_strict_assign(text_encoder, state_dict, "Wan T5 text encoder")
|
| 275 |
+
del state_dict
|
| 276 |
+
return text_encoder.to(dtype=dtype).eval().requires_grad_(False)
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
def onload_all(model: Any) -> None:
|
| 280 |
+
"""Move every VRAM-managed submodule of ``model`` to its compute device."""
|
| 281 |
+
|
| 282 |
+
for module in model.modules():
|
| 283 |
+
if hasattr(module, "onload"):
|
| 284 |
+
module.onload()
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
def offload_all(model: Any) -> None:
|
| 288 |
+
"""Move every VRAM-managed submodule of ``model`` back to host memory."""
|
| 289 |
+
|
| 290 |
+
for module in model.modules():
|
| 291 |
+
if hasattr(module, "offload"):
|
| 292 |
+
module.offload()
|
| 293 |
+
|
| 294 |
+
|
| 295 |
+
def _enable_dit_streaming(
|
| 296 |
+
pipe: Any,
|
| 297 |
+
*,
|
| 298 |
+
device: str,
|
| 299 |
+
persistent_parameters: int,
|
| 300 |
+
) -> None:
|
| 301 |
+
"""Keep a bounded prefix of DiT weights on the GPU and stream the rest."""
|
| 302 |
+
|
| 303 |
+
import torch
|
| 304 |
+
|
| 305 |
+
from fdanyone.vendor.diffsynth.models.wan_video_dit import RMSNorm
|
| 306 |
+
from fdanyone.vendor.diffsynth.vram_management import (
|
| 307 |
+
AutoWrappedLinear,
|
| 308 |
+
AutoWrappedModule,
|
| 309 |
+
enable_vram_management,
|
| 310 |
+
)
|
| 311 |
+
|
| 312 |
+
# The vendored wrapper only owns registered child modules. Stage the full
|
| 313 |
+
# model first so standalone parameters (for example block modulation) stay
|
| 314 |
+
# on the compute device when wrapped modules are moved back to the CPU.
|
| 315 |
+
pipe.dit.to(device=device)
|
| 316 |
+
dtype: torch.dtype = next(iter(pipe.dit.parameters())).dtype
|
| 317 |
+
module_config: dict[str, Any] = {
|
| 318 |
+
"offload_dtype": dtype,
|
| 319 |
+
"offload_device": "cpu",
|
| 320 |
+
"onload_dtype": dtype,
|
| 321 |
+
"onload_device": device,
|
| 322 |
+
"computation_dtype": pipe.torch_dtype,
|
| 323 |
+
"computation_device": device,
|
| 324 |
+
}
|
| 325 |
+
enable_vram_management(
|
| 326 |
+
pipe.dit,
|
| 327 |
+
module_map={
|
| 328 |
+
torch.nn.Linear: AutoWrappedLinear,
|
| 329 |
+
# Every DiT ``Conv3d`` together (patch embedding, view-pack
|
| 330 |
+
# projections, pose encoder) is 15.9M parameters -- 0.3% of the
|
| 331 |
+
# model, 30 MiB in bf16. Streaming them buys nothing and forces
|
| 332 |
+
# callers to read layer metadata such as ``kernel_size`` through
|
| 333 |
+
# the wrapper, so they stay resident and unwrapped.
|
| 334 |
+
torch.nn.LayerNorm: AutoWrappedModule,
|
| 335 |
+
RMSNorm: AutoWrappedModule,
|
| 336 |
+
},
|
| 337 |
+
module_config=module_config,
|
| 338 |
+
max_num_param=persistent_parameters,
|
| 339 |
+
overflow_module_config={**module_config, "onload_device": "cpu"},
|
| 340 |
+
)
|
| 341 |
+
onload_all(pipe.dit)
|
| 342 |
+
torch.cuda.empty_cache()
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
def _enable_regional_compile(dit: Any, *, mode: str = "default") -> None:
|
| 346 |
+
"""Compile repeated transformer blocks in place with a static-shape contract.
|
| 347 |
+
|
| 348 |
+
``Module.compile`` swaps only the block's call implementation, so module
|
| 349 |
+
identity, class, and fully qualified parameter names survive. A rebuilt
|
| 350 |
+
``ModuleList`` of wrappers would rename every DiT weight.
|
| 351 |
+
"""
|
| 352 |
+
|
| 353 |
+
for block in dit.blocks:
|
| 354 |
+
block.compile(dynamic=False, mode=mode)
|
| 355 |
+
|
| 356 |
+
|
| 357 |
+
def load_pipeline(
|
| 358 |
+
*,
|
| 359 |
+
checkpoint_path: str | Path,
|
| 360 |
+
assets: BaseAssets,
|
| 361 |
+
device: str,
|
| 362 |
+
settings: ModeSettings,
|
| 363 |
+
load_text_encoder: bool = True,
|
| 364 |
+
) -> LoadedPipeline:
|
| 365 |
+
"""Load exactly the runtime models required by the released checkpoint."""
|
| 366 |
+
|
| 367 |
+
import torch
|
| 368 |
+
|
| 369 |
+
from fdanyone.vendor.diffsynth.pipelines.wan_video_spatem import (
|
| 370 |
+
WanVideoSpaTemPipeline,
|
| 371 |
+
)
|
| 372 |
+
|
| 373 |
+
dtype = torch.bfloat16
|
| 374 |
+
if load_text_encoder and assets.text_encoder is None:
|
| 375 |
+
raise AssetError(
|
| 376 |
+
"Text encoder does not exist. Supply an exported prompt embedding through "
|
| 377 |
+
"prompt_embedding_path, or run `python scripts/download_model.py` to install it."
|
| 378 |
+
)
|
| 379 |
+
tokenizer_path: str | None = None if assets.tokenizer is None else str(assets.tokenizer)
|
| 380 |
+
pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=dtype, tokenizer_path=tokenizer_path)
|
| 381 |
+
pipe.dit = _load_dit(
|
| 382 |
+
Path(checkpoint_path),
|
| 383 |
+
dtype,
|
| 384 |
+
settings=settings,
|
| 385 |
+
turbo_lora_path=assets.turbo_lora,
|
| 386 |
+
)
|
| 387 |
+
if settings.fp8_w8a8:
|
| 388 |
+
from fdanyone.model.quantization import quantize_dit_fp8_w8a8
|
| 389 |
+
|
| 390 |
+
pipe.dit.to(device=device)
|
| 391 |
+
w8a8_report = quantize_dit_fp8_w8a8(pipe.dit)
|
| 392 |
+
LOGGER.info(
|
| 393 |
+
"FP8 W8A8 DiT: modules=%d granularity=PerTensor",
|
| 394 |
+
w8a8_report.module_count,
|
| 395 |
+
)
|
| 396 |
+
if settings.regional_compile:
|
| 397 |
+
_enable_regional_compile(
|
| 398 |
+
pipe.dit,
|
| 399 |
+
mode="max-autotune-no-cudagraphs",
|
| 400 |
+
)
|
| 401 |
+
LOGGER.info(
|
| 402 |
+
"Regional DiT compile: blocks=%d dynamic=False mode=%s",
|
| 403 |
+
len(pipe.dit.blocks),
|
| 404 |
+
"max-autotune-no-cudagraphs",
|
| 405 |
+
)
|
| 406 |
+
pipe.vae = _load_vae(assets.vae, dtype)
|
| 407 |
+
target_decoder: object | None = None
|
| 408 |
+
if settings.tiny_decoders:
|
| 409 |
+
from fdanyone.model.tiny_decoder import load_tiny_wan_decoder
|
| 410 |
+
|
| 411 |
+
if assets.tiny_decoder is None:
|
| 412 |
+
raise ConfigurationError("TAEW2.2 checkpoint path was not resolved.")
|
| 413 |
+
target_decoder = load_tiny_wan_decoder(assets.tiny_decoder)
|
| 414 |
+
if load_text_encoder:
|
| 415 |
+
pipe.text_encoder = _load_text_encoder(assets.text_encoder, dtype)
|
| 416 |
+
pipe.prompter.fetch_models(pipe.text_encoder)
|
| 417 |
+
pipe.height_division_factor = pipe.vae.upsampling_factor * 2
|
| 418 |
+
pipe.width_division_factor = pipe.vae.upsampling_factor * 2
|
| 419 |
+
persistent_parameters: int | None = 0 if settings.stream_dit_weights else None
|
| 420 |
+
if persistent_parameters is not None:
|
| 421 |
+
_enable_dit_streaming(
|
| 422 |
+
pipe,
|
| 423 |
+
device=device,
|
| 424 |
+
persistent_parameters=persistent_parameters,
|
| 425 |
+
)
|
| 426 |
+
return LoadedPipeline(pipe=pipe, target_decoder=target_decoder)
|
|
@@ -0,0 +1,247 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Stage-static conditioning for the Wan DiT.
|
| 2 |
+
|
| 3 |
+
``prepare_static_conditioning`` computes every tensor that does not change
|
| 4 |
+
between denoising steps of one generation stage; ``forward_dynamic`` then runs
|
| 5 |
+
only the noisy-latent patching, the route gather, the transformer blocks, and
|
| 6 |
+
the head. Both are exact refactors of ``WanModel.forward`` for the released
|
| 7 |
+
packed-source configuration, and ``tests/test_prepared_parity.py`` pins that
|
| 8 |
+
equivalence bit for bit.
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
from __future__ import annotations
|
| 12 |
+
|
| 13 |
+
from contextlib import nullcontext
|
| 14 |
+
from dataclasses import dataclass
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
from einops import rearrange, repeat
|
| 18 |
+
|
| 19 |
+
from fdanyone.vendor.diffsynth.models.wan_video_dit import (
|
| 20 |
+
WanModel,
|
| 21 |
+
pack_viewpack_tokens,
|
| 22 |
+
sinusoidal_embedding_1d,
|
| 23 |
+
split_source_views,
|
| 24 |
+
)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
@dataclass(frozen=True)
|
| 28 |
+
class PreparedCrossAttention:
|
| 29 |
+
"""Prompt keys and values projected for one transformer block."""
|
| 30 |
+
|
| 31 |
+
key: torch.Tensor
|
| 32 |
+
value: torch.Tensor
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
@dataclass(frozen=True)
|
| 36 |
+
class PreparedWanConditioning:
|
| 37 |
+
"""All stage-static inputs consumed by ``forward_dynamic``."""
|
| 38 |
+
|
| 39 |
+
cross_attention: tuple[PreparedCrossAttention, ...]
|
| 40 |
+
source_tokens: torch.Tensor
|
| 41 |
+
pose_tokens: torch.Tensor | None
|
| 42 |
+
null_pose_tokens: torch.Tensor | None
|
| 43 |
+
temporal_freqs: torch.Tensor
|
| 44 |
+
multiview_freqs: torch.Tensor
|
| 45 |
+
time_embeddings: tuple[torch.Tensor, ...]
|
| 46 |
+
time_modulations: tuple[torch.Tensor, ...]
|
| 47 |
+
grid_size: tuple[int, int, int]
|
| 48 |
+
group_size: int
|
| 49 |
+
packed_views: int
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def _prepare_source_tokens(
|
| 53 |
+
model: WanModel,
|
| 54 |
+
x_src: torch.Tensor,
|
| 55 |
+
) -> tuple[torch.Tensor, tuple[int, int, int], int]:
|
| 56 |
+
"""Patch and pack fixed source/reference latents once."""
|
| 57 |
+
|
| 58 |
+
x_src, x_src_2x, x_src_4x = split_source_views(x_src)
|
| 59 |
+
source_tokens, grid = model.patchify(x_src)
|
| 60 |
+
packed = [source_tokens]
|
| 61 |
+
# ``WanModel.forward`` crops the packed tile to the noisy-latent grid;
|
| 62 |
+
# ``forward_dynamic`` rejects any group whose grid differs from this source
|
| 63 |
+
# grid, so the two crops are the same.
|
| 64 |
+
grid_size = tuple(int(size) for size in grid)
|
| 65 |
+
if model.use_viewpack and x_src_2x is not None:
|
| 66 |
+
x_pack = pack_viewpack_tokens(
|
| 67 |
+
model.viewpack_embedding, x_src_2x, x_src_4x, grid_size
|
| 68 |
+
)
|
| 69 |
+
packed.append(x_pack.to(source_tokens.dtype))
|
| 70 |
+
return torch.cat(packed, dim=0), grid_size, len(packed)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def _prepare_prompt_kv(
|
| 74 |
+
model: WanModel,
|
| 75 |
+
context: torch.Tensor,
|
| 76 |
+
) -> tuple[PreparedCrossAttention, ...]:
|
| 77 |
+
"""Project the fixed prompt into per-block keys and values."""
|
| 78 |
+
|
| 79 |
+
return tuple(
|
| 80 |
+
PreparedCrossAttention(
|
| 81 |
+
key=block.cross_attn.norm_k(block.cross_attn.k(context)),
|
| 82 |
+
value=block.cross_attn.v(context),
|
| 83 |
+
)
|
| 84 |
+
for block in model.blocks
|
| 85 |
+
)
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def _prepare_pose_tokens(
|
| 89 |
+
model: WanModel,
|
| 90 |
+
*,
|
| 91 |
+
skeletons: torch.Tensor | tuple[torch.Tensor, ...],
|
| 92 |
+
packed_views: int,
|
| 93 |
+
group_size: int,
|
| 94 |
+
pose_batch_size: int | None,
|
| 95 |
+
device: torch.device,
|
| 96 |
+
stage_timer,
|
| 97 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 98 |
+
"""Encode every skeleton view once, in bounded activation batches.
|
| 99 |
+
|
| 100 |
+
``skeletons`` is either one stacked ``[v, c, f, h, w]`` tensor or a tuple
|
| 101 |
+
of per-view ``[1, c, f, h, w]`` tensors; the tuple form lets callers feed
|
| 102 |
+
cached host views without materializing the whole set as one copy.
|
| 103 |
+
"""
|
| 104 |
+
|
| 105 |
+
views = (
|
| 106 |
+
tuple(skeletons.split(1, dim=0))
|
| 107 |
+
if isinstance(skeletons, torch.Tensor)
|
| 108 |
+
else tuple(skeletons)
|
| 109 |
+
)
|
| 110 |
+
null_view = -torch.ones_like(views[0])
|
| 111 |
+
num_real_views = len(views)
|
| 112 |
+
num_pose_views = num_real_views + packed_views
|
| 113 |
+
batch_size = group_size + packed_views if pose_batch_size is None else pose_batch_size
|
| 114 |
+
encoded_pose = None
|
| 115 |
+
timer = nullcontext() if stage_timer is None else stage_timer.measure("pose_encoder")
|
| 116 |
+
with torch.profiler.record_function("dit.pose_encoder"), timer:
|
| 117 |
+
for start in range(0, num_pose_views, batch_size):
|
| 118 |
+
end = min(start + batch_size, num_pose_views)
|
| 119 |
+
parts = list(views[start:min(end, num_real_views)])
|
| 120 |
+
parts.extend([null_view] * (end - max(start, num_real_views)))
|
| 121 |
+
pose_batch = parts[0] if len(parts) == 1 else torch.cat(parts, dim=0)
|
| 122 |
+
encoded_batch = rearrange(
|
| 123 |
+
model.pose_encoder(pose_batch.to(device=device)), "v c f h w -> v (f h w) c"
|
| 124 |
+
)
|
| 125 |
+
if encoded_pose is None:
|
| 126 |
+
encoded_pose = encoded_batch.new_empty(
|
| 127 |
+
(num_pose_views, *encoded_batch.shape[1:])
|
| 128 |
+
)
|
| 129 |
+
encoded_pose[start:end].copy_(encoded_batch)
|
| 130 |
+
if encoded_pose is None:
|
| 131 |
+
raise ValueError("Prepared pose conditioning requires at least one view.")
|
| 132 |
+
return encoded_pose[:-packed_views], encoded_pose[-packed_views:]
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def prepare_static_conditioning(
|
| 136 |
+
model: WanModel,
|
| 137 |
+
*,
|
| 138 |
+
x_src: torch.Tensor,
|
| 139 |
+
context: torch.Tensor,
|
| 140 |
+
skeletons: torch.Tensor | tuple[torch.Tensor, ...] | None,
|
| 141 |
+
timesteps: torch.Tensor,
|
| 142 |
+
group_size: int,
|
| 143 |
+
pose_batch_size: int | None = None,
|
| 144 |
+
stage_timer=None,
|
| 145 |
+
) -> PreparedWanConditioning:
|
| 146 |
+
"""Compute every stage-invariant tensor once for dynamic denoising."""
|
| 147 |
+
|
| 148 |
+
if group_size <= 0:
|
| 149 |
+
raise ValueError(f"group_size must be positive, got {group_size}.")
|
| 150 |
+
if pose_batch_size is not None and pose_batch_size <= 0:
|
| 151 |
+
raise ValueError(f"pose_batch_size must be positive or None, got {pose_batch_size}.")
|
| 152 |
+
if model.has_image_input:
|
| 153 |
+
raise ValueError("Prepared conditioning supports the released packed-source path only.")
|
| 154 |
+
|
| 155 |
+
context = model.text_embedding(context)
|
| 156 |
+
source_tokens, grid_size, packed_views = _prepare_source_tokens(model, x_src)
|
| 157 |
+
f, h, w = grid_size
|
| 158 |
+
effective_views = group_size + packed_views
|
| 159 |
+
device = source_tokens.device
|
| 160 |
+
pose_tokens = None
|
| 161 |
+
null_pose_tokens = None
|
| 162 |
+
if model.use_pose_encoder:
|
| 163 |
+
if skeletons is None:
|
| 164 |
+
raise ValueError("Prepared pose conditioning requires skeleton tensors.")
|
| 165 |
+
pose_tokens, null_pose_tokens = _prepare_pose_tokens(
|
| 166 |
+
model,
|
| 167 |
+
skeletons=skeletons,
|
| 168 |
+
packed_views=packed_views,
|
| 169 |
+
group_size=group_size,
|
| 170 |
+
pose_batch_size=pose_batch_size,
|
| 171 |
+
device=device,
|
| 172 |
+
stage_timer=stage_timer,
|
| 173 |
+
)
|
| 174 |
+
|
| 175 |
+
time_embeddings = []
|
| 176 |
+
time_modulations = []
|
| 177 |
+
for timestep in timesteps.flatten():
|
| 178 |
+
batched = timestep.reshape(1).to(dtype=source_tokens.dtype, device=device)
|
| 179 |
+
batched = torch.cat([batched] * group_size, dim=0)
|
| 180 |
+
batched = torch.cat(
|
| 181 |
+
[batched, torch.zeros(packed_views, dtype=batched.dtype, device=batched.device)]
|
| 182 |
+
)
|
| 183 |
+
embedding = model.time_embedding(
|
| 184 |
+
sinusoidal_embedding_1d(model.freq_dim, batched).to(source_tokens.dtype)
|
| 185 |
+
)
|
| 186 |
+
time_embeddings.append(embedding)
|
| 187 |
+
time_modulations.append(model.time_projection(embedding).unflatten(1, (6, model.dim)))
|
| 188 |
+
|
| 189 |
+
return PreparedWanConditioning(
|
| 190 |
+
cross_attention=_prepare_prompt_kv(
|
| 191 |
+
model, repeat(context, "1 l c -> v l c", v=effective_views)
|
| 192 |
+
),
|
| 193 |
+
source_tokens=source_tokens,
|
| 194 |
+
pose_tokens=pose_tokens,
|
| 195 |
+
null_pose_tokens=null_pose_tokens,
|
| 196 |
+
temporal_freqs=model._rope_table(f, h, w, device),
|
| 197 |
+
multiview_freqs=model._rope_table(effective_views, h, w, device),
|
| 198 |
+
time_embeddings=tuple(time_embeddings),
|
| 199 |
+
time_modulations=tuple(time_modulations),
|
| 200 |
+
grid_size=grid_size,
|
| 201 |
+
group_size=group_size,
|
| 202 |
+
packed_views=packed_views,
|
| 203 |
+
)
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
def forward_dynamic(
|
| 207 |
+
model: WanModel,
|
| 208 |
+
*,
|
| 209 |
+
x: torch.Tensor,
|
| 210 |
+
prepared: PreparedWanConditioning,
|
| 211 |
+
view_indices: torch.Tensor,
|
| 212 |
+
step_index: int,
|
| 213 |
+
) -> torch.Tensor:
|
| 214 |
+
"""Run only noisy-latent patching, route gather, blocks, and head."""
|
| 215 |
+
|
| 216 |
+
if x.shape[0] != prepared.group_size or view_indices.numel() != prepared.group_size:
|
| 217 |
+
raise ValueError("Dynamic group shape does not match prepared conditioning.")
|
| 218 |
+
x, grid_size = model.patchify(x)
|
| 219 |
+
if tuple(int(size) for size in grid_size) != prepared.grid_size:
|
| 220 |
+
raise ValueError("Dynamic latent grid does not match prepared conditioning.")
|
| 221 |
+
x = torch.cat([x, prepared.source_tokens], dim=0)
|
| 222 |
+
if prepared.pose_tokens is not None:
|
| 223 |
+
if prepared.null_pose_tokens is None:
|
| 224 |
+
raise ValueError("Prepared null-pose tokens are missing.")
|
| 225 |
+
x[: prepared.group_size].add_(
|
| 226 |
+
torch.index_select(prepared.pose_tokens, 0, view_indices)
|
| 227 |
+
)
|
| 228 |
+
x[prepared.group_size :].add_(prepared.null_pose_tokens)
|
| 229 |
+
|
| 230 |
+
t = prepared.time_embeddings[step_index]
|
| 231 |
+
t_mod = prepared.time_modulations[step_index]
|
| 232 |
+
v = x.shape[0]
|
| 233 |
+
for block, cross_attention in zip(model.blocks, prepared.cross_attention):
|
| 234 |
+
x = block(
|
| 235 |
+
x,
|
| 236 |
+
None,
|
| 237 |
+
t_mod,
|
| 238 |
+
prepared.temporal_freqs,
|
| 239 |
+
prepared.multiview_freqs,
|
| 240 |
+
(v, *prepared.grid_size),
|
| 241 |
+
cross_attention_key=cross_attention.key,
|
| 242 |
+
cross_attention_value=cross_attention.value,
|
| 243 |
+
)
|
| 244 |
+
|
| 245 |
+
x = x[: -prepared.packed_views]
|
| 246 |
+
t = t[: -prepared.packed_views]
|
| 247 |
+
return model.unpatchify(model.head(x, t), prepared.grid_size)
|
|
@@ -0,0 +1,147 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Low-overhead CUDA stage timing and one-step DiT profiling."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import os
|
| 7 |
+
import time
|
| 8 |
+
from collections.abc import Iterator
|
| 9 |
+
from contextlib import contextmanager
|
| 10 |
+
from dataclasses import dataclass, field
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
from typing import Protocol, cast
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class _CudaEvent(Protocol):
|
| 16 |
+
def record(self) -> None: ...
|
| 17 |
+
|
| 18 |
+
def elapsed_time(self, end_event: _CudaEvent) -> float: ...
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
@dataclass
|
| 22 |
+
class CudaStageTimer:
|
| 23 |
+
"""Accumulate named CUDA event pairs and synchronize only when read."""
|
| 24 |
+
|
| 25 |
+
device: str
|
| 26 |
+
_events: dict[str, list[tuple[_CudaEvent, _CudaEvent]]] = field(default_factory=dict)
|
| 27 |
+
|
| 28 |
+
@contextmanager
|
| 29 |
+
def measure(self, name: str) -> Iterator[None]:
|
| 30 |
+
"""Record one asynchronous CUDA interval under ``name``."""
|
| 31 |
+
|
| 32 |
+
import torch
|
| 33 |
+
|
| 34 |
+
start = cast(_CudaEvent, torch.cuda.Event(enable_timing=True))
|
| 35 |
+
end = cast(_CudaEvent, torch.cuda.Event(enable_timing=True))
|
| 36 |
+
start.record()
|
| 37 |
+
try:
|
| 38 |
+
yield
|
| 39 |
+
finally:
|
| 40 |
+
end.record()
|
| 41 |
+
self._events.setdefault(name, []).append((start, end))
|
| 42 |
+
|
| 43 |
+
def elapsed_seconds(self, name: str) -> float:
|
| 44 |
+
"""Return the sum of all intervals recorded under ``name``."""
|
| 45 |
+
|
| 46 |
+
import torch
|
| 47 |
+
|
| 48 |
+
intervals = self._events.get(name, [])
|
| 49 |
+
if not intervals:
|
| 50 |
+
return 0.0
|
| 51 |
+
torch.cuda.synchronize(self.device)
|
| 52 |
+
return sum(start.elapsed_time(end) for start, end in intervals) / 1000.0
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
@dataclass
|
| 56 |
+
class StageClock:
|
| 57 |
+
"""One run's named wall-clock stage totals beside its CUDA stage timer."""
|
| 58 |
+
|
| 59 |
+
cuda: CudaStageTimer
|
| 60 |
+
elapsed: dict[str, float] = field(default_factory=dict)
|
| 61 |
+
|
| 62 |
+
@contextmanager
|
| 63 |
+
def stage(self, name: str) -> Iterator[None]:
|
| 64 |
+
"""Record the wall-clock duration of the enclosed stage under ``name``."""
|
| 65 |
+
|
| 66 |
+
started = time.monotonic()
|
| 67 |
+
try:
|
| 68 |
+
yield
|
| 69 |
+
finally:
|
| 70 |
+
self.elapsed[name] = time.monotonic() - started
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
_PROFILE_LABELS = (
|
| 74 |
+
"dit.temporal_attention",
|
| 75 |
+
"dit.multiview_attention",
|
| 76 |
+
"dit.cross_attention",
|
| 77 |
+
"dit.ffn",
|
| 78 |
+
"dit.rope",
|
| 79 |
+
"dit.pose_encoder",
|
| 80 |
+
"dit.route_copies",
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def _event_time(event: object, device: bool) -> float:
|
| 85 |
+
"""Read a profiler event duration in microseconds across torch versions."""
|
| 86 |
+
|
| 87 |
+
candidates = ("device_time_total", "cuda_time_total") if device else ("cpu_time_total",)
|
| 88 |
+
for attribute in candidates:
|
| 89 |
+
value = getattr(event, attribute, None)
|
| 90 |
+
if value is not None:
|
| 91 |
+
return float(value)
|
| 92 |
+
return 0.0
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def _stage_summary(event: object | None) -> dict[str, float | int]:
|
| 96 |
+
"""Summarize one profiled named range, or report an unrecorded stage."""
|
| 97 |
+
|
| 98 |
+
if event is None:
|
| 99 |
+
return {"calls": 0, "cpu_seconds": 0.0, "device_seconds": 0.0}
|
| 100 |
+
return {
|
| 101 |
+
"calls": int(getattr(event, "count", 0)),
|
| 102 |
+
"cpu_seconds": _event_time(event, device=False) / 1_000_000.0,
|
| 103 |
+
"device_seconds": _event_time(event, device=True) / 1_000_000.0,
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
@contextmanager
|
| 108 |
+
def profile_dit_step(*, enabled: bool = True) -> Iterator[None]:
|
| 109 |
+
"""Profile one call when enabled and ``FDANYONE_DIT_PROFILE_DIR`` is set."""
|
| 110 |
+
|
| 111 |
+
output_dir: str | None = os.environ.get("FDANYONE_DIT_PROFILE_DIR")
|
| 112 |
+
if not enabled or output_dir is None:
|
| 113 |
+
yield
|
| 114 |
+
return
|
| 115 |
+
|
| 116 |
+
import torch
|
| 117 |
+
|
| 118 |
+
destination = Path(output_dir).expanduser().resolve()
|
| 119 |
+
destination.mkdir(parents=True, exist_ok=True)
|
| 120 |
+
activities = [torch.profiler.ProfilerActivity.CPU]
|
| 121 |
+
if torch.cuda.is_available():
|
| 122 |
+
activities.append(torch.profiler.ProfilerActivity.CUDA)
|
| 123 |
+
with torch.profiler.profile(
|
| 124 |
+
activities=activities,
|
| 125 |
+
record_shapes=True,
|
| 126 |
+
profile_memory=True,
|
| 127 |
+
with_stack=False,
|
| 128 |
+
) as profile:
|
| 129 |
+
yield
|
| 130 |
+
|
| 131 |
+
profile.export_chrome_trace(str(destination / "dit-step-trace.json"))
|
| 132 |
+
averages = profile.key_averages()
|
| 133 |
+
(destination / "dit-step-table.txt").write_text(
|
| 134 |
+
averages.table(
|
| 135 |
+
sort_by="self_cuda_time_total" if torch.cuda.is_available() else "self_cpu_time_total",
|
| 136 |
+
row_limit=-1,
|
| 137 |
+
)
|
| 138 |
+
)
|
| 139 |
+
events = list(averages)
|
| 140 |
+
by_key = {event.key: event for event in events}
|
| 141 |
+
stages = {label: _stage_summary(by_key.get(label)) for label in _PROFILE_LABELS}
|
| 142 |
+
(destination / "dit-step-summary.json").write_text(
|
| 143 |
+
json.dumps(
|
| 144 |
+
{"profiled_steps": 1, "stages": stages},
|
| 145 |
+
indent=2,
|
| 146 |
+
)
|
| 147 |
+
)
|
|
@@ -0,0 +1,87 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Selective quantization policies for the 4DAnyone DiT."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import re
|
| 6 |
+
from dataclasses import dataclass
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
_FP8_WEIGHT_ONLY_PATTERN: re.Pattern[str] = re.compile(
|
| 11 |
+
r"blocks\.(?P<block>\d+)\."
|
| 12 |
+
r"(?:(?:self_attn|self_attn_mvs|cross_attn)\.(?:q|k|v|o)|"
|
| 13 |
+
r"ffn\.(?:0|2))"
|
| 14 |
+
)
|
| 15 |
+
|
| 16 |
+
QUANTIZED_BLOCK_RANGE: tuple[int, int] = (3, 26)
|
| 17 |
+
"""Inclusive interior DiT block range the guide approves for FP8 projections."""
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def is_quantizable_dit_projection(fqn: str) -> bool:
|
| 21 |
+
"""Return whether an FQN is on the safe interior-projection allowlist.
|
| 22 |
+
|
| 23 |
+
W8A8 uses the same boundary-protected projection set established by the
|
| 24 |
+
earlier weight-only experiment; the name describes the surviving policy.
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
match: re.Match[str] | None = _FP8_WEIGHT_ONLY_PATTERN.fullmatch(fqn)
|
| 28 |
+
if match is None:
|
| 29 |
+
return False
|
| 30 |
+
block_index: int = int(match.group("block"))
|
| 31 |
+
first_block, last_block = QUANTIZED_BLOCK_RANGE
|
| 32 |
+
return first_block <= block_index <= last_block
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
@dataclass(frozen=True, slots=True)
|
| 36 |
+
class FP8W8A8Report:
|
| 37 |
+
"""Summary of one TorchAO dynamic W8A8 conversion."""
|
| 38 |
+
|
| 39 |
+
module_count: int
|
| 40 |
+
"""Number of guide-approved Linear modules passed to TorchAO."""
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _cast_w8a8_activation_to_bf16(
|
| 44 |
+
module: torch.nn.Module,
|
| 45 |
+
args: tuple[object, ...],
|
| 46 |
+
) -> tuple[object, ...] | None:
|
| 47 |
+
"""Restore the BF16 autocast input contract at TorchAO's subclass seam."""
|
| 48 |
+
|
| 49 |
+
del module
|
| 50 |
+
if not args:
|
| 51 |
+
return None
|
| 52 |
+
activation = args[0]
|
| 53 |
+
if (
|
| 54 |
+
not isinstance(activation, torch.Tensor)
|
| 55 |
+
or not activation.is_floating_point()
|
| 56 |
+
or activation.dtype is torch.bfloat16
|
| 57 |
+
):
|
| 58 |
+
return None
|
| 59 |
+
return (activation.to(dtype=torch.bfloat16), *args[1:])
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def quantize_dit_fp8_w8a8(dit: torch.nn.Module) -> FP8W8A8Report:
|
| 63 |
+
"""Apply dynamic per-tensor W8A8 to the approved projections."""
|
| 64 |
+
|
| 65 |
+
from torchao.quantization import (
|
| 66 |
+
Float8DynamicActivationFloat8WeightConfig,
|
| 67 |
+
PerTensor,
|
| 68 |
+
quantize_,
|
| 69 |
+
)
|
| 70 |
+
|
| 71 |
+
selected: set[str] = {
|
| 72 |
+
fqn
|
| 73 |
+
for fqn, module in dit.named_modules()
|
| 74 |
+
if isinstance(module, torch.nn.Linear) and is_quantizable_dit_projection(fqn)
|
| 75 |
+
}
|
| 76 |
+
quantize_(
|
| 77 |
+
dit,
|
| 78 |
+
Float8DynamicActivationFloat8WeightConfig(
|
| 79 |
+
granularity=PerTensor(), activation_value_ub=None
|
| 80 |
+
),
|
| 81 |
+
filter_fn=lambda _, fqn: fqn in selected,
|
| 82 |
+
)
|
| 83 |
+
for fqn in selected:
|
| 84 |
+
dit.get_submodule(fqn).register_forward_pre_hook(
|
| 85 |
+
_cast_w8a8_activation_to_bf16
|
| 86 |
+
)
|
| 87 |
+
return FP8W8A8Report(module_count=len(selected))
|
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Target-context routing across view groups."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from itertools import pairwise
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def _validate_grouping(num_views: int, group_size: int) -> int:
|
| 9 |
+
if num_views <= 0 or group_size <= 0 or num_views % group_size:
|
| 10 |
+
raise ValueError(f"num_views={num_views} must be divisible by positive group_size={group_size}.")
|
| 11 |
+
return num_views // group_size
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def view_groups(
|
| 15 |
+
num_views: int,
|
| 16 |
+
group_size: int,
|
| 17 |
+
offset: int = 0,
|
| 18 |
+
*,
|
| 19 |
+
circular: bool = True,
|
| 20 |
+
) -> tuple[tuple[int, ...], ...]:
|
| 21 |
+
"""Partition one camera layer, optionally without joining its endpoints."""
|
| 22 |
+
|
| 23 |
+
num_groups = _validate_grouping(num_views, group_size)
|
| 24 |
+
if circular:
|
| 25 |
+
return tuple(
|
| 26 |
+
tuple((group_index * group_size + offset + local_index) % num_views for local_index in range(group_size))
|
| 27 |
+
for group_index in range(num_groups)
|
| 28 |
+
)
|
| 29 |
+
|
| 30 |
+
# Shifting an open sequence creates smaller boundary groups instead of a
|
| 31 |
+
# false neighborhood between the two ends of a partial yaw span.
|
| 32 |
+
offset %= group_size
|
| 33 |
+
boundaries = [0]
|
| 34 |
+
if offset:
|
| 35 |
+
boundaries.append(offset)
|
| 36 |
+
boundaries.extend(range(offset + group_size, num_views, group_size))
|
| 37 |
+
boundaries.append(num_views)
|
| 38 |
+
return tuple(tuple(range(start, end)) for start, end in pairwise(boundaries))
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def routing_steps(
|
| 42 |
+
*,
|
| 43 |
+
views_per_layer: int,
|
| 44 |
+
num_layers: int,
|
| 45 |
+
group_size: int,
|
| 46 |
+
num_steps: int,
|
| 47 |
+
enable_tcr: bool,
|
| 48 |
+
circular: bool,
|
| 49 |
+
) -> tuple[tuple[tuple[int, ...], ...], ...]:
|
| 50 |
+
"""Return layer-local target groups for every denoising step."""
|
| 51 |
+
|
| 52 |
+
if num_steps <= 0:
|
| 53 |
+
raise ValueError(f"num_steps must be positive, got {num_steps}.")
|
| 54 |
+
if num_layers <= 0:
|
| 55 |
+
raise ValueError(f"num_layers must be positive, got {num_layers}.")
|
| 56 |
+
_validate_grouping(views_per_layer, group_size)
|
| 57 |
+
return tuple(
|
| 58 |
+
tuple(
|
| 59 |
+
tuple(layer_index * views_per_layer + view_index for view_index in group)
|
| 60 |
+
for layer_index in range(num_layers)
|
| 61 |
+
for group in view_groups(
|
| 62 |
+
views_per_layer,
|
| 63 |
+
group_size,
|
| 64 |
+
step_index if enable_tcr else 0,
|
| 65 |
+
circular=circular,
|
| 66 |
+
)
|
| 67 |
+
)
|
| 68 |
+
for step_index in range(num_steps)
|
| 69 |
+
)
|
|
@@ -0,0 +1,60 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Target-only adapter for the pinned TAEW2.2 decoder experiment."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Any
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from fdanyone.errors import AssetError, ConfigurationError
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def validate_tiny_wan_decoder(decoder: Any) -> None:
|
| 14 |
+
"""Reject tiny decoders that do not match Wan 2.2 TI2V-5B latents."""
|
| 15 |
+
|
| 16 |
+
latent_channels: int = int(getattr(decoder, "latent_channels", 0))
|
| 17 |
+
patch_size: int = int(getattr(decoder, "patch_size", 0))
|
| 18 |
+
if latent_channels != 48 or patch_size != 2:
|
| 19 |
+
raise ConfigurationError(
|
| 20 |
+
"The 4DAnyone Wan 2.2 5B VAE requires 48 latent channels and "
|
| 21 |
+
f"patch size 2; got {latent_channels} channels and patch size "
|
| 22 |
+
f"{patch_size}. Use taew2_2, not taew2_1."
|
| 23 |
+
)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def load_tiny_wan_decoder(checkpoint_path: str | Path) -> torch.nn.Module:
|
| 27 |
+
"""Load the target-only TAEW2.2 decoder in its documented FP16 dtype."""
|
| 28 |
+
|
| 29 |
+
path: Path = Path(checkpoint_path).expanduser().resolve()
|
| 30 |
+
if not path.is_file():
|
| 31 |
+
raise AssetError(f"TAEW2.2 checkpoint does not exist: {path}")
|
| 32 |
+
try:
|
| 33 |
+
from taehv import TAEHV # pyrefly: ignore [missing-import]
|
| 34 |
+
except ImportError as exc:
|
| 35 |
+
raise AssetError(
|
| 36 |
+
"TAEHV is required by target_decoder='taew2_2'; use the pixi "
|
| 37 |
+
"fast environment."
|
| 38 |
+
) from exc
|
| 39 |
+
decoder: torch.nn.Module = TAEHV(checkpoint_path=str(path))
|
| 40 |
+
validate_tiny_wan_decoder(decoder)
|
| 41 |
+
return decoder.to(dtype=torch.float16).eval().requires_grad_(False)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def decode_tiny_target_video(
|
| 45 |
+
decoder: Any,
|
| 46 |
+
latents: torch.Tensor,
|
| 47 |
+
) -> torch.Tensor:
|
| 48 |
+
"""Decode normalized ``NCTHW`` latents to ``NCTHW`` pixels in [-1, 1]."""
|
| 49 |
+
|
| 50 |
+
if latents.ndim != 5 or latents.shape[1] != 48:
|
| 51 |
+
raise ConfigurationError(
|
| 52 |
+
"TAEW2.2 target decode expects NCTHW latents with 48 channels; "
|
| 53 |
+
f"got {tuple(latents.shape)}."
|
| 54 |
+
)
|
| 55 |
+
frame_first: torch.Tensor = decoder.decode_video(
|
| 56 |
+
latents.transpose(1, 2),
|
| 57 |
+
parallel=False,
|
| 58 |
+
show_progress_bar=False,
|
| 59 |
+
)
|
| 60 |
+
return frame_first.transpose(1, 2).mul(2.0).sub(1.0)
|
|
@@ -0,0 +1,125 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Strict streaming merge for Kijai's Wan2.2 5B Turbo-LoRA format."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from dataclasses import dataclass
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from fdanyone.errors import AssetError
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
@dataclass(frozen=True, slots=True)
|
| 14 |
+
class TurboLoraReport:
|
| 15 |
+
"""Auditable summary of one completed adapter merge."""
|
| 16 |
+
|
| 17 |
+
lora_modules: int
|
| 18 |
+
direct_parameters: int
|
| 19 |
+
max_rank: int
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def _target_name(adapter_name: str, suffix: str, parameter: str) -> str:
|
| 23 |
+
"""Map one publisher key to the corresponding Wan parameter name."""
|
| 24 |
+
|
| 25 |
+
prefix = "diffusion_model."
|
| 26 |
+
if not adapter_name.startswith(prefix) or not adapter_name.endswith(suffix):
|
| 27 |
+
raise AssetError(f"Unsupported Turbo-LoRA key {adapter_name!r}.")
|
| 28 |
+
return adapter_name[len(prefix) : -len(suffix)] + parameter
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def merge_wan_turbo_lora(
|
| 32 |
+
model: torch.nn.Module,
|
| 33 |
+
path: str | Path,
|
| 34 |
+
) -> TurboLoraReport:
|
| 35 |
+
"""Merge all low-rank, bias, and normalization deltas without fallback."""
|
| 36 |
+
|
| 37 |
+
checkpoint = Path(path).expanduser().resolve()
|
| 38 |
+
if not checkpoint.is_file():
|
| 39 |
+
raise AssetError(f"Turbo-LoRA checkpoint is missing: {checkpoint}")
|
| 40 |
+
try:
|
| 41 |
+
from safetensors import safe_open
|
| 42 |
+
except ImportError as exc:
|
| 43 |
+
raise AssetError("safetensors is required for Turbo-LoRA.") from exc
|
| 44 |
+
|
| 45 |
+
parameters = dict(model.named_parameters())
|
| 46 |
+
with safe_open(checkpoint, framework="pt", device="cpu") as adapter:
|
| 47 |
+
keys = set(adapter.keys())
|
| 48 |
+
down_keys = sorted(key for key in keys if key.endswith(".lora_down.weight"))
|
| 49 |
+
up_keys = {key for key in keys if key.endswith(".lora_up.weight")}
|
| 50 |
+
direct_keys = sorted(
|
| 51 |
+
key for key in keys if key.endswith((".diff", ".diff_b"))
|
| 52 |
+
)
|
| 53 |
+
consumed_up: set[str] = set()
|
| 54 |
+
max_rank = 0
|
| 55 |
+
|
| 56 |
+
for down_key in down_keys:
|
| 57 |
+
up_key = down_key.removesuffix(".lora_down.weight") + ".lora_up.weight"
|
| 58 |
+
if up_key not in keys:
|
| 59 |
+
raise AssetError(f"Turbo-LoRA has no paired key {up_key!r}.")
|
| 60 |
+
consumed_up.add(up_key)
|
| 61 |
+
target_name = _target_name(down_key, ".lora_down.weight", ".weight")
|
| 62 |
+
if target_name not in parameters:
|
| 63 |
+
raise AssetError(f"Turbo-LoRA target is missing: {target_name}")
|
| 64 |
+
down_shape = tuple(adapter.get_slice(down_key).get_shape())
|
| 65 |
+
up_shape = tuple(adapter.get_slice(up_key).get_shape())
|
| 66 |
+
target_shape = tuple(parameters[target_name].shape)
|
| 67 |
+
if (
|
| 68 |
+
len(down_shape) != 2
|
| 69 |
+
or len(up_shape) != 2
|
| 70 |
+
or up_shape[1] != down_shape[0]
|
| 71 |
+
or (up_shape[0], down_shape[1]) != target_shape
|
| 72 |
+
):
|
| 73 |
+
raise AssetError(
|
| 74 |
+
f"Turbo-LoRA shape mismatch for {target_name}: "
|
| 75 |
+
f"down={down_shape}, up={up_shape}, target={target_shape}."
|
| 76 |
+
)
|
| 77 |
+
max_rank = max(max_rank, down_shape[0])
|
| 78 |
+
if consumed_up != up_keys:
|
| 79 |
+
raise AssetError(
|
| 80 |
+
f"Turbo-LoRA has unpaired up keys: {sorted(up_keys - consumed_up)}"
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
direct_targets: dict[str, str] = {}
|
| 84 |
+
for key in direct_keys:
|
| 85 |
+
if key.endswith(".diff_b"):
|
| 86 |
+
target_name = _target_name(key, ".diff_b", ".bias")
|
| 87 |
+
else:
|
| 88 |
+
target_name = _target_name(key, ".diff", ".weight")
|
| 89 |
+
if target_name not in parameters:
|
| 90 |
+
raise AssetError(f"Turbo-LoRA target is missing: {target_name}")
|
| 91 |
+
patch_shape = tuple(adapter.get_slice(key).get_shape())
|
| 92 |
+
if patch_shape != tuple(parameters[target_name].shape):
|
| 93 |
+
raise AssetError(
|
| 94 |
+
f"Turbo-LoRA shape mismatch for {target_name}: "
|
| 95 |
+
f"patch={patch_shape}, target={tuple(parameters[target_name].shape)}."
|
| 96 |
+
)
|
| 97 |
+
direct_targets[key] = target_name
|
| 98 |
+
|
| 99 |
+
recognized = set(down_keys) | up_keys | set(direct_keys)
|
| 100 |
+
if recognized != keys:
|
| 101 |
+
raise AssetError(f"Unsupported Turbo-LoRA keys: {sorted(keys - recognized)}")
|
| 102 |
+
|
| 103 |
+
with torch.no_grad():
|
| 104 |
+
for down_key in down_keys:
|
| 105 |
+
up_key = down_key.removesuffix(".lora_down.weight") + ".lora_up.weight"
|
| 106 |
+
target_name = _target_name(
|
| 107 |
+
down_key,
|
| 108 |
+
".lora_down.weight",
|
| 109 |
+
".weight",
|
| 110 |
+
)
|
| 111 |
+
parameter = parameters[target_name]
|
| 112 |
+
down = adapter.get_tensor(down_key).float()
|
| 113 |
+
up = adapter.get_tensor(up_key).float()
|
| 114 |
+
delta = torch.mm(up, down)
|
| 115 |
+
parameter.add_(delta.to(device=parameter.device, dtype=parameter.dtype))
|
| 116 |
+
for key, target_name in direct_targets.items():
|
| 117 |
+
parameter = parameters[target_name]
|
| 118 |
+
delta = adapter.get_tensor(key)
|
| 119 |
+
parameter.add_(delta.to(device=parameter.device, dtype=parameter.dtype))
|
| 120 |
+
|
| 121 |
+
return TurboLoraReport(
|
| 122 |
+
lora_modules=len(down_keys),
|
| 123 |
+
direct_parameters=len(direct_keys),
|
| 124 |
+
max_rank=max_rank,
|
| 125 |
+
)
|
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""GVHMR motion recovery."""
|
| 2 |
+
|
| 3 |
+
from fdanyone.motion.result import MotionResult
|
| 4 |
+
|
| 5 |
+
__all__ = ["MotionResult"]
|
|
@@ -0,0 +1,143 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Rebuild the posed SMPL-X body of a finished GVHMR motion result.
|
| 2 |
+
|
| 3 |
+
The heavy dependencies (``numpy``, ``smplx``, ``torch``) stay function-local so
|
| 4 |
+
importing this module costs nothing until a body is actually evaluated.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import logging
|
| 10 |
+
from dataclasses import dataclass
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
from typing import TYPE_CHECKING
|
| 13 |
+
|
| 14 |
+
from fdanyone.errors import FourDAnyoneError
|
| 15 |
+
from fdanyone.runs import discover_run
|
| 16 |
+
|
| 17 |
+
if TYPE_CHECKING:
|
| 18 |
+
from fractions import Fraction
|
| 19 |
+
|
| 20 |
+
import numpy as np
|
| 21 |
+
|
| 22 |
+
LOGGER = logging.getLogger("fdanyone.motion.body")
|
| 23 |
+
|
| 24 |
+
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
| 25 |
+
|
| 26 |
+
# SMPL-X body model roots tried in order, relative to the repository root.
|
| 27 |
+
SMPLX_MODEL_ROOTS = (
|
| 28 |
+
Path("models"),
|
| 29 |
+
Path("third_party/GVHMR/inputs/checkpoints"),
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
# SMPL-X keeps SMPL's body-joint order, so the first 55 entries of the joint
|
| 33 |
+
# tensor are the body, jaw, and hand skeleton addressed by ``parents``.
|
| 34 |
+
NUM_SMPLX_SKELETON_JOINTS = 55
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
@dataclass(frozen=True)
|
| 38 |
+
class BodyMotion:
|
| 39 |
+
"""SMPL-X geometry in the canonical 4DAnyone human world."""
|
| 40 |
+
|
| 41 |
+
vertices: np.ndarray
|
| 42 |
+
joints: np.ndarray
|
| 43 |
+
faces: np.ndarray
|
| 44 |
+
parents: tuple[int, ...]
|
| 45 |
+
keypoints_2d: np.ndarray | None
|
| 46 |
+
image_size: tuple[int, int]
|
| 47 |
+
fps: Fraction
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def _smplx_model_path() -> Path:
|
| 51 |
+
from fdanyone.assets import SMPLX_MODEL
|
| 52 |
+
|
| 53 |
+
for relative in SMPLX_MODEL_ROOTS:
|
| 54 |
+
model = _REPO_ROOT / relative / SMPLX_MODEL
|
| 55 |
+
if model.is_file():
|
| 56 |
+
return model.parents[1]
|
| 57 |
+
raise FourDAnyoneError(
|
| 58 |
+
"The licensed SMPL-X body model is missing. Run `python scripts/download_smplx.py`; "
|
| 59 |
+
f"expected {SMPLX_MODEL} under one of: {[str(_REPO_ROOT / path) for path in SMPLX_MODEL_ROOTS]}."
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def _canonical_rotation(joints: np.ndarray) -> np.ndarray:
|
| 64 |
+
"""Rotate the first frame onto the canonical yaw-zero human world.
|
| 65 |
+
|
| 66 |
+
Mirrors GVHMR's ``compute_T_ayfz2ay``: canonical ``+x`` is the subject's
|
| 67 |
+
anatomical left at frame zero, ``+y`` is up, and ``+z`` completes the
|
| 68 |
+
right-handed frame.
|
| 69 |
+
"""
|
| 70 |
+
|
| 71 |
+
import numpy as np
|
| 72 |
+
|
| 73 |
+
first = joints[0]
|
| 74 |
+
left = (first[1, [0, 2]] - first[2, [0, 2]]) + (first[16, [0, 2]] - first[17, [0, 2]])
|
| 75 |
+
norm = float(np.linalg.norm(left))
|
| 76 |
+
if norm <= 1e-4:
|
| 77 |
+
LOGGER.warning("Cannot determine the facing direction; leaving the motion world unrotated.")
|
| 78 |
+
return np.eye(3, dtype=np.float64)
|
| 79 |
+
x_dir = np.array([left[0] / norm, 0.0, left[1] / norm], dtype=np.float64)
|
| 80 |
+
y_dir = np.array([0.0, 1.0, 0.0], dtype=np.float64)
|
| 81 |
+
z_dir = np.cross(x_dir, y_dir)
|
| 82 |
+
return np.stack([x_dir, y_dir, z_dir], axis=-1)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def load_body_motion(motion_dir: Path, device: str = "cpu") -> BodyMotion:
|
| 86 |
+
"""Rebuild SMPL-X vertices and joints from a saved GVHMR motion result."""
|
| 87 |
+
|
| 88 |
+
import numpy as np
|
| 89 |
+
import smplx
|
| 90 |
+
import torch
|
| 91 |
+
|
| 92 |
+
from fdanyone.motion.result import MotionResult
|
| 93 |
+
|
| 94 |
+
motion = MotionResult.load(motion_dir)
|
| 95 |
+
model_path = _smplx_model_path()
|
| 96 |
+
parameters = motion.smpl_params_global
|
| 97 |
+
num_frames = motion.num_frames
|
| 98 |
+
# GVHMR's "supermotion" body model is plain neutral SMPL-X with ten shape
|
| 99 |
+
# coefficients, twelve hand PCA components, and a non-flat hand mean.
|
| 100 |
+
body_model = smplx.create(
|
| 101 |
+
model_path=str(model_path),
|
| 102 |
+
model_type="smplx",
|
| 103 |
+
gender="neutral",
|
| 104 |
+
num_betas=10,
|
| 105 |
+
num_pca_comps=12,
|
| 106 |
+
flat_hand_mean=False,
|
| 107 |
+
use_pca=True,
|
| 108 |
+
batch_size=num_frames,
|
| 109 |
+
).to(device)
|
| 110 |
+
with torch.inference_mode():
|
| 111 |
+
output = body_model(
|
| 112 |
+
betas=parameters["betas"].to(device),
|
| 113 |
+
global_orient=parameters["global_orient"].to(device),
|
| 114 |
+
body_pose=parameters["body_pose"].to(device),
|
| 115 |
+
transl=parameters["transl"].to(device),
|
| 116 |
+
)
|
| 117 |
+
vertices = output.vertices.detach().cpu().numpy().astype(np.float64)
|
| 118 |
+
joints = output.joints.detach().cpu().numpy().astype(np.float64)[:, :NUM_SMPLX_SKELETON_JOINTS]
|
| 119 |
+
|
| 120 |
+
# Canonicalize exactly like the conditioning stage: drop the first-frame
|
| 121 |
+
# root to the ground plane origin, then align the initial facing yaw.
|
| 122 |
+
offset = joints[0, 0].copy()
|
| 123 |
+
offset[1] = float(vertices[..., 1].min())
|
| 124 |
+
rotation = _canonical_rotation(joints - offset)
|
| 125 |
+
vertices = (vertices - offset) @ rotation
|
| 126 |
+
joints = (joints - offset) @ rotation
|
| 127 |
+
|
| 128 |
+
keypoints = motion.observed_keypoints_2d.detach().cpu().numpy().astype(np.float32)
|
| 129 |
+
return BodyMotion(
|
| 130 |
+
vertices=vertices.astype(np.float32),
|
| 131 |
+
joints=joints.astype(np.float32),
|
| 132 |
+
faces=np.asarray(body_model.faces, dtype=np.uint32),
|
| 133 |
+
parents=tuple(int(value) for value in body_model.parents.detach().cpu().numpy()),
|
| 134 |
+
keypoints_2d=keypoints,
|
| 135 |
+
image_size=(motion.image_width, motion.image_height),
|
| 136 |
+
fps=motion.fps,
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def posed_smplx(data_dir: str | Path, clip: str, device: str = "cpu") -> BodyMotion:
|
| 141 |
+
"""Rebuild the posed body of one finished run, found by clip name."""
|
| 142 |
+
|
| 143 |
+
return load_body_motion(discover_run(Path(data_dir), clip).motion_dir, device=device)
|
|
@@ -0,0 +1,332 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Classic GVHMR inference used by 4DAnyone.
|
| 2 |
+
|
| 3 |
+
The official demo imports training, evaluation, visualization, and
|
| 4 |
+
moving-camera modules eagerly. This file keeps the released static-camera path
|
| 5 |
+
in one place without exposing Hydra or backend abstractions to 4DAnyone users.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import contextlib
|
| 11 |
+
import importlib
|
| 12 |
+
import importlib.util
|
| 13 |
+
import itertools
|
| 14 |
+
import json
|
| 15 |
+
import os
|
| 16 |
+
import subprocess
|
| 17 |
+
import sys
|
| 18 |
+
import types
|
| 19 |
+
import warnings
|
| 20 |
+
from collections.abc import Callable, Iterator
|
| 21 |
+
from contextlib import contextmanager
|
| 22 |
+
from pathlib import Path
|
| 23 |
+
from typing import TypeAlias
|
| 24 |
+
|
| 25 |
+
from fdanyone.errors import AssetError, VideoContractError
|
| 26 |
+
from fdanyone.motion.result import SMPL_PARAMETER_NAMES, MotionResult
|
| 27 |
+
from fdanyone.vendor.pytorch3d_compat import install_if_needed as install_pytorch3d_compat
|
| 28 |
+
from fdanyone.video import CanonicalClip
|
| 29 |
+
|
| 30 |
+
MotionStageHook: TypeAlias = Callable[[str, dict[str, object]], None]
|
| 31 |
+
"""Receives ``(stage_name, payload)`` as each motion stage completes."""
|
| 32 |
+
|
| 33 |
+
GVHMR_ASSETS = (
|
| 34 |
+
"inputs/checkpoints/gvhmr/gvhmr_siga24_release.ckpt",
|
| 35 |
+
"inputs/checkpoints/hmr2/epoch=10-step=25000.ckpt",
|
| 36 |
+
"inputs/checkpoints/vitpose/vitpose-h-multi-coco.pth",
|
| 37 |
+
"inputs/checkpoints/yolo/yolov8x.pt",
|
| 38 |
+
"inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz",
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def validate_gvhmr(root: str | Path) -> tuple[Path, str]:
|
| 43 |
+
"""Locate the GVHMR checkout and files consumed by inference."""
|
| 44 |
+
|
| 45 |
+
path = Path(root).expanduser().resolve()
|
| 46 |
+
required = ("hmr4d/__init__.py", "tools/demo/demo.py", *GVHMR_ASSETS)
|
| 47 |
+
missing = [relative for relative in required if not (path / relative).is_file()]
|
| 48 |
+
if missing:
|
| 49 |
+
formatted = "\n - ".join(missing)
|
| 50 |
+
raise AssetError(
|
| 51 |
+
f"GVHMR is incomplete under {path}. Run `git submodule update --init third_party/GVHMR`, "
|
| 52 |
+
f"`python scripts/download_model.py`, and `python scripts/download_smplx.py`; missing:\n"
|
| 53 |
+
f" - {formatted}"
|
| 54 |
+
)
|
| 55 |
+
try:
|
| 56 |
+
revision = subprocess.check_output(
|
| 57 |
+
["git", "-C", str(path), "rev-parse", "HEAD"],
|
| 58 |
+
text=True,
|
| 59 |
+
stderr=subprocess.DEVNULL,
|
| 60 |
+
).strip()
|
| 61 |
+
except (OSError, subprocess.CalledProcessError) as exc:
|
| 62 |
+
raise AssetError(f"GVHMR must be a git checkout: {path}") from exc
|
| 63 |
+
if len(revision) != 40:
|
| 64 |
+
raise AssetError(f"Cannot identify the GVHMR revision at {path}.")
|
| 65 |
+
return path, revision
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def hydra_override(name: str, value: str | Path) -> str:
|
| 69 |
+
"""Quote a path for the internal GVHMR Hydra config."""
|
| 70 |
+
|
| 71 |
+
if not name.isidentifier():
|
| 72 |
+
raise ValueError(f"Invalid Hydra field name: {name!r}.")
|
| 73 |
+
return f"{name}={json.dumps(str(value), ensure_ascii=False)}"
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
@contextmanager
|
| 77 |
+
def gvhmr_imports(root: Path) -> Iterator[None]:
|
| 78 |
+
"""Temporarily import GVHMR as if its checkout were the working tree."""
|
| 79 |
+
|
| 80 |
+
old_cwd = Path.cwd()
|
| 81 |
+
root_text = str(root)
|
| 82 |
+
already_present = root_text in sys.path
|
| 83 |
+
if not already_present:
|
| 84 |
+
sys.path.insert(0, root_text)
|
| 85 |
+
os.chdir(root)
|
| 86 |
+
try:
|
| 87 |
+
yield
|
| 88 |
+
finally:
|
| 89 |
+
os.chdir(old_cwd)
|
| 90 |
+
if not already_present:
|
| 91 |
+
with contextlib.suppress(ValueError):
|
| 92 |
+
sys.path.remove(root_text)
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
@contextmanager
|
| 96 |
+
def _legacy_checkpoint_loading():
|
| 97 |
+
"""Restore pre-2.6 ``torch.load`` behavior for trusted GVHMR assets."""
|
| 98 |
+
|
| 99 |
+
name = "TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD"
|
| 100 |
+
previous = os.environ.get(name)
|
| 101 |
+
os.environ[name] = "1"
|
| 102 |
+
try:
|
| 103 |
+
with warnings.catch_warnings():
|
| 104 |
+
warnings.filterwarnings(
|
| 105 |
+
"ignore",
|
| 106 |
+
message=r"Environment variable TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD detected.*",
|
| 107 |
+
category=UserWarning,
|
| 108 |
+
)
|
| 109 |
+
yield
|
| 110 |
+
finally:
|
| 111 |
+
if previous is None:
|
| 112 |
+
os.environ.pop(name, None)
|
| 113 |
+
else:
|
| 114 |
+
os.environ[name] = previous
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def _install_optional_import_stubs() -> None:
|
| 118 |
+
"""Avoid visualization and moving-camera dependencies we never call."""
|
| 119 |
+
|
| 120 |
+
def unavailable(*_args, **_kwargs):
|
| 121 |
+
raise RuntimeError("Wis3D visualization is not part of 4DAnyone inference.")
|
| 122 |
+
|
| 123 |
+
if importlib.util.find_spec("wis3d") is None:
|
| 124 |
+
module = types.ModuleType("hmr4d.utils.wis3d_utils")
|
| 125 |
+
module.make_wis3d = unavailable
|
| 126 |
+
module.add_motion_as_lines = unavailable
|
| 127 |
+
sys.modules[module.__name__] = module
|
| 128 |
+
|
| 129 |
+
def moving_camera_unavailable(*_args, **_kwargs):
|
| 130 |
+
raise RuntimeError("SimpleVO is unavailable in the static-camera 4DAnyone runtime.")
|
| 131 |
+
|
| 132 |
+
module = types.ModuleType("hmr4d.utils.preproc.relpose.simple_vo")
|
| 133 |
+
module.SimpleVO = moving_camera_unavailable
|
| 134 |
+
sys.modules[module.__name__] = module
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def _register_inference_store() -> None:
|
| 138 |
+
"""Register only the Hydra groups referenced by GVHMR's demo config."""
|
| 139 |
+
|
| 140 |
+
_install_optional_import_stubs()
|
| 141 |
+
for module in (
|
| 142 |
+
"hmr4d.model.gvhmr.gvhmr_pl_demo",
|
| 143 |
+
"hmr4d.model.gvhmr.utils.endecoder",
|
| 144 |
+
"hmr4d.network.gvhmr.relative_transformer",
|
| 145 |
+
):
|
| 146 |
+
importlib.import_module(module)
|
| 147 |
+
|
| 148 |
+
# GVHMR installs a colored handler on the root logger. Remove only that
|
| 149 |
+
# duplicate because the public pipeline already owns a handler.
|
| 150 |
+
logger_module = sys.modules.get("hmr4d.utils.pylogger")
|
| 151 |
+
logger = getattr(logger_module, "Log", None)
|
| 152 |
+
handler = getattr(logger_module, "ch", None)
|
| 153 |
+
if (
|
| 154 |
+
logger is not None
|
| 155 |
+
and handler in logger.handlers
|
| 156 |
+
and any(candidate is not handler for candidate in logger.handlers)
|
| 157 |
+
):
|
| 158 |
+
logger.removeHandler(handler)
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
def _run_preprocess(cfg, *, on_stage: MotionStageHook | None = None) -> None:
|
| 162 |
+
"""Run tracker, ViTPose, and image-feature extraction."""
|
| 163 |
+
|
| 164 |
+
import torch
|
| 165 |
+
from hmr4d.utils.geo.hmr_cam import get_bbx_xys_from_xyxy
|
| 166 |
+
from hmr4d.utils.preproc.tracker import Tracker
|
| 167 |
+
from hmr4d.utils.preproc.vitfeat_extractor import Extractor
|
| 168 |
+
from hmr4d.utils.preproc.vitpose import VitPoseExtractor
|
| 169 |
+
from hmr4d.utils.pylogger import Log
|
| 170 |
+
|
| 171 |
+
if not bool(cfg.static_cam):
|
| 172 |
+
raise ValueError("4DAnyone requires GVHMR static_cam=true.")
|
| 173 |
+
|
| 174 |
+
Log.info("[Preprocess] Start!")
|
| 175 |
+
started = Log.time()
|
| 176 |
+
video_path = cfg.video_path
|
| 177 |
+
paths = cfg.paths
|
| 178 |
+
|
| 179 |
+
if not Path(paths.bbx).exists():
|
| 180 |
+
tracker = Tracker()
|
| 181 |
+
bbx_xyxy = tracker.get_one_track(video_path).float()
|
| 182 |
+
bbx_xys = get_bbx_xys_from_xyxy(bbx_xyxy, base_enlarge=1.2).float()
|
| 183 |
+
torch.save({"bbx_xyxy": bbx_xyxy, "bbx_xys": bbx_xys}, paths.bbx)
|
| 184 |
+
del tracker
|
| 185 |
+
else:
|
| 186 |
+
bbx_xys = torch.load(paths.bbx, weights_only=True)["bbx_xys"]
|
| 187 |
+
Log.info("[Preprocess] bbx (xyxy, xys) from %s", paths.bbx)
|
| 188 |
+
if on_stage is not None:
|
| 189 |
+
# Only ``bbx_xys`` reaches the model, so read the boxes back from the
|
| 190 |
+
# file both branches guarantee instead of widening the cached branch.
|
| 191 |
+
tracked_boxes = torch.load(paths.bbx, weights_only=True)["bbx_xyxy"]
|
| 192 |
+
on_stage("bboxes", {"bbx_xyxy": tracked_boxes.detach().cpu()})
|
| 193 |
+
|
| 194 |
+
if not Path(paths.vitpose).exists():
|
| 195 |
+
extractor = VitPoseExtractor()
|
| 196 |
+
torch.save(extractor.extract(video_path, bbx_xys), paths.vitpose)
|
| 197 |
+
del extractor
|
| 198 |
+
else:
|
| 199 |
+
Log.info("[Preprocess] vitpose from %s", paths.vitpose)
|
| 200 |
+
if on_stage is not None:
|
| 201 |
+
keypoints_2d = torch.load(paths.vitpose, weights_only=True)
|
| 202 |
+
on_stage("keypoints_2d", {"kp2d": keypoints_2d.detach().cpu()})
|
| 203 |
+
|
| 204 |
+
if not Path(paths.vit_features).exists():
|
| 205 |
+
extractor = Extractor()
|
| 206 |
+
torch.save(extractor.extract_video_features(video_path, bbx_xys), paths.vit_features)
|
| 207 |
+
del extractor
|
| 208 |
+
else:
|
| 209 |
+
Log.info("[Preprocess] vit_features from %s", paths.vit_features)
|
| 210 |
+
if on_stage is not None:
|
| 211 |
+
on_stage("features", {})
|
| 212 |
+
|
| 213 |
+
Log.info("[Preprocess] End. Time elapsed: %.2fs", Log.time() - started)
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def _load_data(cfg):
|
| 217 |
+
"""Build the static-camera tensors consumed by GVHMR."""
|
| 218 |
+
|
| 219 |
+
import torch
|
| 220 |
+
from hmr4d.utils.geo.hmr_cam import estimate_K
|
| 221 |
+
from hmr4d.utils.geo_transform import compute_cam_angvel
|
| 222 |
+
from hmr4d.utils.video_io_utils import get_video_lwh
|
| 223 |
+
|
| 224 |
+
if not bool(cfg.static_cam):
|
| 225 |
+
raise ValueError("4DAnyone requires GVHMR static_cam=true.")
|
| 226 |
+
paths = cfg.paths
|
| 227 |
+
length, width, height = get_video_lwh(cfg.video_path)
|
| 228 |
+
rotation_world_to_camera = torch.eye(3).repeat(length, 1, 1)
|
| 229 |
+
intrinsics = estimate_K(width, height).repeat(length, 1, 1)
|
| 230 |
+
return {
|
| 231 |
+
"length": torch.tensor(length),
|
| 232 |
+
"bbx_xys": torch.load(paths.bbx, weights_only=True)["bbx_xys"],
|
| 233 |
+
"kp2d": torch.load(paths.vitpose, weights_only=True),
|
| 234 |
+
"K_fullimg": intrinsics,
|
| 235 |
+
"cam_angvel": compute_cam_angvel(rotation_world_to_camera),
|
| 236 |
+
"f_imgseq": torch.load(paths.vit_features, weights_only=True),
|
| 237 |
+
}
|
| 238 |
+
|
| 239 |
+
|
| 240 |
+
def _verify_gvhmr_decode(clip: CanonicalClip, working_video: Path, reader_factory) -> None:
|
| 241 |
+
"""Ensure GVHMR's own video reader sees the canonical RGB frames."""
|
| 242 |
+
|
| 243 |
+
import numpy as np
|
| 244 |
+
|
| 245 |
+
reader = reader_factory(str(working_video))
|
| 246 |
+
sentinel = object()
|
| 247 |
+
try:
|
| 248 |
+
for index, (actual, expected) in enumerate(itertools.zip_longest(reader, clip.rgb_frames, fillvalue=sentinel)):
|
| 249 |
+
if actual is sentinel or expected is sentinel or not np.array_equal(actual, expected):
|
| 250 |
+
raise VideoContractError(f"GVHMR decoded a different canonical frame at index {index}.")
|
| 251 |
+
finally:
|
| 252 |
+
close = getattr(reader, "close", None)
|
| 253 |
+
if close is not None:
|
| 254 |
+
close()
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
def run_gvhmr(
|
| 258 |
+
*,
|
| 259 |
+
clip: CanonicalClip,
|
| 260 |
+
working_video: str | Path,
|
| 261 |
+
output_dir: str | Path,
|
| 262 |
+
gvhmr_root: str | Path,
|
| 263 |
+
device: str,
|
| 264 |
+
on_stage: MotionStageHook | None = None,
|
| 265 |
+
) -> MotionResult:
|
| 266 |
+
"""Recover static-camera human motion from the canonical source clip."""
|
| 267 |
+
|
| 268 |
+
root, revision = validate_gvhmr(gvhmr_root)
|
| 269 |
+
working_video = Path(working_video).expanduser().resolve()
|
| 270 |
+
output_root = Path(output_dir).expanduser().resolve()
|
| 271 |
+
output_root.mkdir(parents=True, exist_ok=True)
|
| 272 |
+
|
| 273 |
+
with gvhmr_imports(root), _legacy_checkpoint_loading():
|
| 274 |
+
install_pytorch3d_compat()
|
| 275 |
+
import hydra
|
| 276 |
+
import torch
|
| 277 |
+
from hmr4d.model.gvhmr.gvhmr_pl_demo import DemoPL
|
| 278 |
+
from hmr4d.utils.net_utils import detach_to_cpu
|
| 279 |
+
from hmr4d.utils.video_io_utils import get_video_reader
|
| 280 |
+
from hydra import compose, initialize_config_module
|
| 281 |
+
from omegaconf import open_dict
|
| 282 |
+
|
| 283 |
+
_register_inference_store()
|
| 284 |
+
with initialize_config_module(version_base="1.3", config_module="hmr4d.configs"):
|
| 285 |
+
cfg = compose(
|
| 286 |
+
config_name="demo",
|
| 287 |
+
overrides=[
|
| 288 |
+
hydra_override("video_name", working_video.stem),
|
| 289 |
+
"static_cam=true",
|
| 290 |
+
"verbose=false",
|
| 291 |
+
"use_dpvo=false",
|
| 292 |
+
hydra_override("output_root", output_root),
|
| 293 |
+
],
|
| 294 |
+
)
|
| 295 |
+
Path(cfg.output_dir).mkdir(parents=True, exist_ok=True)
|
| 296 |
+
Path(cfg.preprocess_dir).mkdir(parents=True, exist_ok=True)
|
| 297 |
+
with open_dict(cfg):
|
| 298 |
+
cfg.video_path = str(working_video)
|
| 299 |
+
|
| 300 |
+
_verify_gvhmr_decode(clip, working_video, get_video_reader)
|
| 301 |
+
_run_preprocess(cfg, on_stage=on_stage)
|
| 302 |
+
data = _load_data(cfg)
|
| 303 |
+
if int(data["length"]) != len(clip.frames):
|
| 304 |
+
raise RuntimeError(f"GVHMR decoded {int(data['length'])} frames, expected {len(clip.frames)}.")
|
| 305 |
+
observed_keypoints_2d = data["kp2d"].detach().cpu()
|
| 306 |
+
model: DemoPL = hydra.utils.instantiate(cfg.model, _recursive_=False)
|
| 307 |
+
model.load_pretrained_model(cfg.ckpt_path)
|
| 308 |
+
model = model.eval().to(device)
|
| 309 |
+
with torch.inference_mode():
|
| 310 |
+
prediction = detach_to_cpu(model.predict(data, static_cam=True))
|
| 311 |
+
del model, data
|
| 312 |
+
torch.cuda.empty_cache()
|
| 313 |
+
|
| 314 |
+
result = MotionResult(
|
| 315 |
+
gvhmr_revision=revision,
|
| 316 |
+
fps=clip.fps,
|
| 317 |
+
frame_timestamps_sec=tuple(float(frame.canonical_timestamp) for frame in clip.frames),
|
| 318 |
+
source_frame_indices=tuple(frame.source_index for frame in clip.frames),
|
| 319 |
+
source_pts=tuple(frame.source_pts for frame in clip.frames),
|
| 320 |
+
source_size_bytes=clip.source_size_bytes,
|
| 321 |
+
source_mtime_ns=clip.source_mtime_ns,
|
| 322 |
+
image_height=clip.height,
|
| 323 |
+
image_width=clip.width,
|
| 324 |
+
smpl_params_global={name: prediction["smpl_params_global"][name] for name in SMPL_PARAMETER_NAMES},
|
| 325 |
+
smpl_params_incam={name: prediction["smpl_params_incam"][name] for name in SMPL_PARAMETER_NAMES},
|
| 326 |
+
K_fullimg=prediction["K_fullimg"],
|
| 327 |
+
observed_keypoints_2d=observed_keypoints_2d,
|
| 328 |
+
)
|
| 329 |
+
result.validate(expected_frames=len(clip.frames))
|
| 330 |
+
if on_stage is not None:
|
| 331 |
+
on_stage("smplx", {"result": result})
|
| 332 |
+
return result
|
|
@@ -0,0 +1,188 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Reusable GVHMR result stored as JSON plus safetensors."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
from dataclasses import dataclass
|
| 7 |
+
from fractions import Fraction
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
from typing import Any
|
| 10 |
+
|
| 11 |
+
from fdanyone.errors import FourDAnyoneError
|
| 12 |
+
|
| 13 |
+
SMPL_PARAMETER_NAMES = ("body_pose", "betas", "global_orient", "transl")
|
| 14 |
+
SMPL_PARAMETER_WIDTHS = {"body_pose": 63, "betas": 10, "global_orient": 3, "transl": 3}
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
@dataclass(frozen=True)
|
| 18 |
+
class MotionResult:
|
| 19 |
+
gvhmr_revision: str
|
| 20 |
+
fps: Fraction
|
| 21 |
+
frame_timestamps_sec: tuple[float, ...]
|
| 22 |
+
source_frame_indices: tuple[int, ...]
|
| 23 |
+
source_pts: tuple[int | None, ...]
|
| 24 |
+
source_size_bytes: int
|
| 25 |
+
source_mtime_ns: int
|
| 26 |
+
image_height: int
|
| 27 |
+
image_width: int
|
| 28 |
+
smpl_params_global: dict[str, Any]
|
| 29 |
+
smpl_params_incam: dict[str, Any]
|
| 30 |
+
K_fullimg: Any
|
| 31 |
+
observed_keypoints_2d: Any
|
| 32 |
+
motion_world: str = "gvhmr_gravity_aligned_y_up"
|
| 33 |
+
|
| 34 |
+
@property
|
| 35 |
+
def num_frames(self) -> int:
|
| 36 |
+
return len(self.frame_timestamps_sec)
|
| 37 |
+
|
| 38 |
+
def validate(self, expected_frames: int = 121) -> None:
|
| 39 |
+
try:
|
| 40 |
+
import torch
|
| 41 |
+
except ImportError as exc:
|
| 42 |
+
raise FourDAnyoneError("PyTorch is required to validate motion tensors.") from exc
|
| 43 |
+
if self.motion_world != "gvhmr_gravity_aligned_y_up":
|
| 44 |
+
raise FourDAnyoneError(f"Unknown motion world convention: {self.motion_world!r}.")
|
| 45 |
+
if (
|
| 46 |
+
not isinstance(self.gvhmr_revision, str)
|
| 47 |
+
or len(self.gvhmr_revision) != 40
|
| 48 |
+
or any(character not in "0123456789abcdef" for character in self.gvhmr_revision.lower())
|
| 49 |
+
):
|
| 50 |
+
raise FourDAnyoneError("MotionResult requires a 40-character GVHMR git revision.")
|
| 51 |
+
if self.fps <= 0:
|
| 52 |
+
raise FourDAnyoneError(f"MotionResult FPS must be positive, got {self.fps}.")
|
| 53 |
+
if self.num_frames != expected_frames:
|
| 54 |
+
raise FourDAnyoneError(f"MotionResult has {self.num_frames} frames, expected {expected_frames}.")
|
| 55 |
+
for name, values in (
|
| 56 |
+
("source_frame_indices", self.source_frame_indices),
|
| 57 |
+
("source_pts", self.source_pts),
|
| 58 |
+
):
|
| 59 |
+
if len(values) != self.num_frames:
|
| 60 |
+
raise FourDAnyoneError(f"MotionResult {name} has {len(values)} values, expected {self.num_frames}.")
|
| 61 |
+
expected_timestamps = tuple(float(Fraction(index, 1) / self.fps) for index in range(self.num_frames))
|
| 62 |
+
if self.frame_timestamps_sec != expected_timestamps:
|
| 63 |
+
raise FourDAnyoneError("MotionResult timestamps are not the exact zero-based CFR timeline.")
|
| 64 |
+
if any(index < 0 for index in self.source_frame_indices) or any(
|
| 65 |
+
right < left for left, right in zip(self.source_frame_indices, self.source_frame_indices[1:], strict=False)
|
| 66 |
+
):
|
| 67 |
+
raise FourDAnyoneError("MotionResult source-frame indices must be non-negative and monotonic.")
|
| 68 |
+
if self.source_size_bytes <= 0 or self.source_mtime_ns <= 0:
|
| 69 |
+
raise FourDAnyoneError("MotionResult has an invalid source-file identity.")
|
| 70 |
+
if self.image_height <= 0 or self.image_width <= 0:
|
| 71 |
+
raise FourDAnyoneError("MotionResult image dimensions must be positive.")
|
| 72 |
+
for group_name, parameters in (
|
| 73 |
+
("smpl_params_global", self.smpl_params_global),
|
| 74 |
+
("smpl_params_incam", self.smpl_params_incam),
|
| 75 |
+
):
|
| 76 |
+
if set(parameters) != set(SMPL_PARAMETER_NAMES):
|
| 77 |
+
raise FourDAnyoneError(
|
| 78 |
+
f"MotionResult {group_name} must contain {SMPL_PARAMETER_NAMES}, got {tuple(parameters)}."
|
| 79 |
+
)
|
| 80 |
+
for name, tensor in parameters.items():
|
| 81 |
+
expected_shape = (self.num_frames, SMPL_PARAMETER_WIDTHS[name])
|
| 82 |
+
if not isinstance(tensor, torch.Tensor) or tuple(tensor.shape) != expected_shape:
|
| 83 |
+
raise FourDAnyoneError(f"{group_name}.{name} must have shape {expected_shape}.")
|
| 84 |
+
if not bool(torch.isfinite(tensor).all()):
|
| 85 |
+
raise FourDAnyoneError(f"{group_name}.{name} contains non-finite values.")
|
| 86 |
+
if not isinstance(self.K_fullimg, torch.Tensor) or tuple(self.K_fullimg.shape) != (self.num_frames, 3, 3):
|
| 87 |
+
raise FourDAnyoneError(f"K_fullimg must have shape ({self.num_frames}, 3, 3).")
|
| 88 |
+
if not bool(torch.isfinite(self.K_fullimg).all()):
|
| 89 |
+
raise FourDAnyoneError("K_fullimg contains non-finite values.")
|
| 90 |
+
expected_keypoint_shape = (self.num_frames, 17, 3)
|
| 91 |
+
if (
|
| 92 |
+
not isinstance(self.observed_keypoints_2d, torch.Tensor)
|
| 93 |
+
or tuple(self.observed_keypoints_2d.shape) != expected_keypoint_shape
|
| 94 |
+
):
|
| 95 |
+
raise FourDAnyoneError(f"observed_keypoints_2d must have shape {expected_keypoint_shape}.")
|
| 96 |
+
if not bool(torch.isfinite(self.observed_keypoints_2d).all()):
|
| 97 |
+
raise FourDAnyoneError("observed_keypoints_2d contains non-finite values.")
|
| 98 |
+
|
| 99 |
+
def validate_against_clip(self, clip) -> None:
|
| 100 |
+
"""Reject a cached result produced from a different video timeline."""
|
| 101 |
+
|
| 102 |
+
self.validate(expected_frames=len(clip.frames))
|
| 103 |
+
expected_timestamps = tuple(float(frame.canonical_timestamp) for frame in clip.frames)
|
| 104 |
+
expected_indices = tuple(frame.source_index for frame in clip.frames)
|
| 105 |
+
expected_pts = tuple(frame.source_pts for frame in clip.frames)
|
| 106 |
+
if self.fps != clip.fps:
|
| 107 |
+
raise FourDAnyoneError(f"Motion FPS {self.fps} does not match canonical FPS {clip.fps}.")
|
| 108 |
+
if self.frame_timestamps_sec != expected_timestamps:
|
| 109 |
+
raise FourDAnyoneError("Motion timestamps do not match the canonical clip.")
|
| 110 |
+
if self.source_frame_indices != expected_indices or self.source_pts != expected_pts:
|
| 111 |
+
raise FourDAnyoneError("Motion source-frame identity does not match the canonical clip.")
|
| 112 |
+
if (self.source_size_bytes, self.source_mtime_ns) != (
|
| 113 |
+
clip.source_size_bytes,
|
| 114 |
+
clip.source_mtime_ns,
|
| 115 |
+
):
|
| 116 |
+
raise FourDAnyoneError("Cached GVHMR motion belongs to a different source file.")
|
| 117 |
+
|
| 118 |
+
def save(self, directory: str | Path) -> Path:
|
| 119 |
+
from safetensors.torch import save_file
|
| 120 |
+
|
| 121 |
+
self.validate()
|
| 122 |
+
root = Path(directory).expanduser().resolve()
|
| 123 |
+
root.mkdir(parents=True, exist_ok=True)
|
| 124 |
+
tensor_path = root / "motion.safetensors"
|
| 125 |
+
tensors = {
|
| 126 |
+
**{
|
| 127 |
+
f"smpl_params_global.{name}": self.smpl_params_global[name].detach().cpu().contiguous()
|
| 128 |
+
for name in SMPL_PARAMETER_NAMES
|
| 129 |
+
},
|
| 130 |
+
**{
|
| 131 |
+
f"smpl_params_incam.{name}": self.smpl_params_incam[name].detach().cpu().contiguous()
|
| 132 |
+
for name in SMPL_PARAMETER_NAMES
|
| 133 |
+
},
|
| 134 |
+
"K_fullimg": self.K_fullimg.detach().cpu().contiguous(),
|
| 135 |
+
"observed_keypoints_2d": self.observed_keypoints_2d.detach().cpu().contiguous(),
|
| 136 |
+
}
|
| 137 |
+
save_file(tensors, str(tensor_path))
|
| 138 |
+
metadata = {
|
| 139 |
+
"gvhmr_revision": self.gvhmr_revision,
|
| 140 |
+
"num_frames": self.num_frames,
|
| 141 |
+
"fps_num": self.fps.numerator,
|
| 142 |
+
"fps_den": self.fps.denominator,
|
| 143 |
+
"frame_timestamps_sec": self.frame_timestamps_sec,
|
| 144 |
+
"source_frame_indices": self.source_frame_indices,
|
| 145 |
+
"source_pts": self.source_pts,
|
| 146 |
+
"source_size_bytes": self.source_size_bytes,
|
| 147 |
+
"source_mtime_ns": self.source_mtime_ns,
|
| 148 |
+
"image_height": self.image_height,
|
| 149 |
+
"image_width": self.image_width,
|
| 150 |
+
"motion_world": self.motion_world,
|
| 151 |
+
"tensor_file": tensor_path.name,
|
| 152 |
+
}
|
| 153 |
+
(root / "motion.json").write_text(json.dumps(metadata, indent=2, sort_keys=True) + "\n")
|
| 154 |
+
return root
|
| 155 |
+
|
| 156 |
+
@classmethod
|
| 157 |
+
def load(cls, directory: str | Path) -> MotionResult:
|
| 158 |
+
from safetensors.torch import load_file
|
| 159 |
+
|
| 160 |
+
root = Path(directory).expanduser().resolve()
|
| 161 |
+
metadata = json.loads((root / "motion.json").read_text())
|
| 162 |
+
tensors = load_file(str(root / metadata["tensor_file"]), device="cpu")
|
| 163 |
+
# ``safetensors`` may keep the source file memory-mapped for as long as
|
| 164 |
+
# returned tensors own that storage. Motion results live until final
|
| 165 |
+
# export, which in turn can leave an in-use ``.efc_*`` tombstone on
|
| 166 |
+
# CPFS/NFS when scratch cleanup unlinks the backing file. The payload
|
| 167 |
+
# is tiny compared with the generated videos, so take ownership here
|
| 168 |
+
# and close the file mapping at this explicit process boundary.
|
| 169 |
+
owned = {name: tensor.clone() for name, tensor in tensors.items()}
|
| 170 |
+
del tensors
|
| 171 |
+
result = cls(
|
| 172 |
+
gvhmr_revision=str(metadata["gvhmr_revision"]),
|
| 173 |
+
fps=Fraction(int(metadata["fps_num"]), int(metadata["fps_den"])),
|
| 174 |
+
frame_timestamps_sec=tuple(float(value) for value in metadata["frame_timestamps_sec"]),
|
| 175 |
+
source_frame_indices=tuple(int(value) for value in metadata["source_frame_indices"]),
|
| 176 |
+
source_pts=tuple(None if value is None else int(value) for value in metadata["source_pts"]),
|
| 177 |
+
source_size_bytes=int(metadata["source_size_bytes"]),
|
| 178 |
+
source_mtime_ns=int(metadata["source_mtime_ns"]),
|
| 179 |
+
image_height=int(metadata["image_height"]),
|
| 180 |
+
image_width=int(metadata["image_width"]),
|
| 181 |
+
smpl_params_global={name: owned[f"smpl_params_global.{name}"] for name in SMPL_PARAMETER_NAMES},
|
| 182 |
+
smpl_params_incam={name: owned[f"smpl_params_incam.{name}"] for name in SMPL_PARAMETER_NAMES},
|
| 183 |
+
K_fullimg=owned["K_fullimg"],
|
| 184 |
+
observed_keypoints_2d=owned["observed_keypoints_2d"],
|
| 185 |
+
motion_world=str(metadata["motion_world"]),
|
| 186 |
+
)
|
| 187 |
+
result.validate()
|
| 188 |
+
return result
|
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Private subprocess entry point for motion inference."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import sys
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
|
| 9 |
+
from fdanyone.device import select_cuda_device
|
| 10 |
+
from fdanyone.motion.gvhmr import MotionStageHook, run_gvhmr
|
| 11 |
+
from fdanyone.video import load_canonical_working_clip
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def run_motion_worker(
|
| 15 |
+
*,
|
| 16 |
+
gvhmr_root: str | Path,
|
| 17 |
+
working_video: str | Path,
|
| 18 |
+
clip_metadata: str | Path,
|
| 19 |
+
output_dir: str | Path,
|
| 20 |
+
result_dir: str | Path,
|
| 21 |
+
device: str,
|
| 22 |
+
on_stage: MotionStageHook | None = None,
|
| 23 |
+
) -> None:
|
| 24 |
+
"""Recover motion for one request and publish it under ``result_dir``."""
|
| 25 |
+
|
| 26 |
+
device, _ = select_cuda_device(device)
|
| 27 |
+
clip = load_canonical_working_clip(working_video, clip_metadata)
|
| 28 |
+
common = {
|
| 29 |
+
"clip": clip,
|
| 30 |
+
"working_video": working_video,
|
| 31 |
+
"output_dir": output_dir,
|
| 32 |
+
"device": device,
|
| 33 |
+
}
|
| 34 |
+
result = run_gvhmr(gvhmr_root=gvhmr_root, on_stage=on_stage, **common)
|
| 35 |
+
result.save(result_dir)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def main(request_path: str) -> None:
|
| 39 |
+
request = json.loads(Path(request_path).read_text())
|
| 40 |
+
run_motion_worker(
|
| 41 |
+
gvhmr_root=request["gvhmr_root"],
|
| 42 |
+
working_video=request["working_video"],
|
| 43 |
+
clip_metadata=request["clip_metadata"],
|
| 44 |
+
output_dir=request["output_dir"],
|
| 45 |
+
result_dir=request["result_dir"],
|
| 46 |
+
device=request["device"],
|
| 47 |
+
)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
if __name__ == "__main__":
|
| 51 |
+
if len(sys.argv) != 2:
|
| 52 |
+
raise SystemExit("Usage: python -m fdanyone.motion.worker REQUEST.json")
|
| 53 |
+
main(sys.argv[1])
|
|
@@ -0,0 +1,218 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Publish generated videos and their camera metadata."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import platform
|
| 7 |
+
import shutil
|
| 8 |
+
import sys
|
| 9 |
+
import time
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
from typing import TYPE_CHECKING
|
| 12 |
+
|
| 13 |
+
from fdanyone.config import INFERENCE, ModeSettings
|
| 14 |
+
from fdanyone.errors import FourDAnyoneError
|
| 15 |
+
from fdanyone.io import write_json
|
| 16 |
+
|
| 17 |
+
if TYPE_CHECKING:
|
| 18 |
+
from fdanyone.model.inference import GeneratedViews
|
| 19 |
+
from fdanyone.motion.result import MotionResult
|
| 20 |
+
from fdanyone.skeleton.pipeline import Conditioning
|
| 21 |
+
from fdanyone.video import CanonicalClip
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def _copy_file(source: Path, destination: Path) -> Path:
|
| 25 |
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
| 26 |
+
shutil.copy2(source, destination)
|
| 27 |
+
return destination
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _runtime_metadata(device: str) -> dict:
|
| 31 |
+
import torch
|
| 32 |
+
|
| 33 |
+
cuda = {
|
| 34 |
+
"available": torch.cuda.is_available(),
|
| 35 |
+
"torch_cuda": torch.version.cuda,
|
| 36 |
+
"cudnn": torch.backends.cudnn.version(),
|
| 37 |
+
}
|
| 38 |
+
if torch.cuda.is_available():
|
| 39 |
+
torch_device = torch.device(device)
|
| 40 |
+
properties = torch.cuda.get_device_properties(torch_device)
|
| 41 |
+
cuda.update(
|
| 42 |
+
{
|
| 43 |
+
"device": device,
|
| 44 |
+
"device_name": torch.cuda.get_device_name(torch_device),
|
| 45 |
+
"device_capability": list(torch.cuda.get_device_capability(torch_device)),
|
| 46 |
+
"device_total_memory_bytes": properties.total_memory,
|
| 47 |
+
}
|
| 48 |
+
)
|
| 49 |
+
return {
|
| 50 |
+
"python": sys.version.split()[0],
|
| 51 |
+
"platform": platform.platform(),
|
| 52 |
+
"torch": torch.__version__,
|
| 53 |
+
"cuda": cuda,
|
| 54 |
+
}
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _camera_rig_payload(payload: dict, cameras: list[dict]) -> dict:
|
| 58 |
+
"""Keep the final OpenCV camera rig needed by downstream tools."""
|
| 59 |
+
|
| 60 |
+
records = []
|
| 61 |
+
for camera in cameras:
|
| 62 |
+
camera_id = int(camera["camera_id"])
|
| 63 |
+
records.append(
|
| 64 |
+
{
|
| 65 |
+
"camera_id": camera_id,
|
| 66 |
+
"layer_index": int(camera["layer_index"]),
|
| 67 |
+
"pitch": int(camera["pitch_degrees"]),
|
| 68 |
+
"yaw": float(camera["yaw_degrees"]),
|
| 69 |
+
"K": camera["K"],
|
| 70 |
+
"camera_to_world": camera["camera_to_world"],
|
| 71 |
+
"image_width": int(camera["image_width"]),
|
| 72 |
+
"image_height": int(camera["image_height"]),
|
| 73 |
+
"video": f"videos/dense/{camera_id:02d}.mp4",
|
| 74 |
+
"skeleton_video": f"skeletons/{camera_id:02d}.mp4",
|
| 75 |
+
}
|
| 76 |
+
)
|
| 77 |
+
return {
|
| 78 |
+
"camera_model": "OPENCV",
|
| 79 |
+
"world_frame": payload["world_frame"],
|
| 80 |
+
"camera_frame": payload["camera_frame"],
|
| 81 |
+
"front_camera_ids": payload["front_camera_ids"],
|
| 82 |
+
"framing": payload["framing"],
|
| 83 |
+
"cameras": records,
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def _target_cameras(payload: object, expected_count: int) -> list[dict]:
|
| 88 |
+
"""Read the camera records produced by the conditioning stage."""
|
| 89 |
+
|
| 90 |
+
if not isinstance(payload, dict) or payload.get("camera_model") != "OPENCV":
|
| 91 |
+
raise FourDAnyoneError("Conditioning did not produce an OpenCV camera rig.")
|
| 92 |
+
cameras = payload.get("cameras")
|
| 93 |
+
if not isinstance(cameras, list) or len(cameras) != expected_count:
|
| 94 |
+
raise FourDAnyoneError(f"Conditioning must contain {expected_count} target cameras.")
|
| 95 |
+
if [camera.get("camera_id") for camera in cameras if isinstance(camera, dict)] != list(range(expected_count)):
|
| 96 |
+
raise FourDAnyoneError("Target cameras are not in canonical order.")
|
| 97 |
+
return cameras
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def export_result(
|
| 101 |
+
*,
|
| 102 |
+
clip: CanonicalClip,
|
| 103 |
+
conditioning: Conditioning,
|
| 104 |
+
generated: GeneratedViews,
|
| 105 |
+
destination: str | Path,
|
| 106 |
+
motion: MotionResult,
|
| 107 |
+
model_identity: dict,
|
| 108 |
+
pipeline_started: float,
|
| 109 |
+
settings: ModeSettings,
|
| 110 |
+
) -> dict:
|
| 111 |
+
"""Publish proposal, target, skeleton, camera, and metadata artifacts."""
|
| 112 |
+
|
| 113 |
+
root = Path(destination).expanduser().resolve()
|
| 114 |
+
attention_backend = "sdpa" if settings.exact_attention else "sageattention"
|
| 115 |
+
view_plan = generated.view_plan
|
| 116 |
+
if conditioning.view_plan != view_plan:
|
| 117 |
+
raise FourDAnyoneError("Conditioning and generation resolved different view plans.")
|
| 118 |
+
if len(generated.rcp_videos) != len(view_plan.rcp_camera_ids):
|
| 119 |
+
raise FourDAnyoneError(
|
| 120 |
+
f"Generation returned {len(generated.rcp_videos)} RCP videos, expected {len(view_plan.rcp_camera_ids)}."
|
| 121 |
+
)
|
| 122 |
+
if len(generated.target_videos) != view_plan.num_target_views:
|
| 123 |
+
raise FourDAnyoneError(
|
| 124 |
+
f"Generation returned {len(generated.target_videos)} target videos, expected {view_plan.num_target_views}."
|
| 125 |
+
)
|
| 126 |
+
if len(conditioning.target_skeletons) != view_plan.num_target_views:
|
| 127 |
+
raise FourDAnyoneError(
|
| 128 |
+
f"Conditioning returned {len(conditioning.target_skeletons)} target skeletons, "
|
| 129 |
+
f"expected {view_plan.num_target_views}."
|
| 130 |
+
)
|
| 131 |
+
|
| 132 |
+
sparse_root = root / "videos" / "sparse"
|
| 133 |
+
dense_root = root / "videos" / "dense"
|
| 134 |
+
skeletons_root = root / "skeletons"
|
| 135 |
+
dense_root.mkdir(parents=True, exist_ok=False)
|
| 136 |
+
skeletons_root.mkdir(exist_ok=False)
|
| 137 |
+
if generated.rcp_videos:
|
| 138 |
+
sparse_root.mkdir(exist_ok=False)
|
| 139 |
+
|
| 140 |
+
output_sparse = tuple(
|
| 141 |
+
_copy_file(source, sparse_root / f"{camera_id:02d}.mp4")
|
| 142 |
+
for camera_id, source in zip(view_plan.rcp_camera_ids, generated.rcp_videos, strict=True)
|
| 143 |
+
)
|
| 144 |
+
output_dense = tuple(
|
| 145 |
+
_copy_file(source, dense_root / f"{camera_id:02d}.mp4")
|
| 146 |
+
for camera_id, source in enumerate(generated.target_videos)
|
| 147 |
+
)
|
| 148 |
+
for camera_id, skeleton in enumerate(conditioning.target_skeletons):
|
| 149 |
+
_copy_file(skeleton.path, skeletons_root / f"{camera_id:02d}.mp4")
|
| 150 |
+
|
| 151 |
+
camera_payload = json.loads((conditioning.root / "cameras.json").read_text())
|
| 152 |
+
conditioning_metadata = json.loads((conditioning.root / "metadata.json").read_text())
|
| 153 |
+
camera_records = _target_cameras(camera_payload, view_plan.num_target_views)
|
| 154 |
+
total_elapsed = time.monotonic() - pipeline_started
|
| 155 |
+
|
| 156 |
+
metadata = {
|
| 157 |
+
"input": {
|
| 158 |
+
"filename": clip.source_path.name,
|
| 159 |
+
"fps": f"{clip.fps_num}/{clip.fps_den}",
|
| 160 |
+
"start_time_seconds": float(clip.start_time),
|
| 161 |
+
"num_frames": len(clip.frames),
|
| 162 |
+
"width": clip.width,
|
| 163 |
+
"height": clip.height,
|
| 164 |
+
},
|
| 165 |
+
"motion": {
|
| 166 |
+
"method": "GVHMR",
|
| 167 |
+
"revision": motion.gvhmr_revision,
|
| 168 |
+
},
|
| 169 |
+
"preprocessing": {
|
| 170 |
+
"source_crop_policy": conditioning_metadata["source_crop_policy"],
|
| 171 |
+
"foreground_model": conditioning_metadata["foreground_model"],
|
| 172 |
+
"framing": conditioning_metadata["framing"],
|
| 173 |
+
"skeleton_draw_scale": conditioning_metadata["skeleton_draw_scale"],
|
| 174 |
+
"target_render_deferred": bool(
|
| 175 |
+
conditioning_metadata.get("target_render_deferred", False)
|
| 176 |
+
),
|
| 177 |
+
"target_render_overlap": conditioning_metadata.get("target_render_overlap"),
|
| 178 |
+
},
|
| 179 |
+
"model": dict(model_identity),
|
| 180 |
+
"generation": {
|
| 181 |
+
"mode": settings.mode,
|
| 182 |
+
"seed": generated.seed,
|
| 183 |
+
"view_plan": {
|
| 184 |
+
**view_plan.to_dict(),
|
| 185 |
+
"num_layers": view_plan.num_layers,
|
| 186 |
+
"num_target_views": view_plan.num_target_views,
|
| 187 |
+
"groups_per_layer": view_plan.groups_per_layer,
|
| 188 |
+
"tcr_active": view_plan.tcr_active,
|
| 189 |
+
"routing_topology": "circular" if view_plan.closed_yaw else "open",
|
| 190 |
+
},
|
| 191 |
+
"attention_backend": attention_backend,
|
| 192 |
+
"inference_steps": settings.num_inference_steps,
|
| 193 |
+
"elapsed_seconds": generated.elapsed_seconds,
|
| 194 |
+
"total_elapsed_seconds": total_elapsed,
|
| 195 |
+
"peak_vram_allocated_bytes": generated.peak_vram_allocated_bytes,
|
| 196 |
+
"peak_vram_reserved_bytes": generated.peak_vram_reserved_bytes,
|
| 197 |
+
},
|
| 198 |
+
"output": {
|
| 199 |
+
"rcp_views": len(output_sparse),
|
| 200 |
+
"target_views": len(output_dense),
|
| 201 |
+
"frames_per_video": INFERENCE.num_frames,
|
| 202 |
+
"width": INFERENCE.width,
|
| 203 |
+
"height": INFERENCE.height,
|
| 204 |
+
"fps": f"{clip.fps_num}/{clip.fps_den}",
|
| 205 |
+
},
|
| 206 |
+
"runtime": _runtime_metadata(generated.device),
|
| 207 |
+
}
|
| 208 |
+
write_json(root / "cameras.json", _camera_rig_payload(camera_payload, camera_records))
|
| 209 |
+
write_json(root / "metadata.json", metadata)
|
| 210 |
+
return {
|
| 211 |
+
"attention_backend": attention_backend,
|
| 212 |
+
"num_rcp_videos": len(output_sparse),
|
| 213 |
+
"num_target_videos": len(output_dense),
|
| 214 |
+
"fps": f"{clip.fps_num}/{clip.fps_den}",
|
| 215 |
+
"peak_vram_allocated_bytes": generated.peak_vram_allocated_bytes,
|
| 216 |
+
"peak_vram_reserved_bytes": generated.peak_vram_reserved_bytes,
|
| 217 |
+
"total_pipeline_elapsed_seconds": total_elapsed,
|
| 218 |
+
}
|
|
@@ -0,0 +1,591 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Top-level inference orchestration."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import logging
|
| 7 |
+
import os
|
| 8 |
+
import subprocess
|
| 9 |
+
import sys
|
| 10 |
+
import tempfile
|
| 11 |
+
import time
|
| 12 |
+
from dataclasses import dataclass
|
| 13 |
+
from pathlib import Path
|
| 14 |
+
from typing import TYPE_CHECKING
|
| 15 |
+
|
| 16 |
+
from fdanyone.assets import (
|
| 17 |
+
CHECKPOINT,
|
| 18 |
+
HF_REPO_ID,
|
| 19 |
+
HF_REVISION,
|
| 20 |
+
BaseAssets,
|
| 21 |
+
resolve_base_assets,
|
| 22 |
+
resolve_checkpoint,
|
| 23 |
+
resolve_foreground_model,
|
| 24 |
+
resolve_regressor,
|
| 25 |
+
)
|
| 26 |
+
from fdanyone.config import INFERENCE, ModeSettings
|
| 27 |
+
from fdanyone.device import select_cuda_device
|
| 28 |
+
from fdanyone.download import ensure_example_video, ensure_models, ensure_smplx
|
| 29 |
+
from fdanyone.errors import ConfigurationError, FourDAnyoneError
|
| 30 |
+
from fdanyone.io import AtomicResultDirectory, remove_tree, write_json
|
| 31 |
+
from fdanyone.motion.gvhmr import MotionStageHook, validate_gvhmr
|
| 32 |
+
from fdanyone.motion.result import MotionResult
|
| 33 |
+
from fdanyone.video import (
|
| 34 |
+
CanonicalClip,
|
| 35 |
+
decode_canonical_clip,
|
| 36 |
+
validate_required_video_codecs,
|
| 37 |
+
verify_lossless_video,
|
| 38 |
+
write_gvhmr_video,
|
| 39 |
+
)
|
| 40 |
+
from fdanyone.views import ViewPlan, resolve_view_plan
|
| 41 |
+
|
| 42 |
+
LOGGER = logging.getLogger("fdanyone")
|
| 43 |
+
|
| 44 |
+
if TYPE_CHECKING:
|
| 45 |
+
from fdanyone.model.inference import DenoiseStepHook
|
| 46 |
+
from fdanyone.skeleton.pipeline import Conditioning
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def _data_paths(data_dir: str, video_path: str) -> tuple[Path, Path, Path]:
|
| 50 |
+
data_root = Path(data_dir).expanduser().resolve()
|
| 51 |
+
run_name = Path(video_path).stem
|
| 52 |
+
return (
|
| 53 |
+
data_root,
|
| 54 |
+
data_root / "gvhmr" / "results" / run_name,
|
| 55 |
+
data_root / "fdanyone" / run_name,
|
| 56 |
+
)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def _discard_scratch(path: Path) -> None:
|
| 60 |
+
"""Best-effort cleanup that can never invalidate a published result.
|
| 61 |
+
|
| 62 |
+
Some network filesystems keep an open, hidden tombstone after a file is
|
| 63 |
+
unlinked. Such a tombstone may remain ``EBUSY`` until this process exits,
|
| 64 |
+
so cleanup must not be part of the atomic publication transaction.
|
| 65 |
+
"""
|
| 66 |
+
|
| 67 |
+
try:
|
| 68 |
+
remove_tree(path)
|
| 69 |
+
except OSError as exc:
|
| 70 |
+
LOGGER.warning(
|
| 71 |
+
"Could not remove temporary files at %s (%s). "
|
| 72 |
+
"The result is unaffected; the hidden scratch directory can be removed after this process exits.",
|
| 73 |
+
path,
|
| 74 |
+
exc,
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def _worker_environment() -> dict[str, str]:
|
| 79 |
+
"""Give the short-lived GVHMR workers this checkout and stable CUDA flags."""
|
| 80 |
+
|
| 81 |
+
environment = os.environ.copy()
|
| 82 |
+
environment.update(
|
| 83 |
+
{
|
| 84 |
+
"TORCH_CUDNN_V8_API_DISABLED": "1",
|
| 85 |
+
"CUDNN_FRONTEND_DISABLE": "1",
|
| 86 |
+
"CUDNN_LOGINFO_DBG": "0",
|
| 87 |
+
"CUDNN_LOGDEST_DBG": "stderr",
|
| 88 |
+
"CUDA_DEVICE_MAX_CONNECTIONS": "1",
|
| 89 |
+
"NVIDIA_TF32_OVERRIDE": "0",
|
| 90 |
+
}
|
| 91 |
+
)
|
| 92 |
+
environment.pop("PYTHONHOME", None)
|
| 93 |
+
environment["PYTHONPATH"] = str(Path(__file__).resolve().parent.parent)
|
| 94 |
+
return environment
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def _cpu_worker_environment() -> dict[str, str]:
|
| 98 |
+
"""Run pure CPU overlap workers without exposing a CUDA device."""
|
| 99 |
+
|
| 100 |
+
environment: dict[str, str] = _worker_environment()
|
| 101 |
+
environment["CUDA_VISIBLE_DEVICES"] = ""
|
| 102 |
+
return environment
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def _run_motion(
|
| 106 |
+
*,
|
| 107 |
+
working_video: Path,
|
| 108 |
+
output_dir: Path,
|
| 109 |
+
gvhmr_root: Path,
|
| 110 |
+
device: str,
|
| 111 |
+
worker_python: str,
|
| 112 |
+
clip_metadata: Path,
|
| 113 |
+
inline_workers: bool = False,
|
| 114 |
+
on_motion_stage: MotionStageHook | None = None,
|
| 115 |
+
):
|
| 116 |
+
output_dir.mkdir(parents=True, exist_ok=True)
|
| 117 |
+
request_path = output_dir / ".motion-worker-request.json"
|
| 118 |
+
result_dir = output_dir / "result"
|
| 119 |
+
if inline_workers:
|
| 120 |
+
from fdanyone.motion.worker import run_motion_worker
|
| 121 |
+
|
| 122 |
+
run_motion_worker(
|
| 123 |
+
gvhmr_root=gvhmr_root,
|
| 124 |
+
working_video=working_video,
|
| 125 |
+
clip_metadata=clip_metadata,
|
| 126 |
+
output_dir=output_dir / "runtime",
|
| 127 |
+
result_dir=result_dir,
|
| 128 |
+
device=device,
|
| 129 |
+
on_stage=on_motion_stage,
|
| 130 |
+
)
|
| 131 |
+
return MotionResult.load(result_dir)
|
| 132 |
+
write_json(
|
| 133 |
+
request_path,
|
| 134 |
+
{
|
| 135 |
+
"gvhmr_root": str(gvhmr_root),
|
| 136 |
+
"working_video": str(working_video),
|
| 137 |
+
"clip_metadata": str(clip_metadata),
|
| 138 |
+
"output_dir": str(output_dir / "runtime"),
|
| 139 |
+
"result_dir": str(result_dir),
|
| 140 |
+
"device": device,
|
| 141 |
+
},
|
| 142 |
+
)
|
| 143 |
+
try:
|
| 144 |
+
subprocess.run(
|
| 145 |
+
[worker_python, "-m", "fdanyone.motion.worker", str(request_path)],
|
| 146 |
+
check=True,
|
| 147 |
+
env=_worker_environment(),
|
| 148 |
+
)
|
| 149 |
+
finally:
|
| 150 |
+
request_path.unlink(missing_ok=True)
|
| 151 |
+
return MotionResult.load(result_dir)
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
def _build_conditioning(
|
| 155 |
+
*,
|
| 156 |
+
regressor: Path,
|
| 157 |
+
foreground_model: Path,
|
| 158 |
+
gvhmr_root: Path,
|
| 159 |
+
output_dir: Path,
|
| 160 |
+
device: str,
|
| 161 |
+
worker_python: str,
|
| 162 |
+
working_video: Path,
|
| 163 |
+
clip_metadata: Path,
|
| 164 |
+
motion_result_dir: Path,
|
| 165 |
+
view_plan: ViewPlan,
|
| 166 |
+
settings: ModeSettings,
|
| 167 |
+
inline_workers: bool = False,
|
| 168 |
+
) -> Conditioning:
|
| 169 |
+
from fdanyone.skeleton.pipeline import Conditioning
|
| 170 |
+
|
| 171 |
+
if inline_workers:
|
| 172 |
+
from fdanyone.skeleton.worker import run_skeleton_worker
|
| 173 |
+
|
| 174 |
+
run_skeleton_worker(
|
| 175 |
+
working_video=working_video,
|
| 176 |
+
clip_metadata=clip_metadata,
|
| 177 |
+
motion_result_dir=motion_result_dir,
|
| 178 |
+
regressor_path=regressor,
|
| 179 |
+
foreground_model_path=foreground_model,
|
| 180 |
+
gvhmr_root=gvhmr_root,
|
| 181 |
+
output_dir=output_dir,
|
| 182 |
+
device=device,
|
| 183 |
+
view_plan=view_plan,
|
| 184 |
+
defer_target_skeletons=settings.overlap_target_skeletons,
|
| 185 |
+
)
|
| 186 |
+
else:
|
| 187 |
+
request_path = output_dir.parent / ".skeleton-worker-request.json"
|
| 188 |
+
write_json(
|
| 189 |
+
request_path,
|
| 190 |
+
{
|
| 191 |
+
"working_video": str(working_video),
|
| 192 |
+
"clip_metadata": str(clip_metadata),
|
| 193 |
+
"motion_result_dir": str(motion_result_dir),
|
| 194 |
+
"regressor_path": str(regressor),
|
| 195 |
+
"foreground_model_path": str(foreground_model),
|
| 196 |
+
"gvhmr_root": str(gvhmr_root),
|
| 197 |
+
"output_dir": str(output_dir),
|
| 198 |
+
"device": device,
|
| 199 |
+
"view_plan": view_plan.to_dict(),
|
| 200 |
+
"defer_target_skeletons": settings.overlap_target_skeletons,
|
| 201 |
+
},
|
| 202 |
+
)
|
| 203 |
+
try:
|
| 204 |
+
subprocess.run(
|
| 205 |
+
[
|
| 206 |
+
worker_python,
|
| 207 |
+
"-m",
|
| 208 |
+
"fdanyone.skeleton.worker",
|
| 209 |
+
str(request_path),
|
| 210 |
+
],
|
| 211 |
+
check=True,
|
| 212 |
+
env=_worker_environment(),
|
| 213 |
+
)
|
| 214 |
+
finally:
|
| 215 |
+
request_path.unlink(missing_ok=True)
|
| 216 |
+
conditioning: Conditioning = Conditioning.load(
|
| 217 |
+
output_dir,
|
| 218 |
+
allow_pending_targets=settings.overlap_target_skeletons,
|
| 219 |
+
skeleton_video_decoder=settings.skeleton_video_decoder,
|
| 220 |
+
)
|
| 221 |
+
render_request: Path = output_dir / "target-render-request.json"
|
| 222 |
+
if not render_request.is_file():
|
| 223 |
+
return conditioning
|
| 224 |
+
|
| 225 |
+
# The renderer stays a subprocess even under ``inline_workers``: it needs no
|
| 226 |
+
# GPU, and both it and its completion evidence fail closed unless
|
| 227 |
+
# CUDA_VISIBLE_DEVICES is disabled, which the parent process cannot offer.
|
| 228 |
+
render_log: Path = output_dir / "target-render.log"
|
| 229 |
+
with render_log.open("wb") as log_handle:
|
| 230 |
+
render_process: subprocess.Popen[bytes] = subprocess.Popen(
|
| 231 |
+
[worker_python, "-m", "fdanyone.skeleton.render_worker", str(render_request)],
|
| 232 |
+
env=_cpu_worker_environment(),
|
| 233 |
+
stdout=log_handle,
|
| 234 |
+
stderr=subprocess.STDOUT,
|
| 235 |
+
)
|
| 236 |
+
completed: bool = False
|
| 237 |
+
failure: FourDAnyoneError | None = None
|
| 238 |
+
|
| 239 |
+
def _validate_target_render() -> None:
|
| 240 |
+
return_code: int = render_process.wait()
|
| 241 |
+
if return_code != 0:
|
| 242 |
+
details: str = render_log.read_text(errors="replace")[-4000:]
|
| 243 |
+
raise FourDAnyoneError(
|
| 244 |
+
"CPU target skeleton renderer failed with "
|
| 245 |
+
f"exit {return_code}. Log tail:\n{details}"
|
| 246 |
+
)
|
| 247 |
+
done_path: Path = output_dir / "target-render.done.json"
|
| 248 |
+
if not done_path.is_file():
|
| 249 |
+
raise FourDAnyoneError(
|
| 250 |
+
f"CPU target skeleton renderer exited without its completion record: {done_path}."
|
| 251 |
+
)
|
| 252 |
+
try:
|
| 253 |
+
render_evidence: object = json.loads(done_path.read_text())
|
| 254 |
+
except json.JSONDecodeError as exc:
|
| 255 |
+
raise FourDAnyoneError(
|
| 256 |
+
f"CPU target skeleton completion record is invalid: {done_path}."
|
| 257 |
+
) from exc
|
| 258 |
+
if (
|
| 259 |
+
not isinstance(render_evidence, dict)
|
| 260 |
+
or int(render_evidence.get("schema_version", 0)) != 1
|
| 261 |
+
or int(render_evidence.get("rendered_views", -1))
|
| 262 |
+
!= conditioning.view_plan.num_target_views
|
| 263 |
+
or render_evidence.get("cuda_visible_devices") not in ("", "-1")
|
| 264 |
+
):
|
| 265 |
+
raise FourDAnyoneError(
|
| 266 |
+
"CPU target skeleton completion evidence failed its authenticity contract."
|
| 267 |
+
)
|
| 268 |
+
metadata_path: Path = output_dir / "metadata.json"
|
| 269 |
+
conditioning_metadata: object = json.loads(metadata_path.read_text())
|
| 270 |
+
if not isinstance(conditioning_metadata, dict):
|
| 271 |
+
raise FourDAnyoneError(f"Conditioning metadata must be an object: {metadata_path}.")
|
| 272 |
+
conditioning_metadata["target_render_overlap"] = render_evidence
|
| 273 |
+
write_json(metadata_path, conditioning_metadata)
|
| 274 |
+
|
| 275 |
+
def wait_for_targets() -> None:
|
| 276 |
+
# Re-entry (the pipeline-level cleanup) must repeat the original
|
| 277 |
+
# verdict, not mask an informative failure with a generic one.
|
| 278 |
+
nonlocal completed, failure
|
| 279 |
+
if completed:
|
| 280 |
+
if failure is not None:
|
| 281 |
+
raise failure
|
| 282 |
+
return
|
| 283 |
+
completed = True
|
| 284 |
+
try:
|
| 285 |
+
_validate_target_render()
|
| 286 |
+
except FourDAnyoneError as exc:
|
| 287 |
+
failure = exc
|
| 288 |
+
raise
|
| 289 |
+
|
| 290 |
+
return conditioning.with_target_waiter(wait_for_targets)
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
@dataclass(frozen=True)
|
| 294 |
+
class PreparedRun:
|
| 295 |
+
"""Everything the generation phase needs from one prepared source clip."""
|
| 296 |
+
|
| 297 |
+
settings: ModeSettings
|
| 298 |
+
"""Inference policy shared by both phases."""
|
| 299 |
+
clip: CanonicalClip
|
| 300 |
+
"""Decoded canonical clip on the frozen frame contract."""
|
| 301 |
+
motion: MotionResult
|
| 302 |
+
"""Published GVHMR result validated against the clip."""
|
| 303 |
+
conditioning: Conditioning
|
| 304 |
+
"""Source, proposal, and target conditioning on the resolved camera grid."""
|
| 305 |
+
checkpoint: Path
|
| 306 |
+
"""Resolved 4DAnyone DiT checkpoint."""
|
| 307 |
+
base_assets: BaseAssets
|
| 308 |
+
"""Resolved VAE, text-encoder, and tokenizer locations."""
|
| 309 |
+
prompt_embedding_path: Path | None
|
| 310 |
+
"""Exported prompt context standing in for the T5 encoder, when supplied."""
|
| 311 |
+
model_identity: dict
|
| 312 |
+
"""Model coordinates recorded in the published metadata."""
|
| 313 |
+
device: str
|
| 314 |
+
"""Selected CUDA device."""
|
| 315 |
+
scratch: Path
|
| 316 |
+
"""Hidden working directory that ``release_run`` removes."""
|
| 317 |
+
motion_dir: Path
|
| 318 |
+
"""Published GVHMR result directory."""
|
| 319 |
+
result_dir: Path
|
| 320 |
+
"""Destination the generation phase publishes to."""
|
| 321 |
+
pipeline_started: float
|
| 322 |
+
"""``time.monotonic`` reading taken when preparation began."""
|
| 323 |
+
|
| 324 |
+
|
| 325 |
+
def prepare_run(
|
| 326 |
+
*,
|
| 327 |
+
settings: ModeSettings,
|
| 328 |
+
video_path: str,
|
| 329 |
+
data_dir: str,
|
| 330 |
+
model_dir: str,
|
| 331 |
+
checkpoint_path: str | None,
|
| 332 |
+
mhr70_regressor_path: str | None,
|
| 333 |
+
gvhmr_root: str,
|
| 334 |
+
device: str,
|
| 335 |
+
start_time: float,
|
| 336 |
+
target_fps: str | float,
|
| 337 |
+
views_per_layer: int,
|
| 338 |
+
layer_pitches: list[int],
|
| 339 |
+
start_yaw: int,
|
| 340 |
+
yaw_span: int,
|
| 341 |
+
views_per_group: int | str,
|
| 342 |
+
enable_rcp: bool,
|
| 343 |
+
enable_tcr: bool,
|
| 344 |
+
inline_workers: bool = False,
|
| 345 |
+
on_motion_stage: MotionStageHook | None = None,
|
| 346 |
+
prompt_embedding_path: Path | None = None,
|
| 347 |
+
) -> PreparedRun:
|
| 348 |
+
"""Recover motion and build conditioning for one source video."""
|
| 349 |
+
|
| 350 |
+
pipeline_started = time.monotonic()
|
| 351 |
+
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
|
| 352 |
+
view_plan = resolve_view_plan(
|
| 353 |
+
views_per_layer=views_per_layer,
|
| 354 |
+
layer_pitches=layer_pitches,
|
| 355 |
+
start_yaw=start_yaw,
|
| 356 |
+
yaw_span=yaw_span,
|
| 357 |
+
views_per_group=views_per_group,
|
| 358 |
+
enable_rcp=enable_rcp,
|
| 359 |
+
enable_tcr=enable_tcr,
|
| 360 |
+
)
|
| 361 |
+
data_root, motion_dir, result_dir = _data_paths(data_dir, video_path)
|
| 362 |
+
run_name = Path(video_path).stem
|
| 363 |
+
atomic = AtomicResultDirectory(result_dir)
|
| 364 |
+
# Fail before asset resolution or video decode; the context manager
|
| 365 |
+
# checks again later in case another process creates the path.
|
| 366 |
+
if os.path.lexists(atomic.destination):
|
| 367 |
+
raise ConfigurationError(
|
| 368 |
+
f"4DAnyone result already exists: {atomic.destination}. Choose a new --data_dir or input filename."
|
| 369 |
+
)
|
| 370 |
+
validate_required_video_codecs()
|
| 371 |
+
device, _ = select_cuda_device(device)
|
| 372 |
+
|
| 373 |
+
ensure_example_video(video_path)
|
| 374 |
+
# Resolve the licensed body model before starting the much larger public
|
| 375 |
+
# model download. Interactive use continues automatically after setup;
|
| 376 |
+
# background jobs receive an actionable error instead of hanging.
|
| 377 |
+
ensure_smplx(model_dir, gvhmr_root)
|
| 378 |
+
ensure_models(model_dir, gvhmr_root)
|
| 379 |
+
gvhmr_root, gvhmr_revision = validate_gvhmr(gvhmr_root)
|
| 380 |
+
worker_python = os.path.abspath(sys.executable)
|
| 381 |
+
|
| 382 |
+
regressor = resolve_regressor(mhr70_regressor_path, model_dir=model_dir)
|
| 383 |
+
foreground_model = resolve_foreground_model(model_dir)
|
| 384 |
+
canonical_fps = None if str(target_fps).lower() == "auto" else target_fps
|
| 385 |
+
clip = decode_canonical_clip(
|
| 386 |
+
video_path,
|
| 387 |
+
num_frames=INFERENCE.num_frames,
|
| 388 |
+
start_time=start_time,
|
| 389 |
+
fps=canonical_fps,
|
| 390 |
+
)
|
| 391 |
+
data_root.mkdir(parents=True, exist_ok=True)
|
| 392 |
+
scratch = Path(tempfile.mkdtemp(prefix=f".{run_name}.scratch-", dir=data_root))
|
| 393 |
+
conditioning: Conditioning | None = None
|
| 394 |
+
try:
|
| 395 |
+
clip_metadata = scratch / "canonical_clip.json"
|
| 396 |
+
clip.write_metadata(clip_metadata)
|
| 397 |
+
working_video = write_gvhmr_video(clip, scratch / "canonical_clip.mp4")
|
| 398 |
+
|
| 399 |
+
if os.path.lexists(motion_dir):
|
| 400 |
+
if motion_dir.is_symlink() or not motion_dir.is_dir():
|
| 401 |
+
raise ConfigurationError(f"GVHMR result path is not a regular directory: {motion_dir}")
|
| 402 |
+
motion = MotionResult.load(motion_dir)
|
| 403 |
+
if motion.gvhmr_revision != gvhmr_revision:
|
| 404 |
+
raise ConfigurationError(
|
| 405 |
+
f"Existing GVHMR result at {motion_dir} was produced by "
|
| 406 |
+
f"GVHMR@{motion.gvhmr_revision}, not GVHMR@{gvhmr_revision}."
|
| 407 |
+
)
|
| 408 |
+
motion.validate_against_clip(clip)
|
| 409 |
+
LOGGER.info("Reusing validated GVHMR result at %s", motion_dir)
|
| 410 |
+
else:
|
| 411 |
+
with AtomicResultDirectory(motion_dir) as motion_work:
|
| 412 |
+
motion = _run_motion(
|
| 413 |
+
working_video=working_video,
|
| 414 |
+
output_dir=scratch / "gvhmr",
|
| 415 |
+
gvhmr_root=gvhmr_root,
|
| 416 |
+
device=device,
|
| 417 |
+
worker_python=worker_python,
|
| 418 |
+
clip_metadata=clip_metadata,
|
| 419 |
+
inline_workers=inline_workers,
|
| 420 |
+
on_motion_stage=on_motion_stage,
|
| 421 |
+
)
|
| 422 |
+
motion.validate_against_clip(clip)
|
| 423 |
+
motion.save(motion_work)
|
| 424 |
+
|
| 425 |
+
checkpoint = resolve_checkpoint(checkpoint_path, model_dir=model_dir)
|
| 426 |
+
base_assets = resolve_base_assets(
|
| 427 |
+
model_dir,
|
| 428 |
+
settings,
|
| 429 |
+
have_prompt_embedding=prompt_embedding_path is not None,
|
| 430 |
+
)
|
| 431 |
+
# Record the published identity only for the published checkpoint; an
|
| 432 |
+
# explicit override must not claim the frozen Hugging Face coordinates.
|
| 433 |
+
if checkpoint_path is None:
|
| 434 |
+
model_identity = {"checkpoint": CHECKPOINT, "repo_id": HF_REPO_ID, "revision": HF_REVISION}
|
| 435 |
+
else:
|
| 436 |
+
model_identity = {"checkpoint": checkpoint.name, "source": "local_override"}
|
| 437 |
+
|
| 438 |
+
conditioning = _build_conditioning(
|
| 439 |
+
regressor=regressor,
|
| 440 |
+
foreground_model=foreground_model,
|
| 441 |
+
gvhmr_root=gvhmr_root,
|
| 442 |
+
output_dir=scratch / "conditioning",
|
| 443 |
+
device=device,
|
| 444 |
+
worker_python=worker_python,
|
| 445 |
+
working_video=working_video,
|
| 446 |
+
clip_metadata=clip_metadata,
|
| 447 |
+
motion_result_dir=motion_dir,
|
| 448 |
+
view_plan=view_plan,
|
| 449 |
+
settings=settings,
|
| 450 |
+
inline_workers=inline_workers,
|
| 451 |
+
)
|
| 452 |
+
if conditioning.num_frames != len(clip.frames) or (
|
| 453 |
+
conditioning.fps_num,
|
| 454 |
+
conditioning.fps_den,
|
| 455 |
+
) != (
|
| 456 |
+
clip.fps_num,
|
| 457 |
+
clip.fps_den,
|
| 458 |
+
):
|
| 459 |
+
raise ConfigurationError("Skeleton conditioning does not match the canonical clip timeline.")
|
| 460 |
+
# Re-decode the worker-produced source before it becomes a model tensor.
|
| 461 |
+
verify_lossless_video(clip, conditioning.source_video)
|
| 462 |
+
except BaseException:
|
| 463 |
+
if conditioning is not None and conditioning.target_waiter is not None:
|
| 464 |
+
conditioning.wait_for_target_skeletons()
|
| 465 |
+
_discard_scratch(scratch)
|
| 466 |
+
raise
|
| 467 |
+
return PreparedRun(
|
| 468 |
+
settings=settings,
|
| 469 |
+
clip=clip,
|
| 470 |
+
motion=motion,
|
| 471 |
+
conditioning=conditioning,
|
| 472 |
+
checkpoint=checkpoint,
|
| 473 |
+
base_assets=base_assets,
|
| 474 |
+
prompt_embedding_path=prompt_embedding_path,
|
| 475 |
+
model_identity=model_identity,
|
| 476 |
+
device=device,
|
| 477 |
+
scratch=scratch,
|
| 478 |
+
motion_dir=motion_dir,
|
| 479 |
+
result_dir=result_dir,
|
| 480 |
+
pipeline_started=pipeline_started,
|
| 481 |
+
)
|
| 482 |
+
|
| 483 |
+
|
| 484 |
+
def generate_run(
|
| 485 |
+
prepared: PreparedRun,
|
| 486 |
+
*,
|
| 487 |
+
seed: int,
|
| 488 |
+
on_denoise_step: DenoiseStepHook | None = None,
|
| 489 |
+
) -> dict:
|
| 490 |
+
"""Generate and publish every view for one prepared run.
|
| 491 |
+
|
| 492 |
+
The prompt embedding is fixed by ``prepare_run``, because whether it exists
|
| 493 |
+
decides whether the T5 encoder had to be resolved at all.
|
| 494 |
+
"""
|
| 495 |
+
|
| 496 |
+
# Heavy rendering and generation are imported only after the motion
|
| 497 |
+
# contract has been materialized, keeping CLI/help and CPU tests light.
|
| 498 |
+
from fdanyone.model.inference import generate_views
|
| 499 |
+
from fdanyone.output import export_result
|
| 500 |
+
|
| 501 |
+
with AtomicResultDirectory(prepared.result_dir) as work:
|
| 502 |
+
generated = generate_views(
|
| 503 |
+
clip=prepared.clip,
|
| 504 |
+
conditioning=prepared.conditioning,
|
| 505 |
+
checkpoint_path=prepared.checkpoint,
|
| 506 |
+
assets=prepared.base_assets,
|
| 507 |
+
output_dir=prepared.scratch / "generation",
|
| 508 |
+
device=prepared.device,
|
| 509 |
+
seed=seed,
|
| 510 |
+
settings=prepared.settings,
|
| 511 |
+
on_denoise_step=on_denoise_step,
|
| 512 |
+
prompt_embedding_path=prepared.prompt_embedding_path,
|
| 513 |
+
)
|
| 514 |
+
summary = export_result(
|
| 515 |
+
clip=prepared.clip,
|
| 516 |
+
conditioning=prepared.conditioning,
|
| 517 |
+
generated=generated,
|
| 518 |
+
destination=work,
|
| 519 |
+
motion=prepared.motion,
|
| 520 |
+
model_identity=prepared.model_identity,
|
| 521 |
+
pipeline_started=prepared.pipeline_started,
|
| 522 |
+
settings=prepared.settings,
|
| 523 |
+
)
|
| 524 |
+
summary["result_dir"] = str(prepared.result_dir)
|
| 525 |
+
summary["motion_dir"] = str(prepared.motion_dir)
|
| 526 |
+
return summary
|
| 527 |
+
|
| 528 |
+
|
| 529 |
+
def release_run(prepared: PreparedRun) -> None:
|
| 530 |
+
"""Settle deferred target rendering and drop the run's scratch directory."""
|
| 531 |
+
|
| 532 |
+
if prepared.conditioning.target_waiter is not None:
|
| 533 |
+
prepared.conditioning.wait_for_target_skeletons()
|
| 534 |
+
_discard_scratch(prepared.scratch)
|
| 535 |
+
|
| 536 |
+
|
| 537 |
+
def run_pipeline(
|
| 538 |
+
*,
|
| 539 |
+
settings: ModeSettings,
|
| 540 |
+
video_path: str,
|
| 541 |
+
data_dir: str,
|
| 542 |
+
model_dir: str,
|
| 543 |
+
checkpoint_path: str | None,
|
| 544 |
+
mhr70_regressor_path: str | None,
|
| 545 |
+
gvhmr_root: str,
|
| 546 |
+
device: str,
|
| 547 |
+
start_time: float,
|
| 548 |
+
target_fps: str | float,
|
| 549 |
+
seed: int,
|
| 550 |
+
views_per_layer: int,
|
| 551 |
+
layer_pitches: list[int],
|
| 552 |
+
start_yaw: int,
|
| 553 |
+
yaw_span: int,
|
| 554 |
+
views_per_group: int | str,
|
| 555 |
+
enable_rcp: bool,
|
| 556 |
+
enable_tcr: bool,
|
| 557 |
+
inline_workers: bool = False,
|
| 558 |
+
on_motion_stage: MotionStageHook | None = None,
|
| 559 |
+
on_denoise_step: DenoiseStepHook | None = None,
|
| 560 |
+
prompt_embedding_path: Path | None = None,
|
| 561 |
+
) -> dict:
|
| 562 |
+
"""Execute inference and publish reusable GVHMR plus 4DAnyone results."""
|
| 563 |
+
|
| 564 |
+
if seed < 0:
|
| 565 |
+
raise ConfigurationError(f"seed must be non-negative, got {seed}.")
|
| 566 |
+
prepared = prepare_run(
|
| 567 |
+
settings=settings,
|
| 568 |
+
video_path=video_path,
|
| 569 |
+
data_dir=data_dir,
|
| 570 |
+
model_dir=model_dir,
|
| 571 |
+
checkpoint_path=checkpoint_path,
|
| 572 |
+
mhr70_regressor_path=mhr70_regressor_path,
|
| 573 |
+
gvhmr_root=gvhmr_root,
|
| 574 |
+
device=device,
|
| 575 |
+
start_time=start_time,
|
| 576 |
+
target_fps=target_fps,
|
| 577 |
+
views_per_layer=views_per_layer,
|
| 578 |
+
layer_pitches=layer_pitches,
|
| 579 |
+
start_yaw=start_yaw,
|
| 580 |
+
yaw_span=yaw_span,
|
| 581 |
+
views_per_group=views_per_group,
|
| 582 |
+
enable_rcp=enable_rcp,
|
| 583 |
+
enable_tcr=enable_tcr,
|
| 584 |
+
inline_workers=inline_workers,
|
| 585 |
+
on_motion_stage=on_motion_stage,
|
| 586 |
+
prompt_embedding_path=prompt_embedding_path,
|
| 587 |
+
)
|
| 588 |
+
try:
|
| 589 |
+
return generate_run(prepared, seed=seed, on_denoise_step=on_denoise_step)
|
| 590 |
+
finally:
|
| 591 |
+
release_run(prepared)
|
|
@@ -0,0 +1,161 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Read what one finished 4DAnyone inference run left on disk.
|
| 2 |
+
|
| 3 |
+
Two readers live here. ``discover_run`` collects every artifact of a clip --
|
| 4 |
+
motion, generated videos, skeletons, and the source clip -- into a
|
| 5 |
+
``RunLayout``. The rig readers parse the ``cameras.json`` that
|
| 6 |
+
``fdanyone.output`` writes next to the generated videos.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from __future__ import annotations
|
| 10 |
+
|
| 11 |
+
import json
|
| 12 |
+
import logging
|
| 13 |
+
from dataclasses import dataclass
|
| 14 |
+
from pathlib import Path
|
| 15 |
+
|
| 16 |
+
from fdanyone.errors import FourDAnyoneError
|
| 17 |
+
from fdanyone.io import read_json
|
| 18 |
+
|
| 19 |
+
LOGGER = logging.getLogger("fdanyone.runs")
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
# ---------------------------------------------------------------------------
|
| 23 |
+
# Run layout
|
| 24 |
+
# ---------------------------------------------------------------------------
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
@dataclass(frozen=True)
|
| 28 |
+
class RunLayout:
|
| 29 |
+
"""Files that one inference run left on disk for a single clip."""
|
| 30 |
+
|
| 31 |
+
clip: str
|
| 32 |
+
data_dir: Path
|
| 33 |
+
result_dir: Path
|
| 34 |
+
motion_dir: Path
|
| 35 |
+
cameras_json: Path | None
|
| 36 |
+
metadata_json: Path | None
|
| 37 |
+
source_video: Path | None
|
| 38 |
+
dense_videos: tuple[tuple[int, Path], ...]
|
| 39 |
+
skeleton_videos: tuple[tuple[int, Path], ...]
|
| 40 |
+
sparse_videos: tuple[tuple[int, Path], ...]
|
| 41 |
+
|
| 42 |
+
@property
|
| 43 |
+
def has_motion(self) -> bool:
|
| 44 |
+
return (self.motion_dir / "motion.json").is_file() and (self.motion_dir / "motion.safetensors").is_file()
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def numbered_videos(directory: Path) -> tuple[tuple[int, Path], ...]:
|
| 48 |
+
"""Return ``(camera_id, path)`` pairs for ``NN.mp4`` files, ID-ordered."""
|
| 49 |
+
|
| 50 |
+
if not directory.is_dir():
|
| 51 |
+
return ()
|
| 52 |
+
found: list[tuple[int, Path]] = []
|
| 53 |
+
for path in sorted(directory.glob("*.mp4")):
|
| 54 |
+
try:
|
| 55 |
+
found.append((int(path.stem), path))
|
| 56 |
+
except ValueError:
|
| 57 |
+
LOGGER.warning("Ignoring video with a non-numeric name: %s", path)
|
| 58 |
+
return tuple(sorted(found))
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def find_source_video(data_dir: Path, clip: str, filename: str | None) -> Path | None:
|
| 62 |
+
"""Locate the input clip that produced the run."""
|
| 63 |
+
|
| 64 |
+
candidates: list[Path] = []
|
| 65 |
+
if filename:
|
| 66 |
+
candidates.append(data_dir / "source" / "pexels" / filename)
|
| 67 |
+
for suffix in (".mp4", ".mov", ".mkv", ".webm", ".MP4", ".MOV"):
|
| 68 |
+
candidates.append(data_dir / "source" / "pexels" / f"{clip}{suffix}")
|
| 69 |
+
for candidate in candidates:
|
| 70 |
+
if candidate.is_file():
|
| 71 |
+
return candidate
|
| 72 |
+
# Only the input and result trees can hold a clip; the rest of the data
|
| 73 |
+
# root is model output that a recursive walk would scan for nothing.
|
| 74 |
+
roots = tuple(root for root in (data_dir / "source", data_dir / "fdanyone") if root.is_dir())
|
| 75 |
+
if filename:
|
| 76 |
+
matches = sorted(path for root in roots for path in root.rglob(filename) if path.is_file())
|
| 77 |
+
if matches:
|
| 78 |
+
return matches[0]
|
| 79 |
+
matches = sorted(
|
| 80 |
+
path for root in roots for path in root.rglob(f"{clip}.*") if path.is_file() and path.suffix != ".rrd"
|
| 81 |
+
)
|
| 82 |
+
return matches[0] if matches else None
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def discover_run(data_dir: Path, clip: str) -> RunLayout:
|
| 86 |
+
"""Collect every artifact of ``clip`` and reject an empty run."""
|
| 87 |
+
|
| 88 |
+
result_dir = data_dir / "fdanyone" / clip
|
| 89 |
+
motion_dir = data_dir / "gvhmr" / "results" / clip
|
| 90 |
+
metadata_json = result_dir / "metadata.json"
|
| 91 |
+
cameras_json = result_dir / "cameras.json"
|
| 92 |
+
metadata = read_json(metadata_json)
|
| 93 |
+
filename = None
|
| 94 |
+
if metadata is not None:
|
| 95 |
+
filename = str(metadata.get("input", {}).get("filename") or "") or None
|
| 96 |
+
|
| 97 |
+
layout = RunLayout(
|
| 98 |
+
clip=clip,
|
| 99 |
+
data_dir=data_dir,
|
| 100 |
+
result_dir=result_dir,
|
| 101 |
+
motion_dir=motion_dir,
|
| 102 |
+
cameras_json=cameras_json if cameras_json.is_file() else None,
|
| 103 |
+
metadata_json=metadata_json if metadata_json.is_file() else None,
|
| 104 |
+
source_video=find_source_video(data_dir, clip, filename),
|
| 105 |
+
dense_videos=numbered_videos(result_dir / "videos" / "dense"),
|
| 106 |
+
skeleton_videos=numbered_videos(result_dir / "skeletons"),
|
| 107 |
+
sparse_videos=numbered_videos(result_dir / "videos" / "sparse"),
|
| 108 |
+
)
|
| 109 |
+
if not layout.has_motion and not layout.dense_videos and layout.source_video is None:
|
| 110 |
+
raise FourDAnyoneError(
|
| 111 |
+
f"Nothing to visualize for clip {clip!r}. Expected at least one of:\n"
|
| 112 |
+
f" {motion_dir / 'motion.safetensors'} (GVHMR motion)\n"
|
| 113 |
+
f" {result_dir / 'videos' / 'dense'} (generated views)\n"
|
| 114 |
+
f" a source clip named {clip}.* under {data_dir}\n"
|
| 115 |
+
"Run `python inference.py --video_path <clip>` first."
|
| 116 |
+
)
|
| 117 |
+
return layout
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
# ---------------------------------------------------------------------------
|
| 121 |
+
# Camera rig
|
| 122 |
+
# ---------------------------------------------------------------------------
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def read_cameras(result: Path) -> dict:
|
| 126 |
+
"""Read the ``cameras.json`` a result directory must carry."""
|
| 127 |
+
|
| 128 |
+
path = result / "cameras.json"
|
| 129 |
+
try:
|
| 130 |
+
cameras = json.loads(path.read_text())
|
| 131 |
+
except (FileNotFoundError, json.JSONDecodeError) as exc:
|
| 132 |
+
raise FourDAnyoneError(f"Cannot read {path}: {exc}") from exc
|
| 133 |
+
if not isinstance(cameras, dict):
|
| 134 |
+
raise FourDAnyoneError("cameras.json must contain a JSON object.")
|
| 135 |
+
return cameras
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def camera_records(rig: dict) -> list[dict]:
|
| 139 |
+
"""Validate an OPENCV rig and return its ID-ordered camera records."""
|
| 140 |
+
|
| 141 |
+
if not isinstance(rig, dict) or rig.get("camera_model") != "OPENCV":
|
| 142 |
+
raise FourDAnyoneError("cameras.json has no supported OPENCV camera rig.")
|
| 143 |
+
cameras = rig.get("cameras")
|
| 144 |
+
if not isinstance(cameras, list) or not cameras:
|
| 145 |
+
raise FourDAnyoneError("Camera rig must contain at least one camera.")
|
| 146 |
+
if [camera.get("camera_id") for camera in cameras if isinstance(camera, dict)] != list(range(len(cameras))):
|
| 147 |
+
raise FourDAnyoneError("Camera rig must be ordered by camera ID.")
|
| 148 |
+
return cameras
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
def dense_video_paths(result: Path, cameras: list[dict]) -> tuple[Path, ...]:
|
| 152 |
+
"""Return the generated video of every camera, in rig order."""
|
| 153 |
+
|
| 154 |
+
paths = []
|
| 155 |
+
for camera in cameras:
|
| 156 |
+
camera_id = int(camera["camera_id"])
|
| 157 |
+
relative = f"videos/dense/{camera_id:02d}.mp4"
|
| 158 |
+
if camera.get("video") != relative:
|
| 159 |
+
raise FourDAnyoneError(f"Camera {camera_id:02d} points to the wrong target video.")
|
| 160 |
+
paths.append(result / relative)
|
| 161 |
+
return tuple(paths)
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""MHR70 projection and Goliath40 conditioning renderer."""
|
|
@@ -0,0 +1,170 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Minimal MHR70/Goliath40 conditioning schema.
|
| 2 |
+
|
| 3 |
+
The exact keypoint names/order and no-finger link closure were extracted from
|
| 4 |
+
``facebookresearch/sapiens2`` revision
|
| 5 |
+
``0e51c12d7c7257d88431b2d50e523a7b03004854`` and remain subject to the
|
| 6 |
+
Sapiens2 License. A copy is kept at
|
| 7 |
+
``third_party/licenses/SAPIENS2_LICENSE.md``. The RGB palette, per-keypoint
|
| 8 |
+
color policy, and renderer implementation are original 4DAnyone material
|
| 9 |
+
licensed under Apache-2.0.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
from __future__ import annotations
|
| 13 |
+
|
| 14 |
+
# The exact names/order below are derived from Sapiens2, Copyright (c) Meta
|
| 15 |
+
# Platforms, Inc. and affiliates, under the Sapiens2 License noted above.
|
| 16 |
+
|
| 17 |
+
KEYPOINT_NAMES = (
|
| 18 |
+
"nose",
|
| 19 |
+
"left-eye",
|
| 20 |
+
"right-eye",
|
| 21 |
+
"left-ear",
|
| 22 |
+
"right-ear",
|
| 23 |
+
"left-shoulder",
|
| 24 |
+
"right-shoulder",
|
| 25 |
+
"left-elbow",
|
| 26 |
+
"right-elbow",
|
| 27 |
+
"left-hip",
|
| 28 |
+
"right-hip",
|
| 29 |
+
"left-knee",
|
| 30 |
+
"right-knee",
|
| 31 |
+
"left-ankle",
|
| 32 |
+
"right-ankle",
|
| 33 |
+
"left-big-toe-tip",
|
| 34 |
+
"left-small-toe-tip",
|
| 35 |
+
"left-heel",
|
| 36 |
+
"right-big-toe-tip",
|
| 37 |
+
"right-small-toe-tip",
|
| 38 |
+
"right-heel",
|
| 39 |
+
"right-thumb-tip",
|
| 40 |
+
"right-thumb-first-joint",
|
| 41 |
+
"right-thumb-second-joint",
|
| 42 |
+
"right-thumb-third-joint",
|
| 43 |
+
"right-index-tip",
|
| 44 |
+
"right-index-first-joint",
|
| 45 |
+
"right-index-second-joint",
|
| 46 |
+
"right-index-third-joint",
|
| 47 |
+
"right-middle-tip",
|
| 48 |
+
"right-middle-first-joint",
|
| 49 |
+
"right-middle-second-joint",
|
| 50 |
+
"right-middle-third-joint",
|
| 51 |
+
"right-ring-tip",
|
| 52 |
+
"right-ring-first-joint",
|
| 53 |
+
"right-ring-second-joint",
|
| 54 |
+
"right-ring-third-joint",
|
| 55 |
+
"right-pinky-tip",
|
| 56 |
+
"right-pinky-first-joint",
|
| 57 |
+
"right-pinky-second-joint",
|
| 58 |
+
"right-pinky-third-joint",
|
| 59 |
+
"right-wrist",
|
| 60 |
+
"left-thumb-tip",
|
| 61 |
+
"left-thumb-first-joint",
|
| 62 |
+
"left-thumb-second-joint",
|
| 63 |
+
"left-thumb-third-joint",
|
| 64 |
+
"left-index-tip",
|
| 65 |
+
"left-index-first-joint",
|
| 66 |
+
"left-index-second-joint",
|
| 67 |
+
"left-index-third-joint",
|
| 68 |
+
"left-middle-tip",
|
| 69 |
+
"left-middle-first-joint",
|
| 70 |
+
"left-middle-second-joint",
|
| 71 |
+
"left-middle-third-joint",
|
| 72 |
+
"left-ring-tip",
|
| 73 |
+
"left-ring-first-joint",
|
| 74 |
+
"left-ring-second-joint",
|
| 75 |
+
"left-ring-third-joint",
|
| 76 |
+
"left-pinky-tip",
|
| 77 |
+
"left-pinky-first-joint",
|
| 78 |
+
"left-pinky-second-joint",
|
| 79 |
+
"left-pinky-third-joint",
|
| 80 |
+
"left-wrist",
|
| 81 |
+
"left-olecranon",
|
| 82 |
+
"right-olecranon",
|
| 83 |
+
"left-cubital-fossa",
|
| 84 |
+
"right-cubital-fossa",
|
| 85 |
+
"left-acromion",
|
| 86 |
+
"right-acromion",
|
| 87 |
+
"neck",
|
| 88 |
+
)
|
| 89 |
+
|
| 90 |
+
# First-party 4DAnyone color palette, licensed under Apache-2.0.
|
| 91 |
+
BLUE = (116, 192, 252)
|
| 92 |
+
GREEN = (130, 186, 129)
|
| 93 |
+
ORANGE = (248, 129, 81)
|
| 94 |
+
TEAL = (99, 230, 190)
|
| 95 |
+
YELLOW = (255, 212, 59)
|
| 96 |
+
PINK = (229, 153, 247)
|
| 97 |
+
PURPLE = (177, 151, 252)
|
| 98 |
+
RED = (255, 135, 135)
|
| 99 |
+
|
| 100 |
+
# The schema ids and link endpoints are Sapiens2-derived. The RGB assignments
|
| 101 |
+
# are first-party 4DAnyone material.
|
| 102 |
+
# (schema id, first keypoint, second keypoint, RGB, major)
|
| 103 |
+
LINKS = (
|
| 104 |
+
(0, 13, 11, TEAL, True),
|
| 105 |
+
(1, 11, 9, TEAL, True),
|
| 106 |
+
(2, 14, 12, YELLOW, True),
|
| 107 |
+
(3, 12, 10, YELLOW, True),
|
| 108 |
+
(4, 9, 10, BLUE, True),
|
| 109 |
+
(5, 5, 9, GREEN, True),
|
| 110 |
+
(6, 6, 10, ORANGE, True),
|
| 111 |
+
(7, 5, 6, BLUE, True),
|
| 112 |
+
(8, 5, 7, TEAL, True),
|
| 113 |
+
(9, 6, 8, YELLOW, True),
|
| 114 |
+
(10, 7, 62, TEAL, True),
|
| 115 |
+
(11, 8, 41, YELLOW, True),
|
| 116 |
+
(12, 1, 2, BLUE, True),
|
| 117 |
+
(13, 0, 1, GREEN, True),
|
| 118 |
+
(14, 0, 2, ORANGE, True),
|
| 119 |
+
(15, 1, 3, GREEN, True),
|
| 120 |
+
(16, 2, 4, ORANGE, True),
|
| 121 |
+
(17, 3, 5, GREEN, True),
|
| 122 |
+
(18, 4, 6, ORANGE, True),
|
| 123 |
+
(19, 13, 15, TEAL, True),
|
| 124 |
+
(20, 13, 16, TEAL, True),
|
| 125 |
+
(21, 13, 17, TEAL, True),
|
| 126 |
+
(22, 14, 18, YELLOW, True),
|
| 127 |
+
(23, 14, 19, YELLOW, True),
|
| 128 |
+
(24, 14, 20, YELLOW, True),
|
| 129 |
+
(25, 62, 45, YELLOW, False),
|
| 130 |
+
(29, 62, 49, PINK, False),
|
| 131 |
+
(33, 62, 53, PURPLE, False),
|
| 132 |
+
(37, 62, 57, RED, False),
|
| 133 |
+
(41, 62, 61, TEAL, False),
|
| 134 |
+
(45, 41, 24, YELLOW, False),
|
| 135 |
+
(49, 41, 28, PINK, False),
|
| 136 |
+
(53, 41, 32, PURPLE, False),
|
| 137 |
+
(57, 41, 36, RED, False),
|
| 138 |
+
(61, 41, 40, TEAL, False),
|
| 139 |
+
(65, 5, 10, PURPLE, True),
|
| 140 |
+
(66, 6, 9, PURPLE, True),
|
| 141 |
+
(67, 3, 4, PURPLE, True),
|
| 142 |
+
)
|
| 143 |
+
|
| 144 |
+
EXTRA_KEYPOINT_IDS = frozenset(range(63, 70))
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def keypoint_color(keypoint_id: int) -> tuple[int, int, int]:
|
| 148 |
+
name = KEYPOINT_NAMES[keypoint_id]
|
| 149 |
+
if name == "nose" or name == "neck":
|
| 150 |
+
return BLUE
|
| 151 |
+
if name in {"left-eye", "left-ear"}:
|
| 152 |
+
return GREEN
|
| 153 |
+
if name in {"right-eye", "right-ear"}:
|
| 154 |
+
return ORANGE
|
| 155 |
+
if "thumb" in name:
|
| 156 |
+
return YELLOW
|
| 157 |
+
if "forefinger" in name or "index" in name:
|
| 158 |
+
return PINK
|
| 159 |
+
if "middle" in name:
|
| 160 |
+
return PURPLE
|
| 161 |
+
if "ring" in name:
|
| 162 |
+
return RED
|
| 163 |
+
if "pinky" in name:
|
| 164 |
+
return TEAL
|
| 165 |
+
return TEAL if name.startswith("left-") else YELLOW if name.startswith("right-") else BLUE
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
VISIBLE_KEYPOINT_IDS = (
|
| 169 |
+
frozenset({index for _, first, second, _, _ in LINKS for index in (first, second)}) | EXTRA_KEYPOINT_IDS
|
| 170 |
+
)
|
|
@@ -0,0 +1,839 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Build source-aware Goliath40 conditioning for multi-view generation."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import logging
|
| 7 |
+
from collections.abc import Callable, Iterable
|
| 8 |
+
from contextlib import contextmanager
|
| 9 |
+
from dataclasses import asdict, dataclass, field, replace
|
| 10 |
+
from fractions import Fraction
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
|
| 13 |
+
import numpy as np
|
| 14 |
+
|
| 15 |
+
from fdanyone.assets import BIREFNET_REPO_ID, BIREFNET_REVISION
|
| 16 |
+
from fdanyone.config import CAMERA, CROP, FOREGROUND, INFERENCE, SKELETON, CameraConfig
|
| 17 |
+
from fdanyone.errors import AssetError, FourDAnyoneError
|
| 18 |
+
from fdanyone.foreground import predict_foreground_masks
|
| 19 |
+
from fdanyone.geometry.cameras import (
|
| 20 |
+
CAMERA_FRAME,
|
| 21 |
+
WORLD_FRAME,
|
| 22 |
+
Camera,
|
| 23 |
+
camera_grid,
|
| 24 |
+
camera_ring,
|
| 25 |
+
project_points,
|
| 26 |
+
reference_intrinsics,
|
| 27 |
+
)
|
| 28 |
+
from fdanyone.geometry.crop import (
|
| 29 |
+
Crop,
|
| 30 |
+
center_crop,
|
| 31 |
+
crop_from_bounds,
|
| 32 |
+
mask_bounds,
|
| 33 |
+
transform_intrinsics,
|
| 34 |
+
)
|
| 35 |
+
from fdanyone.geometry.framing import analyze_input_framing, solve_sequence_framing
|
| 36 |
+
from fdanyone.io import write_json
|
| 37 |
+
from fdanyone.motion.gvhmr import gvhmr_imports, validate_gvhmr
|
| 38 |
+
from fdanyone.motion.result import MotionResult
|
| 39 |
+
from fdanyone.skeleton.keypoints import KEYPOINT_NAMES
|
| 40 |
+
from fdanyone.skeleton.renderer import (
|
| 41 |
+
estimate_body_height,
|
| 42 |
+
projected_body_scales,
|
| 43 |
+
render_goliath40,
|
| 44 |
+
)
|
| 45 |
+
from fdanyone.vendor.pytorch3d_compat import (
|
| 46 |
+
install_if_needed as install_pytorch3d_compat,
|
| 47 |
+
)
|
| 48 |
+
from fdanyone.video import (
|
| 49 |
+
CanonicalClip,
|
| 50 |
+
iter_rgb_video,
|
| 51 |
+
write_lossless_video,
|
| 52 |
+
write_video,
|
| 53 |
+
)
|
| 54 |
+
from fdanyone.views import ViewPlan
|
| 55 |
+
|
| 56 |
+
LOGGER = logging.getLogger("fdanyone")
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
@dataclass(frozen=True)
|
| 60 |
+
class SkeletonVideo:
|
| 61 |
+
path: Path
|
| 62 |
+
crop: Crop
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
@dataclass(frozen=True)
|
| 66 |
+
class Conditioning:
|
| 67 |
+
root: Path
|
| 68 |
+
source_video: Path
|
| 69 |
+
source_crop: Crop
|
| 70 |
+
target_skeletons: tuple[SkeletonVideo, ...]
|
| 71 |
+
rcp_skeletons: tuple[SkeletonVideo, ...]
|
| 72 |
+
view_plan: ViewPlan
|
| 73 |
+
fps_num: int
|
| 74 |
+
fps_den: int
|
| 75 |
+
num_frames: int
|
| 76 |
+
skeleton_video_decoder: str = "pyav"
|
| 77 |
+
"""Backend used when loading the rendered skeleton videos."""
|
| 78 |
+
target_waiter: Callable[[], None] | None = field(default=None, repr=False, compare=False)
|
| 79 |
+
"""Optional completion barrier for deferred target skeleton videos."""
|
| 80 |
+
|
| 81 |
+
def load_source_tensor(self):
|
| 82 |
+
return _video_tensor(self.source_video, self.num_frames, crop=self.source_crop)
|
| 83 |
+
|
| 84 |
+
def load_skeleton_tensor(
|
| 85 |
+
self,
|
| 86 |
+
skeletons: Iterable[SkeletonVideo],
|
| 87 |
+
*,
|
| 88 |
+
device: str = "cuda",
|
| 89 |
+
):
|
| 90 |
+
import torch
|
| 91 |
+
|
| 92 |
+
videos = [
|
| 93 |
+
_video_tensor(
|
| 94 |
+
item.path,
|
| 95 |
+
self.num_frames,
|
| 96 |
+
crop=item.crop,
|
| 97 |
+
decoder=self.skeleton_video_decoder,
|
| 98 |
+
device=device,
|
| 99 |
+
)
|
| 100 |
+
for item in skeletons
|
| 101 |
+
]
|
| 102 |
+
return torch.cat(videos, dim=0)
|
| 103 |
+
|
| 104 |
+
def wait_for_target_skeletons(self) -> None:
|
| 105 |
+
"""Wait for a deferred CPU renderer and validate every target artifact."""
|
| 106 |
+
|
| 107 |
+
if self.target_waiter is not None:
|
| 108 |
+
self.target_waiter()
|
| 109 |
+
missing: list[str] = [str(item.path) for item in self.target_skeletons if not item.path.is_file()]
|
| 110 |
+
if missing:
|
| 111 |
+
raise FourDAnyoneError(f"Target skeleton rendering is incomplete: {missing}.")
|
| 112 |
+
|
| 113 |
+
def with_target_waiter(self, waiter: Callable[[], None]) -> Conditioning:
|
| 114 |
+
"""Return this file contract with one process-completion barrier attached."""
|
| 115 |
+
|
| 116 |
+
return replace(self, target_waiter=waiter)
|
| 117 |
+
|
| 118 |
+
@classmethod
|
| 119 |
+
def load(
|
| 120 |
+
cls,
|
| 121 |
+
directory: str | Path,
|
| 122 |
+
*,
|
| 123 |
+
allow_pending_targets: bool = False,
|
| 124 |
+
skeleton_video_decoder: str = "pyav",
|
| 125 |
+
) -> Conditioning:
|
| 126 |
+
"""Load conditioning artifacts, optionally before deferred targets finish."""
|
| 127 |
+
|
| 128 |
+
root = Path(directory).expanduser().resolve()
|
| 129 |
+
camera_payload = json.loads((root / "cameras.json").read_text())
|
| 130 |
+
metadata = json.loads((root / "metadata.json").read_text())
|
| 131 |
+
try:
|
| 132 |
+
view_plan = ViewPlan.from_dict(metadata["view_plan"])
|
| 133 |
+
except (KeyError, TypeError) as exc:
|
| 134 |
+
raise FourDAnyoneError("Conditioning artifacts have no valid view plan.") from exc
|
| 135 |
+
records = camera_payload["cameras"]
|
| 136 |
+
if [int(record["camera_id"]) for record in records] != list(range(view_plan.num_target_views)):
|
| 137 |
+
raise FourDAnyoneError("Target conditioning cameras are not in canonical order.")
|
| 138 |
+
for record, view in zip(records, view_plan.target_views, strict=True):
|
| 139 |
+
if (
|
| 140 |
+
int(record.get("layer_index", -1)) != view.layer_index
|
| 141 |
+
or int(record.get("pitch_degrees", 1000)) != view.pitch
|
| 142 |
+
or abs(float(record.get("yaw_degrees", 1000.0)) - view.yaw) > 1e-8
|
| 143 |
+
):
|
| 144 |
+
raise FourDAnyoneError("Target conditioning cameras do not match the resolved view layout.")
|
| 145 |
+
if camera_payload.get("front_camera_ids") != list(view_plan.front_camera_ids):
|
| 146 |
+
raise FourDAnyoneError("Target conditioning has the wrong frontal-camera IDs.")
|
| 147 |
+
rcp_records = camera_payload.get("rcp_cameras", [])
|
| 148 |
+
if [int(record["camera_id"]) for record in rcp_records] != list(view_plan.rcp_camera_ids):
|
| 149 |
+
raise FourDAnyoneError("RCP conditioning cameras do not match the resolved view plan.")
|
| 150 |
+
|
| 151 |
+
def skeletons(camera_records: list[dict]) -> tuple[SkeletonVideo, ...]:
|
| 152 |
+
return tuple(
|
| 153 |
+
SkeletonVideo(root / record["skeleton_video"], Crop(**record["crop"])) for record in camera_records
|
| 154 |
+
)
|
| 155 |
+
|
| 156 |
+
target_skeletons = skeletons(records)
|
| 157 |
+
rcp_skeletons = skeletons(rcp_records)
|
| 158 |
+
required: tuple[Path, ...] = (
|
| 159 |
+
root / metadata["source_video"],
|
| 160 |
+
*(item.path for item in rcp_skeletons),
|
| 161 |
+
*((item.path for item in target_skeletons) if not allow_pending_targets else ()),
|
| 162 |
+
)
|
| 163 |
+
missing = [str(path) for path in required if not path.is_file()]
|
| 164 |
+
if missing:
|
| 165 |
+
raise FourDAnyoneError(f"Conditioning artifacts are incomplete: {missing}.")
|
| 166 |
+
return cls(
|
| 167 |
+
root=root,
|
| 168 |
+
source_video=root / metadata["source_video"],
|
| 169 |
+
source_crop=Crop(**metadata["source_crop"]),
|
| 170 |
+
target_skeletons=target_skeletons,
|
| 171 |
+
rcp_skeletons=rcp_skeletons,
|
| 172 |
+
view_plan=view_plan,
|
| 173 |
+
fps_num=int(metadata["fps_num"]),
|
| 174 |
+
fps_den=int(metadata["fps_den"]),
|
| 175 |
+
num_frames=int(metadata["num_frames"]),
|
| 176 |
+
skeleton_video_decoder=skeleton_video_decoder,
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
@dataclass(frozen=True)
|
| 181 |
+
class _BodyGeometry:
|
| 182 |
+
vertices_world: np.ndarray
|
| 183 |
+
joints_world: np.ndarray
|
| 184 |
+
keypoints_world: np.ndarray
|
| 185 |
+
keypoints_incam: np.ndarray
|
| 186 |
+
motion_world_to_canonical_world: np.ndarray
|
| 187 |
+
regressor_metadata: dict[str, int | str]
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def _assert_nvdec_engaged(
|
| 191 |
+
decoder,
|
| 192 |
+
*,
|
| 193 |
+
decoded_device_type: str,
|
| 194 |
+
declared_frames: int,
|
| 195 |
+
decoded_frames: int,
|
| 196 |
+
expected_frames: int,
|
| 197 |
+
) -> None:
|
| 198 |
+
"""Fail closed unless TorchCodec proves an exact, GPU-native NVDEC decode."""
|
| 199 |
+
|
| 200 |
+
if declared_frames != expected_frames:
|
| 201 |
+
raise FourDAnyoneError(
|
| 202 |
+
f"NVDEC video declares {declared_frames} frames, expected {expected_frames}."
|
| 203 |
+
)
|
| 204 |
+
if decoded_frames != expected_frames:
|
| 205 |
+
raise FourDAnyoneError(
|
| 206 |
+
f"NVDEC produced {decoded_frames} frames, expected {expected_frames}."
|
| 207 |
+
)
|
| 208 |
+
fallback = decoder.cpu_fallback
|
| 209 |
+
if not fallback.status_known:
|
| 210 |
+
raise FourDAnyoneError(f"NVDEC CPU fallback status is unknown: {fallback}")
|
| 211 |
+
if fallback:
|
| 212 |
+
raise FourDAnyoneError(f"NVDEC fell back to CPU: {fallback}")
|
| 213 |
+
if decoded_device_type != "cuda":
|
| 214 |
+
raise FourDAnyoneError(
|
| 215 |
+
f"NVDEC returned a {decoded_device_type.upper()} tensor instead of CUDA."
|
| 216 |
+
)
|
| 217 |
+
|
| 218 |
+
|
| 219 |
+
def _crop_and_normalize(tensor, crop: Crop | None):
|
| 220 |
+
"""Apply the scaled crop, restore the output size, and map to [-1, 1]."""
|
| 221 |
+
|
| 222 |
+
import torchvision.transforms.functional as transform
|
| 223 |
+
from torchvision.transforms import InterpolationMode
|
| 224 |
+
|
| 225 |
+
if crop is not None:
|
| 226 |
+
height, width = tensor.shape[-2:]
|
| 227 |
+
scaled = _scale_crop(
|
| 228 |
+
crop,
|
| 229 |
+
scale_y=height / crop.original_height,
|
| 230 |
+
scale_x=width / crop.original_width,
|
| 231 |
+
)
|
| 232 |
+
tensor = transform.crop(tensor, scaled.top, scaled.left, scaled.height, scaled.width)
|
| 233 |
+
if tensor.shape[-2:] != (crop.output_height, crop.output_width):
|
| 234 |
+
tensor = transform.resize(
|
| 235 |
+
tensor,
|
| 236 |
+
(crop.output_height, crop.output_width),
|
| 237 |
+
interpolation=InterpolationMode.BICUBIC,
|
| 238 |
+
antialias=True,
|
| 239 |
+
)
|
| 240 |
+
tensor = tensor.clamp_(0.0, 1.0)
|
| 241 |
+
return tensor.mul_(2.0).sub_(1.0)
|
| 242 |
+
|
| 243 |
+
|
| 244 |
+
def _nvdec_video_tensor(
|
| 245 |
+
path: Path,
|
| 246 |
+
num_frames: int,
|
| 247 |
+
*,
|
| 248 |
+
crop: Crop | None,
|
| 249 |
+
device: str,
|
| 250 |
+
):
|
| 251 |
+
"""Decode one complete skeleton video with TorchCodec's NVDEC backend."""
|
| 252 |
+
|
| 253 |
+
import av # noqa: F401
|
| 254 |
+
import torch
|
| 255 |
+
from torchcodec.decoders import VideoDecoder
|
| 256 |
+
|
| 257 |
+
decoder = VideoDecoder(
|
| 258 |
+
path,
|
| 259 |
+
dimension_order="NCHW",
|
| 260 |
+
device=device,
|
| 261 |
+
output_dtype=torch.uint8,
|
| 262 |
+
seek_mode="exact",
|
| 263 |
+
)
|
| 264 |
+
frames = decoder.get_frames_in_range(0, num_frames).data
|
| 265 |
+
_assert_nvdec_engaged(
|
| 266 |
+
decoder,
|
| 267 |
+
decoded_device_type=frames.device.type,
|
| 268 |
+
declared_frames=len(decoder),
|
| 269 |
+
decoded_frames=int(frames.shape[0]),
|
| 270 |
+
expected_frames=num_frames,
|
| 271 |
+
)
|
| 272 |
+
tensor = _crop_and_normalize(frames.float().div_(255.0), crop)
|
| 273 |
+
return tensor.permute(1, 0, 2, 3).unsqueeze(0).contiguous()
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
def _video_tensor(
|
| 277 |
+
path: Path,
|
| 278 |
+
num_frames: int,
|
| 279 |
+
*,
|
| 280 |
+
crop: Crop | None = None,
|
| 281 |
+
decoder: str = "pyav",
|
| 282 |
+
device: str = "cuda",
|
| 283 |
+
):
|
| 284 |
+
if decoder == "torchcodec_cuda":
|
| 285 |
+
if not device.startswith("cuda"):
|
| 286 |
+
raise FourDAnyoneError(
|
| 287 |
+
f"TorchCodec skeleton decoding requires a CUDA device, got {device!r}."
|
| 288 |
+
)
|
| 289 |
+
return _nvdec_video_tensor(
|
| 290 |
+
path,
|
| 291 |
+
num_frames,
|
| 292 |
+
crop=crop,
|
| 293 |
+
device=device,
|
| 294 |
+
)
|
| 295 |
+
if decoder != "pyav":
|
| 296 |
+
raise FourDAnyoneError(f"Unknown skeleton video decoder {decoder!r}.")
|
| 297 |
+
|
| 298 |
+
import torch
|
| 299 |
+
|
| 300 |
+
output_frames = []
|
| 301 |
+
for frame in iter_rgb_video(path):
|
| 302 |
+
tensor = torch.from_numpy(frame).permute(2, 0, 1).float().div_(255.0)
|
| 303 |
+
output_frames.append(_crop_and_normalize(tensor, crop))
|
| 304 |
+
if len(output_frames) != num_frames:
|
| 305 |
+
raise FourDAnyoneError(f"Video {path} has {len(output_frames)} decoded frames, expected {num_frames}.")
|
| 306 |
+
return torch.stack(output_frames, dim=1).unsqueeze(0).contiguous()
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
def _safe_regressor_metadata(raw_metadata: object, support_shape: tuple[int, ...]) -> dict[str, int | str]:
|
| 310 |
+
del raw_metadata
|
| 311 |
+
return {
|
| 312 |
+
"format": "sparse_vertex_regressor",
|
| 313 |
+
"num_keypoints": int(support_shape[0]),
|
| 314 |
+
"support_vertices_per_keypoint": int(support_shape[1]),
|
| 315 |
+
}
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
def _load_regressor(path: Path, device):
|
| 319 |
+
import torch
|
| 320 |
+
|
| 321 |
+
data = torch.load(path, map_location="cpu", weights_only=True)
|
| 322 |
+
support = data["support_vertex_ids"].detach().long().to(device)
|
| 323 |
+
weights = data["weights"].detach().float().to(device)
|
| 324 |
+
names = tuple(str(value) for value in data["keypoint_names"])
|
| 325 |
+
if support.shape != weights.shape or support.shape[0] != 70:
|
| 326 |
+
raise AssetError(f"Unexpected MHR70 regressor shapes: support={support.shape}, weights={weights.shape}.")
|
| 327 |
+
if names != KEYPOINT_NAMES:
|
| 328 |
+
raise AssetError("MHR70 regressor keypoint order does not match the frozen Goliath70 schema.")
|
| 329 |
+
return support, weights, _safe_regressor_metadata(data.get("metadata"), tuple(support.shape))
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
@contextmanager
|
| 333 |
+
def _gvhmr_geometry_context(gvhmr_root: Path):
|
| 334 |
+
install_pytorch3d_compat()
|
| 335 |
+
with gvhmr_imports(gvhmr_root):
|
| 336 |
+
yield
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
def _body_geometry(
|
| 340 |
+
motion: MotionResult,
|
| 341 |
+
regressor_path: Path,
|
| 342 |
+
gvhmr_root: Path,
|
| 343 |
+
device: str,
|
| 344 |
+
) -> _BodyGeometry:
|
| 345 |
+
import torch
|
| 346 |
+
|
| 347 |
+
body_model = gvhmr_root / "inputs/checkpoints/body_models/smplx/SMPLX_NEUTRAL.npz"
|
| 348 |
+
if not body_model.is_file():
|
| 349 |
+
raise AssetError(
|
| 350 |
+
"The licensed SMPL-X body model is missing. Run `python scripts/download_smplx.py`; "
|
| 351 |
+
f"expected the GVHMR compatibility link at {body_model}."
|
| 352 |
+
)
|
| 353 |
+
utility_root = gvhmr_root / "hmr4d/utils/body_model"
|
| 354 |
+
smplx_to_smpl_path = utility_root / "smplx2smpl_sparse.pt"
|
| 355 |
+
joint_regressor_path = utility_root / "smpl_neutral_J_regressor.pt"
|
| 356 |
+
for path in (smplx_to_smpl_path, joint_regressor_path):
|
| 357 |
+
if not path.is_file():
|
| 358 |
+
raise AssetError(f"GVHMR body-model utility is missing: {path}")
|
| 359 |
+
|
| 360 |
+
torch_device = torch.device(device)
|
| 361 |
+
support, weights, regressor_metadata = _load_regressor(regressor_path, torch_device)
|
| 362 |
+
with _gvhmr_geometry_context(gvhmr_root):
|
| 363 |
+
from hmr4d.utils.geo_transform import apply_T_on_points, compute_T_ayfz2ay
|
| 364 |
+
from hmr4d.utils.smplx_utils import make_smplx
|
| 365 |
+
|
| 366 |
+
smplx = make_smplx("supermotion").to(torch_device).eval()
|
| 367 |
+
global_parameters = {name: value.to(torch_device) for name, value in motion.smpl_params_global.items()}
|
| 368 |
+
incam_parameters = {name: value.to(torch_device) for name, value in motion.smpl_params_incam.items()}
|
| 369 |
+
with torch.inference_mode():
|
| 370 |
+
vertices_global = smplx(**global_parameters).vertices.detach()
|
| 371 |
+
vertices_incam = smplx(**incam_parameters).vertices.detach()
|
| 372 |
+
if tuple(vertices_global.shape[1:]) != (10475, 3) or vertices_incam.shape != vertices_global.shape:
|
| 373 |
+
raise FourDAnyoneError(
|
| 374 |
+
"Expected matching global/incam SMPL-X vertices [frames,10475,3], got "
|
| 375 |
+
f"{tuple(vertices_global.shape)} and {tuple(vertices_incam.shape)}."
|
| 376 |
+
)
|
| 377 |
+
keypoints_global = (vertices_global[:, support] * weights[None, :, :, None]).sum(dim=2)
|
| 378 |
+
keypoints_incam = (vertices_incam[:, support] * weights[None, :, :, None]).sum(dim=2)
|
| 379 |
+
smplx_to_smpl = torch.load(smplx_to_smpl_path, map_location=torch_device, weights_only=True)
|
| 380 |
+
joint_regressor = torch.load(joint_regressor_path, map_location=torch_device, weights_only=True)
|
| 381 |
+
vertices_smpl = torch.stack([torch.matmul(smplx_to_smpl, frame) for frame in vertices_global])
|
| 382 |
+
offset = torch.einsum("jv,vi->ji", joint_regressor, vertices_smpl[0])[0]
|
| 383 |
+
offset = offset.clone()
|
| 384 |
+
offset[1] = vertices_smpl[..., 1].min()
|
| 385 |
+
vertices_offset = vertices_smpl - offset
|
| 386 |
+
first_joints = torch.einsum("jv,lvi->lji", joint_regressor, vertices_offset[[0]])
|
| 387 |
+
transform = compute_T_ayfz2ay(first_joints, inverse=True)
|
| 388 |
+
vertices_world = apply_T_on_points(vertices_offset, transform)
|
| 389 |
+
keypoints_world = apply_T_on_points(keypoints_global - offset, transform)
|
| 390 |
+
joints_world = torch.einsum("jv,lvi->lji", joint_regressor, vertices_world)
|
| 391 |
+
|
| 392 |
+
world_transform = transform[0].detach().clone()
|
| 393 |
+
world_transform[:3, 3] -= world_transform[:3, :3] @ offset
|
| 394 |
+
result = _BodyGeometry(
|
| 395 |
+
vertices_world.detach().cpu().numpy().astype(np.float32),
|
| 396 |
+
joints_world.detach().cpu().numpy().astype(np.float32),
|
| 397 |
+
keypoints_world.detach().cpu().numpy().astype(np.float32),
|
| 398 |
+
keypoints_incam.detach().cpu().numpy().astype(np.float32),
|
| 399 |
+
world_transform.detach().cpu().numpy().astype(np.float64),
|
| 400 |
+
regressor_metadata,
|
| 401 |
+
)
|
| 402 |
+
del smplx, vertices_global, vertices_incam, vertices_smpl, vertices_world, keypoints_global, keypoints_world
|
| 403 |
+
torch.cuda.empty_cache()
|
| 404 |
+
return result
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
def _front_direction(joints: np.ndarray) -> np.ndarray:
|
| 408 |
+
first = joints[0]
|
| 409 |
+
left = first[1, [0, 2]] - first[2, [0, 2]] + first[16, [0, 2]] - first[17, [0, 2]]
|
| 410 |
+
norm = float(np.linalg.norm(left))
|
| 411 |
+
if norm <= 1e-8:
|
| 412 |
+
return np.array([0.0, 0.0, -1.0], dtype=np.float64)
|
| 413 |
+
left /= norm
|
| 414 |
+
return np.array([left[1], 0.0, -left[0]], dtype=np.float64)
|
| 415 |
+
|
| 416 |
+
|
| 417 |
+
def _projection_shape(height: int, width: int, max_render_height: int = 1280) -> tuple[int, int]:
|
| 418 |
+
divisor = 2
|
| 419 |
+
while height / divisor > max_render_height:
|
| 420 |
+
divisor += 1
|
| 421 |
+
return height // divisor, width // divisor
|
| 422 |
+
|
| 423 |
+
|
| 424 |
+
def _output_skeleton_shape(height: int, width: int) -> tuple[int, int]:
|
| 425 |
+
while max(height, width) > INFERENCE.skeleton_max_dimension:
|
| 426 |
+
height //= 2
|
| 427 |
+
width //= 2
|
| 428 |
+
return max(2, height - height % 2), max(2, width - width % 2)
|
| 429 |
+
|
| 430 |
+
|
| 431 |
+
def _scale_crop(crop: Crop, *, scale_y: float, scale_x: float) -> Crop:
|
| 432 |
+
original_height = max(1, round(crop.original_height * scale_y))
|
| 433 |
+
original_width = max(1, round(crop.original_width * scale_x))
|
| 434 |
+
top = min(max(0, round(crop.top * scale_y)), original_height - 1)
|
| 435 |
+
left = min(max(0, round(crop.left * scale_x)), original_width - 1)
|
| 436 |
+
height = min(max(1, round(crop.height * scale_y)), original_height - top)
|
| 437 |
+
width = min(max(1, round(crop.width * scale_x)), original_width - left)
|
| 438 |
+
return Crop(
|
| 439 |
+
top,
|
| 440 |
+
left,
|
| 441 |
+
height,
|
| 442 |
+
width,
|
| 443 |
+
original_height,
|
| 444 |
+
original_width,
|
| 445 |
+
crop.output_height,
|
| 446 |
+
crop.output_width,
|
| 447 |
+
)
|
| 448 |
+
|
| 449 |
+
|
| 450 |
+
def _cropped_camera(camera: Camera, crop: Crop) -> Camera:
|
| 451 |
+
intrinsic = transform_intrinsics(np.asarray(camera.K), crop)
|
| 452 |
+
return replace(
|
| 453 |
+
camera,
|
| 454 |
+
K=tuple(tuple(float(value) for value in row) for row in intrinsic),
|
| 455 |
+
image_width=crop.output_width,
|
| 456 |
+
image_height=crop.output_height,
|
| 457 |
+
)
|
| 458 |
+
|
| 459 |
+
|
| 460 |
+
def _resized_camera(camera: Camera, image_height: int, image_width: int) -> Camera:
|
| 461 |
+
intrinsic = np.asarray(camera.K, dtype=np.float64).copy()
|
| 462 |
+
intrinsic[0] *= image_width / camera.image_width
|
| 463 |
+
intrinsic[1] *= image_height / camera.image_height
|
| 464 |
+
intrinsic[2, 2] = 1.0
|
| 465 |
+
return replace(
|
| 466 |
+
camera,
|
| 467 |
+
K=tuple(tuple(float(value) for value in row) for row in intrinsic),
|
| 468 |
+
image_width=image_width,
|
| 469 |
+
image_height=image_height,
|
| 470 |
+
)
|
| 471 |
+
|
| 472 |
+
|
| 473 |
+
def render_skeleton_video(
|
| 474 |
+
*,
|
| 475 |
+
keypoints_world_fkc: np.ndarray,
|
| 476 |
+
camera: Camera,
|
| 477 |
+
output_path: str | Path,
|
| 478 |
+
num_frames: int,
|
| 479 |
+
canvas_height: int,
|
| 480 |
+
canvas_width: int,
|
| 481 |
+
output_height: int,
|
| 482 |
+
output_width: int,
|
| 483 |
+
body_height: float,
|
| 484 |
+
focal_pixels: float,
|
| 485 |
+
fps: Fraction,
|
| 486 |
+
crf: int,
|
| 487 |
+
preset: str,
|
| 488 |
+
) -> Path:
|
| 489 |
+
"""Project and render one prepared CPU-only skeleton video.
|
| 490 |
+
|
| 491 |
+
Args:
|
| 492 |
+
keypoints_world_fkc: Float32 world keypoints shaped ``[frames, keypoints, 3]``.
|
| 493 |
+
camera: Frozen camera used for every frame.
|
| 494 |
+
output_path: MP4 destination.
|
| 495 |
+
num_frames: Required frame count.
|
| 496 |
+
canvas_height: Projection-space image height.
|
| 497 |
+
canvas_width: Projection-space image width.
|
| 498 |
+
output_height: Encoded skeleton height.
|
| 499 |
+
output_width: Encoded skeleton width.
|
| 500 |
+
body_height: Robust 3D body height in metres.
|
| 501 |
+
focal_pixels: Geometric-mean focal length in projection pixels.
|
| 502 |
+
fps: Exact encoded frame rate.
|
| 503 |
+
crf: H.264 constant-rate-factor value.
|
| 504 |
+
preset: libx264 speed preset.
|
| 505 |
+
|
| 506 |
+
Returns:
|
| 507 |
+
The rendered MP4 path.
|
| 508 |
+
"""
|
| 509 |
+
|
| 510 |
+
if keypoints_world_fkc.shape[0] != num_frames:
|
| 511 |
+
raise FourDAnyoneError(
|
| 512 |
+
f"Prepared target render has {keypoints_world_fkc.shape[0]} frames, expected {num_frames}."
|
| 513 |
+
)
|
| 514 |
+
projected_frames_fkc: list[np.ndarray] = []
|
| 515 |
+
depth_frames_fk: list[np.ndarray] = []
|
| 516 |
+
frame_points_kc: np.ndarray
|
| 517 |
+
for frame_points_kc in keypoints_world_fkc:
|
| 518 |
+
projected_kc: np.ndarray
|
| 519 |
+
depth_k: np.ndarray
|
| 520 |
+
projected_kc, depth_k, _ = project_points(frame_points_kc, camera)
|
| 521 |
+
projected_frames_fkc.append(projected_kc)
|
| 522 |
+
depth_frames_fk.append(depth_k)
|
| 523 |
+
keypoints_2d_fkc: np.ndarray = np.stack(projected_frames_fkc)
|
| 524 |
+
keypoint_depths_fk: np.ndarray = np.stack(depth_frames_fk)
|
| 525 |
+
body_scales_f: np.ndarray = projected_body_scales(
|
| 526 |
+
keypoint_depths_fk,
|
| 527 |
+
KEYPOINT_NAMES,
|
| 528 |
+
body_height,
|
| 529 |
+
focal_pixels,
|
| 530 |
+
)
|
| 531 |
+
|
| 532 |
+
def rendered_frames() -> Iterable[np.ndarray]:
|
| 533 |
+
scores_k: np.ndarray = np.ones(len(KEYPOINT_NAMES), dtype=np.float32)
|
| 534 |
+
frame_index: int
|
| 535 |
+
for frame_index in range(num_frames):
|
| 536 |
+
yield render_goliath40(
|
| 537 |
+
keypoints_2d_fkc[frame_index],
|
| 538 |
+
keypoint_depths_fk[frame_index],
|
| 539 |
+
scores_k,
|
| 540 |
+
canvas_height=canvas_height,
|
| 541 |
+
canvas_width=canvas_width,
|
| 542 |
+
output_height=output_height,
|
| 543 |
+
output_width=output_width,
|
| 544 |
+
body_scale_px=float(body_scales_f[frame_index]),
|
| 545 |
+
)
|
| 546 |
+
|
| 547 |
+
return write_video(
|
| 548 |
+
iter(rendered_frames()),
|
| 549 |
+
output_path,
|
| 550 |
+
fps,
|
| 551 |
+
crf=crf,
|
| 552 |
+
preset=preset,
|
| 553 |
+
)
|
| 554 |
+
|
| 555 |
+
|
| 556 |
+
def build_skeleton_conditioning(
|
| 557 |
+
*,
|
| 558 |
+
motion: MotionResult,
|
| 559 |
+
clip: CanonicalClip,
|
| 560 |
+
regressor_path: str | Path,
|
| 561 |
+
foreground_model_path: str | Path,
|
| 562 |
+
gvhmr_root: str | Path,
|
| 563 |
+
output_dir: str | Path,
|
| 564 |
+
device: str,
|
| 565 |
+
view_plan: ViewPlan,
|
| 566 |
+
defer_target_skeletons: bool = False,
|
| 567 |
+
) -> Conditioning:
|
| 568 |
+
"""Build source, RCP, and target conditioning on one camera grid."""
|
| 569 |
+
|
| 570 |
+
gvhmr_root, _ = validate_gvhmr(gvhmr_root)
|
| 571 |
+
root = Path(output_dir).expanduser().resolve()
|
| 572 |
+
root.mkdir(parents=True, exist_ok=True)
|
| 573 |
+
|
| 574 |
+
LOGGER.info("Estimating source foreground masks with BiRefNet")
|
| 575 |
+
masks = predict_foreground_masks(clip.rgb_frames, foreground_model_path, device)
|
| 576 |
+
geometry = _body_geometry(motion, Path(regressor_path), gvhmr_root, device)
|
| 577 |
+
input_framing = analyze_input_framing(
|
| 578 |
+
geometry.keypoints_incam,
|
| 579 |
+
KEYPOINT_NAMES,
|
| 580 |
+
motion.K_fullimg.detach().cpu().numpy(),
|
| 581 |
+
motion.observed_keypoints_2d.detach().cpu().numpy(),
|
| 582 |
+
masks,
|
| 583 |
+
)
|
| 584 |
+
|
| 585 |
+
targets = geometry.vertices_world.mean(axis=1)
|
| 586 |
+
targets[:, 1] = 0.0
|
| 587 |
+
center = targets.mean(axis=0)
|
| 588 |
+
front_direction = _front_direction(geometry.joints_world)
|
| 589 |
+
reference_K = reference_intrinsics(clip.height, clip.width)
|
| 590 |
+
framing_pitches = tuple(dict.fromkeys((*view_plan.layer_pitches, int(CAMERA.pitch_degrees))))
|
| 591 |
+
|
| 592 |
+
def camera_factory(radius: float, target_height: float) -> tuple[Camera, ...]:
|
| 593 |
+
# A full safety ring at every requested pitch keeps framing independent
|
| 594 |
+
# of view density while covering all target and canonical RCP cameras.
|
| 595 |
+
candidate_center = center.copy()
|
| 596 |
+
candidate_center[1] = target_height
|
| 597 |
+
return tuple(
|
| 598 |
+
camera
|
| 599 |
+
for layer_index, pitch in enumerate(framing_pitches)
|
| 600 |
+
for camera in camera_ring(
|
| 601 |
+
center=candidate_center,
|
| 602 |
+
front_direction=front_direction,
|
| 603 |
+
K=reference_K,
|
| 604 |
+
image_height=clip.height,
|
| 605 |
+
image_width=clip.width,
|
| 606 |
+
radius=radius,
|
| 607 |
+
target_height=target_height,
|
| 608 |
+
spec=CameraConfig(count=CAMERA.count, pitch_degrees=float(pitch)),
|
| 609 |
+
layer_index=layer_index,
|
| 610 |
+
camera_id_offset=layer_index * CAMERA.count,
|
| 611 |
+
)
|
| 612 |
+
)
|
| 613 |
+
|
| 614 |
+
projection_height, projection_width = _projection_shape(clip.height, clip.width)
|
| 615 |
+
framing = solve_sequence_framing(
|
| 616 |
+
geometry.keypoints_world,
|
| 617 |
+
KEYPOINT_NAMES,
|
| 618 |
+
input_framing,
|
| 619 |
+
camera_factory,
|
| 620 |
+
projection_width / projection_height,
|
| 621 |
+
)
|
| 622 |
+
LOGGER.info(
|
| 623 |
+
"Adaptive framing: input=%s confidence=%.3f radius=%.3f target_height=%.3f f/H=%.3f",
|
| 624 |
+
input_framing.label,
|
| 625 |
+
input_framing.confidence,
|
| 626 |
+
framing.radius,
|
| 627 |
+
framing.target_height,
|
| 628 |
+
framing.focal_normalized,
|
| 629 |
+
)
|
| 630 |
+
center[1] = framing.target_height
|
| 631 |
+
raw_intrinsic = reference_intrinsics(
|
| 632 |
+
clip.height,
|
| 633 |
+
clip.width,
|
| 634 |
+
focal_normalized=framing.focal_normalized,
|
| 635 |
+
)
|
| 636 |
+
raw_target_cameras = camera_grid(
|
| 637 |
+
center=center,
|
| 638 |
+
front_direction=front_direction,
|
| 639 |
+
K=raw_intrinsic,
|
| 640 |
+
image_height=clip.height,
|
| 641 |
+
image_width=clip.width,
|
| 642 |
+
radius=framing.radius,
|
| 643 |
+
target_height=framing.target_height,
|
| 644 |
+
views_per_layer=view_plan.views_per_layer,
|
| 645 |
+
layer_pitches=view_plan.layer_pitches,
|
| 646 |
+
start_yaw=view_plan.start_yaw,
|
| 647 |
+
yaw_span=view_plan.yaw_span,
|
| 648 |
+
)
|
| 649 |
+
|
| 650 |
+
source_crop = crop_from_bounds(
|
| 651 |
+
bounds=mask_bounds(masks, CROP.mask_threshold),
|
| 652 |
+
image_height=clip.height,
|
| 653 |
+
image_width=clip.width,
|
| 654 |
+
output_height=INFERENCE.height,
|
| 655 |
+
output_width=INFERENCE.width,
|
| 656 |
+
margins=CROP.margins,
|
| 657 |
+
allow_upscale=CROP.allow_upscale,
|
| 658 |
+
)
|
| 659 |
+
target_crop = center_crop(clip.height, clip.width, INFERENCE.height, INFERENCE.width)
|
| 660 |
+
source_video = write_lossless_video(clip, root / "source.mkv")
|
| 661 |
+
skeleton_root = root / "goliath40"
|
| 662 |
+
skeleton_root.mkdir()
|
| 663 |
+
skeleton_height, skeleton_width = _output_skeleton_shape(clip.height, clip.width)
|
| 664 |
+
body_height = estimate_body_height(geometry.keypoints_world, KEYPOINT_NAMES)
|
| 665 |
+
keypoints_path: Path = root / "keypoints_3d.npy"
|
| 666 |
+
np.save(keypoints_path, geometry.keypoints_world)
|
| 667 |
+
focal_pixels: float = float(np.sqrt(raw_intrinsic[0, 0] * raw_intrinsic[1, 1]))
|
| 668 |
+
|
| 669 |
+
def render_skeleton(camera: Camera, path: Path) -> Path:
|
| 670 |
+
return render_skeleton_video(
|
| 671 |
+
keypoints_world_fkc=geometry.keypoints_world,
|
| 672 |
+
camera=camera,
|
| 673 |
+
output_path=path,
|
| 674 |
+
num_frames=motion.num_frames,
|
| 675 |
+
canvas_height=clip.height,
|
| 676 |
+
canvas_width=clip.width,
|
| 677 |
+
output_height=skeleton_height,
|
| 678 |
+
output_width=skeleton_width,
|
| 679 |
+
body_height=body_height,
|
| 680 |
+
focal_pixels=focal_pixels,
|
| 681 |
+
fps=clip.fps,
|
| 682 |
+
crf=INFERENCE.skeleton_h264_crf,
|
| 683 |
+
preset=INFERENCE.h264_preset,
|
| 684 |
+
)
|
| 685 |
+
|
| 686 |
+
target_paths: tuple[Path, ...] = tuple(
|
| 687 |
+
skeleton_root / f"{camera.camera_id:02d}.mp4" for camera in raw_target_cameras
|
| 688 |
+
)
|
| 689 |
+
cropped_target_cameras: tuple[Camera, ...] = tuple(
|
| 690 |
+
_cropped_camera(camera, target_crop) for camera in raw_target_cameras
|
| 691 |
+
)
|
| 692 |
+
|
| 693 |
+
if not view_plan.enable_rcp:
|
| 694 |
+
raw_rcp_cameras: tuple[Camera, ...] = ()
|
| 695 |
+
cropped_rcp_cameras: tuple[Camera, ...] = ()
|
| 696 |
+
elif view_plan.is_canonical_target_ring:
|
| 697 |
+
raw_rcp_cameras = tuple(raw_target_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids)
|
| 698 |
+
cropped_rcp_cameras = tuple(cropped_target_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids)
|
| 699 |
+
else:
|
| 700 |
+
canonical_cameras: tuple[Camera, ...] = camera_ring(
|
| 701 |
+
center=center,
|
| 702 |
+
front_direction=front_direction,
|
| 703 |
+
K=raw_intrinsic,
|
| 704 |
+
image_height=clip.height,
|
| 705 |
+
image_width=clip.width,
|
| 706 |
+
radius=framing.radius,
|
| 707 |
+
target_height=framing.target_height,
|
| 708 |
+
layer_index=-1,
|
| 709 |
+
)
|
| 710 |
+
raw_rcp_cameras = tuple(canonical_cameras[camera_id] for camera_id in view_plan.rcp_camera_ids)
|
| 711 |
+
cropped_rcp_cameras = tuple(_cropped_camera(camera, target_crop) for camera in raw_rcp_cameras)
|
| 712 |
+
|
| 713 |
+
if defer_target_skeletons:
|
| 714 |
+
if raw_rcp_cameras:
|
| 715 |
+
rcp_root: Path = root / "rcp_goliath40"
|
| 716 |
+
rcp_root.mkdir()
|
| 717 |
+
rcp_paths: tuple[Path, ...] = tuple(
|
| 718 |
+
render_skeleton(camera, rcp_root / f"{camera.camera_id:02d}.mp4")
|
| 719 |
+
for camera in raw_rcp_cameras
|
| 720 |
+
)
|
| 721 |
+
else:
|
| 722 |
+
rcp_paths = ()
|
| 723 |
+
else:
|
| 724 |
+
target_paths = tuple(
|
| 725 |
+
render_skeleton(camera, path)
|
| 726 |
+
for camera, path in zip(raw_target_cameras, target_paths, strict=True)
|
| 727 |
+
)
|
| 728 |
+
if not raw_rcp_cameras:
|
| 729 |
+
rcp_paths = ()
|
| 730 |
+
elif view_plan.is_canonical_target_ring:
|
| 731 |
+
rcp_paths = tuple(target_paths[camera_id] for camera_id in view_plan.rcp_camera_ids)
|
| 732 |
+
else:
|
| 733 |
+
rcp_root = root / "rcp_goliath40"
|
| 734 |
+
rcp_root.mkdir()
|
| 735 |
+
rcp_paths = tuple(
|
| 736 |
+
render_skeleton(camera, rcp_root / f"{camera.camera_id:02d}.mp4")
|
| 737 |
+
for camera in raw_rcp_cameras
|
| 738 |
+
)
|
| 739 |
+
|
| 740 |
+
def camera_records(
|
| 741 |
+
raw_cameras: tuple[Camera, ...],
|
| 742 |
+
cropped_cameras: tuple[Camera, ...],
|
| 743 |
+
paths: tuple[Path, ...],
|
| 744 |
+
) -> list[dict]:
|
| 745 |
+
return [
|
| 746 |
+
{
|
| 747 |
+
**camera.to_dict(),
|
| 748 |
+
"crop": asdict(target_crop),
|
| 749 |
+
"raw_camera": raw.to_dict(),
|
| 750 |
+
"skeleton_camera": _resized_camera(raw, skeleton_height, skeleton_width).to_dict(),
|
| 751 |
+
"skeleton_video": path.relative_to(root).as_posix(),
|
| 752 |
+
}
|
| 753 |
+
for raw, camera, path in zip(raw_cameras, cropped_cameras, paths, strict=True)
|
| 754 |
+
]
|
| 755 |
+
|
| 756 |
+
framing_payload = framing.to_dict()
|
| 757 |
+
camera_payload = {
|
| 758 |
+
"camera_model": "OPENCV",
|
| 759 |
+
"world_frame": WORLD_FRAME,
|
| 760 |
+
"camera_frame": CAMERA_FRAME,
|
| 761 |
+
"front_camera_ids": list(view_plan.front_camera_ids),
|
| 762 |
+
"motion_world": motion.motion_world,
|
| 763 |
+
"motion_world_to_canonical_world": geometry.motion_world_to_canonical_world.tolist(),
|
| 764 |
+
"ring_center": center.tolist(),
|
| 765 |
+
"framing": framing_payload,
|
| 766 |
+
"cameras": camera_records(raw_target_cameras, cropped_target_cameras, target_paths),
|
| 767 |
+
"rcp_cameras": camera_records(raw_rcp_cameras, cropped_rcp_cameras, rcp_paths),
|
| 768 |
+
}
|
| 769 |
+
write_json(root / "cameras.json", camera_payload)
|
| 770 |
+
write_json(
|
| 771 |
+
root / "metadata.json",
|
| 772 |
+
{
|
| 773 |
+
"view_plan": view_plan.to_dict(),
|
| 774 |
+
"num_frames": motion.num_frames,
|
| 775 |
+
"fps_num": clip.fps_num,
|
| 776 |
+
"fps_den": clip.fps_den,
|
| 777 |
+
"source_video": source_video.name,
|
| 778 |
+
"source_crop": asdict(source_crop),
|
| 779 |
+
"source_crop_policy": {
|
| 780 |
+
"subject": "fmask",
|
| 781 |
+
"threshold": CROP.mask_threshold,
|
| 782 |
+
"margins": CROP.margins,
|
| 783 |
+
"allow_upscale": CROP.allow_upscale,
|
| 784 |
+
},
|
| 785 |
+
"foreground_model": {
|
| 786 |
+
"repo_id": BIREFNET_REPO_ID,
|
| 787 |
+
"revision": BIREFNET_REVISION,
|
| 788 |
+
"image_size": FOREGROUND.image_size,
|
| 789 |
+
"batch_size": FOREGROUND.batch_size,
|
| 790 |
+
},
|
| 791 |
+
"framing": framing_payload,
|
| 792 |
+
"regressor_metadata": geometry.regressor_metadata,
|
| 793 |
+
"keypoint_names": KEYPOINT_NAMES,
|
| 794 |
+
"visible_keypoint_set": "goliath40",
|
| 795 |
+
"skeleton_codec_boundary": f"libx264_crf{INFERENCE.skeleton_h264_crf}",
|
| 796 |
+
"skeleton_canvas": {"height": skeleton_height, "width": skeleton_width},
|
| 797 |
+
"skeleton_draw_scale": {
|
| 798 |
+
"mode": "kp3d",
|
| 799 |
+
"body_height_3d": body_height,
|
| 800 |
+
"body_reference_px": SKELETON.draw_body_reference_px,
|
| 801 |
+
},
|
| 802 |
+
"target_render_deferred": defer_target_skeletons,
|
| 803 |
+
},
|
| 804 |
+
)
|
| 805 |
+
if defer_target_skeletons:
|
| 806 |
+
write_json(
|
| 807 |
+
root / "target-render-request.json",
|
| 808 |
+
{
|
| 809 |
+
"schema_version": 1,
|
| 810 |
+
"keypoints_3d_path": str(keypoints_path),
|
| 811 |
+
"targets": [
|
| 812 |
+
{"camera": camera.to_dict(), "output_path": str(path)}
|
| 813 |
+
for camera, path in zip(raw_target_cameras, target_paths, strict=True)
|
| 814 |
+
],
|
| 815 |
+
"num_frames": motion.num_frames,
|
| 816 |
+
"canvas_height": clip.height,
|
| 817 |
+
"canvas_width": clip.width,
|
| 818 |
+
"output_height": skeleton_height,
|
| 819 |
+
"output_width": skeleton_width,
|
| 820 |
+
"body_height": body_height,
|
| 821 |
+
"focal_pixels": focal_pixels,
|
| 822 |
+
"fps_num": clip.fps_num,
|
| 823 |
+
"fps_den": clip.fps_den,
|
| 824 |
+
"crf": INFERENCE.skeleton_h264_crf,
|
| 825 |
+
"preset": INFERENCE.h264_preset,
|
| 826 |
+
"done_path": str(root / "target-render.done.json"),
|
| 827 |
+
},
|
| 828 |
+
)
|
| 829 |
+
return Conditioning(
|
| 830 |
+
root=root,
|
| 831 |
+
source_video=source_video,
|
| 832 |
+
source_crop=source_crop,
|
| 833 |
+
target_skeletons=tuple(SkeletonVideo(path, target_crop) for path in target_paths),
|
| 834 |
+
rcp_skeletons=tuple(SkeletonVideo(path, target_crop) for path in rcp_paths),
|
| 835 |
+
view_plan=view_plan,
|
| 836 |
+
fps_num=clip.fps_num,
|
| 837 |
+
fps_den=clip.fps_den,
|
| 838 |
+
num_frames=motion.num_frames,
|
| 839 |
+
)
|
|
@@ -0,0 +1,100 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""CPU-only target skeleton renderer for conditioning overlap."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import os
|
| 7 |
+
import sys
|
| 8 |
+
import time
|
| 9 |
+
from fractions import Fraction
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
|
| 14 |
+
from fdanyone.errors import FourDAnyoneError
|
| 15 |
+
from fdanyone.geometry.cameras import Camera
|
| 16 |
+
from fdanyone.io import write_json
|
| 17 |
+
from fdanyone.skeleton.pipeline import render_skeleton_video
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def _camera_from_payload(payload: object) -> Camera:
|
| 21 |
+
"""Decode a camera record from the conditioning file protocol."""
|
| 22 |
+
|
| 23 |
+
if not isinstance(payload, dict):
|
| 24 |
+
raise FourDAnyoneError("Target render camera must be a JSON object.")
|
| 25 |
+
return Camera.from_dict(payload)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def render_target_skeletons(request_path: str | Path) -> tuple[Path, ...]:
|
| 29 |
+
"""Render every target in one prepared, CUDA-hidden file request."""
|
| 30 |
+
|
| 31 |
+
started: float = time.monotonic()
|
| 32 |
+
request_file: Path = Path(request_path).expanduser().resolve()
|
| 33 |
+
request: object = json.loads(request_file.read_text())
|
| 34 |
+
if not isinstance(request, dict) or int(request.get("schema_version", 0)) != 1:
|
| 35 |
+
raise FourDAnyoneError("Unsupported target skeleton render request.")
|
| 36 |
+
cuda_visible_devices: str | None = os.environ.get("CUDA_VISIBLE_DEVICES")
|
| 37 |
+
if cuda_visible_devices not in ("", "-1"):
|
| 38 |
+
raise FourDAnyoneError(
|
| 39 |
+
"Target skeleton rendering must run with CUDA_VISIBLE_DEVICES disabled."
|
| 40 |
+
)
|
| 41 |
+
|
| 42 |
+
keypoints_path: Path = Path(str(request["keypoints_3d_path"])).expanduser().resolve()
|
| 43 |
+
keypoints_world_fkc: np.ndarray = np.load(keypoints_path, allow_pickle=False)
|
| 44 |
+
if keypoints_world_fkc.dtype != np.float32 or keypoints_world_fkc.ndim != 3:
|
| 45 |
+
raise FourDAnyoneError(
|
| 46 |
+
"Prepared target keypoints must be float32 [frames,keypoints,3], got "
|
| 47 |
+
f"dtype={keypoints_world_fkc.dtype}, shape={keypoints_world_fkc.shape}."
|
| 48 |
+
)
|
| 49 |
+
targets: object = request.get("targets")
|
| 50 |
+
if not isinstance(targets, list) or not targets:
|
| 51 |
+
raise FourDAnyoneError("Target skeleton render request has no targets.")
|
| 52 |
+
|
| 53 |
+
rendered_paths: list[Path] = []
|
| 54 |
+
target: object
|
| 55 |
+
for target in targets:
|
| 56 |
+
if not isinstance(target, dict):
|
| 57 |
+
raise FourDAnyoneError("Target skeleton entry must be a JSON object.")
|
| 58 |
+
camera: Camera = _camera_from_payload(target["camera"])
|
| 59 |
+
output_path: Path = Path(str(target["output_path"])).expanduser().resolve()
|
| 60 |
+
rendered_path: Path = render_skeleton_video(
|
| 61 |
+
keypoints_world_fkc=keypoints_world_fkc,
|
| 62 |
+
camera=camera,
|
| 63 |
+
output_path=output_path,
|
| 64 |
+
num_frames=int(request["num_frames"]),
|
| 65 |
+
canvas_height=int(request["canvas_height"]),
|
| 66 |
+
canvas_width=int(request["canvas_width"]),
|
| 67 |
+
output_height=int(request["output_height"]),
|
| 68 |
+
output_width=int(request["output_width"]),
|
| 69 |
+
body_height=float(request["body_height"]),
|
| 70 |
+
focal_pixels=float(request["focal_pixels"]),
|
| 71 |
+
fps=Fraction(int(request["fps_num"]), int(request["fps_den"])),
|
| 72 |
+
crf=int(request["crf"]),
|
| 73 |
+
preset=str(request["preset"]),
|
| 74 |
+
)
|
| 75 |
+
rendered_paths.append(rendered_path)
|
| 76 |
+
|
| 77 |
+
done_path: Path = Path(str(request["done_path"])).expanduser().resolve()
|
| 78 |
+
write_json(
|
| 79 |
+
done_path,
|
| 80 |
+
{
|
| 81 |
+
"schema_version": 1,
|
| 82 |
+
"rendered_views": len(rendered_paths),
|
| 83 |
+
"elapsed_seconds": time.monotonic() - started,
|
| 84 |
+
"cuda_visible_devices": cuda_visible_devices,
|
| 85 |
+
"outputs": [str(path) for path in rendered_paths],
|
| 86 |
+
},
|
| 87 |
+
)
|
| 88 |
+
return tuple(rendered_paths)
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def main(request_path: str) -> None:
|
| 92 |
+
"""Run one target skeleton request."""
|
| 93 |
+
|
| 94 |
+
render_target_skeletons(request_path)
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
if __name__ == "__main__":
|
| 98 |
+
if len(sys.argv) != 2:
|
| 99 |
+
raise SystemExit("Usage: python -m fdanyone.skeleton.render_worker REQUEST.json")
|
| 100 |
+
main(sys.argv[1])
|
|
@@ -0,0 +1,230 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Dependency-light, depth-aware Goliath40 rasterization."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import math
|
| 6 |
+
|
| 7 |
+
import cv2
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
from fdanyone.config import INFERENCE, SKELETON
|
| 11 |
+
from fdanyone.skeleton.keypoints import EXTRA_KEYPOINT_IDS, LINKS, VISIBLE_KEYPOINT_IDS, keypoint_color
|
| 12 |
+
|
| 13 |
+
_BODY_SCALE_SEGMENTS = (
|
| 14 |
+
(("left-shoulder",), ("right-shoulder",), 0.26),
|
| 15 |
+
(("left-hip",), ("right-hip",), 0.18),
|
| 16 |
+
(("left-shoulder", "right-shoulder"), ("left-hip", "right-hip"), 0.32),
|
| 17 |
+
(("left-shoulder",), ("left-hip",), 0.32),
|
| 18 |
+
(("right-shoulder",), ("right-hip",), 0.32),
|
| 19 |
+
(("left-shoulder",), ("left-elbow",), 0.19),
|
| 20 |
+
(("right-shoulder",), ("right-elbow",), 0.19),
|
| 21 |
+
(("left-elbow",), ("left-wrist",), 0.16),
|
| 22 |
+
(("right-elbow",), ("right-wrist",), 0.16),
|
| 23 |
+
(("left-hip",), ("left-knee",), 0.245),
|
| 24 |
+
(("right-hip",), ("right-knee",), 0.245),
|
| 25 |
+
(("left-knee",), ("left-ankle",), 0.245),
|
| 26 |
+
(("right-knee",), ("right-ankle",), 0.245),
|
| 27 |
+
)
|
| 28 |
+
_BODY_CENTER_NAMES = ("left-shoulder", "right-shoulder", "left-hip", "right-hip")
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def _name_index(names) -> dict[str, int]:
|
| 32 |
+
normalized = [str(name).strip().lower().replace("_", "-") for name in names]
|
| 33 |
+
mapping = dict(zip(normalized, range(len(normalized)), strict=True))
|
| 34 |
+
if len(mapping) != len(normalized):
|
| 35 |
+
raise ValueError("Keypoint names must be unique after normalization.")
|
| 36 |
+
return mapping
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def estimate_body_height(keypoints_3d: np.ndarray, names) -> float:
|
| 40 |
+
"""Robustly infer physical body height from stable anatomical segments."""
|
| 41 |
+
|
| 42 |
+
points = np.asarray(keypoints_3d, dtype=np.float64)
|
| 43 |
+
if points.ndim != 3 or points.shape[1:] != (len(names), 3):
|
| 44 |
+
raise ValueError(f"Expected keypoints [frames,{len(names)},3], got {points.shape}.")
|
| 45 |
+
mapping = _name_index(names)
|
| 46 |
+
required = {name for first, second, _ in _BODY_SCALE_SEGMENTS for group in (first, second) for name in group}
|
| 47 |
+
missing = sorted(required - mapping.keys())
|
| 48 |
+
if missing:
|
| 49 |
+
raise ValueError(f"Missing body-scale keypoints: {missing}.")
|
| 50 |
+
|
| 51 |
+
frame_heights = []
|
| 52 |
+
for frame in points:
|
| 53 |
+
candidates = []
|
| 54 |
+
for first, second, height_ratio in _BODY_SCALE_SEGMENTS:
|
| 55 |
+
point_a = frame[[mapping[name] for name in first]].mean(axis=0)
|
| 56 |
+
point_b = frame[[mapping[name] for name in second]].mean(axis=0)
|
| 57 |
+
length = float(np.linalg.norm(point_a - point_b))
|
| 58 |
+
if np.isfinite(length) and length > 0:
|
| 59 |
+
candidates.append(length / height_ratio)
|
| 60 |
+
if candidates:
|
| 61 |
+
frame_heights.append(float(np.median(candidates)))
|
| 62 |
+
heights = np.asarray(frame_heights, dtype=np.float64)
|
| 63 |
+
heights = heights[np.isfinite(heights) & (heights > 0)]
|
| 64 |
+
if not heights.size:
|
| 65 |
+
raise ValueError("Cannot estimate body height from the supplied keypoints.")
|
| 66 |
+
if heights.size >= 5:
|
| 67 |
+
lower, upper = np.percentile(heights, [10.0, 90.0])
|
| 68 |
+
trimmed = heights[(heights >= lower) & (heights <= upper)]
|
| 69 |
+
if trimmed.size:
|
| 70 |
+
heights = trimmed
|
| 71 |
+
return float(np.median(heights))
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def projected_body_scales(
|
| 75 |
+
keypoint_depths: np.ndarray,
|
| 76 |
+
names,
|
| 77 |
+
body_height_3d: float,
|
| 78 |
+
focal_px: float,
|
| 79 |
+
) -> np.ndarray:
|
| 80 |
+
"""Convert physical body height and camera depth to per-frame pixel scale."""
|
| 81 |
+
|
| 82 |
+
depths = np.asarray(keypoint_depths, dtype=np.float64)
|
| 83 |
+
if depths.ndim != 2 or depths.shape[1] != len(names):
|
| 84 |
+
raise ValueError(f"Expected keypoint depths [frames,{len(names)}], got {depths.shape}.")
|
| 85 |
+
if not np.isfinite(body_height_3d) or body_height_3d <= 0:
|
| 86 |
+
raise ValueError("body_height_3d must be positive and finite.")
|
| 87 |
+
if not np.isfinite(focal_px) or focal_px <= 0:
|
| 88 |
+
raise ValueError("focal_px must be positive and finite.")
|
| 89 |
+
mapping = _name_index(names)
|
| 90 |
+
missing = [name for name in _BODY_CENTER_NAMES if name not in mapping]
|
| 91 |
+
if missing:
|
| 92 |
+
raise ValueError(f"Missing body-center keypoints: {missing}.")
|
| 93 |
+
center_depths = depths[:, [mapping[name] for name in _BODY_CENTER_NAMES]]
|
| 94 |
+
valid = np.isfinite(center_depths) & (center_depths > 1e-6)
|
| 95 |
+
counts = valid.sum(axis=1)
|
| 96 |
+
mean_depths = np.full(depths.shape[0], np.nan, dtype=np.float64)
|
| 97 |
+
enough = counts >= 2
|
| 98 |
+
mean_depths[enough] = np.where(valid, center_depths, 0.0).sum(axis=1)[enough] / counts[enough]
|
| 99 |
+
scales = np.full(depths.shape[0], np.nan, dtype=np.float32)
|
| 100 |
+
scales[enough] = (body_height_3d * focal_px / mean_depths[enough]).astype(np.float32)
|
| 101 |
+
return scales
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
def _draw_point(z_buffer, canvas, point, depth, radius, color) -> None:
|
| 105 |
+
if not np.isfinite(depth) or depth <= 0:
|
| 106 |
+
return
|
| 107 |
+
height, width = z_buffer.shape
|
| 108 |
+
x, y = point
|
| 109 |
+
radius = max(1, int(radius))
|
| 110 |
+
x0, x1 = max(0, x - radius), min(width, x + radius + 1)
|
| 111 |
+
y0, y1 = max(0, y - radius), min(height, y + radius + 1)
|
| 112 |
+
if x0 >= x1 or y0 >= y1:
|
| 113 |
+
return
|
| 114 |
+
yy, xx = np.ogrid[y0:y1, x0:x1]
|
| 115 |
+
mask = (xx - x) ** 2 + (yy - y) ** 2 <= radius**2
|
| 116 |
+
view = z_buffer[y0:y1, x0:x1]
|
| 117 |
+
update = mask & (depth <= view)
|
| 118 |
+
view[update] = depth
|
| 119 |
+
canvas[y0:y1, x0:x1][update] = color
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
def _draw_line(z_buffer, canvas, p1, p2, d1, d2, thickness, color) -> None:
|
| 123 |
+
if not np.isfinite(d1) or not np.isfinite(d2) or max(d1, d2) <= 0:
|
| 124 |
+
return
|
| 125 |
+
height, width = z_buffer.shape
|
| 126 |
+
x1, y1 = p1
|
| 127 |
+
x2, y2 = p2
|
| 128 |
+
dx, dy = float(x2 - x1), float(y2 - y1)
|
| 129 |
+
length_sq = dx * dx + dy * dy
|
| 130 |
+
if length_sq < 1e-6:
|
| 131 |
+
_draw_point(z_buffer, canvas, p1, (d1 + d2) / 2.0, max(1, thickness // 2), color)
|
| 132 |
+
return
|
| 133 |
+
radius = max(0.5, float(thickness) / 2.0)
|
| 134 |
+
pad = int(math.ceil(radius)) + 1
|
| 135 |
+
x0, x3 = max(0, min(x1, x2) - pad), min(width, max(x1, x2) + pad + 1)
|
| 136 |
+
y0, y3 = max(0, min(y1, y2) - pad), min(height, max(y1, y2) + pad + 1)
|
| 137 |
+
if x0 >= x3 or y0 >= y3:
|
| 138 |
+
return
|
| 139 |
+
yy, xx = np.ogrid[y0:y3, x0:x3]
|
| 140 |
+
t = np.clip(((xx - x1) * dx + (yy - y1) * dy) / length_sq, 0.0, 1.0)
|
| 141 |
+
mask = (xx - (x1 + t * dx)) ** 2 + (yy - (y1 + t * dy)) ** 2 <= radius**2
|
| 142 |
+
depth = d1 + t * (d2 - d1)
|
| 143 |
+
view = z_buffer[y0:y3, x0:x3]
|
| 144 |
+
update = mask & (depth > 0) & (depth <= view)
|
| 145 |
+
view[update] = depth[update]
|
| 146 |
+
canvas[y0:y3, x0:x3][update] = color
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def render_goliath40(
|
| 150 |
+
keypoints: np.ndarray,
|
| 151 |
+
depths: np.ndarray,
|
| 152 |
+
scores: np.ndarray,
|
| 153 |
+
*,
|
| 154 |
+
canvas_height: int,
|
| 155 |
+
canvas_width: int,
|
| 156 |
+
output_height: int,
|
| 157 |
+
output_width: int,
|
| 158 |
+
score_threshold: float = 0.3,
|
| 159 |
+
body_scale_px: float | None = None,
|
| 160 |
+
) -> np.ndarray:
|
| 161 |
+
"""Render one RGB frame using the reference Sapiens2 sizing rules."""
|
| 162 |
+
|
| 163 |
+
canvas_scale = max(1.0, INFERENCE.skeleton_max_dimension / max(output_height, output_width))
|
| 164 |
+
render_height = int(round(output_height * canvas_scale))
|
| 165 |
+
render_width = int(round(output_width * canvas_scale))
|
| 166 |
+
points = np.asarray(keypoints, dtype=np.float32).copy()
|
| 167 |
+
points[:, 0] *= render_width / canvas_width
|
| 168 |
+
points[:, 1] *= render_height / canvas_height
|
| 169 |
+
depths = np.asarray(depths, dtype=np.float32)
|
| 170 |
+
scores = np.asarray(scores, dtype=np.float32)
|
| 171 |
+
# Keep the reference renderer's BGR draw -> resize -> RGB conversion order.
|
| 172 |
+
# Resizing a channel-permuted uint8 image can differ by one LSB in OpenCV's
|
| 173 |
+
# optimized interpolation kernels, so drawing directly in RGB is not quite
|
| 174 |
+
# byte-exact even though the colors are semantically identical.
|
| 175 |
+
canvas = np.zeros((render_height, render_width, 3), dtype=np.uint8)
|
| 176 |
+
z_buffer = np.full((render_height, render_width), np.inf, dtype=np.float32)
|
| 177 |
+
line_scale = render_height / 1024.0
|
| 178 |
+
if body_scale_px is not None and np.isfinite(body_scale_px) and body_scale_px > 0:
|
| 179 |
+
keypoint_scale = render_height / canvas_height
|
| 180 |
+
line_scale = max(
|
| 181 |
+
0.25,
|
| 182 |
+
float(body_scale_px) * keypoint_scale / SKELETON.draw_body_reference_px,
|
| 183 |
+
)
|
| 184 |
+
base_radius = max(1, int(round(2 * line_scale)))
|
| 185 |
+
base_thickness = max(1, int(round(2 * line_scale)))
|
| 186 |
+
point_items: dict[int, tuple[tuple[int, int], float, int, tuple[int, int, int]]] = {}
|
| 187 |
+
|
| 188 |
+
def remember(index: int, radius: int) -> None:
|
| 189 |
+
if scores[index] < score_threshold or not np.isfinite(points[index]).all():
|
| 190 |
+
return
|
| 191 |
+
item = (
|
| 192 |
+
(int(round(points[index, 0])), int(round(points[index, 1]))),
|
| 193 |
+
float(depths[index]),
|
| 194 |
+
radius,
|
| 195 |
+
keypoint_color(index)[::-1],
|
| 196 |
+
)
|
| 197 |
+
if index not in point_items or radius > point_items[index][2]:
|
| 198 |
+
point_items[index] = item
|
| 199 |
+
|
| 200 |
+
for _, first, second, color, major in LINKS:
|
| 201 |
+
if min(float(scores[first]), float(scores[second])) < score_threshold:
|
| 202 |
+
continue
|
| 203 |
+
if not np.isfinite(points[[first, second]]).all():
|
| 204 |
+
continue
|
| 205 |
+
p1 = (int(round(points[first, 0])), int(round(points[first, 1])))
|
| 206 |
+
p2 = (int(round(points[second, 0])), int(round(points[second, 1])))
|
| 207 |
+
thickness = base_thickness * (2 if major else 1)
|
| 208 |
+
_draw_line(
|
| 209 |
+
z_buffer,
|
| 210 |
+
canvas,
|
| 211 |
+
p1,
|
| 212 |
+
p2,
|
| 213 |
+
float(depths[first]),
|
| 214 |
+
float(depths[second]),
|
| 215 |
+
thickness,
|
| 216 |
+
color[::-1],
|
| 217 |
+
)
|
| 218 |
+
radius = max(1, int(round(base_radius * (1.75 if major else 1.0))))
|
| 219 |
+
remember(first, radius)
|
| 220 |
+
remember(second, radius)
|
| 221 |
+
|
| 222 |
+
for index in EXTRA_KEYPOINT_IDS:
|
| 223 |
+
remember(index, max(2, int(round(base_radius * 1.5))))
|
| 224 |
+
for index, (point, depth, radius, color) in point_items.items():
|
| 225 |
+
if index in VISIBLE_KEYPOINT_IDS:
|
| 226 |
+
_draw_point(z_buffer, canvas, point, depth, radius, color)
|
| 227 |
+
if (render_height, render_width) != (output_height, output_width):
|
| 228 |
+
canvas = cv2.resize(canvas, (output_width, output_height), interpolation=cv2.INTER_AREA)
|
| 229 |
+
canvas = cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB)
|
| 230 |
+
return np.ascontiguousarray(canvas)
|
|
@@ -0,0 +1,66 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Private subprocess entry point for licensed body-model conditioning."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import json
|
| 6 |
+
import sys
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
|
| 9 |
+
from fdanyone.device import select_cuda_device
|
| 10 |
+
from fdanyone.motion.result import MotionResult
|
| 11 |
+
from fdanyone.skeleton.pipeline import build_skeleton_conditioning
|
| 12 |
+
from fdanyone.video import load_canonical_working_clip
|
| 13 |
+
from fdanyone.views import ViewPlan
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def run_skeleton_worker(
|
| 17 |
+
*,
|
| 18 |
+
working_video: str | Path,
|
| 19 |
+
clip_metadata: str | Path,
|
| 20 |
+
motion_result_dir: str | Path,
|
| 21 |
+
regressor_path: str | Path,
|
| 22 |
+
foreground_model_path: str | Path,
|
| 23 |
+
gvhmr_root: str | Path,
|
| 24 |
+
output_dir: str | Path,
|
| 25 |
+
device: str,
|
| 26 |
+
view_plan: ViewPlan,
|
| 27 |
+
defer_target_skeletons: bool = False,
|
| 28 |
+
) -> None:
|
| 29 |
+
"""Build conditioning for one request and publish it under ``output_dir``."""
|
| 30 |
+
|
| 31 |
+
device, _ = select_cuda_device(device)
|
| 32 |
+
clip = load_canonical_working_clip(working_video, clip_metadata)
|
| 33 |
+
motion = MotionResult.load(motion_result_dir)
|
| 34 |
+
build_skeleton_conditioning(
|
| 35 |
+
motion=motion,
|
| 36 |
+
clip=clip,
|
| 37 |
+
regressor_path=regressor_path,
|
| 38 |
+
foreground_model_path=foreground_model_path,
|
| 39 |
+
gvhmr_root=gvhmr_root,
|
| 40 |
+
output_dir=output_dir,
|
| 41 |
+
device=device,
|
| 42 |
+
view_plan=view_plan,
|
| 43 |
+
defer_target_skeletons=defer_target_skeletons,
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def main(request_path: str) -> None:
|
| 48 |
+
request = json.loads(Path(request_path).read_text())
|
| 49 |
+
run_skeleton_worker(
|
| 50 |
+
working_video=request["working_video"],
|
| 51 |
+
clip_metadata=request["clip_metadata"],
|
| 52 |
+
motion_result_dir=request["motion_result_dir"],
|
| 53 |
+
regressor_path=request["regressor_path"],
|
| 54 |
+
foreground_model_path=request["foreground_model_path"],
|
| 55 |
+
gvhmr_root=request["gvhmr_root"],
|
| 56 |
+
output_dir=request["output_dir"],
|
| 57 |
+
device=request["device"],
|
| 58 |
+
view_plan=ViewPlan.from_dict(request["view_plan"]),
|
| 59 |
+
defer_target_skeletons=bool(request.get("defer_target_skeletons", False)),
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
if __name__ == "__main__":
|
| 64 |
+
if len(sys.argv) != 2:
|
| 65 |
+
raise SystemExit("Usage: python -m fdanyone.skeleton.worker REQUEST.json")
|
| 66 |
+
main(sys.argv[1])
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Third-party inference code redistributed with its upstream notices."""
|
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding those notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 179 |
+
|
| 180 |
+
To apply the Apache License to your work, attach the following
|
| 181 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 182 |
+
replaced with your own identifying information. (Don't include
|
| 183 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 184 |
+
comment syntax for the file format. We also recommend that a
|
| 185 |
+
file or class name and description of purpose be included on the
|
| 186 |
+
same "printed page" as the copyright notice for easier
|
| 187 |
+
identification within third-party archives.
|
| 188 |
+
|
| 189 |
+
Copyright [2023] [Zhongjie Duan]
|
| 190 |
+
|
| 191 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 192 |
+
you may not use this file except in compliance with the License.
|
| 193 |
+
You may obtain a copy of the License at
|
| 194 |
+
|
| 195 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 196 |
+
|
| 197 |
+
Unless required by applicable law or agreed to in writing, software
|
| 198 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 199 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 200 |
+
See the License for the specific language governing permissions and
|
| 201 |
+
limitations under the License.
|
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DiffSynth-Studio provenance
|
| 2 |
+
|
| 3 |
+
This directory is a deliberately small inference-only extract from [modelscope/DiffSynth-Studio](https://github.com/modelscope/DiffSynth-Studio), licensed under Apache-2.0.
|
| 4 |
+
|
| 5 |
+
- Public base revision: `04e39f7de53df7276a7b40ca1791c2a393e05ff3`
|
| 6 |
+
- Research fork revision used by the original experiment: `c00782d90c872c97bda4745a9e6a41a0a4a7c4db`
|
| 7 |
+
- `UPSTREAM.patch` SHA-256: `178e6035e451f94a2122fa4d2c876a488546964768b90275549cbf609f97daba`
|
| 8 |
+
- Extracted: 2026-07-16
|
| 9 |
+
|
| 10 |
+
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.
|
| 11 |
+
|
| 12 |
+
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.
|
| 13 |
+
|
| 14 |
+
To reproduce the retained research sources before pruning:
|
| 15 |
+
|
| 16 |
+
```bash
|
| 17 |
+
git clone https://github.com/modelscope/DiffSynth-Studio.git
|
| 18 |
+
git -C DiffSynth-Studio checkout 04e39f7de53df7276a7b40ca1791c2a393e05ff3
|
| 19 |
+
git -C DiffSynth-Studio apply --check /path/to/UPSTREAM.patch
|
| 20 |
+
git -C DiffSynth-Studio apply /path/to/UPSTREAM.patch
|
| 21 |
+
```
|
| 22 |
+
|
| 23 |
+
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.
|
|
@@ -0,0 +1,1334 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
diff --git a/diffsynth/models/wan_video_dit.py b/diffsynth/models/wan_video_dit.py
|
| 2 |
+
index 1a54728..b663722 100644
|
| 3 |
+
--- a/diffsynth/models/wan_video_dit.py
|
| 4 |
+
+++ b/diffsynth/models/wan_video_dit.py
|
| 5 |
+
@@ -3,7 +3,7 @@ import torch.nn as nn
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
import math
|
| 8 |
+
from typing import Tuple, Optional
|
| 9 |
+
-from einops import rearrange
|
| 10 |
+
+from einops import rearrange, repeat
|
| 11 |
+
from .utils import hash_state_dict_keys
|
| 12 |
+
from .wan_video_camera_controller import SimpleAdapter
|
| 13 |
+
try:
|
| 14 |
+
@@ -23,7 +23,12 @@ try:
|
| 15 |
+
SAGE_ATTN_AVAILABLE = True
|
| 16 |
+
except ModuleNotFoundError:
|
| 17 |
+
SAGE_ATTN_AVAILABLE = False
|
| 18 |
+
-
|
| 19 |
+
+
|
| 20 |
+
+try:
|
| 21 |
+
+ from spas_sage_attn import spas_sage2_attn_meansim_topk_cuda
|
| 22 |
+
+ SPARGE_ATTN_AVAILABLE = True
|
| 23 |
+
+except ModuleNotFoundError:
|
| 24 |
+
+ SPARGE_ATTN_AVAILABLE = False
|
| 25 |
+
|
| 26 |
+
def flash_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, num_heads: int, compatibility_mode=False):
|
| 27 |
+
if compatibility_mode:
|
| 28 |
+
@@ -46,6 +51,13 @@ def flash_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, num_heads
|
| 29 |
+
v = rearrange(v, "b s (n d) -> b s n d", n=num_heads)
|
| 30 |
+
x = flash_attn.flash_attn_func(q, k, v)
|
| 31 |
+
x = rearrange(x, "b s n d -> b s (n d)", n=num_heads)
|
| 32 |
+
+ # TODO: test spas_sage2_attn_meansim_topk_cuda
|
| 33 |
+
+ # elif SPARGE_ATTN_AVAILABLE:
|
| 34 |
+
+ # q = rearrange(q, "b s (n d) -> b n s d", n=num_heads)
|
| 35 |
+
+ # k = rearrange(k, "b s (n d) -> b n s d", n=num_heads)
|
| 36 |
+
+ # v = rearrange(v, "b s (n d) -> b n s d", n=num_heads)
|
| 37 |
+
+ # x = spas_sage2_attn_meansim_topk_cuda(q, k, v, simthreshd1=-0.1, topk=0.5, pvthreshd=15, is_causal=False)
|
| 38 |
+
+ # x = rearrange(x, "b n s d -> b s (n d)", n=num_heads)
|
| 39 |
+
elif SAGE_ATTN_AVAILABLE:
|
| 40 |
+
q = rearrange(q, "b s (n d) -> b n s d", n=num_heads)
|
| 41 |
+
k = rearrange(k, "b s (n d) -> b n s d", n=num_heads)
|
| 42 |
+
@@ -186,6 +198,32 @@ class CrossAttention(nn.Module):
|
| 43 |
+
return self.o(x)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
+class CrossAttentionSrcCam(nn.Module):
|
| 47 |
+
+ def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
|
| 48 |
+
+ super().__init__()
|
| 49 |
+
+ self.dim = dim
|
| 50 |
+
+ self.num_heads = num_heads
|
| 51 |
+
+ self.head_dim = dim // num_heads
|
| 52 |
+
+
|
| 53 |
+
+ self.q = nn.Linear(dim, dim)
|
| 54 |
+
+ self.k = nn.Linear(dim, dim)
|
| 55 |
+
+ self.v = nn.Linear(dim, dim)
|
| 56 |
+
+ self.o = nn.Linear(dim, dim)
|
| 57 |
+
+ self.norm_q = RMSNorm(dim, eps=eps)
|
| 58 |
+
+ self.norm_k = RMSNorm(dim, eps=eps)
|
| 59 |
+
+
|
| 60 |
+
+ self.attn = AttentionModule(self.num_heads)
|
| 61 |
+
+
|
| 62 |
+
+ def forward(self, x, x_src, freqs, freqs_src):
|
| 63 |
+
+ q = self.norm_q(self.q(x))
|
| 64 |
+
+ k = self.norm_k(self.k(x_src))
|
| 65 |
+
+ v = self.v(x_src)
|
| 66 |
+
+ q = rope_apply(q, freqs, self.num_heads)
|
| 67 |
+
+ k = rope_apply(k, freqs_src, self.num_heads)
|
| 68 |
+
+ x = self.attn(q, k, v)
|
| 69 |
+
+ return self.o(x)
|
| 70 |
+
+
|
| 71 |
+
+
|
| 72 |
+
class GateModule(nn.Module):
|
| 73 |
+
def __init__(self,):
|
| 74 |
+
super().__init__()
|
| 75 |
+
@@ -200,6 +238,13 @@ class DiTBlock(nn.Module):
|
| 76 |
+
self.num_heads = num_heads
|
| 77 |
+
self.ffn_dim = ffn_dim
|
| 78 |
+
|
| 79 |
+
+ self.disable_video_attn = False
|
| 80 |
+
+ self.use_4d_attn = False
|
| 81 |
+
+ self.use_mvs_attn = False
|
| 82 |
+
+ self.use_src_self_attn = False
|
| 83 |
+
+ self.use_src_cross_attn = False
|
| 84 |
+
+ self.use_cam_encoder = False
|
| 85 |
+
+
|
| 86 |
+
self.self_attn = SelfAttention(dim, num_heads, eps)
|
| 87 |
+
self.cross_attn = CrossAttention(
|
| 88 |
+
dim, num_heads, eps, has_image_input=has_image_input)
|
| 89 |
+
@@ -211,7 +256,10 @@ class DiTBlock(nn.Module):
|
| 90 |
+
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
| 91 |
+
self.gate = GateModule()
|
| 92 |
+
|
| 93 |
+
- def forward(self, x, context, t_mod, freqs):
|
| 94 |
+
+ def forward(self, x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, shape):
|
| 95 |
+
+ # breakpoint()
|
| 96 |
+
+ v, f, h, w = shape
|
| 97 |
+
+
|
| 98 |
+
has_seq = len(t_mod.shape) == 4
|
| 99 |
+
chunk_dim = 2 if has_seq else 1
|
| 100 |
+
# msa: multi-head self-attention mlp: multi-layer perceptron
|
| 101 |
+
@@ -222,12 +270,80 @@ class DiTBlock(nn.Module):
|
| 102 |
+
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
|
| 103 |
+
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
|
| 104 |
+
)
|
| 105 |
+
- input_x = modulate(self.norm1(x), shift_msa, scale_msa)
|
| 106 |
+
- x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
|
| 107 |
+
+
|
| 108 |
+
+ # video self-attention
|
| 109 |
+
+ if not self.disable_video_attn:
|
| 110 |
+
+ input_x = modulate(self.norm1(x), shift_msa, scale_msa)
|
| 111 |
+
+ if self.use_4d_attn:
|
| 112 |
+
+ input_x = rearrange(input_x, "v fhw c -> (v fhw) c").unsqueeze(0)
|
| 113 |
+
+ freqs = repeat(freqs, "fhw 1 c -> (v fhw) 1 c", v=v)
|
| 114 |
+
+ input_x = self.self_attn(input_x, freqs)
|
| 115 |
+
+ if self.use_4d_attn:
|
| 116 |
+
+ input_x = rearrange(input_x.squeeze(0), "(v fhw) c -> v fhw c", v=v)
|
| 117 |
+
+ x = self.gate(x, gate_msa, input_x)
|
| 118 |
+
+
|
| 119 |
+
+ # source-view self-attention
|
| 120 |
+
+ if self.use_src_self_attn:
|
| 121 |
+
+ x_cat = torch.cat([x, x_src], dim=1)
|
| 122 |
+
+ shift_src, scale_src, gate_src = (
|
| 123 |
+
+ self.modulation_src.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1)
|
| 124 |
+
+ input_x_cat = modulate(self.norm1_src(x_cat), shift_src, scale_src)
|
| 125 |
+
+
|
| 126 |
+
+ freqs_cat = torch.cat([freqs, freqs_src], dim=0)
|
| 127 |
+
+ input_x_cat = self.self_attn_src(input_x_cat, freqs_cat)
|
| 128 |
+
+ x_cat = self.gate(x_cat, gate_src, input_x_cat)
|
| 129 |
+
+ len_src = x_src.shape[1]
|
| 130 |
+
+ x, x_src = x_cat[:, :-len_src, ...], x_cat[:, -len_src:, ...]
|
| 131 |
+
+
|
| 132 |
+
+ # source-view cross-attention
|
| 133 |
+
+ if self.use_src_cross_attn:
|
| 134 |
+
+ shift_src, scale_src, gate_src = (
|
| 135 |
+
+ self.modulation_src.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1)
|
| 136 |
+
+ input_x = modulate(self.norm1_src(x), shift_src, scale_src)
|
| 137 |
+
+ input_x_src = self.norm1_src(x_src)
|
| 138 |
+
+
|
| 139 |
+
+ input_x = self.cross_attn_src(input_x, input_x_src, freqs, freqs_src)
|
| 140 |
+
+ x = self.gate(x, gate_src, input_x)
|
| 141 |
+
+
|
| 142 |
+
+ # multiview self-attention
|
| 143 |
+
+ if self.use_mvs_attn:
|
| 144 |
+
+ shift_mvs, scale_mvs, gate_mvs = (
|
| 145 |
+
+ self.modulation_mvs.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod[:, :3, :]).chunk(3, dim=1)
|
| 146 |
+
+ input_x = modulate(self.norm1_mvs(x), shift_mvs, scale_mvs)
|
| 147 |
+
+
|
| 148 |
+
+ # add camera embedding before multiview attention
|
| 149 |
+
+ if self.use_cam_encoder and cam_emb is not None:
|
| 150 |
+
+ cam_proj = self.cam_encoder(cam_emb) # (v, 1, dim)
|
| 151 |
+
+ cam_proj = cam_proj.unsqueeze(2).unsqueeze(3).expand(-1, f, h, w, -1) # (v, f, h, w, dim)
|
| 152 |
+
+ cam_proj = rearrange(cam_proj, "v f h w d -> v (f h w) d")
|
| 153 |
+
+ input_x = input_x + cam_proj
|
| 154 |
+
+
|
| 155 |
+
+ input_x = rearrange(input_x, "v (f h w) c -> f (v h w) c", v=v, f=f, h=h, w=w)
|
| 156 |
+
+ input_x = self.self_attn_mvs(input_x, freqs_mvs)
|
| 157 |
+
+ input_x = rearrange(input_x, "f (v h w) c -> v (f h w) c", v=v, f=f, h=h, w=w)
|
| 158 |
+
+
|
| 159 |
+
+ # projector wraps multiview attention output
|
| 160 |
+
+ if self.use_cam_encoder:
|
| 161 |
+
+ input_x = self.projector(input_x)
|
| 162 |
+
+
|
| 163 |
+
+ x = self.gate(x, gate_mvs, input_x)
|
| 164 |
+
+
|
| 165 |
+
+ # prompt cross-attention
|
| 166 |
+
+ context = repeat(context, "1 l c -> v l c", v=x.shape[0])
|
| 167 |
+
x = x + self.cross_attn(self.norm3(x), context)
|
| 168 |
+
+
|
| 169 |
+
+ # feed-forward network
|
| 170 |
+
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
|
| 171 |
+
x = self.gate(x, gate_mlp, self.ffn(input_x))
|
| 172 |
+
- return x
|
| 173 |
+
+
|
| 174 |
+
+ if self.use_src_self_attn:
|
| 175 |
+
+ # for src self-attention, the x_src is updated as well
|
| 176 |
+
+ x_src = x_src + self.cross_attn(self.norm3(x_src), context)
|
| 177 |
+
+
|
| 178 |
+
+ input_x_src = modulate(self.norm2(x_src), shift_mlp, scale_mlp)
|
| 179 |
+
+ x_src = self.gate(x_src, gate_mlp, self.ffn(input_x_src))
|
| 180 |
+
+
|
| 181 |
+
+ return x, x_src
|
| 182 |
+
|
| 183 |
+
|
| 184 |
+
class MLP(torch.nn.Module):
|
| 185 |
+
@@ -264,11 +380,57 @@ class Head(nn.Module):
|
| 186 |
+
shift, scale = (self.modulation.unsqueeze(0).to(dtype=t_mod.dtype, device=t_mod.device) + t_mod.unsqueeze(2)).chunk(2, dim=2)
|
| 187 |
+
x = (self.head(self.norm(x) * (1 + scale.squeeze(2)) + shift.squeeze(2)))
|
| 188 |
+
else:
|
| 189 |
+
+ if t_mod.shape[0] != 1:
|
| 190 |
+
+ t_mod = t_mod[:, None, :] # [b, d] -> [b, 1, d]
|
| 191 |
+
shift, scale = (self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(2, dim=1)
|
| 192 |
+
x = (self.head(self.norm(x) * (1 + scale) + shift))
|
| 193 |
+
return x
|
| 194 |
+
|
| 195 |
+
|
| 196 |
+
+def pad_for_3d_conv(x, kernel_size):
|
| 197 |
+
+ """Pad to be divisible by kernel_size. From FramePack."""
|
| 198 |
+
+ _, _, t, h, w = x.shape
|
| 199 |
+
+ pt, ph, pw = kernel_size
|
| 200 |
+
+ pad_t = (pt - (t % pt)) % pt
|
| 201 |
+
+ pad_h = (ph - (h % ph)) % ph
|
| 202 |
+
+ pad_w = (pw - (w % pw)) % pw
|
| 203 |
+
+ if pad_t == 0 and pad_h == 0 and pad_w == 0:
|
| 204 |
+
+ return x
|
| 205 |
+
+ return F.pad(x, (0, pad_w, 0, pad_h, 0, pad_t), mode='replicate')
|
| 206 |
+
+
|
| 207 |
+
+
|
| 208 |
+
+class ViewPackEmbedding(nn.Module):
|
| 209 |
+
+ """Multi-resolution spatial patch embedding for clean source views.
|
| 210 |
+
+ Spatial-only downsampling: temporal dim preserved.
|
| 211 |
+
+ Ref: HunyuanVideoPatchEmbedForCleanLatents in FramePack.
|
| 212 |
+
+
|
| 213 |
+
+ Supported configurations (all exactly fill 4 quadrants):
|
| 214 |
+
+ - 4×2x views (each 2x view → 1 quadrant)
|
| 215 |
+
+ - 3×2x + 4×4x (3 quadrants from 2x + 1 quadrant from 4×4x tile)
|
| 216 |
+
+ """
|
| 217 |
+
+
|
| 218 |
+
+ def __init__(self, in_dim, dim, patch_size):
|
| 219 |
+
+ super().__init__()
|
| 220 |
+
+ pt, ph, pw = patch_size # (1, 2, 2) for Wan
|
| 221 |
+
+ # 2x: spatial 2x downsample relative to 1x -> kernel (1, 4, 4)
|
| 222 |
+
+ self.proj_2x = nn.Conv3d(in_dim, dim, kernel_size=(pt, ph*2, pw*2), stride=(pt, ph*2, pw*2))
|
| 223 |
+
+ # 4x: spatial 4x downsample relative to 1x -> kernel (1, 8, 8)
|
| 224 |
+
+ self.proj_4x = nn.Conv3d(in_dim, dim, kernel_size=(pt, ph*4, pw*4), stride=(pt, ph*4, pw*4))
|
| 225 |
+
+
|
| 226 |
+
+ @torch.no_grad()
|
| 227 |
+
+ def initialize_from_patch_embedding(self, patch_embedding: nn.Conv3d):
|
| 228 |
+
+ """FramePack-style init: tile spatial dims and scale by 1/area_ratio."""
|
| 229 |
+
+ weight = patch_embedding.weight.detach().clone() # (dim, in_dim, 1, 2, 2)
|
| 230 |
+
+ bias = patch_embedding.bias.detach().clone()
|
| 231 |
+
+ sd = {
|
| 232 |
+
+ 'proj_2x.weight': repeat(weight, 'b c t h w -> b c t (h 2) (w 2)') / 4.0,
|
| 233 |
+
+ 'proj_2x.bias': bias.clone(),
|
| 234 |
+
+ 'proj_4x.weight': repeat(weight, 'b c t h w -> b c t (h 4) (w 4)') / 16.0,
|
| 235 |
+
+ 'proj_4x.bias': bias.clone(),
|
| 236 |
+
+ }
|
| 237 |
+
+ self.load_state_dict(sd)
|
| 238 |
+
+
|
| 239 |
+
+
|
| 240 |
+
class WanModel(torch.nn.Module):
|
| 241 |
+
def __init__(
|
| 242 |
+
self,
|
| 243 |
+
@@ -355,32 +517,121 @@ class WanModel(torch.nn.Module):
|
| 244 |
+
|
| 245 |
+
def forward(self,
|
| 246 |
+
x: torch.Tensor,
|
| 247 |
+
+ x_src: torch.Tensor,
|
| 248 |
+
timestep: torch.Tensor,
|
| 249 |
+
context: torch.Tensor,
|
| 250 |
+
+ skeletons: Optional[torch.Tensor] = None,
|
| 251 |
+
+ cam_emb: Optional[torch.Tensor] = None,
|
| 252 |
+
+ drop_viewpack_tokens: bool = False,
|
| 253 |
+
clip_feature: Optional[torch.Tensor] = None,
|
| 254 |
+
y: Optional[torch.Tensor] = None,
|
| 255 |
+
use_gradient_checkpointing: bool = False,
|
| 256 |
+
use_gradient_checkpointing_offload: bool = False,
|
| 257 |
+
**kwargs,
|
| 258 |
+
):
|
| 259 |
+
- t = self.time_embedding(
|
| 260 |
+
- sinusoidal_embedding_1d(self.freq_dim, timestep))
|
| 261 |
+
- t_mod = self.time_projection(t).unflatten(1, (6, self.dim))
|
| 262 |
+
+ # breakpoint()
|
| 263 |
+
context = self.text_embedding(context)
|
| 264 |
+
-
|
| 265 |
+
+
|
| 266 |
+
if self.has_image_input:
|
| 267 |
+
x = torch.cat([x, y], dim=1) # (b, c_x + c_y, f, h, w)
|
| 268 |
+
clip_embdding = self.img_emb(clip_feature)
|
| 269 |
+
context = torch.cat([clip_embdding, context], dim=1)
|
| 270 |
+
-
|
| 271 |
+
+
|
| 272 |
+
x, (f, h, w) = self.patchify(x)
|
| 273 |
+
-
|
| 274 |
+
+
|
| 275 |
+
freqs = torch.cat([
|
| 276 |
+
self.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
| 277 |
+
self.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
| 278 |
+
self.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
| 279 |
+
], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
|
| 280 |
+
-
|
| 281 |
+
+
|
| 282 |
+
+ # build packed views from 1x/2x/4x source views
|
| 283 |
+
+ v_src = x_src.shape[0]
|
| 284 |
+
+ if v_src == 1:
|
| 285 |
+
+ # 1×1x: no 2x/4x sources
|
| 286 |
+
+ x_src_2x, x_src_4x = None, None
|
| 287 |
+
+ elif v_src == 5:
|
| 288 |
+
+ # 1×1x + 4×2x
|
| 289 |
+
+ x_src, x_src_2x, x_src_4x = x_src[:1], x_src[1:], None
|
| 290 |
+
+ elif v_src == 8:
|
| 291 |
+
+ # 1×1x + 3×2x + 4×4x
|
| 292 |
+
+ x_src, x_src_2x, x_src_4x = x_src[:1], x_src[1:4], x_src[4:]
|
| 293 |
+
+ else:
|
| 294 |
+
+ raise ValueError(f"Unsupported number of source views: {v_src}")
|
| 295 |
+
+
|
| 296 |
+
+ # 1x source: always packed as one extra view
|
| 297 |
+
+ x_src, (f_src, h_src, w_src) = self.patchify(x_src)
|
| 298 |
+
+ if not self.use_src_attn:
|
| 299 |
+
+ # use viewpack for 1x source views
|
| 300 |
+
+ x = torch.cat([x, x_src], dim=0)
|
| 301 |
+
+ x_src = None
|
| 302 |
+
+ freqs_src = None
|
| 303 |
+
+ v_pack = 1
|
| 304 |
+
+ else:
|
| 305 |
+
+ # use src self/cross-attention for 1x source views
|
| 306 |
+
+ x_src = repeat(x_src, "1 fhw_src c -> v fhw_src c", v=x.shape[0])
|
| 307 |
+
+ freqs_src = torch.cat([
|
| 308 |
+
+ 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),
|
| 309 |
+
+ self.freqs[1][:h_src].view(1, h_src, 1, -1).expand(f_src, h_src, w_src, -1),
|
| 310 |
+
+ self.freqs[2][:w_src].view(1, 1, w_src, -1).expand(f_src, h_src, w_src, -1)
|
| 311 |
+
+ ], dim=-1).reshape(f_src * h_src * w_src, 1, -1).to(x.device)
|
| 312 |
+
+ v_pack = 0
|
| 313 |
+
+
|
| 314 |
+
+ # 2x/4x source views: packed after 1x source views
|
| 315 |
+
+ if self.use_viewpack:
|
| 316 |
+
+ if x_src_2x is not None:
|
| 317 |
+
+ x_src_2x = pad_for_3d_conv(x_src_2x, self.viewpack_embedding.proj_2x.kernel_size)
|
| 318 |
+
+ x_src_2x = self.viewpack_embedding.proj_2x(x_src_2x) # (v_2x, dim, f, h//2, w//2)
|
| 319 |
+
+
|
| 320 |
+
+ if x_src_4x is not None:
|
| 321 |
+
+ # tile 4x source views into 2x source views
|
| 322 |
+
+ x_src_4x = pad_for_3d_conv(x_src_4x, self.viewpack_embedding.proj_4x.kernel_size)
|
| 323 |
+
+ x_src_4x = self.viewpack_embedding.proj_4x(x_src_4x) # (4, dim, f, h//4, w//4)
|
| 324 |
+
+ 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)
|
| 325 |
+
+ f_2x, h_2x, w_2x = x_src_2x.shape[2:] # crop padding surplus to match 2x source views (f, h, w)
|
| 326 |
+
+ x_src_4x = x_src_4x[:, :, :f_2x, :h_2x, :w_2x]
|
| 327 |
+
+ x_src_2x = torch.cat([x_src_2x, x_src_4x], dim=0) # (4, dim, f, h//2, w//2)
|
| 328 |
+
+
|
| 329 |
+
+ # tile 2x source views into 1x source views
|
| 330 |
+
+ x_pack = rearrange(x_src_2x, '(g1 g2) c f h w -> 1 c f (g1 h) (g2 w)', g1=2, g2=2)
|
| 331 |
+
+ x_pack = x_pack[:, :, :f, :h, :w] # crop padding surplus to match 1x source views (f, h, w)
|
| 332 |
+
+ x_pack = rearrange(x_pack, '1 c f h w -> 1 (f h w) c')
|
| 333 |
+
+ x_pack = x_pack.to(dtype=x.dtype)
|
| 334 |
+
+ if drop_viewpack_tokens:
|
| 335 |
+
+ # Keep viewpack parameters in the autograd graph across distributed ranks.
|
| 336 |
+
+ zero_dependency = x_pack.float().mean().to(dtype=x.dtype) * 0.0
|
| 337 |
+
+ x = x + zero_dependency
|
| 338 |
+
+ else:
|
| 339 |
+
+ x = torch.cat([x, x_pack], dim=0)
|
| 340 |
+
+ v_pack += 1
|
| 341 |
+
+
|
| 342 |
+
+ timestep = torch.cat([timestep, torch.zeros(v_pack, device=timestep.device, dtype=timestep.dtype)])
|
| 343 |
+
+ if skeletons is not None:
|
| 344 |
+
+ skeletons = torch.cat([skeletons, -torch.ones_like(skeletons[:1]).expand(v_pack, -1, -1, -1, -1)], dim=0)
|
| 345 |
+
+
|
| 346 |
+
+ # Expand cam_emb for viewpack views (zero vectors for packed source views)
|
| 347 |
+
+ if cam_emb is not None:
|
| 348 |
+
+ cam_emb = torch.cat([cam_emb, torch.zeros(v_pack, cam_emb.shape[-1],
|
| 349 |
+
+ device=cam_emb.device, dtype=cam_emb.dtype)], dim=0)
|
| 350 |
+
+ cam_emb = cam_emb.unsqueeze(1) # (v, 1, 12)
|
| 351 |
+
+
|
| 352 |
+
+ # Compute time embeddings (after concat since timestep may have been extended)
|
| 353 |
+
+ t = self.time_embedding(
|
| 354 |
+
+ sinusoidal_embedding_1d(self.freq_dim, timestep).to(x.dtype))
|
| 355 |
+
+ t_mod = self.time_projection(t).unflatten(1, (6, self.dim))
|
| 356 |
+
+
|
| 357 |
+
+ if self.use_pose_encoder:
|
| 358 |
+
+ skeleton_latents = self.pose_encoder(skeletons)
|
| 359 |
+
+ skeleton_tokens = rearrange(skeleton_latents, "v c f h w -> v (f h w) c")
|
| 360 |
+
+ x = x + skeleton_tokens
|
| 361 |
+
+
|
| 362 |
+
+ v = x.shape[0]
|
| 363 |
+
+ freqs_mvs = torch.cat([
|
| 364 |
+
+ self.freqs[0][:v].view(v, 1, 1, -1).expand(v, h, w, -1),
|
| 365 |
+
+ self.freqs[1][:h].view(1, h, 1, -1).expand(v, h, w, -1),
|
| 366 |
+
+ self.freqs[2][:w].view(1, 1, w, -1).expand(v, h, w, -1)
|
| 367 |
+
+ ], dim=-1).reshape(v * h * w, 1, -1).to(x.device)
|
| 368 |
+
+
|
| 369 |
+
def create_custom_forward(module):
|
| 370 |
+
def custom_forward(*inputs):
|
| 371 |
+
return module(*inputs)
|
| 372 |
+
@@ -390,19 +641,24 @@ class WanModel(torch.nn.Module):
|
| 373 |
+
if self.training and use_gradient_checkpointing:
|
| 374 |
+
if use_gradient_checkpointing_offload:
|
| 375 |
+
with torch.autograd.graph.save_on_cpu():
|
| 376 |
+
- x = torch.utils.checkpoint.checkpoint(
|
| 377 |
+
+ x, x_src = torch.utils.checkpoint.checkpoint(
|
| 378 |
+
create_custom_forward(block),
|
| 379 |
+
- x, context, t_mod, freqs,
|
| 380 |
+
+ x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w),
|
| 381 |
+
use_reentrant=False,
|
| 382 |
+
)
|
| 383 |
+
else:
|
| 384 |
+
- x = torch.utils.checkpoint.checkpoint(
|
| 385 |
+
+ x, x_src = torch.utils.checkpoint.checkpoint(
|
| 386 |
+
create_custom_forward(block),
|
| 387 |
+
- x, context, t_mod, freqs,
|
| 388 |
+
+ x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w),
|
| 389 |
+
use_reentrant=False,
|
| 390 |
+
)
|
| 391 |
+
else:
|
| 392 |
+
- x = block(x, context, t_mod, freqs)
|
| 393 |
+
+ x, x_src = block(x, x_src, context, t_mod, freqs, freqs_mvs, freqs_src, cam_emb, (v, f, h, w))
|
| 394 |
+
+
|
| 395 |
+
+ # Strip added views (viewpack)
|
| 396 |
+
+ if v_pack > 0:
|
| 397 |
+
+ x = x[:-v_pack, ...]
|
| 398 |
+
+ t = t[:-v_pack, ...]
|
| 399 |
+
|
| 400 |
+
x = self.head(x, t)
|
| 401 |
+
x = self.unpatchify(x, (f, h, w))
|
| 402 |
+
@@ -411,7 +667,7 @@ class WanModel(torch.nn.Module):
|
| 403 |
+
@staticmethod
|
| 404 |
+
def state_dict_converter():
|
| 405 |
+
return WanModelStateDictConverter()
|
| 406 |
+
-
|
| 407 |
+
+
|
| 408 |
+
|
| 409 |
+
class WanModelStateDictConverter:
|
| 410 |
+
def __init__(self):
|
| 411 |
+
diff --git a/diffsynth/models/wan_video_pose_encoder.py b/diffsynth/models/wan_video_pose_encoder.py
|
| 412 |
+
new file mode 100644
|
| 413 |
+
index 0000000..5180625
|
| 414 |
+
--- /dev/null
|
| 415 |
+
+++ b/diffsynth/models/wan_video_pose_encoder.py
|
| 416 |
+
@@ -0,0 +1,81 @@
|
| 417 |
+
+import torch
|
| 418 |
+
+import torch.nn as nn
|
| 419 |
+
+import numpy as np
|
| 420 |
+
+from torch.nn import init
|
| 421 |
+
+
|
| 422 |
+
+# PoseEncoder is 3D version of PoseNet in MimicMotion:
|
| 423 |
+
+# https://github.com/Tencent/MimicMotion/blob/c053153a1d124abae8c08568925ae88debc63001/mimicmotion/modules/pose_net.py
|
| 424 |
+
+
|
| 425 |
+
+
|
| 426 |
+
+class PoseEncoder(nn.Module):
|
| 427 |
+
+ def __init__(self, out_dim=5120, in_channels=3):
|
| 428 |
+
+ super().__init__()
|
| 429 |
+
+
|
| 430 |
+
+ if out_dim in (5120, 1536):
|
| 431 |
+
+ # Wan2.1-T2V-14B / 1.3B
|
| 432 |
+
+ t_strides = (1, 1, 1, 2, 2) # downsampled by 4
|
| 433 |
+
+ s_strides = (2, 2, 1, 2, 2) # downsampled by 16
|
| 434 |
+
+ kernel_size = (3, 3, 3)
|
| 435 |
+
+ elif out_dim == 3072:
|
| 436 |
+
+ # Wan2.2-TI2V-5B
|
| 437 |
+
+ t_strides = (1, 1, 1, 2, 2) # downsampled by 4
|
| 438 |
+
+ s_strides = (2, 2, 2, 2, 2) # downsampled by 32
|
| 439 |
+
+ kernel_size = (3, 4, 4)
|
| 440 |
+
+ else:
|
| 441 |
+
+ raise ValueError(f"Invalid out_dim: {out_dim}")
|
| 442 |
+
+
|
| 443 |
+
+ strides = [(t, s, s) for t, s in zip(t_strides, s_strides)]
|
| 444 |
+
+
|
| 445 |
+
+ self.conv_layers = nn.Sequential(
|
| 446 |
+
+ nn.Conv3d(in_channels, in_channels, kernel_size=3, stride=1, padding=1),
|
| 447 |
+
+ nn.SiLU(),
|
| 448 |
+
+ nn.Conv3d(in_channels, 16, kernel_size=kernel_size, stride=strides[0], padding=(1, 1, 1)),
|
| 449 |
+
+ nn.SiLU(),
|
| 450 |
+
+ nn.Conv3d(16, 16, kernel_size=3, stride=1, padding=1),
|
| 451 |
+
+ nn.SiLU(),
|
| 452 |
+
+ nn.Conv3d(16, 32, kernel_size=kernel_size, stride=strides[1], padding=(1, 1, 1)),
|
| 453 |
+
+ nn.SiLU(),
|
| 454 |
+
+ nn.Conv3d(32, 32, kernel_size=3, stride=1, padding=1),
|
| 455 |
+
+ nn.SiLU(),
|
| 456 |
+
+ nn.Conv3d(32, 64, kernel_size=kernel_size, stride=strides[2], padding=(1, 1, 1)),
|
| 457 |
+
+ nn.SiLU(),
|
| 458 |
+
+ nn.Conv3d(64, 64, kernel_size=3, stride=1, padding=1),
|
| 459 |
+
+ nn.SiLU(),
|
| 460 |
+
+ nn.Conv3d(64, 128, kernel_size=kernel_size, stride=strides[3], padding=(1, 1, 1)),
|
| 461 |
+
+ nn.SiLU(),
|
| 462 |
+
+ nn.Conv3d(128, 128, kernel_size=3, stride=1, padding=1),
|
| 463 |
+
+ nn.SiLU(),
|
| 464 |
+
+ nn.Conv3d(128, 256, kernel_size=kernel_size, stride=strides[4], padding=(1, 1, 1)),
|
| 465 |
+
+ nn.SiLU(),
|
| 466 |
+
+ )
|
| 467 |
+
+
|
| 468 |
+
+ self.final_proj = nn.Conv3d(256, out_dim, kernel_size=1)
|
| 469 |
+
+
|
| 470 |
+
+ self.scale = nn.Parameter(torch.ones(1) * 2.0)
|
| 471 |
+
+
|
| 472 |
+
+ self._initialize_weights()
|
| 473 |
+
+
|
| 474 |
+
+ def _initialize_weights(self):
|
| 475 |
+
+ for m in self.modules():
|
| 476 |
+
+ if isinstance(m, nn.Conv3d):
|
| 477 |
+
+ # He (Kaiming) initialization in fan‑in mode
|
| 478 |
+
+ receptive = np.prod(m.kernel_size) * m.in_channels
|
| 479 |
+
+ init.normal_(m.weight, mean=0.0, std=np.sqrt(2.0 / receptive))
|
| 480 |
+
+ if m.bias is not None:
|
| 481 |
+
+ init.zeros_(m.bias)
|
| 482 |
+
+ # start with zero output so model behaves like unconditional
|
| 483 |
+
+ init.zeros_(self.final_proj.weight)
|
| 484 |
+
+ if self.final_proj.bias is not None:
|
| 485 |
+
+ init.zeros_(self.final_proj.bias)
|
| 486 |
+
+
|
| 487 |
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 488 |
+
+ """
|
| 489 |
+
+ x: (B, C, F, H, W) -> latent grid matching the DiT patch tokens.
|
| 490 |
+
+ Wan2.1 uses F/4, H/16, W/16; Wan2.2-TI2V-5B uses F/4, H/32, W/32.
|
| 491 |
+
+ """
|
| 492 |
+
+ # Wan pattern: 1 -> 4 -> 4 -> ...
|
| 493 |
+
+ x = torch.cat([x[:, :, :1].repeat(1, 1, 3, 1, 1), x], dim=2)
|
| 494 |
+
+
|
| 495 |
+
+ x = self.conv_layers(x)
|
| 496 |
+
+ x = self.final_proj(x)
|
| 497 |
+
+ return x * self.scale
|
| 498 |
+
diff --git a/diffsynth/models/wan_video_vae.py b/diffsynth/models/wan_video_vae.py
|
| 499 |
+
index 397a2e7..43057ff 100644
|
| 500 |
+
--- a/diffsynth/models/wan_video_vae.py
|
| 501 |
+
+++ b/diffsynth/models/wan_video_vae.py
|
| 502 |
+
@@ -1121,7 +1121,7 @@ class WanVideoVAE(nn.Module):
|
| 503 |
+
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
| 504 |
+
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
| 505 |
+
|
| 506 |
+
- for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
|
| 507 |
+
+ for h, h_, w, w_ in tasks:
|
| 508 |
+
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
|
| 509 |
+
hidden_states_batch = self.model.decode(hidden_states_batch, self.scale).to(data_device)
|
| 510 |
+
|
| 511 |
+
@@ -1173,7 +1173,7 @@ class WanVideoVAE(nn.Module):
|
| 512 |
+
weight = torch.zeros((1, 1, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
|
| 513 |
+
values = torch.zeros((1, self.z_dim, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
|
| 514 |
+
|
| 515 |
+
- for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
|
| 516 |
+
+ for h, h_, w, w_ in tasks:
|
| 517 |
+
hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device)
|
| 518 |
+
hidden_states_batch = self.model.encode(hidden_states_batch, self.scale).to(data_device)
|
| 519 |
+
|
| 520 |
+
@@ -1216,10 +1216,9 @@ class WanVideoVAE(nn.Module):
|
| 521 |
+
|
| 522 |
+
|
| 523 |
+
def encode(self, videos, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 524 |
+
-
|
| 525 |
+
videos = [video.to("cpu") for video in videos]
|
| 526 |
+
hidden_states = []
|
| 527 |
+
- for video in videos:
|
| 528 |
+
+ for video in tqdm(videos, desc="VAE encoding", disable=not tiled):
|
| 529 |
+
video = video.unsqueeze(0)
|
| 530 |
+
if tiled:
|
| 531 |
+
tile_size = (tile_size[0] * self.upsampling_factor, tile_size[1] * self.upsampling_factor)
|
| 532 |
+
@@ -1234,11 +1233,18 @@ class WanVideoVAE(nn.Module):
|
| 533 |
+
|
| 534 |
+
|
| 535 |
+
def decode(self, hidden_states, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 536 |
+
- if tiled:
|
| 537 |
+
- video = self.tiled_decode(hidden_states, device, tile_size, tile_stride)
|
| 538 |
+
- else:
|
| 539 |
+
- video = self.single_decode(hidden_states, device)
|
| 540 |
+
- return video
|
| 541 |
+
+ hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states]
|
| 542 |
+
+ videos = []
|
| 543 |
+
+ for hidden_state in tqdm(hidden_states, desc="VAE decoding", disable=not tiled):
|
| 544 |
+
+ hidden_state = hidden_state.unsqueeze(0)
|
| 545 |
+
+ if tiled:
|
| 546 |
+
+ video = self.tiled_decode(hidden_state, device, tile_size, tile_stride)
|
| 547 |
+
+ else:
|
| 548 |
+
+ video = self.single_decode(hidden_state, device)
|
| 549 |
+
+ video = video.squeeze(0)
|
| 550 |
+
+ videos.append(video)
|
| 551 |
+
+ videos = torch.stack(videos)
|
| 552 |
+
+ return videos
|
| 553 |
+
|
| 554 |
+
|
| 555 |
+
@staticmethod
|
| 556 |
+
diff --git a/diffsynth/pipelines/__init__.py b/diffsynth/pipelines/__init__.py
|
| 557 |
+
index e2ad551..f878ad8 100644
|
| 558 |
+
--- a/diffsynth/pipelines/__init__.py
|
| 559 |
+
+++ b/diffsynth/pipelines/__init__.py
|
| 560 |
+
@@ -12,4 +12,5 @@ from .pipeline_runner import SDVideoPipelineRunner
|
| 561 |
+
from .hunyuan_video import HunyuanVideoPipeline
|
| 562 |
+
from .step_video import StepVideoPipeline
|
| 563 |
+
from .wan_video import WanVideoPipeline
|
| 564 |
+
+from .wan_video_spatem import WanVideoSpaTemPipeline
|
| 565 |
+
KolorsImagePipeline = SDXLImagePipeline
|
| 566 |
+
diff --git a/diffsynth/pipelines/wan_video_spatem.py b/diffsynth/pipelines/wan_video_spatem.py
|
| 567 |
+
new file mode 100644
|
| 568 |
+
index 0000000..85b4f65
|
| 569 |
+
--- /dev/null
|
| 570 |
+
+++ b/diffsynth/pipelines/wan_video_spatem.py
|
| 571 |
+
@@ -0,0 +1,659 @@
|
| 572 |
+
+from ..models import ModelManager
|
| 573 |
+
+from ..models.wan_video_dit import WanModel
|
| 574 |
+
+from ..models.wan_video_pose_encoder import PoseEncoder
|
| 575 |
+
+from ..models.wan_video_text_encoder import WanTextEncoder
|
| 576 |
+
+from ..models.wan_video_vae import WanVideoVAE
|
| 577 |
+
+from ..models.wan_video_image_encoder import WanImageEncoder
|
| 578 |
+
+from ..schedulers.flow_match import FlowMatchScheduler
|
| 579 |
+
+from ..schedulers.bride_match import BridgeMatchScheduler
|
| 580 |
+
+from ..pipelines.base import BasePipeline
|
| 581 |
+
+from ..prompters import WanPrompter
|
| 582 |
+
+import torch, os
|
| 583 |
+
+import torch.nn as nn
|
| 584 |
+
+import numpy as np
|
| 585 |
+
+import torch.nn.functional as F
|
| 586 |
+
+from PIL import Image
|
| 587 |
+
+from tqdm import tqdm
|
| 588 |
+
+from typing import Optional, Union
|
| 589 |
+
+from functools import partial
|
| 590 |
+
+from einops import rearrange
|
| 591 |
+
+
|
| 592 |
+
+from ..vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
|
| 593 |
+
+from ..models.wan_video_text_encoder import T5RelativeEmbedding, T5LayerNorm
|
| 594 |
+
+from ..models.wan_video_dit import RMSNorm, SelfAttention, CrossAttentionSrcCam, ViewPackEmbedding
|
| 595 |
+
+from ..models.wan_video_vae import RMS_norm, CausalConv3d, Upsample
|
| 596 |
+
+from ..utils import ModelConfig
|
| 597 |
+
+
|
| 598 |
+
+
|
| 599 |
+
+class WanVideoSpaTemPipeline(BasePipeline):
|
| 600 |
+
+
|
| 601 |
+
+ def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None):
|
| 602 |
+
+ super().__init__(device=device, torch_dtype=torch_dtype)
|
| 603 |
+
+ self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
|
| 604 |
+
+ self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
|
| 605 |
+
+ self.text_encoder: WanTextEncoder = None
|
| 606 |
+
+ self.image_encoder: WanImageEncoder = None
|
| 607 |
+
+ self.dit: WanModel = None
|
| 608 |
+
+ self.vae: WanVideoVAE = None
|
| 609 |
+
+ self.model_names = ["text_encoder", "dit", "vae"]
|
| 610 |
+
+ self.height_division_factor = 16
|
| 611 |
+
+ self.width_division_factor = 16
|
| 612 |
+
+
|
| 613 |
+
+ def enable_vram_management(self, num_persistent_param_in_dit=None):
|
| 614 |
+
+ dtype = next(iter(self.text_encoder.parameters())).dtype
|
| 615 |
+
+ enable_vram_management(
|
| 616 |
+
+ self.text_encoder,
|
| 617 |
+
+ module_map={
|
| 618 |
+
+ torch.nn.Linear: AutoWrappedLinear,
|
| 619 |
+
+ torch.nn.Embedding: AutoWrappedModule,
|
| 620 |
+
+ T5RelativeEmbedding: AutoWrappedModule,
|
| 621 |
+
+ T5LayerNorm: AutoWrappedModule,
|
| 622 |
+
+ },
|
| 623 |
+
+ module_config=dict(
|
| 624 |
+
+ offload_dtype=dtype,
|
| 625 |
+
+ offload_device="cpu",
|
| 626 |
+
+ onload_dtype=dtype,
|
| 627 |
+
+ onload_device="cpu",
|
| 628 |
+
+ computation_dtype=self.torch_dtype,
|
| 629 |
+
+ computation_device=self.device,
|
| 630 |
+
+ ),
|
| 631 |
+
+ )
|
| 632 |
+
+ dtype = next(iter(self.dit.parameters())).dtype
|
| 633 |
+
+ enable_vram_management(
|
| 634 |
+
+ self.dit,
|
| 635 |
+
+ module_map={
|
| 636 |
+
+ torch.nn.Linear: AutoWrappedLinear,
|
| 637 |
+
+ torch.nn.Conv3d: AutoWrappedModule,
|
| 638 |
+
+ torch.nn.LayerNorm: AutoWrappedModule,
|
| 639 |
+
+ RMSNorm: AutoWrappedModule,
|
| 640 |
+
+ },
|
| 641 |
+
+ module_config=dict(
|
| 642 |
+
+ offload_dtype=dtype,
|
| 643 |
+
+ offload_device="cpu",
|
| 644 |
+
+ onload_dtype=dtype,
|
| 645 |
+
+ onload_device=self.device,
|
| 646 |
+
+ computation_dtype=self.torch_dtype,
|
| 647 |
+
+ computation_device=self.device,
|
| 648 |
+
+ ),
|
| 649 |
+
+ max_num_param=num_persistent_param_in_dit,
|
| 650 |
+
+ overflow_module_config=dict(
|
| 651 |
+
+ offload_dtype=dtype,
|
| 652 |
+
+ offload_device="cpu",
|
| 653 |
+
+ onload_dtype=dtype,
|
| 654 |
+
+ onload_device="cpu",
|
| 655 |
+
+ computation_dtype=self.torch_dtype,
|
| 656 |
+
+ computation_device=self.device,
|
| 657 |
+
+ ),
|
| 658 |
+
+ )
|
| 659 |
+
+ dtype = next(iter(self.vae.parameters())).dtype
|
| 660 |
+
+ enable_vram_management(
|
| 661 |
+
+ self.vae,
|
| 662 |
+
+ module_map={
|
| 663 |
+
+ torch.nn.Linear: AutoWrappedLinear,
|
| 664 |
+
+ torch.nn.Conv2d: AutoWrappedModule,
|
| 665 |
+
+ RMS_norm: AutoWrappedModule,
|
| 666 |
+
+ CausalConv3d: AutoWrappedModule,
|
| 667 |
+
+ Upsample: AutoWrappedModule,
|
| 668 |
+
+ torch.nn.SiLU: AutoWrappedModule,
|
| 669 |
+
+ torch.nn.Dropout: AutoWrappedModule,
|
| 670 |
+
+ },
|
| 671 |
+
+ module_config=dict(
|
| 672 |
+
+ offload_dtype=dtype,
|
| 673 |
+
+ offload_device="cpu",
|
| 674 |
+
+ onload_dtype=dtype,
|
| 675 |
+
+ onload_device=self.device,
|
| 676 |
+
+ computation_dtype=self.torch_dtype,
|
| 677 |
+
+ computation_device=self.device,
|
| 678 |
+
+ ),
|
| 679 |
+
+ )
|
| 680 |
+
+ if self.image_encoder is not None:
|
| 681 |
+
+ dtype = next(iter(self.image_encoder.parameters())).dtype
|
| 682 |
+
+ enable_vram_management(
|
| 683 |
+
+ self.image_encoder,
|
| 684 |
+
+ module_map={
|
| 685 |
+
+ torch.nn.Linear: AutoWrappedLinear,
|
| 686 |
+
+ torch.nn.Conv2d: AutoWrappedModule,
|
| 687 |
+
+ torch.nn.LayerNorm: AutoWrappedModule,
|
| 688 |
+
+ },
|
| 689 |
+
+ module_config=dict(
|
| 690 |
+
+ offload_dtype=dtype,
|
| 691 |
+
+ offload_device="cpu",
|
| 692 |
+
+ onload_dtype=dtype,
|
| 693 |
+
+ onload_device="cpu",
|
| 694 |
+
+ computation_dtype=dtype,
|
| 695 |
+
+ computation_device=self.device,
|
| 696 |
+
+ ),
|
| 697 |
+
+ )
|
| 698 |
+
+ self.enable_cpu_offload()
|
| 699 |
+
+
|
| 700 |
+
+ def fetch_models(self, model_manager: ModelManager):
|
| 701 |
+
+ text_encoder_model_and_path = model_manager.fetch_model("wan_video_text_encoder", require_model_path=True)
|
| 702 |
+
+ if text_encoder_model_and_path is not None:
|
| 703 |
+
+ self.text_encoder, tokenizer_path = text_encoder_model_and_path
|
| 704 |
+
+ self.prompter.fetch_models(self.text_encoder)
|
| 705 |
+
+ self.prompter.fetch_tokenizer(os.path.join(os.path.dirname(tokenizer_path), "google/umt5-xxl"))
|
| 706 |
+
+ self.dit = model_manager.fetch_model("wan_video_dit")
|
| 707 |
+
+ self.vae = model_manager.fetch_model("wan_video_vae")
|
| 708 |
+
+ self.image_encoder = model_manager.fetch_model("wan_video_image_encoder")
|
| 709 |
+
+
|
| 710 |
+
+ @staticmethod
|
| 711 |
+
+ def from_model_manager(model_manager: ModelManager, torch_dtype=None, device=None):
|
| 712 |
+
+ if device is None:
|
| 713 |
+
+ device = model_manager.device
|
| 714 |
+
+ if torch_dtype is None:
|
| 715 |
+
+ torch_dtype = model_manager.torch_dtype
|
| 716 |
+
+ pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=torch_dtype)
|
| 717 |
+
+ pipe.fetch_models(model_manager)
|
| 718 |
+
+ return pipe
|
| 719 |
+
+
|
| 720 |
+
+ @staticmethod
|
| 721 |
+
+ def from_pretrained(
|
| 722 |
+
+ torch_dtype: torch.dtype = torch.bfloat16,
|
| 723 |
+
+ device: Union[str, torch.device] = "cuda",
|
| 724 |
+
+ model_configs: list[ModelConfig] = [],
|
| 725 |
+
+ tokenizer_config: ModelConfig = ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="google/*"),
|
| 726 |
+
+ redirect_common_files: bool = True,
|
| 727 |
+
+ use_usp=False,
|
| 728 |
+
+ ):
|
| 729 |
+
+ # Redirect model path
|
| 730 |
+
+ if redirect_common_files:
|
| 731 |
+
+ redirect_dict = {
|
| 732 |
+
+ "models_t5_umt5-xxl-enc-bf16.pth": "Wan-AI/Wan2.1-T2V-1.3B",
|
| 733 |
+
+ "Wan2.1_VAE.pth": "Wan-AI/Wan2.1-T2V-1.3B",
|
| 734 |
+
+ "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth": "Wan-AI/Wan2.1-I2V-14B-480P",
|
| 735 |
+
+ }
|
| 736 |
+
+ for model_config in model_configs:
|
| 737 |
+
+ if model_config.origin_file_pattern is None or model_config.model_id is None:
|
| 738 |
+
+ continue
|
| 739 |
+
+ if (
|
| 740 |
+
+ model_config.origin_file_pattern in redirect_dict
|
| 741 |
+
+ and model_config.model_id != redirect_dict[model_config.origin_file_pattern]
|
| 742 |
+
+ ):
|
| 743 |
+
+ print(
|
| 744 |
+
+ 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."
|
| 745 |
+
+ )
|
| 746 |
+
+ model_config.model_id = redirect_dict[model_config.origin_file_pattern]
|
| 747 |
+
+
|
| 748 |
+
+ # Initialize pipeline
|
| 749 |
+
+ pipe = WanVideoSpaTemPipeline(device=device, torch_dtype=torch_dtype)
|
| 750 |
+
+ if use_usp:
|
| 751 |
+
+ pipe.initialize_usp()
|
| 752 |
+
+
|
| 753 |
+
+ # Download and load models
|
| 754 |
+
+ model_manager = ModelManager()
|
| 755 |
+
+ for model_config in model_configs:
|
| 756 |
+
+ model_config.download_if_necessary(use_usp=use_usp)
|
| 757 |
+
+ model_manager.load_model(
|
| 758 |
+
+ model_config.path,
|
| 759 |
+
+ device=model_config.offload_device or device,
|
| 760 |
+
+ torch_dtype=model_config.offload_dtype or torch_dtype,
|
| 761 |
+
+ )
|
| 762 |
+
+
|
| 763 |
+
+ # Load models
|
| 764 |
+
+ pipe.text_encoder = model_manager.fetch_model("wan_video_text_encoder")
|
| 765 |
+
+ dit = model_manager.fetch_model("wan_video_dit", index=2)
|
| 766 |
+
+ if isinstance(dit, list):
|
| 767 |
+
+ pipe.dit, pipe.dit2 = dit
|
| 768 |
+
+ else:
|
| 769 |
+
+ pipe.dit = dit
|
| 770 |
+
+ pipe.vae = model_manager.fetch_model("wan_video_vae")
|
| 771 |
+
+ pipe.image_encoder = model_manager.fetch_model("wan_video_image_encoder")
|
| 772 |
+
+ pipe.motion_controller = model_manager.fetch_model("wan_video_motion_controller")
|
| 773 |
+
+ pipe.vace = model_manager.fetch_model("wan_video_vace")
|
| 774 |
+
+
|
| 775 |
+
+ # Size division factor
|
| 776 |
+
+ if pipe.vae is not None:
|
| 777 |
+
+ pipe.height_division_factor = pipe.vae.upsampling_factor * 2
|
| 778 |
+
+ pipe.width_division_factor = pipe.vae.upsampling_factor * 2
|
| 779 |
+
+
|
| 780 |
+
+ # Initialize tokenizer
|
| 781 |
+
+ tokenizer_config.local_model_path = model_configs[0].local_model_path
|
| 782 |
+
+ tokenizer_config.skip_download = model_configs[0].skip_download
|
| 783 |
+
+ tokenizer_config.download_if_necessary(use_usp=use_usp)
|
| 784 |
+
+ pipe.prompter.fetch_models(pipe.text_encoder)
|
| 785 |
+
+ pipe.prompter.fetch_tokenizer(tokenizer_config.path)
|
| 786 |
+
+
|
| 787 |
+
+ # Unified Sequence Parallel
|
| 788 |
+
+ if use_usp:
|
| 789 |
+
+ pipe.enable_usp()
|
| 790 |
+
+
|
| 791 |
+
+ return pipe
|
| 792 |
+
+
|
| 793 |
+
+ def init_spatem_modules(
|
| 794 |
+
+ self,
|
| 795 |
+
+ disable_video_attn: bool = False,
|
| 796 |
+
+ use_4d_attn: bool = False,
|
| 797 |
+
+ use_mvs_attn: bool = False,
|
| 798 |
+
+ use_src_self_attn: bool = False,
|
| 799 |
+
+ use_src_cross_attn: bool = False,
|
| 800 |
+
+ freqs_src_shift: int = 121,
|
| 801 |
+
+ use_viewpack: bool = True,
|
| 802 |
+
+ viewpack_dropout_prob: float = 0.0,
|
| 803 |
+
+ use_pose_encoder: bool = True,
|
| 804 |
+
+ pose_encoder_type: str = "rgb",
|
| 805 |
+
+ use_cam_encoder: bool = False,
|
| 806 |
+
+ range_4d_attn: tuple[int, int, int] = (0, None, 2),
|
| 807 |
+
+ range_mvs_attn: tuple[int, int, int] = (1, None, 2),
|
| 808 |
+
+ range_src_self_attn: tuple[int, int, int] = (0, None, 2),
|
| 809 |
+
+ range_src_cross_attn: tuple[int, int, int] = (0, None, 2),
|
| 810 |
+
+ use_lbm: bool = False,
|
| 811 |
+
+ fill_wpmask_with_noise: bool = False,
|
| 812 |
+
+ ):
|
| 813 |
+
+ # breakpoint()
|
| 814 |
+
+ device, dtype = self.dit.patch_embedding.weight.device, self.dit.patch_embedding.weight.dtype
|
| 815 |
+
+
|
| 816 |
+
+ if disable_video_attn:
|
| 817 |
+
+ # todo: delete self_attn layers from the model
|
| 818 |
+
+ if use_4d_attn:
|
| 819 |
+
+ raise ValueError("Cannot use 4D attention when video attention is disabled")
|
| 820 |
+
+ for block in self.dit.blocks:
|
| 821 |
+
+ block.disable_video_attn = True
|
| 822 |
+
+
|
| 823 |
+
+ if use_4d_attn:
|
| 824 |
+
+ b, e, s = range_4d_attn
|
| 825 |
+
+ for block in self.dit.blocks[b:e:s]:
|
| 826 |
+
+ block.use_4d_attn = True
|
| 827 |
+
+
|
| 828 |
+
+ if use_mvs_attn:
|
| 829 |
+
+ b, e, s = range_mvs_attn
|
| 830 |
+
+ for block in self.dit.blocks[b:e:s]:
|
| 831 |
+
+ block.use_mvs_attn = True
|
| 832 |
+
+
|
| 833 |
+
+ dim = block.self_attn.q.weight.shape[0]
|
| 834 |
+
+ block.modulation_mvs = nn.Parameter(block.modulation[:, :3, :].detach().clone())
|
| 835 |
+
+ block.norm1_mvs = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to(
|
| 836 |
+
+ device=device, dtype=dtype
|
| 837 |
+
+ )
|
| 838 |
+
+ block.self_attn_mvs = SelfAttention(dim, block.self_attn.num_heads, block.self_attn.norm_q.eps).to(
|
| 839 |
+
+ device=device, dtype=dtype
|
| 840 |
+
+ )
|
| 841 |
+
+ block.self_attn_mvs.load_state_dict(block.self_attn.state_dict(), strict=True)
|
| 842 |
+
+
|
| 843 |
+
+ if not 0.0 <= viewpack_dropout_prob <= 1.0:
|
| 844 |
+
+ raise ValueError("viewpack_dropout_prob should be between 0 and 1")
|
| 845 |
+
+ if viewpack_dropout_prob > 0.0 and not use_viewpack:
|
| 846 |
+
+ raise ValueError("viewpack_dropout_prob requires use_viewpack=True")
|
| 847 |
+
+
|
| 848 |
+
+ if use_viewpack:
|
| 849 |
+
+ viewpack_emb = ViewPackEmbedding(
|
| 850 |
+
+ in_dim=self.dit.patch_embedding.weight.shape[1],
|
| 851 |
+
+ dim=self.dit.patch_embedding.weight.shape[0],
|
| 852 |
+
+ patch_size=list(self.dit.patch_embedding.kernel_size),
|
| 853 |
+
+ )
|
| 854 |
+
+ viewpack_emb.initialize_from_patch_embedding(self.dit.patch_embedding)
|
| 855 |
+
+ self.dit.viewpack_embedding = viewpack_emb.to(device=device, dtype=dtype)
|
| 856 |
+
+ elif use_src_self_attn:
|
| 857 |
+
+ if use_src_cross_attn:
|
| 858 |
+
+ raise ValueError("Cannot use both src self-attention and src cross-attention")
|
| 859 |
+
+
|
| 860 |
+
+ b, e, s = range_src_self_attn
|
| 861 |
+
+ for block in self.dit.blocks[b:e:s]:
|
| 862 |
+
+ block.use_src_self_attn = True
|
| 863 |
+
+
|
| 864 |
+
+ dim = block.self_attn.q.weight.shape[0]
|
| 865 |
+
+ block.modulation_src = nn.Parameter(block.modulation[:, :3, :].detach().clone())
|
| 866 |
+
+ block.norm1_src = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to(
|
| 867 |
+
+ device=device, dtype=dtype
|
| 868 |
+
+ )
|
| 869 |
+
+ block.self_attn_src = SelfAttention(dim, block.self_attn.num_heads, block.self_attn.norm_q.eps).to(
|
| 870 |
+
+ device=device, dtype=dtype
|
| 871 |
+
+ )
|
| 872 |
+
+ block.self_attn_src.load_state_dict(block.self_attn.state_dict(), strict=True)
|
| 873 |
+
+ elif use_src_cross_attn:
|
| 874 |
+
+ b, e, s = range_src_cross_attn
|
| 875 |
+
+ for block in self.dit.blocks[b:e:s]:
|
| 876 |
+
+ block.use_src_cross_attn = True
|
| 877 |
+
+
|
| 878 |
+
+ dim = block.self_attn.q.weight.shape[0]
|
| 879 |
+
+ block.modulation_src = nn.Parameter(block.modulation[:, :3, :].detach().clone())
|
| 880 |
+
+ block.norm1_src = nn.LayerNorm(dim, eps=block.norm1.eps, elementwise_affine=False).to(
|
| 881 |
+
+ device=device, dtype=dtype
|
| 882 |
+
+ )
|
| 883 |
+
+ block.cross_attn_src = CrossAttentionSrcCam(
|
| 884 |
+
+ dim, block.self_attn.num_heads, block.self_attn.norm_q.eps
|
| 885 |
+
+ ).to(device=device, dtype=dtype)
|
| 886 |
+
+ block.cross_attn_src.load_state_dict(block.self_attn.state_dict(), strict=True)
|
| 887 |
+
+
|
| 888 |
+
+ if use_pose_encoder:
|
| 889 |
+
+ if pose_encoder_type == "rgb":
|
| 890 |
+
+ in_channels = 3
|
| 891 |
+
+ elif pose_encoder_type == "rgbd":
|
| 892 |
+
+ in_channels = 4
|
| 893 |
+
+ else:
|
| 894 |
+
+ raise ValueError(f"Invalid pose_encoder_type: {pose_encoder_type}")
|
| 895 |
+
+ pose_encoder = PoseEncoder(out_dim=self.dit.patch_embedding.out_channels, in_channels=in_channels)
|
| 896 |
+
+ self.dit.pose_encoder = pose_encoder.to(device=device, dtype=dtype)
|
| 897 |
+
+
|
| 898 |
+
+ if use_cam_encoder:
|
| 899 |
+
+ dim = self.dit.blocks[0].self_attn.q.weight.shape[0]
|
| 900 |
+
+ for block in self.dit.blocks:
|
| 901 |
+
+ block.use_cam_encoder = True
|
| 902 |
+
+ block.cam_encoder = nn.Linear(12, dim).to(device=device, dtype=dtype)
|
| 903 |
+
+ block.projector = nn.Linear(dim, dim).to(device=device, dtype=dtype)
|
| 904 |
+
+ block.cam_encoder.weight.data.zero_()
|
| 905 |
+
+ block.cam_encoder.bias.data.zero_()
|
| 906 |
+
+ block.projector.weight = nn.Parameter(torch.eye(dim, device=device, dtype=dtype))
|
| 907 |
+
+ block.projector.bias = nn.Parameter(torch.zeros(dim, device=device, dtype=dtype))
|
| 908 |
+
+
|
| 909 |
+
+ if use_lbm:
|
| 910 |
+
+ # TODO: hard-code for now
|
| 911 |
+
+ self.scheduler = BridgeMatchScheduler()
|
| 912 |
+
+ self.dit.fill_wpmask_with_noise = fill_wpmask_with_noise
|
| 913 |
+
+
|
| 914 |
+
+ self.dit.use_pose_encoder = use_pose_encoder
|
| 915 |
+
+ self.dit.use_cam_encoder = use_cam_encoder
|
| 916 |
+
+ self.dit.use_viewpack = use_viewpack
|
| 917 |
+
+ self.dit.viewpack_dropout_prob = viewpack_dropout_prob
|
| 918 |
+
+ self.dit.use_src_attn = use_src_self_attn or use_src_cross_attn
|
| 919 |
+
+ self.dit.use_lbm = use_lbm
|
| 920 |
+
+ self.dit.freqs_src_shift = freqs_src_shift
|
| 921 |
+
+
|
| 922 |
+
+ def denoising_model(self):
|
| 923 |
+
+ return self.dit
|
| 924 |
+
+
|
| 925 |
+
+ def encode_prompt(self, prompt, positive=True):
|
| 926 |
+
+ prompt_emb = self.prompter.encode_prompt(prompt, positive=positive)
|
| 927 |
+
+ return {"context": prompt_emb}
|
| 928 |
+
+
|
| 929 |
+
+ def encode_image(self, image, num_frames, height, width):
|
| 930 |
+
+ image = self.preprocess_image(image.resize((width, height))).to(self.device)
|
| 931 |
+
+ clip_context = self.image_encoder.encode_image([image])
|
| 932 |
+
+ msk = torch.ones(1, num_frames, height // 8, width // 8, device=self.device)
|
| 933 |
+
+ msk[:, 1:] = 0
|
| 934 |
+
+ msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
| 935 |
+
+ msk = msk.view(1, msk.shape[1] // 4, 4, height // 8, width // 8)
|
| 936 |
+
+ msk = msk.transpose(1, 2)[0]
|
| 937 |
+
+
|
| 938 |
+
+ vae_input = torch.concat(
|
| 939 |
+
+ [image.transpose(0, 1), torch.zeros(3, num_frames - 1, height, width).to(image.device)], dim=1
|
| 940 |
+
+ )
|
| 941 |
+
+ y = self.vae.encode([vae_input.to(dtype=self.torch_dtype, device=self.device)], device=self.device)[0]
|
| 942 |
+
+ y = torch.concat([msk, y])
|
| 943 |
+
+ y = y.unsqueeze(0)
|
| 944 |
+
+ clip_context = clip_context.to(dtype=self.torch_dtype, device=self.device)
|
| 945 |
+
+ y = y.to(dtype=self.torch_dtype, device=self.device)
|
| 946 |
+
+ return {"clip_feature": clip_context, "y": y}
|
| 947 |
+
+
|
| 948 |
+
+ def tensor2video(self, frames):
|
| 949 |
+
+ frames = rearrange(frames, "c f h w -> f h w c")
|
| 950 |
+
+ frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
|
| 951 |
+
+ frames = [Image.fromarray(frame) for frame in frames]
|
| 952 |
+
+ return frames
|
| 953 |
+
+
|
| 954 |
+
+ def prepare_extra_input(self, latents=None):
|
| 955 |
+
+ return {}
|
| 956 |
+
+
|
| 957 |
+
+ def encode_video(self, input_video, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 958 |
+
+ latents = self.vae.encode(
|
| 959 |
+
+ input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride
|
| 960 |
+
+ )
|
| 961 |
+
+ return latents
|
| 962 |
+
+
|
| 963 |
+
+ def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 964 |
+
+ frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
|
| 965 |
+
+ return frames
|
| 966 |
+
+
|
| 967 |
+
+ def encode_fmask(self, mask, size):
|
| 968 |
+
+ f_, h_, w_ = size
|
| 969 |
+
+ g_ = (mask.shape[2] - 1) // (f_ - 1)
|
| 970 |
+
+
|
| 971 |
+
+ # union of the frames in each latent (the first frame is encoded independently)
|
| 972 |
+
+ mask = torch.cat([mask[:, :, :1].repeat(1, 1, g_ - 1, 1, 1), mask], dim=2)
|
| 973 |
+
+ mask = rearrange(mask, "v c (f g) h w -> v c f g h w", f=f_, g=g_)
|
| 974 |
+
+ mask = mask.max(dim=3).values
|
| 975 |
+
+
|
| 976 |
+
+ # interpolate along the spatial dimensions
|
| 977 |
+
+ mask = rearrange(mask, "v c f h w -> (v f) c h w")
|
| 978 |
+
+ mask = F.interpolate(mask, size=(h_, w_), mode="area")
|
| 979 |
+
+ mask = rearrange(mask, "(v f) c h w -> v c f h w", f=f_)
|
| 980 |
+
+ return mask
|
| 981 |
+
+
|
| 982 |
+
+ def encode_wpmask(self, mask, size):
|
| 983 |
+
+ f_, h_, w_ = size
|
| 984 |
+
+ g_ = (mask.shape[2] - 1) // (f_ - 1)
|
| 985 |
+
+
|
| 986 |
+
+ # intersection of the frames in each latent (the first frame is encoded independently)
|
| 987 |
+
+ mask = torch.cat([mask[:, :, :1].repeat(1, 1, g_ - 1, 1, 1), mask], dim=2)
|
| 988 |
+
+ mask = rearrange(mask, "v c (f g) h w -> v c f g h w", f=f_, g=g_)
|
| 989 |
+
+ mask = mask.min(dim=3).values
|
| 990 |
+
+
|
| 991 |
+
+ # interpolate along the spatial dimensions
|
| 992 |
+
+ mask = rearrange(mask, "v c f h w -> (v f) c h w")
|
| 993 |
+
+ mask = F.interpolate(mask, size=(h_, w_), mode="area")
|
| 994 |
+
+ mask = rearrange(mask, "(v f) c h w -> v c f h w", f=f_)
|
| 995 |
+
+ return mask
|
| 996 |
+
+
|
| 997 |
+
+ @torch.no_grad()
|
| 998 |
+
+ def __call__(
|
| 999 |
+
+ self,
|
| 1000 |
+
+ prompt,
|
| 1001 |
+
+ negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
| 1002 |
+
+ src_videos: torch.Tensor = None,
|
| 1003 |
+
+ skeletons: torch.Tensor = None,
|
| 1004 |
+
+ wpvideos: torch.Tensor = None,
|
| 1005 |
+
+ wpmasks: torch.Tensor = None,
|
| 1006 |
+
+ cam_emb: torch.Tensor = None,
|
| 1007 |
+
+ input_image: Image.Image = None,
|
| 1008 |
+
+ input_video: torch.Tensor = None,
|
| 1009 |
+
+ denoising_strength: float = 1.0,
|
| 1010 |
+
+ seed: int = None,
|
| 1011 |
+
+ rand_device: str = "cpu",
|
| 1012 |
+
+ height: int = 832,
|
| 1013 |
+
+ width: int = 480,
|
| 1014 |
+
+ num_frames: int = None,
|
| 1015 |
+
+ cfg_scale: float = 5.0,
|
| 1016 |
+
+ num_inference_steps: int = 50,
|
| 1017 |
+
+ sigma_shift: float = 5.0,
|
| 1018 |
+
+ tiled: bool = True,
|
| 1019 |
+
+ tile_size: tuple[int, int] = (52, 30),
|
| 1020 |
+
+ tile_stride: tuple[int, int] = (26, 15),
|
| 1021 |
+
+ tea_cache_l1_thresh: float = None,
|
| 1022 |
+
+ tea_cache_model_id: str = "",
|
| 1023 |
+
+ progress_bar_cmd=partial(tqdm, desc="Denoising"),
|
| 1024 |
+
+ progress_bar_st=None,
|
| 1025 |
+
+ return_tensor=False,
|
| 1026 |
+
+ ):
|
| 1027 |
+
+ # breakpoint()
|
| 1028 |
+
+ assert num_frames is None, "num_frames is not supported for WanVideoSpaTemPipeline"
|
| 1029 |
+
+ assert input_image is None, "input_image is not supported for WanVideoSpaTemPipeline"
|
| 1030 |
+
+ assert input_video is None, "input_video is not supported for WanVideoSpaTemPipeline"
|
| 1031 |
+
+ assert tea_cache_l1_thresh is None, "tea_cache_l1_thresh is not supported for WanVideoSpaTemPipeline"
|
| 1032 |
+
+
|
| 1033 |
+
+ # Parameter check
|
| 1034 |
+
+ height, width = self.check_resize_height_width(height, width)
|
| 1035 |
+
+
|
| 1036 |
+
+ # Tiler parameters
|
| 1037 |
+
+ tiler_kwargs = {"tiled": tiled, "tile_size": tile_size, "tile_stride": tile_stride}
|
| 1038 |
+
+
|
| 1039 |
+
+ # Scheduler
|
| 1040 |
+
+ if self.dit.use_lbm:
|
| 1041 |
+
+ # bridge matching scheduler
|
| 1042 |
+
+ self.scheduler.set_timesteps(num_inference_steps)
|
| 1043 |
+
+ else:
|
| 1044 |
+
+ # flow matching scheduler
|
| 1045 |
+
+ self.scheduler.set_timesteps(num_inference_steps, denoising_strength=denoising_strength, shift=sigma_shift)
|
| 1046 |
+
+
|
| 1047 |
+
+ src_videos = src_videos.to(dtype=self.torch_dtype, device=self.device)
|
| 1048 |
+
+ if skeletons is not None:
|
| 1049 |
+
+ skeletons = skeletons.to(dtype=self.torch_dtype, device=self.device)
|
| 1050 |
+
+ if wpvideos is not None:
|
| 1051 |
+
+ wpvideos = wpvideos.to(dtype=self.torch_dtype, device=self.device)
|
| 1052 |
+
+ if wpmasks is not None:
|
| 1053 |
+
+ wpmasks = wpmasks.to(dtype=self.torch_dtype, device=self.device)
|
| 1054 |
+
+ if cam_emb is not None:
|
| 1055 |
+
+ cam_emb = cam_emb.to(dtype=self.torch_dtype, device=self.device)
|
| 1056 |
+
+
|
| 1057 |
+
+ if skeletons is not None:
|
| 1058 |
+
+ num_cameras = skeletons.shape[0]
|
| 1059 |
+
+ num_frames = skeletons.shape[2]
|
| 1060 |
+
+ elif wpvideos is not None:
|
| 1061 |
+
+ num_cameras = wpvideos.shape[0]
|
| 1062 |
+
+ num_frames = wpvideos.shape[2]
|
| 1063 |
+
+ else:
|
| 1064 |
+
+ raise ValueError("Either skeletons or wpvideos must be provided")
|
| 1065 |
+
+
|
| 1066 |
+
+ # Initialize noise
|
| 1067 |
+
+ noise_shape = (
|
| 1068 |
+
+ num_cameras,
|
| 1069 |
+
+ self.vae.model.z_dim,
|
| 1070 |
+
+ (num_frames - 1) // 4 + 1,
|
| 1071 |
+
+ height // self.vae.upsampling_factor,
|
| 1072 |
+
+ width // self.vae.upsampling_factor,
|
| 1073 |
+
+ )
|
| 1074 |
+
+ noise = self.generate_noise(noise_shape, seed=seed, device=rand_device, dtype=torch.float32)
|
| 1075 |
+
+ noise = noise.to(dtype=self.torch_dtype, device=self.device)
|
| 1076 |
+
+
|
| 1077 |
+
+ if input_video is not None:
|
| 1078 |
+
+ self.load_models_to_device(["vae"])
|
| 1079 |
+
+ input_video = self.preprocess_images(input_video)
|
| 1080 |
+
+ input_video = torch.stack(input_video, dim=2).to(dtype=self.torch_dtype, device=self.device)
|
| 1081 |
+
+ latents = self.encode_video(input_video, **tiler_kwargs).to(dtype=self.torch_dtype, device=self.device)
|
| 1082 |
+
+ latents = self.scheduler.add_noise(latents, noise, timestep=self.scheduler.timesteps[0])
|
| 1083 |
+
+ else:
|
| 1084 |
+
+ latents = noise
|
| 1085 |
+
+
|
| 1086 |
+
+ # Encode source video
|
| 1087 |
+
+ self.load_models_to_device(["vae"])
|
| 1088 |
+
+ src_latents = self.encode_video(src_videos, **tiler_kwargs)
|
| 1089 |
+
+ src_latents = src_latents.to(dtype=self.torch_dtype, device=self.device)
|
| 1090 |
+
+ src_latents_nega = torch.zeros_like(src_latents)
|
| 1091 |
+
+
|
| 1092 |
+
+ # Latent bridge matching
|
| 1093 |
+
+ if self.dit.use_lbm:
|
| 1094 |
+
+ if skeletons is not None:
|
| 1095 |
+
+ # skeleton-based: use primary src_latents as bridge source
|
| 1096 |
+
+ lbm_src_latents = src_latents[:1].expand_as(latents)
|
| 1097 |
+
+ elif wpvideos is not None:
|
| 1098 |
+
+ lbm_src_latents = self.encode_video(wpvideos, **tiler_kwargs).to(
|
| 1099 |
+
+ dtype=self.torch_dtype, device=self.device
|
| 1100 |
+
+ )
|
| 1101 |
+
+ if self.dit.fill_wpmask_with_noise:
|
| 1102 |
+
+ wpmask_latents = self.encode_wpmask(wpmasks, size=lbm_src_latents.shape[-3:])
|
| 1103 |
+
+ lbm_src_latents = lbm_src_latents * wpmask_latents + noise * (1 - wpmask_latents)
|
| 1104 |
+
+
|
| 1105 |
+
+ latents = lbm_src_latents
|
| 1106 |
+
+
|
| 1107 |
+
+ # Encode prompts
|
| 1108 |
+
+ self.load_models_to_device(["text_encoder"])
|
| 1109 |
+
+ prompt_emb_posi = self.encode_prompt(prompt, positive=True)
|
| 1110 |
+
+ if cfg_scale != 1.0:
|
| 1111 |
+
+ prompt_emb_nega = self.encode_prompt(negative_prompt, positive=False)
|
| 1112 |
+
+
|
| 1113 |
+
+ # Encode image
|
| 1114 |
+
+ if input_image is not None and self.image_encoder is not None:
|
| 1115 |
+
+ self.load_models_to_device(["image_encoder", "vae"])
|
| 1116 |
+
+ image_emb = self.encode_image(input_image, num_frames, height, width)
|
| 1117 |
+
+ else:
|
| 1118 |
+
+ image_emb = {}
|
| 1119 |
+
+
|
| 1120 |
+
+ # Extra input
|
| 1121 |
+
+ extra_input = self.prepare_extra_input(latents)
|
| 1122 |
+
+
|
| 1123 |
+
+ # Denoise
|
| 1124 |
+
+ self.load_models_to_device(["dit"])
|
| 1125 |
+
+ for progress_id, timestep in enumerate(progress_bar_cmd(self.scheduler.timesteps)):
|
| 1126 |
+
+ timestep = timestep.unsqueeze(0).to(dtype=self.torch_dtype, device=self.device)
|
| 1127 |
+
+ timestep = torch.cat([timestep] * num_cameras, dim=0)
|
| 1128 |
+
+
|
| 1129 |
+
+ # Inference
|
| 1130 |
+
+ noise_pred_posi = self.denoising_model()(
|
| 1131 |
+
+ x=latents,
|
| 1132 |
+
+ x_src=src_latents,
|
| 1133 |
+
+ timestep=timestep,
|
| 1134 |
+
+ skeletons=skeletons,
|
| 1135 |
+
+ cam_emb=cam_emb,
|
| 1136 |
+
+ **prompt_emb_posi,
|
| 1137 |
+
+ **image_emb,
|
| 1138 |
+
+ **extra_input,
|
| 1139 |
+
+ )
|
| 1140 |
+
+ if cfg_scale != 1.0:
|
| 1141 |
+
+ noise_pred_nega = self.denoising_model()(
|
| 1142 |
+
+ x=latents,
|
| 1143 |
+
+ x_src=src_latents_nega,
|
| 1144 |
+
+ timestep=timestep,
|
| 1145 |
+
+ skeletons=skeletons,
|
| 1146 |
+
+ cam_emb=cam_emb,
|
| 1147 |
+
+ **prompt_emb_nega,
|
| 1148 |
+
+ **image_emb,
|
| 1149 |
+
+ **extra_input,
|
| 1150 |
+
+ )
|
| 1151 |
+
+ noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)
|
| 1152 |
+
+ else:
|
| 1153 |
+
+ noise_pred = noise_pred_posi
|
| 1154 |
+
+
|
| 1155 |
+
+ # Scheduler
|
| 1156 |
+
+ latents = self.scheduler.step(noise_pred, self.scheduler.timesteps[progress_id], latents)
|
| 1157 |
+
+
|
| 1158 |
+
+ # Decode
|
| 1159 |
+
+ self.load_models_to_device(["vae"])
|
| 1160 |
+
+ pred_videos = self.decode_video(latents, **tiler_kwargs)
|
| 1161 |
+
+
|
| 1162 |
+
+ if return_tensor:
|
| 1163 |
+
+ return pred_videos
|
| 1164 |
+
+
|
| 1165 |
+
+ self.load_models_to_device([])
|
| 1166 |
+
+ pred_video_list = []
|
| 1167 |
+
+ for pred_video in pred_videos:
|
| 1168 |
+
+ pred_video_list.append(self.tensor2video(pred_video))
|
| 1169 |
+
+ return pred_video_list
|
| 1170 |
+
+
|
| 1171 |
+
+
|
| 1172 |
+
+class TeaCache:
|
| 1173 |
+
+ def __init__(self, num_inference_steps, rel_l1_thresh, model_id):
|
| 1174 |
+
+ self.num_inference_steps = num_inference_steps
|
| 1175 |
+
+ self.step = 0
|
| 1176 |
+
+ self.accumulated_rel_l1_distance = 0
|
| 1177 |
+
+ self.previous_modulated_input = None
|
| 1178 |
+
+ self.rel_l1_thresh = rel_l1_thresh
|
| 1179 |
+
+ self.previous_residual = None
|
| 1180 |
+
+ self.previous_hidden_states = None
|
| 1181 |
+
+
|
| 1182 |
+
+ self.coefficients_dict = {
|
| 1183 |
+
+ "Wan2.1-T2V-1.3B": [-5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02],
|
| 1184 |
+
+ "Wan2.1-T2V-14B": [-3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01],
|
| 1185 |
+
+ "Wan2.1-I2V-14B-480P": [2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01],
|
| 1186 |
+
+ "Wan2.1-I2V-14B-720P": [8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02],
|
| 1187 |
+
+ }
|
| 1188 |
+
+ if model_id not in self.coefficients_dict:
|
| 1189 |
+
+ supported_model_ids = ", ".join([i for i in self.coefficients_dict])
|
| 1190 |
+
+ raise ValueError(
|
| 1191 |
+
+ f"{model_id} is not a supported TeaCache model id. Please choose a valid model id in ({supported_model_ids})."
|
| 1192 |
+
+ )
|
| 1193 |
+
+ self.coefficients = self.coefficients_dict[model_id]
|
| 1194 |
+
+
|
| 1195 |
+
+ def check(self, dit: WanModel, x, t_mod):
|
| 1196 |
+
+ modulated_inp = t_mod.clone()
|
| 1197 |
+
+ if self.step == 0 or self.step == self.num_inference_steps - 1:
|
| 1198 |
+
+ should_calc = True
|
| 1199 |
+
+ self.accumulated_rel_l1_distance = 0
|
| 1200 |
+
+ else:
|
| 1201 |
+
+ coefficients = self.coefficients
|
| 1202 |
+
+ rescale_func = np.poly1d(coefficients)
|
| 1203 |
+
+ self.accumulated_rel_l1_distance += rescale_func(
|
| 1204 |
+
+ (
|
| 1205 |
+
+ (modulated_inp - self.previous_modulated_input).abs().mean()
|
| 1206 |
+
+ / self.previous_modulated_input.abs().mean()
|
| 1207 |
+
+ )
|
| 1208 |
+
+ .cpu()
|
| 1209 |
+
+ .item()
|
| 1210 |
+
+ )
|
| 1211 |
+
+ if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
|
| 1212 |
+
+ should_calc = False
|
| 1213 |
+
+ else:
|
| 1214 |
+
+ should_calc = True
|
| 1215 |
+
+ self.accumulated_rel_l1_distance = 0
|
| 1216 |
+
+ self.previous_modulated_input = modulated_inp
|
| 1217 |
+
+ self.step += 1
|
| 1218 |
+
+ if self.step == self.num_inference_steps:
|
| 1219 |
+
+ self.step = 0
|
| 1220 |
+
+ if should_calc:
|
| 1221 |
+
+ self.previous_hidden_states = x.clone()
|
| 1222 |
+
+ return not should_calc
|
| 1223 |
+
+
|
| 1224 |
+
+ def store(self, hidden_states):
|
| 1225 |
+
+ self.previous_residual = hidden_states - self.previous_hidden_states
|
| 1226 |
+
+ self.previous_hidden_states = None
|
| 1227 |
+
+
|
| 1228 |
+
+ def update(self, hidden_states):
|
| 1229 |
+
+ hidden_states = hidden_states + self.previous_residual
|
| 1230 |
+
+ return hidden_states
|
| 1231 |
+
diff --git a/diffsynth/schedulers/__init__.py b/diffsynth/schedulers/__init__.py
|
| 1232 |
+
index 0ec4325..03cae00 100644
|
| 1233 |
+
--- a/diffsynth/schedulers/__init__.py
|
| 1234 |
+
+++ b/diffsynth/schedulers/__init__.py
|
| 1235 |
+
@@ -1,3 +1,4 @@
|
| 1236 |
+
from .ddim import EnhancedDDIMScheduler
|
| 1237 |
+
from .continuous_ode import ContinuousODEScheduler
|
| 1238 |
+
from .flow_match import FlowMatchScheduler
|
| 1239 |
+
+from .bride_match import BridgeMatchScheduler
|
| 1240 |
+
diff --git a/diffsynth/schedulers/bride_match.py b/diffsynth/schedulers/bride_match.py
|
| 1241 |
+
new file mode 100644
|
| 1242 |
+
index 0000000..c157702
|
| 1243 |
+
--- /dev/null
|
| 1244 |
+
+++ b/diffsynth/schedulers/bride_match.py
|
| 1245 |
+
@@ -0,0 +1,71 @@
|
| 1246 |
+
+import torch, math
|
| 1247 |
+
+
|
| 1248 |
+
+
|
| 1249 |
+
+class BridgeMatchScheduler:
|
| 1250 |
+
+
|
| 1251 |
+
+ def __init__(
|
| 1252 |
+
+ self,
|
| 1253 |
+
+ num_train_timesteps=1000,
|
| 1254 |
+
+ num_inference_steps=8,
|
| 1255 |
+
+ sigma_max=1.0,
|
| 1256 |
+
+ bridge_noise_sigma=0.005,
|
| 1257 |
+
+ ):
|
| 1258 |
+
+ self.sigma_max = sigma_max
|
| 1259 |
+
+ self.sigma_min = sigma_max / num_train_timesteps
|
| 1260 |
+
+ self.num_train_timesteps = num_train_timesteps # train timesteps for base model
|
| 1261 |
+
+ self.num_inference_steps = num_inference_steps # inference steps for bridge matching
|
| 1262 |
+
+ self.bridge_noise_sigma = bridge_noise_sigma
|
| 1263 |
+
+
|
| 1264 |
+
+ self.set_timesteps(self.num_inference_steps)
|
| 1265 |
+
+
|
| 1266 |
+
+ def set_timesteps(self, num_inference_steps=8, training=False):
|
| 1267 |
+
+ sigma_start = self.sigma_max
|
| 1268 |
+
+ sigma_end = self.sigma_max / num_inference_steps
|
| 1269 |
+
+ self.sigmas = torch.linspace(sigma_start, sigma_end, num_inference_steps)
|
| 1270 |
+
+ self.timesteps = self.sigmas * self.num_train_timesteps
|
| 1271 |
+
+
|
| 1272 |
+
+ self.training = training
|
| 1273 |
+
+
|
| 1274 |
+
+ def retrieve_sigma(self, timestep):
|
| 1275 |
+
+ if isinstance(timestep, torch.Tensor):
|
| 1276 |
+
+ timestep = timestep.cpu()
|
| 1277 |
+
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
|
| 1278 |
+
+ sigma = self.sigmas[timestep_id]
|
| 1279 |
+
+ return sigma
|
| 1280 |
+
+
|
| 1281 |
+
+ def get_noise_term(self, sigma, sample):
|
| 1282 |
+
+ # bridge noise term == 0 when sigma == 1 or 0
|
| 1283 |
+
+ return self.bridge_noise_sigma * (sigma * (1.0 - sigma)) ** 0.5 * torch.randn_like(sample)
|
| 1284 |
+
+
|
| 1285 |
+
+ def step(self, model_output, timestep, sample, to_final=False, **kwargs):
|
| 1286 |
+
+ if isinstance(timestep, torch.Tensor):
|
| 1287 |
+
+ timestep = timestep.cpu()
|
| 1288 |
+
+ timestep_id = torch.argmin((self.timesteps - timestep).abs())
|
| 1289 |
+
+ sigma = self.sigmas[timestep_id]
|
| 1290 |
+
+ if to_final or timestep_id + 1 >= len(self.timesteps):
|
| 1291 |
+
+ sigma_ = 0
|
| 1292 |
+
+ else:
|
| 1293 |
+
+ sigma_ = self.sigmas[timestep_id + 1]
|
| 1294 |
+
+
|
| 1295 |
+
+ prev_sample = sample + model_output * (sigma_ - sigma) + self.get_noise_term(sigma_, sample)
|
| 1296 |
+
+ return prev_sample
|
| 1297 |
+
+
|
| 1298 |
+
+ def add_noise(self, tgt_sample, src_sample, timestep):
|
| 1299 |
+
+ sigma = self.retrieve_sigma(timestep)
|
| 1300 |
+
+ noisy_sample = sigma * src_sample + (1 - sigma) * tgt_sample + self.get_noise_term(sigma, tgt_sample)
|
| 1301 |
+
+ return noisy_sample
|
| 1302 |
+
+
|
| 1303 |
+
+ def training_target(self, tgt_sample, noisy_sample, timestep):
|
| 1304 |
+
+ sigma = self.retrieve_sigma(timestep)
|
| 1305 |
+
+
|
| 1306 |
+
+ target = (noisy_sample - tgt_sample) / sigma
|
| 1307 |
+
+ return target
|
| 1308 |
+
+
|
| 1309 |
+
+ def denoised_sample(self, prediction, noisy_sample, timestep):
|
| 1310 |
+
+ sigma = self.retrieve_sigma(timestep)
|
| 1311 |
+
+
|
| 1312 |
+
+ sample = noisy_sample - prediction * sigma
|
| 1313 |
+
+ return sample
|
| 1314 |
+
+
|
| 1315 |
+
+ def training_weight(self, timestep):
|
| 1316 |
+
+ return 1.0
|
| 1317 |
+
diff --git a/diffsynth/schedulers/flow_match.py b/diffsynth/schedulers/flow_match.py
|
| 1318 |
+
index 6a8e235..0cd3e08 100644
|
| 1319 |
+
--- a/diffsynth/schedulers/flow_match.py
|
| 1320 |
+
+++ b/diffsynth/schedulers/flow_match.py
|
| 1321 |
+
@@ -98,7 +98,12 @@ class FlowMatchScheduler():
|
| 1322 |
+
def training_target(self, sample, noise, timestep):
|
| 1323 |
+
target = noise - sample
|
| 1324 |
+
return target
|
| 1325 |
+
-
|
| 1326 |
+
+
|
| 1327 |
+
+
|
| 1328 |
+
+ def denoised_sample(self, prediction, noise, timestep):
|
| 1329 |
+
+ sample = noise - prediction
|
| 1330 |
+
+ return sample
|
| 1331 |
+
+
|
| 1332 |
+
|
| 1333 |
+
def training_weight(self, timestep):
|
| 1334 |
+
timestep_id = torch.argmin((self.timesteps - timestep.to(self.timesteps.device)).abs())
|
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Paths are relative to fdanyone/vendor/diffsynth.
|
| 2 |
+
LICENSE
|
| 3 |
+
UPSTREAM.md
|
| 4 |
+
UPSTREAM.patch
|
| 5 |
+
VENDORED_FILES.txt
|
| 6 |
+
__init__.py
|
| 7 |
+
models/__init__.py
|
| 8 |
+
models/utils.py
|
| 9 |
+
models/wan_video_dit.py
|
| 10 |
+
models/wan_video_pose_encoder.py
|
| 11 |
+
models/wan_video_text_encoder.py
|
| 12 |
+
models/wan_video_vae.py
|
| 13 |
+
pipelines/__init__.py
|
| 14 |
+
pipelines/base.py
|
| 15 |
+
pipelines/wan_video_spatem.py
|
| 16 |
+
prompters/__init__.py
|
| 17 |
+
prompters/base_prompter.py
|
| 18 |
+
prompters/wan_prompter.py
|
| 19 |
+
schedulers/__init__.py
|
| 20 |
+
schedulers/flow_match.py
|
| 21 |
+
vram_management/__init__.py
|
| 22 |
+
vram_management/layers.py
|
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Minimal DiffSynth-Studio inference closure used by 4DAnyone."""
|
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Vendored inference pipelines."""
|
| 2 |
+
|
| 3 |
+
from .wan_video_spatem import WanVideoSpaTemPipeline
|
| 4 |
+
|
| 5 |
+
__all__ = ["WanVideoSpaTemPipeline"]
|
|
@@ -0,0 +1,127 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
from PIL import Image
|
| 4 |
+
from torchvision.transforms import GaussianBlur
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class BasePipeline(torch.nn.Module):
|
| 9 |
+
|
| 10 |
+
def __init__(self, device="cuda", torch_dtype=torch.float16, height_division_factor=64, width_division_factor=64):
|
| 11 |
+
super().__init__()
|
| 12 |
+
self.device = device
|
| 13 |
+
self.torch_dtype = torch_dtype
|
| 14 |
+
self.height_division_factor = height_division_factor
|
| 15 |
+
self.width_division_factor = width_division_factor
|
| 16 |
+
self.cpu_offload = False
|
| 17 |
+
self.model_names = []
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def check_resize_height_width(self, height, width):
|
| 21 |
+
if height % self.height_division_factor != 0:
|
| 22 |
+
height = (height + self.height_division_factor - 1) // self.height_division_factor * self.height_division_factor
|
| 23 |
+
print(f"The height cannot be evenly divided by {self.height_division_factor}. We round it up to {height}.")
|
| 24 |
+
if width % self.width_division_factor != 0:
|
| 25 |
+
width = (width + self.width_division_factor - 1) // self.width_division_factor * self.width_division_factor
|
| 26 |
+
print(f"The width cannot be evenly divided by {self.width_division_factor}. We round it up to {width}.")
|
| 27 |
+
return height, width
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def preprocess_image(self, image):
|
| 31 |
+
image = torch.Tensor(np.array(image, dtype=np.float32) * (2 / 255) - 1).permute(2, 0, 1).unsqueeze(0)
|
| 32 |
+
return image
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def preprocess_images(self, images):
|
| 36 |
+
return [self.preprocess_image(image) for image in images]
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def vae_output_to_image(self, vae_output):
|
| 40 |
+
image = vae_output[0].cpu().float().permute(1, 2, 0).numpy()
|
| 41 |
+
image = Image.fromarray(((image / 2 + 0.5).clip(0, 1) * 255).astype("uint8"))
|
| 42 |
+
return image
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def vae_output_to_video(self, vae_output):
|
| 46 |
+
video = vae_output.cpu().permute(1, 2, 0).numpy()
|
| 47 |
+
video = [Image.fromarray(((image / 2 + 0.5).clip(0, 1) * 255).astype("uint8")) for image in video]
|
| 48 |
+
return video
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def merge_latents(self, value, latents, masks, scales, blur_kernel_size=33, blur_sigma=10.0):
|
| 52 |
+
if len(latents) > 0:
|
| 53 |
+
blur = GaussianBlur(kernel_size=blur_kernel_size, sigma=blur_sigma)
|
| 54 |
+
height, width = value.shape[-2:]
|
| 55 |
+
weight = torch.ones_like(value)
|
| 56 |
+
for latent, mask, scale in zip(latents, masks, scales):
|
| 57 |
+
mask = self.preprocess_image(mask.resize((width, height))).mean(dim=1, keepdim=True) > 0
|
| 58 |
+
mask = mask.repeat(1, latent.shape[1], 1, 1).to(dtype=latent.dtype, device=latent.device)
|
| 59 |
+
mask = blur(mask)
|
| 60 |
+
value += latent * mask * scale
|
| 61 |
+
weight += mask * scale
|
| 62 |
+
value /= weight
|
| 63 |
+
return value
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
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):
|
| 67 |
+
if special_kwargs is None:
|
| 68 |
+
noise_pred_global = inference_callback(prompt_emb_global)
|
| 69 |
+
else:
|
| 70 |
+
noise_pred_global = inference_callback(prompt_emb_global, special_kwargs)
|
| 71 |
+
if special_local_kwargs_list is None:
|
| 72 |
+
noise_pred_locals = [inference_callback(prompt_emb_local) for prompt_emb_local in prompt_emb_locals]
|
| 73 |
+
else:
|
| 74 |
+
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)]
|
| 75 |
+
noise_pred = self.merge_latents(noise_pred_global, noise_pred_locals, masks, mask_scales)
|
| 76 |
+
return noise_pred
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def extend_prompt(self, prompt, local_prompts, masks, mask_scales):
|
| 80 |
+
local_prompts = local_prompts or []
|
| 81 |
+
masks = masks or []
|
| 82 |
+
mask_scales = mask_scales or []
|
| 83 |
+
extended_prompt_dict = self.prompter.extend_prompt(prompt)
|
| 84 |
+
prompt = extended_prompt_dict.get("prompt", prompt)
|
| 85 |
+
local_prompts += extended_prompt_dict.get("prompts", [])
|
| 86 |
+
masks += extended_prompt_dict.get("masks", [])
|
| 87 |
+
mask_scales += [100.0] * len(extended_prompt_dict.get("masks", []))
|
| 88 |
+
return prompt, local_prompts, masks, mask_scales
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def enable_cpu_offload(self):
|
| 92 |
+
self.cpu_offload = True
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def load_models_to_device(self, loadmodel_names=[]):
|
| 96 |
+
# only load models to device if cpu_offload is enabled
|
| 97 |
+
if not self.cpu_offload:
|
| 98 |
+
return
|
| 99 |
+
# offload the unneeded models to cpu
|
| 100 |
+
for model_name in self.model_names:
|
| 101 |
+
if model_name not in loadmodel_names:
|
| 102 |
+
model = getattr(self, model_name)
|
| 103 |
+
if model is not None:
|
| 104 |
+
if hasattr(model, "vram_management_enabled") and model.vram_management_enabled:
|
| 105 |
+
for module in model.modules():
|
| 106 |
+
if hasattr(module, "offload"):
|
| 107 |
+
module.offload()
|
| 108 |
+
else:
|
| 109 |
+
model.cpu()
|
| 110 |
+
# load the needed models to device
|
| 111 |
+
for model_name in loadmodel_names:
|
| 112 |
+
model = getattr(self, model_name)
|
| 113 |
+
if model is not None:
|
| 114 |
+
if hasattr(model, "vram_management_enabled") and model.vram_management_enabled:
|
| 115 |
+
for module in model.modules():
|
| 116 |
+
if hasattr(module, "onload"):
|
| 117 |
+
module.onload()
|
| 118 |
+
else:
|
| 119 |
+
model.to(self.device)
|
| 120 |
+
# fresh the cuda cache
|
| 121 |
+
torch.cuda.empty_cache()
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
def generate_noise(self, shape, seed=None, device="cpu", dtype=torch.float16):
|
| 125 |
+
generator = None if seed is None else torch.Generator(device).manual_seed(seed)
|
| 126 |
+
noise = torch.randn(shape, generator=generator, device=device, dtype=dtype)
|
| 127 |
+
return noise
|
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
|
| 4 |
+
from ..models.wan_video_dit import SelfAttention, ViewPackEmbedding, WanModel
|
| 5 |
+
from ..models.wan_video_pose_encoder import PoseEncoder
|
| 6 |
+
from ..models.wan_video_text_encoder import WanTextEncoder
|
| 7 |
+
from ..models.wan_video_vae import WanVideoVAE
|
| 8 |
+
from ..pipelines.base import BasePipeline
|
| 9 |
+
from ..prompters.wan_prompter import WanPrompter
|
| 10 |
+
from ..schedulers.flow_match import FlowMatchScheduler
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class WanVideoSpaTemPipeline(BasePipeline):
|
| 14 |
+
def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None):
|
| 15 |
+
super().__init__(device=device, torch_dtype=torch_dtype)
|
| 16 |
+
self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True)
|
| 17 |
+
self.prompter = WanPrompter(tokenizer_path=tokenizer_path)
|
| 18 |
+
self.text_encoder: WanTextEncoder = None
|
| 19 |
+
self.dit: WanModel = None
|
| 20 |
+
self.vae: WanVideoVAE = None
|
| 21 |
+
self.model_names = ["text_encoder", "dit", "vae"]
|
| 22 |
+
self.height_division_factor = 16
|
| 23 |
+
self.width_division_factor = 16
|
| 24 |
+
|
| 25 |
+
def init_spatem_modules(
|
| 26 |
+
self,
|
| 27 |
+
use_mvs_attn: bool = False,
|
| 28 |
+
use_viewpack: bool = True,
|
| 29 |
+
viewpack_dropout_prob: float = 0.0,
|
| 30 |
+
use_pose_encoder: bool = True,
|
| 31 |
+
pose_encoder_type: str = "rgb",
|
| 32 |
+
range_mvs_attn: tuple[int, int, int] = (1, None, 2),
|
| 33 |
+
):
|
| 34 |
+
device = self.dit.patch_embedding.weight.device
|
| 35 |
+
dtype = self.dit.patch_embedding.weight.dtype
|
| 36 |
+
|
| 37 |
+
if use_mvs_attn:
|
| 38 |
+
begin, end, stride = range_mvs_attn
|
| 39 |
+
for block in self.dit.blocks[begin:end:stride]:
|
| 40 |
+
block.use_mvs_attn = True
|
| 41 |
+
dim = block.self_attn.q.weight.shape[0]
|
| 42 |
+
block.modulation_mvs = nn.Parameter(
|
| 43 |
+
block.modulation[:, :3, :].detach().clone()
|
| 44 |
+
)
|
| 45 |
+
block.norm1_mvs = nn.LayerNorm(
|
| 46 |
+
dim,
|
| 47 |
+
eps=block.norm1.eps,
|
| 48 |
+
elementwise_affine=False,
|
| 49 |
+
).to(device=device, dtype=dtype)
|
| 50 |
+
block.self_attn_mvs = SelfAttention(
|
| 51 |
+
dim,
|
| 52 |
+
block.self_attn.num_heads,
|
| 53 |
+
block.self_attn.norm_q.eps,
|
| 54 |
+
).to(device=device, dtype=dtype)
|
| 55 |
+
block.self_attn_mvs.load_state_dict(
|
| 56 |
+
block.self_attn.state_dict(),
|
| 57 |
+
strict=True,
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
if not 0.0 <= viewpack_dropout_prob <= 1.0:
|
| 61 |
+
raise ValueError("viewpack_dropout_prob should be between 0 and 1")
|
| 62 |
+
if viewpack_dropout_prob > 0.0 and not use_viewpack:
|
| 63 |
+
raise ValueError("viewpack_dropout_prob requires use_viewpack=True")
|
| 64 |
+
if use_viewpack:
|
| 65 |
+
viewpack_emb = ViewPackEmbedding(
|
| 66 |
+
in_dim=self.dit.patch_embedding.weight.shape[1],
|
| 67 |
+
dim=self.dit.patch_embedding.weight.shape[0],
|
| 68 |
+
patch_size=list(self.dit.patch_embedding.kernel_size),
|
| 69 |
+
)
|
| 70 |
+
viewpack_emb.initialize_from_patch_embedding(self.dit.patch_embedding)
|
| 71 |
+
self.dit.viewpack_embedding = viewpack_emb.to(
|
| 72 |
+
device=device,
|
| 73 |
+
dtype=dtype,
|
| 74 |
+
)
|
| 75 |
+
|
| 76 |
+
if use_pose_encoder:
|
| 77 |
+
if pose_encoder_type != "rgb":
|
| 78 |
+
raise ValueError(f"Invalid pose_encoder_type: {pose_encoder_type}")
|
| 79 |
+
pose_encoder = PoseEncoder(
|
| 80 |
+
out_dim=self.dit.patch_embedding.out_channels,
|
| 81 |
+
in_channels=3,
|
| 82 |
+
)
|
| 83 |
+
self.dit.pose_encoder = pose_encoder.to(device=device, dtype=dtype)
|
| 84 |
+
|
| 85 |
+
self.dit.use_pose_encoder = use_pose_encoder
|
| 86 |
+
self.dit.use_viewpack = use_viewpack
|
| 87 |
+
self.dit.viewpack_dropout_prob = viewpack_dropout_prob
|
| 88 |
+
|
| 89 |
+
def encode_video(self, input_video):
|
| 90 |
+
return self.vae.encode(input_video, device=self.device)
|
| 91 |
+
|
| 92 |
+
def decode_video(self, latents):
|
| 93 |
+
return self.vae.decode(latents, device=self.device)
|
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Vendored Wan prompt helpers."""
|
| 2 |
+
|
| 3 |
+
from .wan_prompter import WanPrompter
|
| 4 |
+
|
| 5 |
+
__all__ = ["WanPrompter"]
|
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def tokenize_long_prompt(tokenizer, prompt, max_length=None):
|
| 6 |
+
# Get model_max_length from self.tokenizer
|
| 7 |
+
length = tokenizer.model_max_length if max_length is None else max_length
|
| 8 |
+
|
| 9 |
+
# To avoid the warning. set self.tokenizer.model_max_length to +oo.
|
| 10 |
+
tokenizer.model_max_length = 99999999
|
| 11 |
+
|
| 12 |
+
# Tokenize it!
|
| 13 |
+
input_ids = tokenizer(prompt, return_tensors="pt").input_ids
|
| 14 |
+
|
| 15 |
+
# Determine the real length.
|
| 16 |
+
max_length = (input_ids.shape[1] + length - 1) // length * length
|
| 17 |
+
|
| 18 |
+
# Restore tokenizer.model_max_length
|
| 19 |
+
tokenizer.model_max_length = length
|
| 20 |
+
|
| 21 |
+
# Tokenize it again with fixed length.
|
| 22 |
+
input_ids = tokenizer(
|
| 23 |
+
prompt,
|
| 24 |
+
return_tensors="pt",
|
| 25 |
+
padding="max_length",
|
| 26 |
+
max_length=max_length,
|
| 27 |
+
truncation=True
|
| 28 |
+
).input_ids
|
| 29 |
+
|
| 30 |
+
# Reshape input_ids to fit the text encoder.
|
| 31 |
+
num_sentence = input_ids.shape[1] // length
|
| 32 |
+
input_ids = input_ids.reshape((num_sentence, length))
|
| 33 |
+
|
| 34 |
+
return input_ids
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class BasePrompter:
|
| 39 |
+
def __init__(self):
|
| 40 |
+
self.refiners = []
|
| 41 |
+
self.extenders = []
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def load_prompt_refiners(self, model_manager, refiner_classes=[]):
|
| 45 |
+
for refiner_class in refiner_classes:
|
| 46 |
+
refiner = refiner_class.from_model_manager(model_manager)
|
| 47 |
+
self.refiners.append(refiner)
|
| 48 |
+
|
| 49 |
+
def load_prompt_extenders(self, model_manager, extender_classes=[]):
|
| 50 |
+
for extender_class in extender_classes:
|
| 51 |
+
extender = extender_class.from_model_manager(model_manager)
|
| 52 |
+
self.extenders.append(extender)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
@torch.no_grad()
|
| 56 |
+
def process_prompt(self, prompt, positive=True):
|
| 57 |
+
if isinstance(prompt, list):
|
| 58 |
+
prompt = [self.process_prompt(prompt_, positive=positive) for prompt_ in prompt]
|
| 59 |
+
else:
|
| 60 |
+
for refiner in self.refiners:
|
| 61 |
+
prompt = refiner(prompt, positive=positive)
|
| 62 |
+
return prompt
|
| 63 |
+
|
| 64 |
+
@torch.no_grad()
|
| 65 |
+
def extend_prompt(self, prompt:str, positive=True):
|
| 66 |
+
extended_prompt = dict(prompt=prompt)
|
| 67 |
+
for extender in self.extenders:
|
| 68 |
+
extended_prompt = extender(extended_prompt)
|
| 69 |
+
return extended_prompt
|
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .base_prompter import BasePrompter
|
| 2 |
+
from ..models.wan_video_text_encoder import WanTextEncoder
|
| 3 |
+
from transformers import AutoTokenizer
|
| 4 |
+
import os, torch
|
| 5 |
+
import ftfy
|
| 6 |
+
import html
|
| 7 |
+
import string
|
| 8 |
+
import regex as re
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def basic_clean(text):
|
| 12 |
+
text = ftfy.fix_text(text)
|
| 13 |
+
text = html.unescape(html.unescape(text))
|
| 14 |
+
return text.strip()
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def whitespace_clean(text):
|
| 18 |
+
text = re.sub(r'\s+', ' ', text)
|
| 19 |
+
text = text.strip()
|
| 20 |
+
return text
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def canonicalize(text, keep_punctuation_exact_string=None):
|
| 24 |
+
text = text.replace('_', ' ')
|
| 25 |
+
if keep_punctuation_exact_string:
|
| 26 |
+
text = keep_punctuation_exact_string.join(
|
| 27 |
+
part.translate(str.maketrans('', '', string.punctuation))
|
| 28 |
+
for part in text.split(keep_punctuation_exact_string))
|
| 29 |
+
else:
|
| 30 |
+
text = text.translate(str.maketrans('', '', string.punctuation))
|
| 31 |
+
text = text.lower()
|
| 32 |
+
text = re.sub(r'\s+', ' ', text)
|
| 33 |
+
return text.strip()
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class HuggingfaceTokenizer:
|
| 37 |
+
|
| 38 |
+
def __init__(self, name, seq_len=None, clean=None, **kwargs):
|
| 39 |
+
assert clean in (None, 'whitespace', 'lower', 'canonicalize')
|
| 40 |
+
self.name = name
|
| 41 |
+
self.seq_len = seq_len
|
| 42 |
+
self.clean = clean
|
| 43 |
+
|
| 44 |
+
# init tokenizer
|
| 45 |
+
self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
|
| 46 |
+
self.vocab_size = self.tokenizer.vocab_size
|
| 47 |
+
|
| 48 |
+
def __call__(self, sequence, **kwargs):
|
| 49 |
+
return_mask = kwargs.pop('return_mask', False)
|
| 50 |
+
|
| 51 |
+
# arguments
|
| 52 |
+
_kwargs = {'return_tensors': 'pt'}
|
| 53 |
+
if self.seq_len is not None:
|
| 54 |
+
_kwargs.update({
|
| 55 |
+
'padding': 'max_length',
|
| 56 |
+
'truncation': True,
|
| 57 |
+
'max_length': self.seq_len
|
| 58 |
+
})
|
| 59 |
+
_kwargs.update(**kwargs)
|
| 60 |
+
|
| 61 |
+
# tokenization
|
| 62 |
+
if isinstance(sequence, str):
|
| 63 |
+
sequence = [sequence]
|
| 64 |
+
if self.clean:
|
| 65 |
+
sequence = [self._clean(u) for u in sequence]
|
| 66 |
+
ids = self.tokenizer(sequence, **_kwargs)
|
| 67 |
+
|
| 68 |
+
# output
|
| 69 |
+
if return_mask:
|
| 70 |
+
return ids.input_ids, ids.attention_mask
|
| 71 |
+
else:
|
| 72 |
+
return ids.input_ids
|
| 73 |
+
|
| 74 |
+
def _clean(self, text):
|
| 75 |
+
if self.clean == 'whitespace':
|
| 76 |
+
text = whitespace_clean(basic_clean(text))
|
| 77 |
+
elif self.clean == 'lower':
|
| 78 |
+
text = whitespace_clean(basic_clean(text)).lower()
|
| 79 |
+
elif self.clean == 'canonicalize':
|
| 80 |
+
text = canonicalize(basic_clean(text))
|
| 81 |
+
return text
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
class WanPrompter(BasePrompter):
|
| 85 |
+
|
| 86 |
+
def __init__(self, tokenizer_path=None, text_len=512):
|
| 87 |
+
super().__init__()
|
| 88 |
+
self.text_len = text_len
|
| 89 |
+
self.text_encoder = None
|
| 90 |
+
self.fetch_tokenizer(tokenizer_path)
|
| 91 |
+
|
| 92 |
+
def fetch_tokenizer(self, tokenizer_path=None):
|
| 93 |
+
if tokenizer_path is not None:
|
| 94 |
+
self.tokenizer = HuggingfaceTokenizer(name=tokenizer_path, seq_len=self.text_len, clean='whitespace')
|
| 95 |
+
|
| 96 |
+
def fetch_models(self, text_encoder: WanTextEncoder = None):
|
| 97 |
+
self.text_encoder = text_encoder
|
| 98 |
+
|
| 99 |
+
def encode_prompt(self, prompt, positive=True, device="cuda"):
|
| 100 |
+
prompt = self.process_prompt(prompt, positive=positive)
|
| 101 |
+
|
| 102 |
+
ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True)
|
| 103 |
+
ids = ids.to(device)
|
| 104 |
+
mask = mask.to(device)
|
| 105 |
+
seq_lens = mask.gt(0).sum(dim=1).long()
|
| 106 |
+
prompt_emb = self.text_encoder(ids, mask)
|
| 107 |
+
for i, v in enumerate(seq_lens):
|
| 108 |
+
prompt_emb[:, v:] = 0
|
| 109 |
+
return prompt_emb
|