pablovela5620 Claude Fable 5 commited on
Commit
1b2e96f
·
1 Parent(s): 7d62919

Vendor the inference-path subset of the fdanyone package

Browse files

The 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>

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +1 -0
  2. .gitignore +10 -0
  3. PROVENANCE.md +36 -0
  4. fdanyone/__init__.py +1 -0
  5. fdanyone/assets.py +191 -0
  6. fdanyone/config.py +200 -0
  7. fdanyone/device.py +25 -0
  8. fdanyone/download.py +427 -0
  9. fdanyone/errors.py +17 -0
  10. fdanyone/foreground.py +65 -0
  11. fdanyone/geometry/__init__.py +1 -0
  12. fdanyone/geometry/cameras.py +250 -0
  13. fdanyone/geometry/crop.py +144 -0
  14. fdanyone/geometry/framing.py +579 -0
  15. fdanyone/io.py +115 -0
  16. fdanyone/model/__init__.py +1 -0
  17. fdanyone/model/inference.py +689 -0
  18. fdanyone/model/loader.py +426 -0
  19. fdanyone/model/prepared.py +247 -0
  20. fdanyone/model/profiling.py +147 -0
  21. fdanyone/model/quantization.py +87 -0
  22. fdanyone/model/routing.py +69 -0
  23. fdanyone/model/tiny_decoder.py +60 -0
  24. fdanyone/model/turbo_lora.py +125 -0
  25. fdanyone/motion/__init__.py +5 -0
  26. fdanyone/motion/body.py +143 -0
  27. fdanyone/motion/gvhmr.py +332 -0
  28. fdanyone/motion/result.py +188 -0
  29. fdanyone/motion/worker.py +53 -0
  30. fdanyone/output.py +218 -0
  31. fdanyone/pipeline.py +591 -0
  32. fdanyone/runs.py +161 -0
  33. fdanyone/skeleton/__init__.py +1 -0
  34. fdanyone/skeleton/keypoints.py +170 -0
  35. fdanyone/skeleton/pipeline.py +839 -0
  36. fdanyone/skeleton/render_worker.py +100 -0
  37. fdanyone/skeleton/renderer.py +230 -0
  38. fdanyone/skeleton/worker.py +66 -0
  39. fdanyone/vendor/__init__.py +1 -0
  40. fdanyone/vendor/diffsynth/LICENSE +201 -0
  41. fdanyone/vendor/diffsynth/UPSTREAM.md +23 -0
  42. fdanyone/vendor/diffsynth/UPSTREAM.patch +1334 -0
  43. fdanyone/vendor/diffsynth/VENDORED_FILES.txt +22 -0
  44. fdanyone/vendor/diffsynth/__init__.py +1 -0
  45. fdanyone/vendor/diffsynth/pipelines/__init__.py +5 -0
  46. fdanyone/vendor/diffsynth/pipelines/base.py +127 -0
  47. fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py +93 -0
  48. fdanyone/vendor/diffsynth/prompters/__init__.py +5 -0
  49. fdanyone/vendor/diffsynth/prompters/base_prompter.py +69 -0
  50. fdanyone/vendor/diffsynth/prompters/wan_prompter.py +109 -0
.gitattributes CHANGED
@@ -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
.gitignore ADDED
@@ -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/
PROVENANCE.md ADDED
@@ -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` |
fdanyone/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """4DAnyone inference."""
fdanyone/assets.py ADDED
@@ -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
+ )
fdanyone/config.py ADDED
@@ -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()
fdanyone/device.py ADDED
@@ -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
fdanyone/download.py ADDED
@@ -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
fdanyone/errors.py ADDED
@@ -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."""
fdanyone/foreground.py ADDED
@@ -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)
fdanyone/geometry/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Camera and crop geometry."""
fdanyone/geometry/cameras.py ADDED
@@ -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
fdanyone/geometry/crop.py ADDED
@@ -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
fdanyone/geometry/framing.py ADDED
@@ -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
+ )
fdanyone/io.py ADDED
@@ -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
fdanyone/model/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """4DAnyone model loading and stage inference."""
fdanyone/model/inference.py ADDED
@@ -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
+ )
fdanyone/model/loader.py ADDED
@@ -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)
fdanyone/model/prepared.py ADDED
@@ -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)
fdanyone/model/profiling.py ADDED
@@ -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
+ )
fdanyone/model/quantization.py ADDED
@@ -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))
fdanyone/model/routing.py ADDED
@@ -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
+ )
fdanyone/model/tiny_decoder.py ADDED
@@ -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)
fdanyone/model/turbo_lora.py ADDED
@@ -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
+ )
fdanyone/motion/__init__.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ """GVHMR motion recovery."""
2
+
3
+ from fdanyone.motion.result import MotionResult
4
+
5
+ __all__ = ["MotionResult"]
fdanyone/motion/body.py ADDED
@@ -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)
fdanyone/motion/gvhmr.py ADDED
@@ -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
fdanyone/motion/result.py ADDED
@@ -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
fdanyone/motion/worker.py ADDED
@@ -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])
fdanyone/output.py ADDED
@@ -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
+ }
fdanyone/pipeline.py ADDED
@@ -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)
fdanyone/runs.py ADDED
@@ -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)
fdanyone/skeleton/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """MHR70 projection and Goliath40 conditioning renderer."""
fdanyone/skeleton/keypoints.py ADDED
@@ -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
+ )
fdanyone/skeleton/pipeline.py ADDED
@@ -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
+ )
fdanyone/skeleton/render_worker.py ADDED
@@ -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])
fdanyone/skeleton/renderer.py ADDED
@@ -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)
fdanyone/skeleton/worker.py ADDED
@@ -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])
fdanyone/vendor/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Third-party inference code redistributed with its upstream notices."""
fdanyone/vendor/diffsynth/LICENSE ADDED
@@ -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.
fdanyone/vendor/diffsynth/UPSTREAM.md ADDED
@@ -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.
fdanyone/vendor/diffsynth/UPSTREAM.patch ADDED
@@ -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())
fdanyone/vendor/diffsynth/VENDORED_FILES.txt ADDED
@@ -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
fdanyone/vendor/diffsynth/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Minimal DiffSynth-Studio inference closure used by 4DAnyone."""
fdanyone/vendor/diffsynth/pipelines/__init__.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ """Vendored inference pipelines."""
2
+
3
+ from .wan_video_spatem import WanVideoSpaTemPipeline
4
+
5
+ __all__ = ["WanVideoSpaTemPipeline"]
fdanyone/vendor/diffsynth/pipelines/base.py ADDED
@@ -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
fdanyone/vendor/diffsynth/pipelines/wan_video_spatem.py ADDED
@@ -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)
fdanyone/vendor/diffsynth/prompters/__init__.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ """Vendored Wan prompt helpers."""
2
+
3
+ from .wan_prompter import WanPrompter
4
+
5
+ __all__ = ["WanPrompter"]
fdanyone/vendor/diffsynth/prompters/base_prompter.py ADDED
@@ -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
fdanyone/vendor/diffsynth/prompters/wan_prompter.py ADDED
@@ -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