Spaces:
Runtime error
Runtime error
Commit ·
982899c
1
Parent(s): 986d8aa
DreamX-Creator 1.0 on ZeroGPU: vendored videox_fun + dreamx_inference from AMAP-ML upstream; generate(image, prompt)->(mp4, last-frame PNG, seed), neutral keyframe when image empty, DREAMX_CKPT_DIR for persistent checkpoints, diffusers 0.37.1 stack
Browse files- .gitattributes +1 -0
- .gitignore +2 -0
- README.md +22 -6
- app.py +540 -395
- config/config.yaml +54 -0
- dreamx_inference.py +1030 -0
- examples/case1.jpg +3 -0
- examples/case2.jpg +3 -0
- examples/case3.jpg +3 -0
- examples/case4.jpg +3 -0
- examples/case5.jpg +3 -0
- examples/case6.jpg +3 -0
- requirements.txt +12 -12
- videox_fun/__init__.py +1 -0
- videox_fun/dist/__init__.py +3 -0
- videox_fun/dist/sequence_parallel.py +42 -0
- videox_fun/models/__init__.py +7 -0
- videox_fun/models/attention_utils.py +211 -0
- videox_fun/models/creator/__init__.py +1 -0
- videox_fun/models/creator/creator_audio_dit.py +74 -0
- videox_fun/models/creator/creator_video_dit.py +103 -0
- videox_fun/models/creator/dac_vae.py +878 -0
- videox_fun/models/creator_audio.py +353 -0
- videox_fun/models/creator_dac_vae.py +151 -0
- videox_fun/models/creator_gating.py +1286 -0
- videox_fun/models/wan_text_encoder.py +389 -0
- videox_fun/models/wan_transformer3d_prope.py +1035 -0
- videox_fun/models/wan_vae3_8.py +1248 -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 |
+
examples/*.jpg filter=lfs diff=lfs merge=lfs -text
|
.gitignore
CHANGED
|
@@ -1,6 +1,8 @@
|
|
| 1 |
# Local model / cache artifacts
|
| 2 |
*.gguf
|
| 3 |
*.safetensors
|
|
|
|
|
|
|
| 4 |
/tmp/
|
| 5 |
.cache/
|
| 6 |
__pycache__/
|
|
|
|
| 1 |
# Local model / cache artifacts
|
| 2 |
*.gguf
|
| 3 |
*.safetensors
|
| 4 |
+
*.pth
|
| 5 |
+
/checkpoints/
|
| 6 |
/tmp/
|
| 7 |
.cache/
|
| 8 |
__pycache__/
|
README.md
CHANGED
|
@@ -1,14 +1,30 @@
|
|
| 1 |
---
|
| 2 |
-
title:
|
| 3 |
-
emoji:
|
| 4 |
colorFrom: green
|
| 5 |
colorTo: blue
|
| 6 |
sdk: docker
|
| 7 |
app_file: app.py
|
| 8 |
pinned: false
|
|
|
|
|
|
|
| 9 |
---
|
| 10 |
|
| 11 |
-
Docker-hosted
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
title: DreamX Creator
|
| 3 |
+
emoji: 🎬
|
| 4 |
colorFrom: green
|
| 5 |
colorTo: blue
|
| 6 |
sdk: docker
|
| 7 |
app_file: app.py
|
| 8 |
pinned: false
|
| 9 |
+
hardware:
|
| 10 |
+
gpu: zero
|
| 11 |
---
|
| 12 |
|
| 13 |
+
Docker-hosted DreamX-Creator 1.0 Space for the AI Shorts Factory backend
|
| 14 |
+
(engine vendored from `AMAP-ML/DreamX-Creator` / `hugging-apps/gd-ml-dreamx-creator`).
|
| 15 |
+
|
| 16 |
+
One endpoint via Gradio's `/call/generate` protocol:
|
| 17 |
+
|
| 18 |
+
- `generate(image, prompt, seconds, num_inference_steps, resolution_tokens,
|
| 19 |
+
guidance_scale, negative_prompt, seed, randomize_seed)`
|
| 20 |
+
→ `(mp4 with jointly-denoised audio, last-frame PNG, seed)`.
|
| 21 |
+
|
| 22 |
+
`image` may be empty — a neutral keyframe is generated in-app. The last-frame
|
| 23 |
+
PNG is returned so each scene can feed the next one (continuity chaining).
|
| 24 |
+
|
| 25 |
+
Weights (~43 GB peak) download at first boot from the ungated
|
| 26 |
+
`GD-ML/DreamX-Creator` repo into `./checkpoints` (override with
|
| 27 |
+
`DREAMX_CKPT_DIR`, e.g. `/data/dreamx` when persistent storage is enabled).
|
| 28 |
+
|
| 29 |
+
Requires **ZeroGPU** hardware (`hardware: gpu: zero` above and/or the Space
|
| 30 |
+
Settings toggle) — the free CPU tier cannot denoise.
|
app.py
CHANGED
|
@@ -1,443 +1,588 @@
|
|
| 1 |
-
"""
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
|
|
|
| 11 |
|
| 12 |
-
|
| 13 |
-
pipeline call image_to_video with the last frame URL of the previous clip so
|
| 14 |
-
the subject stays continuous (last-frame chaining).
|
| 15 |
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
-
|
|
|
|
| 22 |
|
| 23 |
-
import glob
|
| 24 |
import os
|
| 25 |
-
import subprocess
|
| 26 |
-
import threading
|
| 27 |
-
from pathlib import Path
|
| 28 |
-
|
| 29 |
-
# ZeroGPU rule #1: `import spaces` must precede any CUDA-touching import
|
| 30 |
-
# (torch etc.) — it monkey-patches torch.cuda at import time.
|
| 31 |
-
import spaces # noqa: E402
|
| 32 |
-
|
| 33 |
-
import numpy as np
|
| 34 |
-
import torch
|
| 35 |
-
|
| 36 |
-
# Requires the Space secret HF_TOKEN after accepting the LTX-2.x Community
|
| 37 |
-
# License (free for entities under $10M annual revenue). Needed for every file
|
| 38 |
-
# served from the gated Lightricks/LTX-2.5-Diffusers repo (config, Gemma-4-12B
|
| 39 |
-
# text encoder, VAEs, connectors, audio_vae + vocoder).
|
| 40 |
-
HF_TOKEN = os.environ.get("HF_TOKEN", "").strip()
|
| 41 |
-
|
| 42 |
-
# Direct download URL of the distilled GGUF transformer (~15GB). The backend
|
| 43 |
-
# only needs LTX_SPACE_URL; this key is configured in the Space. A GGUF the
|
| 44 |
-
# user uploads into the repo `model/` folder is preferred over this URL.
|
| 45 |
-
LTX_MODEL_URL = os.environ.get(
|
| 46 |
-
"LTX_MODEL_URL",
|
| 47 |
-
"https://huggingface.co/realrebelai/LTX-2.5_GGUFs/resolve/main/LTX-2.5-Distilled-Q4_K_M.gguf",
|
| 48 |
-
)
|
| 49 |
-
PIPELINE_ID = "Lightricks/LTX-2.5-Diffusers"
|
| 50 |
-
# Overridable repo id (e.g. a locally mirrored copy: LTX_PIPELINE_ID=/app/ltx25).
|
| 51 |
-
LTX_PIPELINE_ID = os.environ.get("LTX_PIPELINE_ID", PIPELINE_ID)
|
| 52 |
-
SEED = int(os.environ.get("LTX_SEED", "0"))
|
| 53 |
-
FPS = int(os.environ.get("LTX_FPS", "24"))
|
| 54 |
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
# Honest worst-case GPU wall-time per generation. ZeroGPU validates this
|
| 60 |
-
# after a 1.5x multiplier and sits a ~60-120s continuous execution cap per
|
| 61 |
-
# call on the free tier — 120 keeps us inside both (120 * 1.5 = 180 < 300
|
| 62 |
-
# cap, and matches the real per-call window after the model is resident).
|
| 63 |
-
GPU_DURATION = int(os.environ.get("LTX_GPU_DURATION", "120"))
|
| 64 |
|
| 65 |
-
|
| 66 |
-
_model = None
|
| 67 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
|
| 69 |
-
|
|
|
|
| 70 |
|
|
|
|
|
|
|
|
|
|
| 71 |
|
| 72 |
-
def _find_local_gguf() -> str | None:
|
| 73 |
-
"""Pick a GGUF the user uploaded into the repo (web UI ``model/`` folder)."""
|
| 74 |
-
for folder in ("model", "models", "weights"):
|
| 75 |
-
hits = [p for p in glob.glob(os.path.join(folder, "**", "*.gguf"), recursive=True) if os.path.isfile(p)]
|
| 76 |
-
if hits:
|
| 77 |
-
return sorted(hits)[0]
|
| 78 |
-
return None
|
| 79 |
|
|
|
|
|
|
|
|
|
|
| 80 |
|
| 81 |
-
def _gguf_fingerprint(path: str) -> None:
|
| 82 |
-
"""Fail fast if the GGUF is not an LTX-2 video transformer.
|
| 83 |
-
|
| 84 |
-
diffusers 0.40 converts GGUFs that keep the ComfyUI/native
|
| 85 |
-
``model.diffusion_model.*`` tensor names (llama.cpp-style ``blk.*`` renames
|
| 86 |
-
are not supported by the LTX-2 GGUF converter).
|
| 87 |
-
"""
|
| 88 |
-
from gguf import GGUFReader
|
| 89 |
|
| 90 |
-
|
| 91 |
-
names = [tensor.name for tensor in reader.tensors]
|
| 92 |
-
print(f"[ltx] gguf tensors={len(names)}")
|
| 93 |
-
|
| 94 |
-
def count(marker: str) -> int:
|
| 95 |
-
return sum(1 for name in names if marker in name)
|
| 96 |
-
|
| 97 |
-
blocks = count("transformer_blocks.")
|
| 98 |
-
ltx2_av_gate = count("model.diffusion_model.av_ca_a2v_gate_adaln_single")
|
| 99 |
-
llama_blk = count("blk.") + count(".attn_qkv.")
|
| 100 |
-
print(
|
| 101 |
-
f"[ltx] gguf fingerprint blocks={blocks} ltx2_av_gate={ltx2_av_gate} llama_blk={llama_blk}"
|
| 102 |
-
)
|
| 103 |
-
if blocks == 0 and llama_blk == 0:
|
| 104 |
-
raise RuntimeError(
|
| 105 |
-
"GGUF does not look like an LTX-2 video transformer (no transformer_blocks found)"
|
| 106 |
-
)
|
| 107 |
-
if blocks == 0 and llama_blk > 0:
|
| 108 |
-
raise RuntimeError(
|
| 109 |
-
"GGUF uses llama.cpp-renamed tensors (blk.*); diffusers 0.40's LTX-2 GGUF "
|
| 110 |
-
"converter expects the native model.diffusion_model.* names. Use a GGUF "
|
| 111 |
-
"converted from the ComfyUI/native checkpoint instead."
|
| 112 |
-
)
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
def _download_gguf(dest: str) -> str:
|
| 116 |
-
import httpx
|
| 117 |
-
|
| 118 |
-
if LTX_MODEL_URL.startswith("http"):
|
| 119 |
-
if os.path.exists(dest):
|
| 120 |
-
return dest
|
| 121 |
-
with httpx.stream("GET", LTX_MODEL_URL, follow_redirects=True, timeout=600) as r:
|
| 122 |
-
r.raise_for_status()
|
| 123 |
-
with open(dest, "wb") as fh:
|
| 124 |
-
for chunk in r.iter_bytes(chunk_size=1 << 20):
|
| 125 |
-
fh.write(chunk)
|
| 126 |
-
return dest
|
| 127 |
-
from huggingface_hub import hf_hub_download
|
| 128 |
-
|
| 129 |
-
# Support "repo_id:filename" shorthand.
|
| 130 |
-
repo, _, filename = LTX_MODEL_URL.partition(":")
|
| 131 |
-
return hf_hub_download(repo, filename or "LTX-2.5-Distilled-Q4_K_M.gguf", token=HF_TOKEN or None)
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
def _ensure_gguf() -> str:
|
| 135 |
-
"""Local vendored GGUF first, then network download."""
|
| 136 |
-
local = _find_local_gguf()
|
| 137 |
-
if local is not None:
|
| 138 |
-
print(f"[ltx] using vendored GGUF: {local}")
|
| 139 |
-
return local
|
| 140 |
-
print("[ltx] no local GGUF found; downloading…")
|
| 141 |
-
return _download_gguf(os.environ.get("LTX_GGUF_CACHE", "/tmp/ltx-model.gguf"))
|
| 142 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 143 |
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
from diffusers import AutoModel as TransformerCls
|
| 151 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
try:
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
kwargs = dict(
|
| 159 |
-
quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16),
|
| 160 |
-
dtype=torch.bfloat16,
|
| 161 |
-
)
|
| 162 |
-
_gguf_fingerprint(gguf_path)
|
| 163 |
-
token = HF_TOKEN or None
|
| 164 |
-
try:
|
| 165 |
-
# `config=` pulls transformer/config.json from the gated repo (token).
|
| 166 |
-
return TransformerCls.from_single_file(gguf_path, config=LTX_PIPELINE_ID, token=token, **kwargs)
|
| 167 |
-
except Exception as exc: # config may live elsewhere; trust the GGUF KV metadata
|
| 168 |
-
print(f"[ltx] from_single_file with config failed ({exc}), retrying without")
|
| 169 |
-
return TransformerCls.from_single_file(gguf_path, token=token, **kwargs)
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
def _load_model() -> dict:
|
| 173 |
-
global _model
|
| 174 |
-
if _model is not None:
|
| 175 |
-
return _model
|
| 176 |
-
|
| 177 |
-
with _lock:
|
| 178 |
-
if _model is not None:
|
| 179 |
-
return _model
|
| 180 |
-
|
| 181 |
-
if not HF_TOKEN and not os.path.isdir(LTX_PIPELINE_ID):
|
| 182 |
-
raise RuntimeError(
|
| 183 |
-
"HF_TOKEN secret is not set. Open the Space → Settings → Variables and "
|
| 184 |
-
"secrets and create a secret named HF_TOKEN: a fine-grained token with "
|
| 185 |
-
"read access to Lightricks/LTX-2.5-Diffusers (after accepting its "
|
| 186 |
-
"LTX-2.x Community License on that repo)."
|
| 187 |
-
)
|
| 188 |
|
| 189 |
-
|
| 190 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 191 |
|
| 192 |
-
|
| 193 |
-
|
|
|
|
|
|
|
|
|
|
| 194 |
|
| 195 |
-
from diffusers import LTX2ImageToVideoPipeline, LTX2Pipeline
|
| 196 |
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
token = HF_TOKEN or None
|
| 201 |
-
try:
|
| 202 |
-
built["t2v"] = LTX2Pipeline.from_pretrained(
|
| 203 |
-
LTX_PIPELINE_ID, transformer=transformer, torch_dtype=torch.bfloat16, token=token
|
| 204 |
-
)
|
| 205 |
-
built["i2v"] = LTX2ImageToVideoPipeline.from_pretrained(
|
| 206 |
-
LTX_PIPELINE_ID, transformer=transformer, torch_dtype=torch.bfloat16, token=token
|
| 207 |
-
)
|
| 208 |
-
except Exception as exc: # pragma: no cover - component layout differences
|
| 209 |
-
msg = str(exc)
|
| 210 |
-
hint = ""
|
| 211 |
-
if "401" in msg or "is not a valid model identifier" in msg or "gated" in msg.lower():
|
| 212 |
-
hint = " (hint: accept the LTX-2.x Community License on the repo + set the HF_TOKEN secret)"
|
| 213 |
-
raise RuntimeError(
|
| 214 |
-
f"[ltx] pipeline build failed - check Lightricks/LTX-2.5-Diffusers "
|
| 215 |
-
f"component layout. {msg}{hint}"
|
| 216 |
-
) from exc
|
| 217 |
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
return _model
|
| 223 |
|
| 224 |
|
| 225 |
-
|
| 226 |
-
# at startup and resident for every forked worker, instead of per-call.
|
| 227 |
-
# `_generate` still guards with `_load_model()` so early calls block on the
|
| 228 |
-
# lock until this background warm-up finishes.
|
| 229 |
-
_warmup = threading.Thread(target=_load_model, name="ltx-warmup", daemon=True)
|
| 230 |
-
_warmup.start()
|
| 231 |
|
| 232 |
|
| 233 |
-
|
|
|
|
| 234 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 235 |
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
|
|
|
|
|
|
| 296 |
)
|
| 297 |
-
return muxed
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
def _load_image(image_url: str):
|
| 301 |
-
if not image_url:
|
| 302 |
-
return None
|
| 303 |
-
import base64
|
| 304 |
-
import io
|
| 305 |
-
|
| 306 |
-
from PIL import Image, ImageOps
|
| 307 |
-
|
| 308 |
-
if image_url.startswith("data:"):
|
| 309 |
-
encoded = image_url.partition(",")[2]
|
| 310 |
-
return Image.open(io.BytesIO(base64.b64decode(encoded))).convert("RGB")
|
| 311 |
-
|
| 312 |
-
import httpx
|
| 313 |
-
|
| 314 |
-
data = httpx.get(image_url, timeout=120).content
|
| 315 |
-
image = Image.open(io.BytesIO(data)).convert("RGB")
|
| 316 |
-
return ImageOps.exif_transpose(image)
|
| 317 |
-
|
| 318 |
-
|
| 319 |
-
def _generate(prompt, image_url, width, height, num_frames, enhance_prompt, with_audio):
|
| 320 |
-
if enhance_prompt:
|
| 321 |
-
# Never let the enhancer rewrite the scripted dialogue.
|
| 322 |
-
enhance_prompt = False
|
| 323 |
|
| 324 |
-
|
| 325 |
-
|
| 326 |
-
|
| 327 |
-
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
|
| 335 |
-
|
| 336 |
-
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
|
| 340 |
-
|
| 341 |
-
generator=generator,
|
| 342 |
)
|
| 343 |
-
|
| 344 |
-
output = _try_call(model["i2v"], image=image, **common)
|
| 345 |
-
else:
|
| 346 |
-
output = _try_call(model["t2v"], **common)
|
| 347 |
-
|
| 348 |
-
frames = _extract_frames(output)
|
| 349 |
-
audio = _extract_audio(output) if with_audio else None
|
| 350 |
-
|
| 351 |
-
video_path = f"/tmp/gradio/ltx_{threading.get_ident()}.mp4"
|
| 352 |
-
muxed = _mux_audio(video_path, frames, FPS, audio)
|
| 353 |
-
|
| 354 |
-
first_frame = f"/tmp/gradio/ltx_{threading.get_ident()}_first.png"
|
| 355 |
-
frames[0].save(first_frame)
|
| 356 |
-
return muxed, first_frame
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
# ------------------------------------------------------------------ gradio UI
|
| 360 |
-
|
| 361 |
|
| 362 |
-
|
| 363 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 364 |
|
|
|
|
| 365 |
|
| 366 |
-
|
| 367 |
-
|
| 368 |
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
|
| 372 |
-
|
| 373 |
-
|
| 374 |
-
def image_to_video_av(prompt, image_url, width, height, num_frames, enhance_prompt=False):
|
| 375 |
-
return _generate(prompt, image_url, width, height, num_frames, enhance_prompt, with_audio=True)
|
| 376 |
-
|
| 377 |
-
|
| 378 |
-
import gradio as gr # noqa: E402
|
| 379 |
-
|
| 380 |
-
t2v_fns = [spaces.GPU(duration=GPU_DURATION)(text_to_video), spaces.GPU(duration=GPU_DURATION)(text_to_video_av)]
|
| 381 |
-
i2v_fns = [spaces.GPU(duration=GPU_DURATION)(image_to_video), spaces.GPU(duration=GPU_DURATION)(image_to_video_av)]
|
| 382 |
-
|
| 383 |
-
with gr.Blocks(title="LTX-2.5 Shorts Space") as demo:
|
| 384 |
-
gr.Markdown(
|
| 385 |
-
"# LTX-2.5 Shorts Space\n"
|
| 386 |
-
"ZeroGPU inference for the AI Shorts Factory backend. "
|
| 387 |
-
"Portrait 9:16 clips (~4s) with optional synchronized narration audio. "
|
| 388 |
-
"The prompt enhancer stays **off** so quoted dialogue is spoken exactly."
|
| 389 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 390 |
|
| 391 |
-
|
| 392 |
-
|
| 393 |
-
|
| 394 |
-
|
| 395 |
-
|
| 396 |
-
|
| 397 |
-
|
| 398 |
-
|
| 399 |
-
|
| 400 |
-
|
| 401 |
-
|
| 402 |
-
t2v_fns[0],
|
| 403 |
-
inputs=[t2v_prompt, width, height, num_frames, enhance],
|
| 404 |
-
outputs=[video_out, frame_out],
|
| 405 |
)
|
| 406 |
|
| 407 |
-
|
| 408 |
-
|
| 409 |
-
|
| 410 |
-
|
| 411 |
-
|
| 412 |
-
|
| 413 |
-
|
| 414 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 415 |
|
| 416 |
-
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
label="First-frame image URL (last frame from the previous clip)",
|
| 420 |
-
value="",
|
| 421 |
-
)
|
| 422 |
-
i2v_btn = gr.Button("Generate (video only)")
|
| 423 |
-
i2v_btn.click(
|
| 424 |
-
i2v_fns[0],
|
| 425 |
-
inputs=[i2v_prompt, i2v_image, width, height, num_frames, enhance],
|
| 426 |
-
outputs=[video_out, frame_out],
|
| 427 |
)
|
| 428 |
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
|
| 434 |
-
|
| 435 |
-
|
| 436 |
-
|
| 437 |
-
i2v_fns[1],
|
| 438 |
-
inputs=[i2va_prompt, i2va_image, width, height, num_frames, enhance],
|
| 439 |
-
outputs=[video_out, frame_out],
|
| 440 |
)
|
| 441 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 442 |
if __name__ == "__main__":
|
| 443 |
-
demo.
|
|
|
|
| 1 |
+
"""DreamX-Creator 1.0 — native joint audio-video generation on ZeroGPU.
|
| 2 |
|
| 3 |
+
Image + prompt -> a video whose soundtrack is denoised jointly with the frames
|
| 4 |
+
by the same model (gated A2V / V2A cross-attention), so speech, foley and
|
| 5 |
+
ambience stay in sync with the picture.
|
| 6 |
|
| 7 |
+
The inference path is the authors' own release code (`videox_fun/` +
|
| 8 |
+
`dreamx_inference.py`, copied verbatim from the reference Space
|
| 9 |
+
`hugging-apps/gd-ml-dreamx-creator` / `AMAP-ML/DreamX-Creator`); this file only
|
| 10 |
+
wires it into Gradio and stages the checkpoint download so the container never
|
| 11 |
+
holds all 43 GB of fp32 weights on disk at once.
|
| 12 |
|
| 13 |
+
Adapted for the AI Shorts Factory backend (Space `text_amon_API`):
|
|
|
|
|
|
|
| 14 |
|
| 15 |
+
- `image` is optional: an empty value yields a neutral keyframe (Option A), so
|
| 16 |
+
scenes without a first-frame image can still generate.
|
| 17 |
+
- Returns ``(mp4, last-frame PNG, seed)`` instead of ``(mp4, seed)``: the last
|
| 18 |
+
frame is what scene 2..N sends back as the first frame for continuity.
|
| 19 |
+
- Checkpoint root can be redirected to persistent storage with
|
| 20 |
+
``DREAMX_CKPT_DIR`` (default: this repo's ``./checkpoints``, the proven
|
| 21 |
+
upstream layout).
|
| 22 |
|
| 23 |
+
Served through Gradio's `/call/generate` protocol (``api_name="generate"``).
|
| 24 |
+
"""
|
| 25 |
|
|
|
|
| 26 |
import os
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
|
| 28 |
+
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 29 |
+
# videox_fun falls back to torch SDPA when flash-attn is absent; the numerics are
|
| 30 |
+
# identical here because every sequence in the batch is full-length (no padding).
|
| 31 |
+
os.environ.setdefault("VIDEOX_ATTENTION_TYPE", "FLASH_ATTENTION")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
|
| 33 |
+
import spaces # noqa: E402 (must precede torch)
|
|
|
|
| 34 |
|
| 35 |
+
import math # noqa: E402
|
| 36 |
+
import random # noqa: E402
|
| 37 |
+
import shutil # noqa: E402
|
| 38 |
+
import subprocess # noqa: E402
|
| 39 |
+
import tempfile # noqa: E402
|
| 40 |
+
import time # noqa: E402
|
| 41 |
+
from pathlib import Path # noqa: E402
|
| 42 |
+
from types import SimpleNamespace # noqa: E402
|
| 43 |
|
| 44 |
+
import numpy as np # noqa: E402
|
| 45 |
+
import torch # noqa: E402
|
| 46 |
|
| 47 |
+
# Trusted upstream .pth checkpoints (UMT5 encoder, Wan2.2 VAE) are plain tensor
|
| 48 |
+
# dicts saved before torch 2.6 flipped `weights_only` to True.
|
| 49 |
+
_torch_load = torch.load
|
| 50 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 51 |
|
| 52 |
+
def _torch_load_compat(*args, **kwargs):
|
| 53 |
+
kwargs.setdefault("weights_only", False)
|
| 54 |
+
return _torch_load(*args, **kwargs)
|
| 55 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
|
| 57 |
+
torch.load = _torch_load_compat
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 58 |
|
| 59 |
+
import gradio as gr # noqa: E402
|
| 60 |
+
from diffusers import FlowMatchEulerDiscreteScheduler # noqa: E402
|
| 61 |
+
from huggingface_hub import snapshot_download # noqa: E402
|
| 62 |
+
from omegaconf import OmegaConf # noqa: E402
|
| 63 |
+
from PIL import Image # noqa: E402
|
| 64 |
+
from transformers import AutoTokenizer # noqa: E402
|
| 65 |
+
|
| 66 |
+
from dreamx_inference import ( # noqa: E402
|
| 67 |
+
DEFAULT_NEGATIVE_PROMPT,
|
| 68 |
+
DirectionalMultimodalCFGAdapter,
|
| 69 |
+
filter_kwargs,
|
| 70 |
+
generate_joint_audio_video,
|
| 71 |
+
)
|
| 72 |
+
from videox_fun.models import AutoencoderKLWan3_8, WanT5EncoderModel # noqa: E402
|
| 73 |
+
from videox_fun.models.creator_dac_vae import CreatorDACVAE # noqa: E402
|
| 74 |
+
from videox_fun.models.creator_gating import WanCreatorGatingAVModel # noqa: E402
|
| 75 |
+
|
| 76 |
+
REPO_ID = "GD-ML/DreamX-Creator"
|
| 77 |
+
APP_DIR = Path(__file__).parent.resolve()
|
| 78 |
+
CKPT_DIR = Path(os.environ.get("DREAMX_CKPT_DIR", str(APP_DIR / "checkpoints")))
|
| 79 |
+
CONFIG_PATH = APP_DIR / "config" / "config.yaml"
|
| 80 |
+
WEIGHT_DTYPE = torch.bfloat16
|
| 81 |
+
FPS = 24
|
| 82 |
+
CACHE_VERSION = 1
|
| 83 |
+
|
| 84 |
+
# Authors' defaults (audio_video_generation/inference.py + inference.sh).
|
| 85 |
+
VIDEO_BRIDGE_GUIDANCE = 3.5
|
| 86 |
+
AUDIO_BRIDGE_GUIDANCE = 3.5
|
| 87 |
+
VIDEO_SHIFT = 5.0
|
| 88 |
+
AUDIO_SHIFT = 5.0
|
| 89 |
+
|
| 90 |
+
RESOLUTION_CHOICES = [
|
| 91 |
+
("Fast — ~360p", 220),
|
| 92 |
+
("Balanced — ~480p", 440),
|
| 93 |
+
("Sharp — ~600p", 660),
|
| 94 |
+
]
|
| 95 |
+
|
| 96 |
+
# Neutral keyframe (Option A) dimensions — portrait 9:16, shorts-first.
|
| 97 |
+
KEYFRAME_WIDTH = int(os.environ.get("DREAMX_KEYFRAME_WIDTH", "720"))
|
| 98 |
+
KEYFRAME_HEIGHT = int(os.environ.get("DREAMX_KEYFRAME_HEIGHT", "1280"))
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def _log(msg: str) -> None:
|
| 102 |
+
print(f"[dreamx] {msg}", flush=True)
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def _disk() -> str:
|
| 106 |
+
total, used, free = shutil.disk_usage("/")
|
| 107 |
+
return f"disk used={used / 2**30:.1f}GB free={free / 2**30:.1f}GB"
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
# --------------------------------------------------------------------------- #
|
| 111 |
+
# Model loading (module scope, eagerly on "cuda" — ZeroGPU packs from here)
|
| 112 |
+
# --------------------------------------------------------------------------- #
|
| 113 |
+
|
| 114 |
+
_cfg = OmegaConf.load(CONFIG_PATH)
|
| 115 |
+
_video_kwargs = OmegaConf.to_container(_cfg["video_transformer_additional_kwargs"], resolve=True)
|
| 116 |
+
_audio_kwargs = OmegaConf.to_container(_cfg["audio_transformer_additional_kwargs"], resolve=True)
|
| 117 |
+
_gating_kwargs = OmegaConf.to_container(_cfg["creator_gating_kwargs"], resolve=True)
|
| 118 |
+
_video_vae_kwargs = OmegaConf.to_container(_cfg["video_vae_kwargs"], resolve=True)
|
| 119 |
+
_text_encoder_kwargs = OmegaConf.to_container(_cfg["text_encoder_kwargs"], resolve=True)
|
| 120 |
+
_scheduler_kwargs = OmegaConf.to_container(_cfg["scheduler_kwargs"], resolve=True)
|
| 121 |
+
MAX_SEQUENCE_LENGTH = int(_text_encoder_kwargs.get("text_length", 512))
|
| 122 |
+
|
| 123 |
+
_log(f"checkpoints in {CKPT_DIR} ({_disk()})")
|
| 124 |
+
_log(f"downloading joint AV generator ... ({_disk()})")
|
| 125 |
+
snapshot_download(REPO_ID, local_dir=str(CKPT_DIR), allow_patterns=["creator/*", "creator/**/*"])
|
| 126 |
+
_log(f"loading joint AV generator ... ({_disk()})")
|
| 127 |
+
|
| 128 |
+
transformer = WanCreatorGatingAVModel.from_pretrained(
|
| 129 |
+
pretrained_model_path=str(CKPT_DIR / "creator"),
|
| 130 |
+
video_pretrained_model_path=str(CKPT_DIR / "creator" / "video_model"),
|
| 131 |
+
audio_pretrained_model_path=str(CKPT_DIR / "creator" / "audio_model"),
|
| 132 |
+
video_subfolder=_video_kwargs.get("transformer_low_noise_model_subpath", None),
|
| 133 |
+
audio_subfolder=_audio_kwargs.get("transformer_low_noise_model_subpath", None),
|
| 134 |
+
video_kwargs=_video_kwargs,
|
| 135 |
+
audio_kwargs=_audio_kwargs,
|
| 136 |
+
low_cpu_mem_usage=True,
|
| 137 |
+
torch_dtype=WEIGHT_DTYPE,
|
| 138 |
+
use_temporal_rope=_gating_kwargs.get("use_temporal_rope", True),
|
| 139 |
+
audio_fps=_gating_kwargs.get("audio_fps", 48000.0 / 960.0),
|
| 140 |
+
vae_temporal_stride=_gating_kwargs.get("vae_temporal_stride", 4),
|
| 141 |
+
a2v_cross_attn_layers=_gating_kwargs.get("a2v_cross_attn_layers", None),
|
| 142 |
+
v2a_cross_attn_layers=_gating_kwargs.get("v2a_cross_attn_layers", None),
|
| 143 |
+
use_gating=_gating_kwargs.get("use_gating", True),
|
| 144 |
+
zero_init_cross_attn=_gating_kwargs.get("zero_init_cross_attn", False),
|
| 145 |
+
zero_init_gating=_gating_kwargs.get("zero_init_gating", True),
|
| 146 |
+
gate_init_value=_gating_kwargs.get("gate_init_value", 0.0),
|
| 147 |
+
a2v_gate_alphas=_gating_kwargs.get("a2v_gate_alphas", None),
|
| 148 |
+
v2a_gate_alphas=_gating_kwargs.get("v2a_gate_alphas", None),
|
| 149 |
+
)
|
| 150 |
+
transformer.eval()
|
| 151 |
+
# fp32 source shards are no longer needed once the bf16 model is in RAM.
|
| 152 |
+
shutil.rmtree(CKPT_DIR / "creator", ignore_errors=True)
|
| 153 |
+
_log(f"joint AV generator loaded ({_disk()})")
|
| 154 |
+
|
| 155 |
+
_log("downloading VAEs + UMT5-xxl text encoder ...")
|
| 156 |
+
snapshot_download(
|
| 157 |
+
REPO_ID,
|
| 158 |
+
local_dir=str(CKPT_DIR),
|
| 159 |
+
allow_patterns=["audio_vae/*", "wan2.2_ti2v_5b/*", "wan2.2_ti2v_5b/**/*"],
|
| 160 |
+
)
|
| 161 |
+
_log(f"loading VAEs + text encoder ... ({_disk()})")
|
| 162 |
|
| 163 |
+
_wan_dir = CKPT_DIR / "wan2.2_ti2v_5b"
|
| 164 |
+
_video_vae_path = _wan_dir / _video_vae_kwargs.get("vae_subpath", "Wan2.2_VAE.pth")
|
| 165 |
+
video_vae = AutoencoderKLWan3_8.from_pretrained(
|
| 166 |
+
str(_video_vae_path), additional_kwargs=_video_vae_kwargs
|
| 167 |
+
).eval()
|
| 168 |
+
audio_vae = CreatorDACVAE.from_pretrained(str(CKPT_DIR / "audio_vae"), strict=False).eval()
|
|
|
|
| 169 |
|
| 170 |
+
tokenizer = AutoTokenizer.from_pretrained(
|
| 171 |
+
str(_wan_dir / _text_encoder_kwargs.get("tokenizer_subpath", "google/umt5-xxl"))
|
| 172 |
+
)
|
| 173 |
+
_text_encoder_path = _wan_dir / _text_encoder_kwargs.get(
|
| 174 |
+
"text_encoder_subpath", "models_t5_umt5-xxl-enc-bf16.pth"
|
| 175 |
+
)
|
| 176 |
+
text_encoder = WanT5EncoderModel.from_pretrained(
|
| 177 |
+
str(_text_encoder_path),
|
| 178 |
+
additional_kwargs=_text_encoder_kwargs,
|
| 179 |
+
low_cpu_mem_usage=True,
|
| 180 |
+
torch_dtype=WEIGHT_DTYPE,
|
| 181 |
+
).eval()
|
| 182 |
+
|
| 183 |
+
for _stale in (_video_vae_path, _text_encoder_path):
|
| 184 |
try:
|
| 185 |
+
os.remove(_stale)
|
| 186 |
+
except OSError:
|
| 187 |
+
pass
|
| 188 |
+
shutil.rmtree(CKPT_DIR / "audio_vae", ignore_errors=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 189 |
|
| 190 |
+
transformer = DirectionalMultimodalCFGAdapter(
|
| 191 |
+
transformer,
|
| 192 |
+
video_scale=VIDEO_BRIDGE_GUIDANCE,
|
| 193 |
+
audio_scale=AUDIO_BRIDGE_GUIDANCE,
|
| 194 |
+
enable_a2v=True,
|
| 195 |
+
enable_v2a=True,
|
| 196 |
+
).eval()
|
| 197 |
|
| 198 |
+
transformer.to("cuda")
|
| 199 |
+
text_encoder.to("cuda")
|
| 200 |
+
video_vae.to("cuda")
|
| 201 |
+
audio_vae.to("cuda")
|
| 202 |
+
_log(f"all models resident on cuda ({_disk()})")
|
| 203 |
|
|
|
|
| 204 |
|
| 205 |
+
# --------------------------------------------------------------------------- #
|
| 206 |
+
# Helpers
|
| 207 |
+
# --------------------------------------------------------------------------- #
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 208 |
|
| 209 |
+
def _latent_frames(duration: float) -> int:
|
| 210 |
+
num_frames = int(duration * FPS)
|
| 211 |
+
num_frames = int((num_frames - 1) // 4 * 4) + 1
|
| 212 |
+
return (num_frames - 1) // 4 + 1
|
|
|
|
| 213 |
|
| 214 |
|
| 215 |
+
DURATION_CAP = 300
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 216 |
|
| 217 |
|
| 218 |
+
def _raw_gpu_seconds(seconds, num_inference_steps, resolution_tokens) -> float:
|
| 219 |
+
"""Predicted GPU seconds, fit to two measured ZeroGPU runs.
|
| 220 |
|
| 221 |
+
Measured on RTX PRO 6000 (sm_120), bf16, 3-branch directional multimodal CFG:
|
| 222 |
+
2640 video tokens x 10 steps -> 10.77s denoise, 16.0s total
|
| 223 |
+
7560 video tokens x 30 steps -> 95.04s denoise, 104.7s total
|
| 224 |
+
"""
|
| 225 |
+
tokens = _latent_frames(float(seconds)) * int(resolution_tokens)
|
| 226 |
+
# per denoising step: linear (FFN/proj) + quadratic (self-attention) terms
|
| 227 |
+
per_step = 4.02e-4 * tokens + 2.256e-9 * tokens * tokens
|
| 228 |
+
denoise = per_step * int(num_inference_steps)
|
| 229 |
+
overhead = 4.5 + 7.0e-4 * tokens # T5 encode, first-frame VAE encode, decode, mux
|
| 230 |
+
return denoise + overhead
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
def _estimate_duration(
|
| 234 |
+
image=None,
|
| 235 |
+
prompt="",
|
| 236 |
+
seconds=3.0,
|
| 237 |
+
num_inference_steps=30,
|
| 238 |
+
resolution_tokens=440,
|
| 239 |
+
*args,
|
| 240 |
+
**kwargs,
|
| 241 |
+
):
|
| 242 |
+
"""GPU seconds to reserve — calibrated against measured runs on ZeroGPU."""
|
| 243 |
+
raw = _raw_gpu_seconds(seconds, num_inference_steps, resolution_tokens)
|
| 244 |
+
return int(min(DURATION_CAP, math.ceil(raw * 1.15) + 5))
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
def _write_mp4(frames: np.ndarray, audio: np.ndarray, sample_rate: int, fps: int) -> str:
|
| 248 |
+
"""Mux uint8 RGB frames + mono float audio into a single H.264/AAC mp4."""
|
| 249 |
+
height, width = frames.shape[1], frames.shape[2]
|
| 250 |
+
out_path = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name
|
| 251 |
+
wav_path = tempfile.NamedTemporaryFile(suffix=".wav", delete=False).name
|
| 252 |
+
|
| 253 |
+
import soundfile as sf
|
| 254 |
+
|
| 255 |
+
sf.write(wav_path, np.clip(audio, -1.0, 1.0), sample_rate)
|
| 256 |
+
|
| 257 |
+
base = [
|
| 258 |
+
"ffmpeg", "-hide_banner", "-loglevel", "error", "-y",
|
| 259 |
+
"-f", "rawvideo", "-pix_fmt", "rgb24",
|
| 260 |
+
"-s", f"{width}x{height}", "-r", str(fps), "-i", "-",
|
| 261 |
+
"-i", wav_path,
|
| 262 |
+
]
|
| 263 |
+
for vcodec in ("libx264", "mpeg4"):
|
| 264 |
+
cmd = base + [
|
| 265 |
+
"-c:v", vcodec, "-pix_fmt", "yuv420p", "-crf", "18",
|
| 266 |
+
"-c:a", "aac", "-b:a", "192k", "-shortest", out_path,
|
| 267 |
+
]
|
| 268 |
+
if vcodec == "mpeg4":
|
| 269 |
+
cmd.remove("-crf")
|
| 270 |
+
cmd.remove("18")
|
| 271 |
+
proc = subprocess.run(cmd, input=frames.tobytes(), capture_output=True)
|
| 272 |
+
if proc.returncode == 0 and os.path.getsize(out_path) > 0:
|
| 273 |
+
os.remove(wav_path)
|
| 274 |
+
return out_path
|
| 275 |
+
_log(f"ffmpeg ({vcodec}) failed: {proc.stderr.decode()[-600:]}")
|
| 276 |
+
raise gr.Error("ffmpeg failed to encode the generated video.")
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
def _neutral_keyframe(width: int = KEYFRAME_WIDTH, height: int = KEYFRAME_HEIGHT) -> str:
|
| 280 |
+
"""Deterministic dark diagonal gradient — the "Option A" neutral first frame.
|
| 281 |
+
|
| 282 |
+
By default a 9:16 portrait so generated clips keep a shorts-friendly aspect
|
| 283 |
+
ratio when no real first-frame image is provided.
|
| 284 |
+
"""
|
| 285 |
+
y, x = np.mgrid[0:height, 0:width]
|
| 286 |
+
norm = np.sqrt((x / max(width - 1, 1)) ** 2 + (y / max(height - 1, 1)) ** 2)
|
| 287 |
+
tone = (0.06 + 0.16 * norm[..., None] * np.array([0.86, 0.95, 1.0])).clip(0, 1)
|
| 288 |
+
frame = (tone * 255.0).astype(np.uint8)
|
| 289 |
+
path = tempfile.NamedTemporaryFile(suffix=".png", delete=False).name
|
| 290 |
+
Image.fromarray(frame).save(path)
|
| 291 |
+
return path
|
| 292 |
+
|
| 293 |
+
|
| 294 |
+
# --------------------------------------------------------------------------- #
|
| 295 |
+
# Inference
|
| 296 |
+
# --------------------------------------------------------------------------- #
|
| 297 |
+
|
| 298 |
+
|
| 299 |
+
@spaces.GPU(duration=_estimate_duration)
|
| 300 |
+
def generate(
|
| 301 |
+
image: str,
|
| 302 |
+
prompt: str,
|
| 303 |
+
seconds: float = 3.0,
|
| 304 |
+
num_inference_steps: int = 30,
|
| 305 |
+
resolution_tokens: int = 440,
|
| 306 |
+
guidance_scale: float = 5.0,
|
| 307 |
+
negative_prompt: str = DEFAULT_NEGATIVE_PROMPT,
|
| 308 |
+
seed: int = 113,
|
| 309 |
+
randomize_seed: bool = False,
|
| 310 |
+
progress=gr.Progress(track_tqdm=True),
|
| 311 |
+
):
|
| 312 |
+
"""Generate a video with a jointly-denoised soundtrack from a first frame and a prompt.
|
| 313 |
+
|
| 314 |
+
Args:
|
| 315 |
+
image: path to the image used as the first frame of the video. Empty
|
| 316 |
+
means "no image": a neutral keyframe is generated in-app instead.
|
| 317 |
+
prompt: description of the action AND the sound to generate; put spoken
|
| 318 |
+
lines in quotes (e.g. `Man says, 'hello there.'`).
|
| 319 |
+
seconds: length of the clip in seconds (24 fps).
|
| 320 |
+
num_inference_steps: number of flow-matching denoising steps.
|
| 321 |
+
resolution_tokens: spatial token budget; higher means higher resolution.
|
| 322 |
+
guidance_scale: classifier-free guidance strength for text.
|
| 323 |
+
negative_prompt: what to avoid in both the video and the audio.
|
| 324 |
+
seed: RNG seed for reproducible sampling.
|
| 325 |
+
randomize_seed: draw a fresh random seed instead of using `seed`.
|
| 326 |
+
|
| 327 |
+
Returns:
|
| 328 |
+
A tuple of (path to the generated mp4 with audio, path to the last-frame
|
| 329 |
+
PNG for chaining into the next scene, the seed actually used).
|
| 330 |
+
"""
|
| 331 |
+
if not prompt or not prompt.strip():
|
| 332 |
+
raise gr.Error("Please provide a prompt describing the motion and the sound.")
|
| 333 |
+
if not image:
|
| 334 |
+
_log("empty first-frame image -> neutral keyframe")
|
| 335 |
+
image = _neutral_keyframe()
|
| 336 |
+
|
| 337 |
+
# ZeroGPU kills a task that outruns its reservation, and the reservation is
|
| 338 |
+
# capped, so refuse the few extreme knob combinations that cannot fit.
|
| 339 |
+
if _raw_gpu_seconds(seconds, num_inference_steps, resolution_tokens) * 1.15 + 5 > DURATION_CAP:
|
| 340 |
+
raise gr.Error(
|
| 341 |
+
f"That combination needs more than the {DURATION_CAP}s GPU budget. "
|
| 342 |
+
"Lower the duration, the resolution, or the number of steps."
|
| 343 |
+
)
|
| 344 |
|
| 345 |
+
if randomize_seed:
|
| 346 |
+
seed = random.randint(0, 2**31 - 1)
|
| 347 |
+
seed = int(seed)
|
| 348 |
+
|
| 349 |
+
device = torch.device("cuda")
|
| 350 |
+
video_scheduler_kwargs = dict(_scheduler_kwargs, shift=VIDEO_SHIFT)
|
| 351 |
+
audio_scheduler_kwargs = dict(_scheduler_kwargs, shift=AUDIO_SHIFT)
|
| 352 |
+
|
| 353 |
+
models = {
|
| 354 |
+
"config": _cfg,
|
| 355 |
+
"transformer": transformer,
|
| 356 |
+
"video_vae": video_vae,
|
| 357 |
+
"audio_vae": audio_vae,
|
| 358 |
+
"tokenizer": tokenizer,
|
| 359 |
+
"text_encoder": text_encoder,
|
| 360 |
+
# fresh schedulers per request: they carry mutable step state
|
| 361 |
+
"video_scheduler": FlowMatchEulerDiscreteScheduler(
|
| 362 |
+
**filter_kwargs(FlowMatchEulerDiscreteScheduler, video_scheduler_kwargs)
|
| 363 |
+
),
|
| 364 |
+
"audio_scheduler": FlowMatchEulerDiscreteScheduler(
|
| 365 |
+
**filter_kwargs(FlowMatchEulerDiscreteScheduler, audio_scheduler_kwargs)
|
| 366 |
+
),
|
| 367 |
+
"max_sequence_length": MAX_SEQUENCE_LENGTH,
|
| 368 |
+
}
|
| 369 |
+
|
| 370 |
+
args = SimpleNamespace(
|
| 371 |
+
config_path=str(CONFIG_PATH),
|
| 372 |
+
image=image,
|
| 373 |
+
output="output.mp4",
|
| 374 |
+
negative_prompt=negative_prompt or DEFAULT_NEGATIVE_PROMPT,
|
| 375 |
+
duration=float(seconds),
|
| 376 |
+
target_spatial_tokens=int(resolution_tokens),
|
| 377 |
+
min_token_ratio=0.95,
|
| 378 |
+
fps=FPS,
|
| 379 |
+
num_inference_steps=int(num_inference_steps),
|
| 380 |
+
guidance_scale=float(guidance_scale),
|
| 381 |
+
cfg_mode="multimodal",
|
| 382 |
+
video_bridge_guidance_scale=VIDEO_BRIDGE_GUIDANCE,
|
| 383 |
+
audio_bridge_guidance_scale=AUDIO_BRIDGE_GUIDANCE,
|
| 384 |
+
seed=seed,
|
| 385 |
+
video_shift=VIDEO_SHIFT,
|
| 386 |
+
audio_shift=AUDIO_SHIFT,
|
| 387 |
+
flow_match_mu=None,
|
| 388 |
+
sampler_name="Flow",
|
| 389 |
+
weight_dtype="bfloat16",
|
| 390 |
+
GPU_memory_mode="model_full_load",
|
| 391 |
+
text_encoder_cpu_offload=False,
|
| 392 |
+
video_vae_cpu_offload=False,
|
| 393 |
+
audio_vae_cpu_offload=False,
|
| 394 |
+
vae_cpu_offload=False,
|
| 395 |
+
use_temporal_rope=True,
|
| 396 |
+
audio_fps=48000.0 / 960.0,
|
| 397 |
+
vae_temporal_stride=4,
|
| 398 |
+
disable_a2v_cross_attn=False,
|
| 399 |
+
disable_v2a_cross_attn=False,
|
| 400 |
+
suppress_aux_writes=True,
|
| 401 |
+
skip_output_decode=False,
|
| 402 |
+
disable_progress=False,
|
| 403 |
+
synchronize_noise=False,
|
| 404 |
+
ulysses_degree=1,
|
| 405 |
+
ring_degree=1,
|
| 406 |
+
fsdp_dit=False,
|
| 407 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 408 |
|
| 409 |
+
item = {
|
| 410 |
+
"prompt": prompt.strip(),
|
| 411 |
+
"video_prompt": prompt.strip(),
|
| 412 |
+
"audio_prompt": prompt.strip(),
|
| 413 |
+
"negative_prompt": args.negative_prompt,
|
| 414 |
+
"audio_negative_prompt": args.negative_prompt,
|
| 415 |
+
"duration": args.duration,
|
| 416 |
+
"guidance_scale": args.guidance_scale,
|
| 417 |
+
"num_inference_steps": args.num_inference_steps,
|
| 418 |
+
"seed": seed,
|
| 419 |
+
"first_frame_path": image,
|
| 420 |
+
"name": "sample",
|
| 421 |
+
}
|
| 422 |
+
|
| 423 |
+
started = time.perf_counter()
|
| 424 |
+
video_decoded, audio_decoded, num_frames = generate_joint_audio_video(
|
| 425 |
+
args, models, device, WEIGHT_DTYPE, item
|
|
|
|
| 426 |
)
|
| 427 |
+
gpu_seconds = time.perf_counter() - started
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 428 |
|
| 429 |
+
frames = (
|
| 430 |
+
video_decoded[0].permute(1, 2, 3, 0).clamp(0, 1).numpy() * 255.0
|
| 431 |
+
).astype(np.uint8)
|
| 432 |
+
waveform = audio_decoded.detach().float().cpu()
|
| 433 |
+
while waveform.ndim > 1:
|
| 434 |
+
waveform = waveform[0]
|
| 435 |
+
audio = waveform.numpy()
|
| 436 |
|
| 437 |
+
out_path = _write_mp4(frames, audio, int(audio_vae.sample_rate), FPS)
|
| 438 |
|
| 439 |
+
last_frame_path = tempfile.NamedTemporaryFile(suffix=".png", delete=False).name
|
| 440 |
+
Image.fromarray(frames[-1]).save(last_frame_path)
|
| 441 |
|
| 442 |
+
_log(
|
| 443 |
+
f"done in {gpu_seconds:.1f}s | {num_frames} frames @ {frames.shape[2]}x{frames.shape[1]} "
|
| 444 |
+
f"| steps={args.num_inference_steps} tokens={args.target_spatial_tokens} "
|
| 445 |
+
f"| reserved={_estimate_duration(image, prompt, seconds, num_inference_steps, resolution_tokens)}s"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 446 |
)
|
| 447 |
+
return out_path, last_frame_path, seed
|
| 448 |
+
|
| 449 |
+
|
| 450 |
+
# --------------------------------------------------------------------------- #
|
| 451 |
+
# UI
|
| 452 |
+
# --------------------------------------------------------------------------- #
|
| 453 |
+
|
| 454 |
+
EXAMPLES = [
|
| 455 |
+
[
|
| 456 |
+
"examples/case1.jpg",
|
| 457 |
+
"A man in a dark suit and white shirt is seated on a yellow couch, speaking about "
|
| 458 |
+
"the language of Americans. He uses hand gestures to emphasize his points, and the "
|
| 459 |
+
"background shows a cityscape with illuminated buildings, suggesting an urban setting, "
|
| 460 |
+
"possibly a studio with a city view. Man says, 'The thing about Americans that I've "
|
| 461 |
+
"thought about the language is that they speak, they say they speak English.'.",
|
| 462 |
+
],
|
| 463 |
+
[
|
| 464 |
+
"examples/case4.jpg",
|
| 465 |
+
"The video shows a wolf standing on its hind legs, howling with its mouth wide open, "
|
| 466 |
+
"showing its teeth and tongue. The wolf's ears are perked up, and its eyes are focused "
|
| 467 |
+
"on something in the distance. The background consists of trees with green and yellow "
|
| 468 |
+
"leaves, indicating it is autumn. The wolf's howl is loud and resonant, filling the air "
|
| 469 |
+
"with its powerful voice. The sound of a dog howls and howls can be heard.",
|
| 470 |
+
],
|
| 471 |
+
[
|
| 472 |
+
"examples/case3.jpg",
|
| 473 |
+
"The video captures a dramatic night scene with a series of lightning strikes "
|
| 474 |
+
"illuminating the dark sky and revealing the city lights below. The clouds move across "
|
| 475 |
+
"the sky, and the lightning strikes again, followed by a thunderclap. The sound of a "
|
| 476 |
+
"thunderstorm and rain falling can be heard.",
|
| 477 |
+
],
|
| 478 |
+
[
|
| 479 |
+
"examples/case2.jpg",
|
| 480 |
+
"A humanoid robot is cooking in a modern kitchen. The robot, with a white body and blue "
|
| 481 |
+
"eyes, is stirring food in a pan on the stove. The kitchen is equipped with dark cabinets "
|
| 482 |
+
"and various utensils. The robot's movements are smooth and precise as it stirs the food, "
|
| 483 |
+
"causing steam to rise. The ambient sound of cooking can be heard throughout the scene.",
|
| 484 |
+
],
|
| 485 |
+
[
|
| 486 |
+
"examples/case5.jpg",
|
| 487 |
+
"A person in a red plaid shirt is typing on a white keyboard placed on a wooden table. "
|
| 488 |
+
"The scene is set in a room with a wooden floor and a part of a white blanket visible in "
|
| 489 |
+
"the background. The person's hands are actively moving across the keyboard, indicating "
|
| 490 |
+
"typing activity. The ambient sound is the distinct sound of keys being pressed.",
|
| 491 |
+
],
|
| 492 |
+
[
|
| 493 |
+
"examples/case6.jpg",
|
| 494 |
+
"A man is playing an acoustic guitar in a modern kitchen setting. He is focused on his "
|
| 495 |
+
"playing, with his hands moving along the strings and fretboard. The background features "
|
| 496 |
+
"a well-lit kitchen with wooden cabinets and hanging lights.",
|
| 497 |
+
],
|
| 498 |
+
]
|
| 499 |
+
|
| 500 |
+
CSS = """
|
| 501 |
+
#col-container { max-width: 1180px; margin: 0 auto; }
|
| 502 |
+
.dark .gradio-container { color: var(--body-text-color); }
|
| 503 |
+
"""
|
| 504 |
|
| 505 |
+
with gr.Blocks(title="DreamX-Creator 1.0") as demo:
|
| 506 |
+
with gr.Column(elem_id="col-container"):
|
| 507 |
+
gr.Markdown(
|
| 508 |
+
"# 🎬 DreamX-Creator 1.0\n"
|
| 509 |
+
"Turn **one image + one prompt** into a video whose **soundtrack is generated "
|
| 510 |
+
"jointly with the frames** — speech, foley and ambience come out of the same "
|
| 511 |
+
"denoiser as the picture, so they stay in sync.\n\n"
|
| 512 |
+
"Describe the *sound* as well as the action. For speech, quote the line: "
|
| 513 |
+
"`Man says, 'hello there.'`\n\n"
|
| 514 |
+
"[Model](https://huggingface.co/GD-ML/DreamX-Creator) · "
|
| 515 |
+
"[Code](https://github.com/AMAP-ML/DreamX-Creator)"
|
|
|
|
|
|
|
|
|
|
| 516 |
)
|
| 517 |
|
| 518 |
+
with gr.Row():
|
| 519 |
+
with gr.Column(scale=1):
|
| 520 |
+
image = gr.Image(label="First frame (optional — leave empty for a neutral keyframe)", type="filepath", height=320)
|
| 521 |
+
prompt = gr.Textbox(
|
| 522 |
+
label="Prompt",
|
| 523 |
+
placeholder="Describe the motion and the sound you want to hear…",
|
| 524 |
+
lines=4,
|
| 525 |
+
)
|
| 526 |
+
run = gr.Button("Generate audio + video", variant="primary")
|
| 527 |
+
with gr.Column(scale=1):
|
| 528 |
+
video_out = gr.Video(label="Result (video + generated audio)", height=380)
|
| 529 |
+
last_frame_out = gr.Image(label="Last frame (feeds the next scene)", height=380)
|
| 530 |
+
|
| 531 |
+
with gr.Accordion("Advanced settings", open=False):
|
| 532 |
+
with gr.Row():
|
| 533 |
+
seconds = gr.Slider(
|
| 534 |
+
label="Duration (seconds)", minimum=2.0, maximum=5.0, step=0.5, value=3.0
|
| 535 |
+
)
|
| 536 |
+
num_inference_steps = gr.Slider(
|
| 537 |
+
label="Denoising steps", minimum=10, maximum=50, step=1, value=30
|
| 538 |
+
)
|
| 539 |
+
with gr.Row():
|
| 540 |
+
resolution_tokens = gr.Dropdown(
|
| 541 |
+
label="Resolution budget (spatial tokens)",
|
| 542 |
+
choices=RESOLUTION_CHOICES,
|
| 543 |
+
value=440,
|
| 544 |
+
)
|
| 545 |
+
guidance_scale = gr.Slider(
|
| 546 |
+
label="Guidance scale", minimum=1.0, maximum=10.0, step=0.1, value=5.0
|
| 547 |
+
)
|
| 548 |
+
negative_prompt = gr.Textbox(
|
| 549 |
+
label="Negative prompt", value=DEFAULT_NEGATIVE_PROMPT, lines=2
|
| 550 |
+
)
|
| 551 |
+
with gr.Row():
|
| 552 |
+
seed = gr.Number(label="Seed", value=113, precision=0)
|
| 553 |
+
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
| 554 |
|
| 555 |
+
gr.Markdown(
|
| 556 |
+
"Longer clips, more steps and a bigger resolution budget all cost GPU time "
|
| 557 |
+
"roughly linearly (and attention grows quadratically with tokens × frames)."
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 558 |
)
|
| 559 |
|
| 560 |
+
gr.Examples(
|
| 561 |
+
examples=EXAMPLES,
|
| 562 |
+
inputs=[image, prompt],
|
| 563 |
+
outputs=[video_out, last_frame_out, seed],
|
| 564 |
+
fn=generate,
|
| 565 |
+
cache_examples=True,
|
| 566 |
+
cache_mode="lazy",
|
| 567 |
+
label=f"Official Verse-Bench cases from the DreamX-Creator repo (v{CACHE_VERSION})",
|
|
|
|
|
|
|
|
|
|
| 568 |
)
|
| 569 |
|
| 570 |
+
run.click(
|
| 571 |
+
fn=generate,
|
| 572 |
+
inputs=[
|
| 573 |
+
image,
|
| 574 |
+
prompt,
|
| 575 |
+
seconds,
|
| 576 |
+
num_inference_steps,
|
| 577 |
+
resolution_tokens,
|
| 578 |
+
guidance_scale,
|
| 579 |
+
negative_prompt,
|
| 580 |
+
seed,
|
| 581 |
+
randomize_seed,
|
| 582 |
+
],
|
| 583 |
+
outputs=[video_out, last_frame_out, seed],
|
| 584 |
+
api_name="generate",
|
| 585 |
+
)
|
| 586 |
+
|
| 587 |
if __name__ == "__main__":
|
| 588 |
+
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)
|
config/config.yaml
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
video_transformer_additional_kwargs:
|
| 2 |
+
transformer_low_noise_model_subpath: ./
|
| 3 |
+
transformer_combination_type: single
|
| 4 |
+
dict_mapping:
|
| 5 |
+
in_dim: in_channels
|
| 6 |
+
dim: hidden_size
|
| 7 |
+
|
| 8 |
+
audio_transformer_additional_kwargs:
|
| 9 |
+
transformer_low_noise_model_subpath: .
|
| 10 |
+
patch_size: [1]
|
| 11 |
+
in_dim: 128
|
| 12 |
+
out_dim: 128
|
| 13 |
+
vae_type: dac
|
| 14 |
+
|
| 15 |
+
video_vae_kwargs:
|
| 16 |
+
vae_type: AutoencoderKLWan3_8
|
| 17 |
+
vae_subpath: Wan2.2_VAE.pth
|
| 18 |
+
temporal_compression_ratio: 4
|
| 19 |
+
spatial_compression_ratio: 16
|
| 20 |
+
|
| 21 |
+
text_encoder_kwargs:
|
| 22 |
+
text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
|
| 23 |
+
tokenizer_subpath: google/umt5-xxl
|
| 24 |
+
text_length: 512
|
| 25 |
+
vocab: 256384
|
| 26 |
+
dim: 4096
|
| 27 |
+
dim_attn: 4096
|
| 28 |
+
dim_ffn: 10240
|
| 29 |
+
num_heads: 64
|
| 30 |
+
num_layers: 24
|
| 31 |
+
num_buckets: 32
|
| 32 |
+
shared_pos: false
|
| 33 |
+
dropout: 0.0
|
| 34 |
+
|
| 35 |
+
scheduler_kwargs:
|
| 36 |
+
scheduler_subpath: null
|
| 37 |
+
num_train_timesteps: 1000
|
| 38 |
+
shift: 5.0
|
| 39 |
+
use_dynamic_shifting: false
|
| 40 |
+
base_shift: 0.5
|
| 41 |
+
max_shift: 1.15
|
| 42 |
+
base_image_seq_len: 256
|
| 43 |
+
max_image_seq_len: 4096
|
| 44 |
+
|
| 45 |
+
creator_gating_kwargs:
|
| 46 |
+
use_temporal_rope: true
|
| 47 |
+
audio_fps: 50.0
|
| 48 |
+
vae_temporal_stride: 4
|
| 49 |
+
a2v_cross_attn_layers: [15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29]
|
| 50 |
+
v2a_cross_attn_layers: [15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29]
|
| 51 |
+
use_gating: true
|
| 52 |
+
zero_init_cross_attn: false
|
| 53 |
+
zero_init_gating: false
|
| 54 |
+
gate_init_value: 0.5
|
dreamx_inference.py
ADDED
|
@@ -0,0 +1,1030 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Single-image Creator audio-video inference.
|
| 2 |
+
|
| 3 |
+
The supplied image is encoded as the first video frame. The prompt conditions
|
| 4 |
+
the joint Creator video/audio denoiser, which writes a video, a WAV, and a merged
|
| 5 |
+
MP4. This entrypoint intentionally has no dataset, Verse-Bench, training, or
|
| 6 |
+
model-download workflow.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
import argparse
|
| 10 |
+
import inspect
|
| 11 |
+
import math
|
| 12 |
+
import os
|
| 13 |
+
import sys
|
| 14 |
+
import time
|
| 15 |
+
import subprocess
|
| 16 |
+
|
| 17 |
+
import numpy as np
|
| 18 |
+
import torch
|
| 19 |
+
import torch.distributed as dist
|
| 20 |
+
from diffusers import FlowMatchEulerDiscreteScheduler
|
| 21 |
+
from omegaconf import OmegaConf
|
| 22 |
+
from PIL import Image
|
| 23 |
+
from transformers import AutoTokenizer as HFAutoTokenizer
|
| 24 |
+
try: # optional: only used by this module's standalone CLI entrypoint
|
| 25 |
+
from torchvision.io import write_video
|
| 26 |
+
except Exception: # pragma: no cover
|
| 27 |
+
write_video = None
|
| 28 |
+
from torch import nn
|
| 29 |
+
from tqdm.auto import tqdm
|
| 30 |
+
|
| 31 |
+
import warnings
|
| 32 |
+
warnings.simplefilter(action='ignore', category=FutureWarning)
|
| 33 |
+
|
| 34 |
+
current_file_path = os.path.abspath(__file__)
|
| 35 |
+
release_root = os.path.dirname(current_file_path)
|
| 36 |
+
if release_root not in sys.path:
|
| 37 |
+
sys.path.insert(0, release_root)
|
| 38 |
+
|
| 39 |
+
from videox_fun.models import AutoencoderKLWan3_8, WanT5EncoderModel
|
| 40 |
+
from videox_fun.models.creator_gating import WanCreatorGatingAVModel
|
| 41 |
+
from videox_fun.models.creator_dac_vae import CreatorDACVAE
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def filter_kwargs(cls, kwargs):
|
| 45 |
+
"""Keep scheduler options compatible across diffusers versions."""
|
| 46 |
+
valid = set(inspect.signature(cls.__init__).parameters)
|
| 47 |
+
return {key: value for key, value in kwargs.items() if key in valid}
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class DirectionalMultimodalCFGAdapter(nn.Module):
|
| 51 |
+
"""Preserve the experiment's three-branch multimodal CFG in one call."""
|
| 52 |
+
|
| 53 |
+
def __init__(self, model, video_scale, audio_scale, enable_a2v, enable_v2a):
|
| 54 |
+
super().__init__()
|
| 55 |
+
self.model = model
|
| 56 |
+
self.video_scale = float(video_scale)
|
| 57 |
+
self.audio_scale = float(audio_scale)
|
| 58 |
+
self.enable_a2v = bool(enable_a2v)
|
| 59 |
+
self.enable_v2a = bool(enable_v2a)
|
| 60 |
+
|
| 61 |
+
@property
|
| 62 |
+
def video_patch_size(self):
|
| 63 |
+
return self.model.video_patch_size
|
| 64 |
+
|
| 65 |
+
@property
|
| 66 |
+
def audio_patch_size(self):
|
| 67 |
+
return self.model.audio_patch_size
|
| 68 |
+
|
| 69 |
+
@staticmethod
|
| 70 |
+
def _expand(value, name):
|
| 71 |
+
if isinstance(value, torch.Tensor):
|
| 72 |
+
if value.ndim == 0 or value.shape[0] != 2:
|
| 73 |
+
raise ValueError(f"Multimodal CFG expects {name} batch size 2")
|
| 74 |
+
return torch.cat((value[0:1], value[0:1], value[1:2]), dim=0)
|
| 75 |
+
if isinstance(value, (list, tuple)) and len(value) == 2:
|
| 76 |
+
return [value[0], value[0], value[1]]
|
| 77 |
+
raise TypeError(f"Unsupported {name} batch value: {type(value)!r}")
|
| 78 |
+
|
| 79 |
+
@classmethod
|
| 80 |
+
def _expand_inputs(cls, inputs):
|
| 81 |
+
expanded = dict(inputs)
|
| 82 |
+
for key in ("x", "t", "context", "y", "clip_fea"):
|
| 83 |
+
if expanded.get(key) is not None:
|
| 84 |
+
expanded[key] = cls._expand(expanded[key], key)
|
| 85 |
+
return expanded
|
| 86 |
+
|
| 87 |
+
@staticmethod
|
| 88 |
+
def _collapse(prediction, bridge_scale):
|
| 89 |
+
if prediction.shape[0] != 3:
|
| 90 |
+
raise ValueError("Creator multimodal CFG expects three model predictions")
|
| 91 |
+
d00, d0b, dtb = prediction[0:1], prediction[1:2], prediction[2:3]
|
| 92 |
+
bridge = d00 + bridge_scale * (d0b - d00)
|
| 93 |
+
return torch.cat((bridge, bridge + (dtb - d0b)), dim=0)
|
| 94 |
+
|
| 95 |
+
def forward(self, video, audio, dtype=torch.bfloat16, return_dict=True, **kwargs):
|
| 96 |
+
if self.model.training:
|
| 97 |
+
raise RuntimeError("Inference adapter received a training model")
|
| 98 |
+
model_video = self._expand_inputs(video)
|
| 99 |
+
model_audio = self._expand_inputs(audio)
|
| 100 |
+
device = model_video["t"].device
|
| 101 |
+
bridge_mask = torch.tensor([False, True, True], device=device, dtype=torch.bool)
|
| 102 |
+
output = self.model(
|
| 103 |
+
video=model_video,
|
| 104 |
+
audio=model_audio,
|
| 105 |
+
dtype=dtype,
|
| 106 |
+
return_dict=True,
|
| 107 |
+
enable_a2v=bridge_mask & self.enable_a2v,
|
| 108 |
+
enable_v2a=bridge_mask & self.enable_v2a,
|
| 109 |
+
)
|
| 110 |
+
result = {
|
| 111 |
+
"video": self._collapse(output["video"], self.video_scale),
|
| 112 |
+
"audio": self._collapse(output["audio"], self.audio_scale),
|
| 113 |
+
}
|
| 114 |
+
return result if return_dict else (result["video"], result["audio"])
|
| 115 |
+
|
| 116 |
+
try:
|
| 117 |
+
import soundfile as sf
|
| 118 |
+
except ImportError:
|
| 119 |
+
sf = None
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
# ============================================================================
|
| 123 |
+
# Single-process helpers
|
| 124 |
+
# ============================================================================
|
| 125 |
+
|
| 126 |
+
def print_info(*args, **kwargs):
|
| 127 |
+
print(*args, **kwargs)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def synchronize_tensor(tensor: torch.Tensor, args) -> torch.Tensor:
|
| 131 |
+
"""Broadcast a rank-0 tensor when the SP launcher requests exact inputs."""
|
| 132 |
+
if getattr(args, "synchronize_noise", False) and dist.is_initialized():
|
| 133 |
+
dist.broadcast(tensor, src=0)
|
| 134 |
+
return tensor
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
# ============================================================================
|
| 138 |
+
# Argument parsing
|
| 139 |
+
# ============================================================================
|
| 140 |
+
|
| 141 |
+
DEFAULT_NEGATIVE_PROMPT = (
|
| 142 |
+
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指"
|
| 143 |
+
)
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def parse_args():
|
| 147 |
+
parser = argparse.ArgumentParser(description="Single-image Creator audio-video inference")
|
| 148 |
+
|
| 149 |
+
# Model paths
|
| 150 |
+
parser.add_argument("--config_path", type=str, default=os.path.join(os.path.dirname(__file__), "config/config.yaml"))
|
| 151 |
+
parser.add_argument("--model_name", type=str, required=True,
|
| 152 |
+
help="Wan2.2-TI2V-5B model directory")
|
| 153 |
+
parser.add_argument("--transformer_path", type=str, required=True,
|
| 154 |
+
help="Creator checkpoint directory with video_model/ and audio_model/")
|
| 155 |
+
parser.add_argument("--audio_vae_path", type=str, required=True,
|
| 156 |
+
help="Path to CreatorDACVAE audio VAE")
|
| 157 |
+
|
| 158 |
+
# Input / Output
|
| 159 |
+
parser.add_argument("--image", type=str, required=True,
|
| 160 |
+
help="Input image used as the first video frame")
|
| 161 |
+
parser.add_argument("--prompt", type=str, required=True, help="Text prompt")
|
| 162 |
+
parser.add_argument("--output", type=str, default="./outputs/output.mp4",
|
| 163 |
+
help="Output MP4; a WAV with the same stem is also written")
|
| 164 |
+
parser.add_argument("--negative_prompt", type=str, default=DEFAULT_NEGATIVE_PROMPT)
|
| 165 |
+
# Generation parameters
|
| 166 |
+
parser.add_argument("--duration", type=float, default=5.0,
|
| 167 |
+
help="Target duration in seconds")
|
| 168 |
+
parser.add_argument("--target_spatial_tokens", type=int, default=880,
|
| 169 |
+
help="Maximum spatial tokens per frame")
|
| 170 |
+
parser.add_argument("--min_token_ratio", type=float, default=0.95,
|
| 171 |
+
help="Minimum fraction of target spatial tokens")
|
| 172 |
+
parser.add_argument("--fps", type=int, default=24)
|
| 173 |
+
parser.add_argument("--num_inference_steps", type=int, default=50)
|
| 174 |
+
parser.add_argument("--guidance_scale", type=float, default=5.0)
|
| 175 |
+
parser.add_argument("--cfg_mode", type=str, default="multimodal",
|
| 176 |
+
choices=["text", "multimodal"],
|
| 177 |
+
help="CFG mode. 'multimodal' uses Creator dual CFG with separate "
|
| 178 |
+
"bridge and text guidance (3 model evaluations per step).")
|
| 179 |
+
parser.add_argument("--video_bridge_guidance_scale", type=float, default=3.5)
|
| 180 |
+
parser.add_argument("--audio_bridge_guidance_scale", type=float, default=3.5)
|
| 181 |
+
parser.add_argument("--seed", type=int, default=42)
|
| 182 |
+
parser.add_argument("--video_shift", type=float, default=5.0,
|
| 183 |
+
help="Noise schedule shift for video denoising")
|
| 184 |
+
parser.add_argument("--audio_shift", type=float, default=5.0,
|
| 185 |
+
help="Noise schedule shift for audio denoising")
|
| 186 |
+
|
| 187 |
+
# Sampler
|
| 188 |
+
parser.add_argument("--sampler_name", type=str, default="Flow", choices=["Flow"])
|
| 189 |
+
|
| 190 |
+
# Memory & compute
|
| 191 |
+
parser.add_argument("--weight_dtype", type=str, default="bfloat16",
|
| 192 |
+
choices=["float16", "bfloat16", "float32"])
|
| 193 |
+
parser.add_argument("--GPU_memory_mode", type=str, default="model_full_load",
|
| 194 |
+
choices=["model_full_load", "model_cpu_offload"])
|
| 195 |
+
offload_action = argparse.BooleanOptionalAction
|
| 196 |
+
parser.add_argument(
|
| 197 |
+
"--text_encoder_cpu_offload",
|
| 198 |
+
action=offload_action,
|
| 199 |
+
default=None,
|
| 200 |
+
help="Keep the text encoder on CPU except while encoding prompts. "
|
| 201 |
+
"Defaults to the GPU_memory_mode setting.",
|
| 202 |
+
)
|
| 203 |
+
parser.add_argument(
|
| 204 |
+
"--video_vae_cpu_offload",
|
| 205 |
+
action=offload_action,
|
| 206 |
+
default=None,
|
| 207 |
+
help="Keep the video VAE on CPU except while encoding/decoding. "
|
| 208 |
+
"Defaults to the GPU_memory_mode setting.",
|
| 209 |
+
)
|
| 210 |
+
parser.add_argument(
|
| 211 |
+
"--audio_vae_cpu_offload",
|
| 212 |
+
action=offload_action,
|
| 213 |
+
default=None,
|
| 214 |
+
help="Keep the audio VAE on CPU except while decoding. "
|
| 215 |
+
"Defaults to the GPU_memory_mode setting.",
|
| 216 |
+
)
|
| 217 |
+
parser.add_argument(
|
| 218 |
+
"--vae_cpu_offload",
|
| 219 |
+
action=offload_action,
|
| 220 |
+
default=None,
|
| 221 |
+
help="Enable or disable CPU offload for both video and audio VAEs. "
|
| 222 |
+
"Individual VAE flags take precedence.",
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
# Temporal RoPE
|
| 226 |
+
parser.add_argument("--use_temporal_rope", type=bool, default=True)
|
| 227 |
+
parser.add_argument("--audio_fps", type=float, default=48000.0 / 960.0,
|
| 228 |
+
help="Audio latent FPS (DAC: 48000/960=50)")
|
| 229 |
+
parser.add_argument("--vae_temporal_stride", type=int, default=4)
|
| 230 |
+
|
| 231 |
+
# Runtime cross-attention controls
|
| 232 |
+
parser.add_argument("--disable_a2v_cross_attn", "--disable-a2v-cross-attn",
|
| 233 |
+
"--disable_a2v", "--disable-a2v", action="store_true",
|
| 234 |
+
help="Disable audio-to-video cross attention at inference time.")
|
| 235 |
+
parser.add_argument("--disable_v2a_cross_attn", "--disable-v2a-cross-attn",
|
| 236 |
+
"--disable_v2a", "--disable-v2a", action="store_true",
|
| 237 |
+
help="Disable video-to-audio cross attention at inference time.")
|
| 238 |
+
|
| 239 |
+
args = parser.parse_args()
|
| 240 |
+
return args
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
# ============================================================================
|
| 244 |
+
# Device / dtype helpers
|
| 245 |
+
# ============================================================================
|
| 246 |
+
|
| 247 |
+
def init_device() -> torch.device:
|
| 248 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 249 |
+
if device.type == "cuda":
|
| 250 |
+
torch.cuda.set_device(0)
|
| 251 |
+
return device
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
def resolve_weight_dtype(dtype_arg: str, device: torch.device) -> torch.dtype:
|
| 255 |
+
if device.type == "cpu":
|
| 256 |
+
return torch.float32
|
| 257 |
+
if dtype_arg == "float16":
|
| 258 |
+
return torch.float16
|
| 259 |
+
if dtype_arg == "bfloat16":
|
| 260 |
+
return torch.bfloat16
|
| 261 |
+
return torch.float32
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
def resolve_cpu_offload_flags(args):
|
| 265 |
+
"""Resolve dedicated offload flags, preserving the legacy memory mode."""
|
| 266 |
+
legacy_offload = args.GPU_memory_mode == "model_cpu_offload"
|
| 267 |
+
|
| 268 |
+
def resolve(value):
|
| 269 |
+
return legacy_offload if value is None else bool(value)
|
| 270 |
+
|
| 271 |
+
vae_override = getattr(args, "vae_cpu_offload", None)
|
| 272 |
+
|
| 273 |
+
def resolve_vae(value):
|
| 274 |
+
return resolve(vae_override if value is None else value)
|
| 275 |
+
|
| 276 |
+
return {
|
| 277 |
+
"transformer": legacy_offload,
|
| 278 |
+
"text_encoder": resolve(getattr(args, "text_encoder_cpu_offload", None)),
|
| 279 |
+
"video_vae": resolve_vae(getattr(args, "video_vae_cpu_offload", None)),
|
| 280 |
+
"audio_vae": resolve_vae(getattr(args, "audio_vae_cpu_offload", None)),
|
| 281 |
+
}
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def clear_cuda_cache(device: torch.device):
|
| 285 |
+
if device.type == "cuda":
|
| 286 |
+
torch.cuda.empty_cache()
|
| 287 |
+
|
| 288 |
+
|
| 289 |
+
# ============================================================================
|
| 290 |
+
# Text encoding
|
| 291 |
+
# ============================================================================
|
| 292 |
+
|
| 293 |
+
def _get_t5_prompt_embeds(tokenizer, text_encoder, prompt, max_sequence_length, device, dtype):
|
| 294 |
+
prompt = [prompt] if isinstance(prompt, str) else prompt
|
| 295 |
+
text_inputs = tokenizer(
|
| 296 |
+
prompt,
|
| 297 |
+
padding="max_length",
|
| 298 |
+
max_length=max_sequence_length,
|
| 299 |
+
truncation=True,
|
| 300 |
+
add_special_tokens=True,
|
| 301 |
+
return_tensors="pt",
|
| 302 |
+
)
|
| 303 |
+
text_input_ids = text_inputs.input_ids
|
| 304 |
+
prompt_attention_mask = text_inputs.attention_mask
|
| 305 |
+
seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long()
|
| 306 |
+
|
| 307 |
+
prompt_embeds = text_encoder(
|
| 308 |
+
text_input_ids.to(device),
|
| 309 |
+
attention_mask=prompt_attention_mask.to(device),
|
| 310 |
+
)[0]
|
| 311 |
+
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
| 312 |
+
return [embed[:seq_len] for embed, seq_len in zip(prompt_embeds, seq_lens.tolist())]
|
| 313 |
+
|
| 314 |
+
|
| 315 |
+
def encode_prompt(tokenizer, text_encoder, prompt, negative_prompt, guidance_scale,
|
| 316 |
+
max_sequence_length, device, dtype):
|
| 317 |
+
prompt_embeds = _get_t5_prompt_embeds(
|
| 318 |
+
tokenizer, text_encoder, prompt, max_sequence_length, device, dtype)
|
| 319 |
+
|
| 320 |
+
if guidance_scale <= 1.0:
|
| 321 |
+
return prompt_embeds
|
| 322 |
+
|
| 323 |
+
negative_prompt = negative_prompt or ""
|
| 324 |
+
negative_prompt_embeds = _get_t5_prompt_embeds(
|
| 325 |
+
tokenizer, text_encoder, negative_prompt, max_sequence_length, device, dtype)
|
| 326 |
+
return negative_prompt_embeds + prompt_embeds
|
| 327 |
+
|
| 328 |
+
|
| 329 |
+
# ============================================================================
|
| 330 |
+
# Latent shape helpers
|
| 331 |
+
# ============================================================================
|
| 332 |
+
|
| 333 |
+
def compute_audio_latent_length(duration: float, audio_vae: CreatorDACVAE) -> int:
|
| 334 |
+
"""Compute audio latent time length from duration.
|
| 335 |
+
|
| 336 |
+
For DAC VAE: latent_T = ceil(duration * sample_rate / hop_length)
|
| 337 |
+
"""
|
| 338 |
+
sample_rate = audio_vae.sample_rate
|
| 339 |
+
hop_length = audio_vae.hop_length
|
| 340 |
+
num_samples = int(duration * sample_rate)
|
| 341 |
+
latent_T = math.ceil(num_samples / hop_length)
|
| 342 |
+
return latent_T
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
def get_audio_num_tokens(latent_T: int, patch_size: tuple) -> int:
|
| 346 |
+
"""Compute number of audio tokens after patching."""
|
| 347 |
+
return latent_T // patch_size[0]
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
def compute_video_latent_shape(duration, fps, vae_temporal_ratio=4, vae_spatial_ratio=8,
|
| 351 |
+
height=480, width=832):
|
| 352 |
+
"""Compute video latent dimensions from duration and resolution."""
|
| 353 |
+
num_frames = int(duration * fps)
|
| 354 |
+
num_frames = int((num_frames - 1) // vae_temporal_ratio * vae_temporal_ratio) + 1
|
| 355 |
+
latent_frames = (num_frames - 1) // vae_temporal_ratio + 1
|
| 356 |
+
latent_height = height // vae_spatial_ratio
|
| 357 |
+
latent_width = width // vae_spatial_ratio
|
| 358 |
+
return num_frames, latent_frames, latent_height, latent_width
|
| 359 |
+
|
| 360 |
+
|
| 361 |
+
# ============================================================================
|
| 362 |
+
# Audio saving
|
| 363 |
+
# ============================================================================
|
| 364 |
+
|
| 365 |
+
def save_audio_wav(waveform: torch.Tensor, sample_rate: int, output_path: str):
|
| 366 |
+
"""Save waveform tensor to wav file.
|
| 367 |
+
|
| 368 |
+
Args:
|
| 369 |
+
waveform: [B, 1, T] or [1, T] or [T] tensor
|
| 370 |
+
sample_rate: audio sample rate
|
| 371 |
+
output_path: path to save wav file
|
| 372 |
+
"""
|
| 373 |
+
waveform = waveform.cpu().float()
|
| 374 |
+
if waveform.ndim == 3:
|
| 375 |
+
waveform = waveform.squeeze(0) # [1, T]
|
| 376 |
+
if waveform.ndim == 1:
|
| 377 |
+
waveform = waveform.unsqueeze(0) # [1, T]
|
| 378 |
+
|
| 379 |
+
if sf is not None:
|
| 380 |
+
sf.write(output_path, waveform.transpose(0, 1).numpy(), sample_rate)
|
| 381 |
+
return
|
| 382 |
+
|
| 383 |
+
try:
|
| 384 |
+
import torchaudio
|
| 385 |
+
torchaudio.save(output_path, waveform, sample_rate)
|
| 386 |
+
return
|
| 387 |
+
except ImportError:
|
| 388 |
+
pass
|
| 389 |
+
|
| 390 |
+
import wave
|
| 391 |
+
waveform = waveform.clamp(-1.0, 1.0)
|
| 392 |
+
waveform_int16 = (waveform * 32767.0).to(torch.int16).transpose(0, 1).contiguous().numpy()
|
| 393 |
+
with wave.open(output_path, "wb") as wav_file:
|
| 394 |
+
wav_file.setnchannels(int(waveform_int16.shape[1]))
|
| 395 |
+
wav_file.setsampwidth(2)
|
| 396 |
+
wav_file.setframerate(sample_rate)
|
| 397 |
+
wav_file.writeframes(waveform_int16.tobytes())
|
| 398 |
+
|
| 399 |
+
|
| 400 |
+
# ============================================================================
|
| 401 |
+
# Scheduler helpers
|
| 402 |
+
# ============================================================================
|
| 403 |
+
|
| 404 |
+
def retrieve_timesteps(scheduler, num_inference_steps=None, device=None,
|
| 405 |
+
timesteps=None, sigmas=None, **kwargs):
|
| 406 |
+
if timesteps is not None and sigmas is not None:
|
| 407 |
+
raise ValueError("Only one of `timesteps` or `sigmas` can be passed.")
|
| 408 |
+
if timesteps is not None:
|
| 409 |
+
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
| 410 |
+
return scheduler.timesteps, len(scheduler.timesteps)
|
| 411 |
+
elif sigmas is not None:
|
| 412 |
+
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
| 413 |
+
return scheduler.timesteps, len(scheduler.timesteps)
|
| 414 |
+
else:
|
| 415 |
+
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
| 416 |
+
return scheduler.timesteps, num_inference_steps
|
| 417 |
+
|
| 418 |
+
|
| 419 |
+
def prepare_extra_step_kwargs(scheduler, eta=0.0):
|
| 420 |
+
extra_step_kwargs = {}
|
| 421 |
+
if "eta" in set(inspect.signature(scheduler.step).parameters.keys()):
|
| 422 |
+
extra_step_kwargs["eta"] = eta
|
| 423 |
+
return extra_step_kwargs
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
def resolve_scheduler_shifts(args):
|
| 427 |
+
"""Resolve separate AV shifts while supporting callers that only define --shift."""
|
| 428 |
+
video_shift = getattr(args, "video_shift", 5.0)
|
| 429 |
+
audio_shift = getattr(args, "audio_shift", 5.0)
|
| 430 |
+
return float(video_shift), float(audio_shift)
|
| 431 |
+
|
| 432 |
+
|
| 433 |
+
# ============================================================================
|
| 434 |
+
# Image to latent encoding
|
| 435 |
+
# ============================================================================
|
| 436 |
+
|
| 437 |
+
def encode_first_frame(image: Image.Image, vae, device, dtype):
|
| 438 |
+
"""Encode first frame image into VAE latent space.
|
| 439 |
+
|
| 440 |
+
Returns latent of shape [1, C, 1, H//8, W//8].
|
| 441 |
+
"""
|
| 442 |
+
image = image.convert("RGB")
|
| 443 |
+
img_tensor = torch.from_numpy(np.array(image)).permute(2, 0, 1).float() / 255.0
|
| 444 |
+
img_tensor = img_tensor * 2.0 - 1.0 # normalize to [-1, 1]
|
| 445 |
+
img_tensor = img_tensor.unsqueeze(0).unsqueeze(2) # [1, 3, 1, H, W]
|
| 446 |
+
img_tensor = img_tensor.to(device=device, dtype=dtype)
|
| 447 |
+
vae = vae.to(dtype=dtype, device=device)
|
| 448 |
+
with torch.no_grad():
|
| 449 |
+
posterior = vae.encode(img_tensor)[0]
|
| 450 |
+
latent = posterior.sample()
|
| 451 |
+
return latent # [1, C, 1, H//8, W//8]
|
| 452 |
+
|
| 453 |
+
|
| 454 |
+
def _budget_spatial_size(
|
| 455 |
+
target_spatial_tokens: int,
|
| 456 |
+
source_height: int,
|
| 457 |
+
source_width: int,
|
| 458 |
+
spatial_divisor_h: int,
|
| 459 |
+
spatial_divisor_w: int,
|
| 460 |
+
min_token_ratio: float = 0.95,
|
| 461 |
+
):
|
| 462 |
+
"""Pick a near-aspect-ratio size within the spatial-token budget."""
|
| 463 |
+
if target_spatial_tokens <= 0:
|
| 464 |
+
raise ValueError("target_spatial_tokens must be positive")
|
| 465 |
+
if source_height <= 0 or source_width <= 0:
|
| 466 |
+
raise ValueError(f"Invalid source size: {(source_height, source_width)}")
|
| 467 |
+
if not 0.0 < min_token_ratio <= 1.0:
|
| 468 |
+
raise ValueError("min_token_ratio must be in (0, 1]")
|
| 469 |
+
|
| 470 |
+
min_spatial_tokens = max(1, int(math.ceil(target_spatial_tokens * min_token_ratio)))
|
| 471 |
+
source_ratio = float(source_height) / float(source_width)
|
| 472 |
+
best = None
|
| 473 |
+
for token_height in range(1, target_spatial_tokens + 1):
|
| 474 |
+
max_token_width = target_spatial_tokens // token_height
|
| 475 |
+
if max_token_width < 1:
|
| 476 |
+
continue
|
| 477 |
+
ideal_token_width = (
|
| 478 |
+
source_width * token_height * spatial_divisor_h
|
| 479 |
+
/ (source_height * spatial_divisor_w)
|
| 480 |
+
)
|
| 481 |
+
for token_width in {
|
| 482 |
+
1,
|
| 483 |
+
max_token_width,
|
| 484 |
+
int(math.floor(ideal_token_width)),
|
| 485 |
+
int(math.ceil(ideal_token_width)),
|
| 486 |
+
}:
|
| 487 |
+
if not 1 <= token_width <= max_token_width:
|
| 488 |
+
continue
|
| 489 |
+
spatial_tokens = token_height * token_width
|
| 490 |
+
height = token_height * spatial_divisor_h
|
| 491 |
+
width = token_width * spatial_divisor_w
|
| 492 |
+
aspect_error = abs(math.log((height / width) / source_ratio))
|
| 493 |
+
score = (
|
| 494 |
+
max(0, min_spatial_tokens - spatial_tokens),
|
| 495 |
+
aspect_error,
|
| 496 |
+
target_spatial_tokens - spatial_tokens,
|
| 497 |
+
height,
|
| 498 |
+
width,
|
| 499 |
+
)
|
| 500 |
+
if best is None or score < best[0]:
|
| 501 |
+
best = (score, height, width, spatial_tokens)
|
| 502 |
+
|
| 503 |
+
if best is None:
|
| 504 |
+
raise ValueError(f"Unable to resolve target_spatial_tokens={target_spatial_tokens}")
|
| 505 |
+
return best[1], best[2], best[3]
|
| 506 |
+
|
| 507 |
+
|
| 508 |
+
def compute_dynamic_resolution(
|
| 509 |
+
image: Image.Image,
|
| 510 |
+
target_spatial_tokens: int = 880,
|
| 511 |
+
min_token_ratio: float = 0.95,
|
| 512 |
+
spatial_divisor_h: int = 32,
|
| 513 |
+
spatial_divisor_w: int = 32,
|
| 514 |
+
):
|
| 515 |
+
"""Preserve input aspect ratio while using 0.95-1.0 of the token budget."""
|
| 516 |
+
source_width, source_height = image.size
|
| 517 |
+
height, width, actual_tokens = _budget_spatial_size(
|
| 518 |
+
int(target_spatial_tokens),
|
| 519 |
+
source_height,
|
| 520 |
+
source_width,
|
| 521 |
+
spatial_divisor_h,
|
| 522 |
+
spatial_divisor_w,
|
| 523 |
+
min_token_ratio,
|
| 524 |
+
)
|
| 525 |
+
return height, width, int(target_spatial_tokens), actual_tokens
|
| 526 |
+
|
| 527 |
+
|
| 528 |
+
# ============================================================================
|
| 529 |
+
# Model setup
|
| 530 |
+
# ============================================================================
|
| 531 |
+
|
| 532 |
+
def setup_models(args, device, weight_dtype):
|
| 533 |
+
config = OmegaConf.load(args.config_path)
|
| 534 |
+
|
| 535 |
+
# Video transformer kwargs
|
| 536 |
+
video_transformer_kwargs = OmegaConf.to_container(
|
| 537 |
+
config.get("video_transformer_additional_kwargs",
|
| 538 |
+
config.get("transformer_additional_kwargs", {})),
|
| 539 |
+
resolve=True,
|
| 540 |
+
)
|
| 541 |
+
# Audio transformer kwargs
|
| 542 |
+
audio_transformer_kwargs = OmegaConf.to_container(
|
| 543 |
+
config.get("audio_transformer_additional_kwargs",
|
| 544 |
+
config.get("transformer_additional_kwargs", {})),
|
| 545 |
+
resolve=True,
|
| 546 |
+
)
|
| 547 |
+
|
| 548 |
+
# Resolve transformer paths
|
| 549 |
+
video_path = args.transformer_path
|
| 550 |
+
audio_path = args.transformer_path
|
| 551 |
+
|
| 552 |
+
# Check if joint checkpoint (has video_model/ and audio_model/ subdirs)
|
| 553 |
+
if args.transformer_path is not None:
|
| 554 |
+
video_sub = os.path.join(args.transformer_path, "video_model")
|
| 555 |
+
audio_sub = os.path.join(args.transformer_path, "audio_model")
|
| 556 |
+
if os.path.isdir(video_sub) and os.path.isdir(audio_sub):
|
| 557 |
+
video_path = video_sub
|
| 558 |
+
audio_path = audio_sub
|
| 559 |
+
|
| 560 |
+
# Creator gating kwargs
|
| 561 |
+
creator_gating_kwargs = OmegaConf.to_container(
|
| 562 |
+
config.get("creator_gating_kwargs", {}),
|
| 563 |
+
resolve=True,
|
| 564 |
+
)
|
| 565 |
+
|
| 566 |
+
print_info(f"Loading Creator gating AV transformer: video={video_path}, audio={audio_path}")
|
| 567 |
+
transformer = WanCreatorGatingAVModel.from_pretrained(
|
| 568 |
+
pretrained_model_path=args.transformer_path,
|
| 569 |
+
video_pretrained_model_path=video_path,
|
| 570 |
+
audio_pretrained_model_path=audio_path,
|
| 571 |
+
video_subfolder=video_transformer_kwargs.get("transformer_low_noise_model_subpath", None),
|
| 572 |
+
audio_subfolder=audio_transformer_kwargs.get("transformer_low_noise_model_subpath", None),
|
| 573 |
+
video_kwargs=video_transformer_kwargs,
|
| 574 |
+
audio_kwargs=audio_transformer_kwargs,
|
| 575 |
+
low_cpu_mem_usage=True,
|
| 576 |
+
torch_dtype=weight_dtype,
|
| 577 |
+
# Pass gating-specific kwargs from config
|
| 578 |
+
use_temporal_rope=creator_gating_kwargs.get("use_temporal_rope", args.use_temporal_rope),
|
| 579 |
+
audio_fps=creator_gating_kwargs.get("audio_fps", args.audio_fps),
|
| 580 |
+
vae_temporal_stride=creator_gating_kwargs.get("vae_temporal_stride", args.vae_temporal_stride),
|
| 581 |
+
a2v_cross_attn_layers=creator_gating_kwargs.get("a2v_cross_attn_layers", None),
|
| 582 |
+
v2a_cross_attn_layers=creator_gating_kwargs.get("v2a_cross_attn_layers", None),
|
| 583 |
+
use_gating=creator_gating_kwargs.get("use_gating", True),
|
| 584 |
+
zero_init_cross_attn=creator_gating_kwargs.get("zero_init_cross_attn", False),
|
| 585 |
+
zero_init_gating=creator_gating_kwargs.get("zero_init_gating", True),
|
| 586 |
+
gate_init_value=creator_gating_kwargs.get("gate_init_value", 0.0),
|
| 587 |
+
a2v_gate_alphas=creator_gating_kwargs.get("a2v_gate_alphas", None),
|
| 588 |
+
v2a_gate_alphas=creator_gating_kwargs.get("v2a_gate_alphas", None),
|
| 589 |
+
)
|
| 590 |
+
|
| 591 |
+
# Load video VAE
|
| 592 |
+
video_vae_path = os.path.join(
|
| 593 |
+
args.model_name, config.get("video_vae_kwargs", {}).get("vae_subpath", "vae"))
|
| 594 |
+
print_info(f"Loading video VAE from: {video_vae_path}")
|
| 595 |
+
vae_kwargs = OmegaConf.to_container(config.get("video_vae_kwargs", {}), resolve=True)
|
| 596 |
+
video_vae = AutoencoderKLWan3_8.from_pretrained(video_vae_path, additional_kwargs=vae_kwargs)
|
| 597 |
+
|
| 598 |
+
# Load audio VAE (CreatorDACVAE)
|
| 599 |
+
print_info(f"Loading audio VAE (CreatorDACVAE) from: {args.audio_vae_path}")
|
| 600 |
+
audio_vae = CreatorDACVAE.from_pretrained(args.audio_vae_path, strict=False)
|
| 601 |
+
|
| 602 |
+
# Load tokenizer and text encoder
|
| 603 |
+
text_encoder_kwargs = OmegaConf.to_container(config.get("text_encoder_kwargs", {}), resolve=True)
|
| 604 |
+
tokenizer_path = os.path.join(args.model_name, text_encoder_kwargs.get("tokenizer_subpath", "tokenizer"))
|
| 605 |
+
text_encoder_path = os.path.join(
|
| 606 |
+
args.model_name, text_encoder_kwargs.get("text_encoder_subpath", "text_encoder"))
|
| 607 |
+
|
| 608 |
+
print_info(f"Loading tokenizer from: {tokenizer_path}")
|
| 609 |
+
tokenizer = HFAutoTokenizer.from_pretrained(tokenizer_path)
|
| 610 |
+
|
| 611 |
+
print_info(f"Loading text encoder from: {text_encoder_path}")
|
| 612 |
+
text_encoder = WanT5EncoderModel.from_pretrained(
|
| 613 |
+
text_encoder_path,
|
| 614 |
+
additional_kwargs=text_encoder_kwargs,
|
| 615 |
+
low_cpu_mem_usage=True,
|
| 616 |
+
torch_dtype=weight_dtype,
|
| 617 |
+
)
|
| 618 |
+
|
| 619 |
+
# Setup schedulers
|
| 620 |
+
scheduler_dict = {"Flow": FlowMatchEulerDiscreteScheduler}
|
| 621 |
+
video_shift, audio_shift = resolve_scheduler_shifts(args)
|
| 622 |
+
scheduler_kwargs = OmegaConf.to_container(config.get("scheduler_kwargs", {}), resolve=True)
|
| 623 |
+
video_scheduler_kwargs = dict(scheduler_kwargs)
|
| 624 |
+
audio_scheduler_kwargs = dict(scheduler_kwargs)
|
| 625 |
+
video_scheduler_kwargs["shift"] = video_shift
|
| 626 |
+
audio_scheduler_kwargs["shift"] = audio_shift
|
| 627 |
+
Chosen_Scheduler = scheduler_dict[args.sampler_name]
|
| 628 |
+
video_scheduler = Chosen_Scheduler(**filter_kwargs(Chosen_Scheduler, video_scheduler_kwargs))
|
| 629 |
+
audio_scheduler = Chosen_Scheduler(**filter_kwargs(Chosen_Scheduler, audio_scheduler_kwargs))
|
| 630 |
+
|
| 631 |
+
# Set models to eval
|
| 632 |
+
transformer.eval()
|
| 633 |
+
if args.cfg_mode == "multimodal":
|
| 634 |
+
transformer = DirectionalMultimodalCFGAdapter(
|
| 635 |
+
transformer,
|
| 636 |
+
video_scale=args.video_bridge_guidance_scale,
|
| 637 |
+
audio_scale=args.audio_bridge_guidance_scale,
|
| 638 |
+
enable_a2v=not args.disable_a2v_cross_attn,
|
| 639 |
+
enable_v2a=not args.disable_v2a_cross_attn,
|
| 640 |
+
).eval()
|
| 641 |
+
text_encoder.eval()
|
| 642 |
+
video_vae.eval()
|
| 643 |
+
audio_vae.eval()
|
| 644 |
+
|
| 645 |
+
offload_flags = resolve_cpu_offload_flags(args)
|
| 646 |
+
|
| 647 |
+
# Keep explicitly offloaded modules on CPU until their short GPU phase.
|
| 648 |
+
if not offload_flags["transformer"]:
|
| 649 |
+
transformer.to(device)
|
| 650 |
+
if not offload_flags["text_encoder"]:
|
| 651 |
+
text_encoder.to(device)
|
| 652 |
+
if not offload_flags["video_vae"]:
|
| 653 |
+
video_vae.to(device)
|
| 654 |
+
if not offload_flags["audio_vae"]:
|
| 655 |
+
audio_vae.to(device)
|
| 656 |
+
|
| 657 |
+
return {
|
| 658 |
+
"config": config,
|
| 659 |
+
"transformer": transformer,
|
| 660 |
+
"video_vae": video_vae,
|
| 661 |
+
"audio_vae": audio_vae,
|
| 662 |
+
"tokenizer": tokenizer,
|
| 663 |
+
"text_encoder": text_encoder,
|
| 664 |
+
"video_scheduler": video_scheduler,
|
| 665 |
+
"audio_scheduler": audio_scheduler,
|
| 666 |
+
"max_sequence_length": int(text_encoder_kwargs.get("text_length", 512)),
|
| 667 |
+
}
|
| 668 |
+
|
| 669 |
+
|
| 670 |
+
# ============================================================================
|
| 671 |
+
# Joint denoising loop
|
| 672 |
+
# ============================================================================
|
| 673 |
+
|
| 674 |
+
@torch.no_grad()
|
| 675 |
+
def generate_joint_audio_video(args, models, device, weight_dtype, item):
|
| 676 |
+
"""Generate audio and video jointly for a single item."""
|
| 677 |
+
transformer = models["transformer"]
|
| 678 |
+
video_vae = models["video_vae"]
|
| 679 |
+
audio_vae = models["audio_vae"]
|
| 680 |
+
tokenizer = models["tokenizer"]
|
| 681 |
+
text_encoder = models["text_encoder"]
|
| 682 |
+
video_scheduler = models["video_scheduler"]
|
| 683 |
+
audio_scheduler = models["audio_scheduler"]
|
| 684 |
+
max_sequence_length = models["max_sequence_length"]
|
| 685 |
+
offload_flags = resolve_cpu_offload_flags(args)
|
| 686 |
+
|
| 687 |
+
prompt = item["prompt"]
|
| 688 |
+
video_prompt = item.get("video_prompt", prompt)
|
| 689 |
+
audio_prompt = item.get("audio_prompt", prompt)
|
| 690 |
+
negative_prompt = item.get("negative_prompt", args.negative_prompt)
|
| 691 |
+
audio_negative_prompt = item.get("audio_negative_prompt", args.negative_prompt)
|
| 692 |
+
duration = float(item.get("duration", args.duration))
|
| 693 |
+
guidance_scale = float(item.get("guidance_scale", args.guidance_scale))
|
| 694 |
+
num_inference_steps = int(item.get("num_inference_steps", args.num_inference_steps))
|
| 695 |
+
seed = int(item.get("seed", args.seed))
|
| 696 |
+
# ---- Step 1: Load the supplied first frame ----
|
| 697 |
+
first_frame_path = item.get("first_frame_path", args.image)
|
| 698 |
+
if not os.path.isfile(first_frame_path):
|
| 699 |
+
raise FileNotFoundError(f"Input image not found: {first_frame_path}")
|
| 700 |
+
print_info(f"Loading first frame from: {first_frame_path}")
|
| 701 |
+
with Image.open(first_frame_path) as source_image:
|
| 702 |
+
source_image = source_image.convert("RGB")
|
| 703 |
+
video_patch_size = tuple(int(v) for v in transformer.video_patch_size)
|
| 704 |
+
spatial_compression = int(getattr(video_vae.config, "spatial_compression_ratio", 16))
|
| 705 |
+
height, width, target_tokens, actual_tokens = compute_dynamic_resolution(
|
| 706 |
+
source_image,
|
| 707 |
+
target_spatial_tokens=args.target_spatial_tokens,
|
| 708 |
+
min_token_ratio=args.min_token_ratio,
|
| 709 |
+
spatial_divisor_h=spatial_compression * video_patch_size[1],
|
| 710 |
+
spatial_divisor_w=spatial_compression * video_patch_size[2],
|
| 711 |
+
)
|
| 712 |
+
print_info(
|
| 713 |
+
f"Dynamic resolution: {height}x{width}, "
|
| 714 |
+
f"spatial_tokens={actual_tokens}/{target_tokens}"
|
| 715 |
+
)
|
| 716 |
+
first_frame_image = source_image.resize((width, height), Image.Resampling.BICUBIC)
|
| 717 |
+
|
| 718 |
+
do_classifier_free_guidance = guidance_scale > 1.0
|
| 719 |
+
|
| 720 |
+
# Save the resized frame once in distributed inference.
|
| 721 |
+
if not getattr(args, "suppress_aux_writes", False):
|
| 722 |
+
first_frame_save_path = os.path.splitext(args.output)[0] + "_first_frame.png"
|
| 723 |
+
first_frame_image.save(first_frame_save_path)
|
| 724 |
+
print_info(f"First frame saved to: {first_frame_save_path}")
|
| 725 |
+
|
| 726 |
+
# ---- Step 2: Encode first frame to video latent ----
|
| 727 |
+
if offload_flags["video_vae"]:
|
| 728 |
+
video_vae.to(device)
|
| 729 |
+
|
| 730 |
+
first_frame_latent = encode_first_frame(first_frame_image, video_vae, device, weight_dtype)
|
| 731 |
+
first_frame_latent = synchronize_tensor(first_frame_latent, args)
|
| 732 |
+
latent_channels = first_frame_latent.shape[1]
|
| 733 |
+
latent_height = first_frame_latent.shape[3]
|
| 734 |
+
latent_width = first_frame_latent.shape[4]
|
| 735 |
+
|
| 736 |
+
if offload_flags["video_vae"]:
|
| 737 |
+
video_vae.to("cpu")
|
| 738 |
+
clear_cuda_cache(device)
|
| 739 |
+
|
| 740 |
+
# ---- Step 3: Compute latent shapes ----
|
| 741 |
+
vae_temporal_ratio = int(getattr(video_vae.config, "temporal_compression_ratio", 4))
|
| 742 |
+
num_frames = int(duration * args.fps)
|
| 743 |
+
num_frames = int((num_frames - 1) // vae_temporal_ratio * vae_temporal_ratio) + 1
|
| 744 |
+
latent_frames = (num_frames - 1) // vae_temporal_ratio + 1
|
| 745 |
+
|
| 746 |
+
# Video sequence length
|
| 747 |
+
video_patch_size = tuple(int(v) for v in transformer.video_patch_size)
|
| 748 |
+
video_seq_len = (latent_frames * latent_height * latent_width) // math.prod(video_patch_size)
|
| 749 |
+
|
| 750 |
+
# Audio latent shape (DAC VAE: latent is [D, T])
|
| 751 |
+
audio_duration = num_frames / args.fps
|
| 752 |
+
audio_latent_T = compute_audio_latent_length(audio_duration, audio_vae)
|
| 753 |
+
audio_latent_dim = audio_vae.latent_dim
|
| 754 |
+
audio_patch_size = tuple(int(v) for v in transformer.audio_patch_size)
|
| 755 |
+
# Align latent_T to patch_size boundary
|
| 756 |
+
if audio_latent_T % audio_patch_size[0] != 0:
|
| 757 |
+
audio_latent_T = math.ceil(audio_latent_T / audio_patch_size[0]) * audio_patch_size[0]
|
| 758 |
+
audio_seq_len = get_audio_num_tokens(audio_latent_T, audio_patch_size)
|
| 759 |
+
|
| 760 |
+
print_info(f"Video: {num_frames} frames, latent [{latent_frames}, {latent_height}, {latent_width}], "
|
| 761 |
+
f"seq_len={video_seq_len}")
|
| 762 |
+
print_info(f"Audio: duration={audio_duration}s, latent [{audio_latent_dim}, {audio_latent_T}], "
|
| 763 |
+
f"seq_len={audio_seq_len}")
|
| 764 |
+
|
| 765 |
+
# ---- Step 4: Encode text prompts ----
|
| 766 |
+
if offload_flags["text_encoder"]:
|
| 767 |
+
text_encoder.to(device)
|
| 768 |
+
|
| 769 |
+
video_prompt_embeds = encode_prompt(
|
| 770 |
+
tokenizer, text_encoder, video_prompt, negative_prompt,
|
| 771 |
+
guidance_scale, max_sequence_length, device, weight_dtype)
|
| 772 |
+
|
| 773 |
+
if audio_prompt != video_prompt or audio_negative_prompt != negative_prompt:
|
| 774 |
+
audio_prompt_embeds = encode_prompt(
|
| 775 |
+
tokenizer, text_encoder, audio_prompt, audio_negative_prompt,
|
| 776 |
+
guidance_scale, max_sequence_length, device, weight_dtype)
|
| 777 |
+
else:
|
| 778 |
+
audio_prompt_embeds = video_prompt_embeds
|
| 779 |
+
|
| 780 |
+
if offload_flags["text_encoder"]:
|
| 781 |
+
text_encoder.to("cpu")
|
| 782 |
+
clear_cuda_cache(device)
|
| 783 |
+
|
| 784 |
+
# ---- Step 5: Initialize noise latents ----
|
| 785 |
+
generator = torch.Generator(device=device).manual_seed(seed)
|
| 786 |
+
|
| 787 |
+
# Video noise: [1, C, latent_frames, H_lat, W_lat]
|
| 788 |
+
video_latents = torch.randn(
|
| 789 |
+
(1, latent_channels, latent_frames, latent_height, latent_width),
|
| 790 |
+
generator=generator, device=device, dtype=weight_dtype,
|
| 791 |
+
)
|
| 792 |
+
# Apply first-frame conditioning
|
| 793 |
+
video_latents[:, :, 0:1, :, :] = first_frame_latent.to(weight_dtype)
|
| 794 |
+
|
| 795 |
+
# Audio noise: [1, D, T] (DAC latent format)
|
| 796 |
+
audio_latents = torch.randn(
|
| 797 |
+
(1, audio_latent_dim, audio_latent_T),
|
| 798 |
+
generator=generator, device=device, dtype=weight_dtype,
|
| 799 |
+
)
|
| 800 |
+
video_latents = synchronize_tensor(video_latents, args)
|
| 801 |
+
audio_latents = synchronize_tensor(audio_latents, args)
|
| 802 |
+
|
| 803 |
+
# ---- Step 6: Setup schedulers ----
|
| 804 |
+
video_shift, audio_shift = resolve_scheduler_shifts(args)
|
| 805 |
+
timestep_kwargs = {}
|
| 806 |
+
if getattr(args, "flow_match_mu", None) is not None:
|
| 807 |
+
timestep_kwargs["mu"] = float(args.flow_match_mu)
|
| 808 |
+
video_timesteps, num_inference_steps = retrieve_timesteps(
|
| 809 |
+
video_scheduler, num_inference_steps, device, **timestep_kwargs)
|
| 810 |
+
audio_timesteps, _ = retrieve_timesteps(
|
| 811 |
+
audio_scheduler, num_inference_steps, device, **timestep_kwargs)
|
| 812 |
+
|
| 813 |
+
timesteps = video_timesteps
|
| 814 |
+
if len(audio_timesteps) != len(video_timesteps):
|
| 815 |
+
raise ValueError(
|
| 816 |
+
f"Video/audio timestep count mismatch: {len(video_timesteps)} vs {len(audio_timesteps)}"
|
| 817 |
+
)
|
| 818 |
+
if hasattr(video_scheduler, "init_noise_sigma"):
|
| 819 |
+
video_latents[:, :, 1:, :, :] = video_latents[:, :, 1:, :, :] * video_scheduler.init_noise_sigma
|
| 820 |
+
if hasattr(audio_scheduler, "init_noise_sigma"):
|
| 821 |
+
audio_latents = audio_latents * audio_scheduler.init_noise_sigma
|
| 822 |
+
|
| 823 |
+
video_extra_step_kwargs = prepare_extra_step_kwargs(video_scheduler)
|
| 824 |
+
audio_extra_step_kwargs = prepare_extra_step_kwargs(audio_scheduler)
|
| 825 |
+
|
| 826 |
+
# ---- Step 7: Per-timestep video token count for first frame ----
|
| 827 |
+
patch_t, patch_h, patch_w = video_patch_size
|
| 828 |
+
first_frame_tokens = (latent_height * latent_width) // (patch_h * patch_w)
|
| 829 |
+
|
| 830 |
+
# ---- Step 8: Denoising loop ----
|
| 831 |
+
if offload_flags["transformer"]:
|
| 832 |
+
transformer.to(device)
|
| 833 |
+
|
| 834 |
+
if device.type == "cuda":
|
| 835 |
+
torch.cuda.synchronize()
|
| 836 |
+
start_time = time.time()
|
| 837 |
+
|
| 838 |
+
enable_a2v = not getattr(args, "disable_a2v_cross_attn", False)
|
| 839 |
+
enable_v2a = not getattr(args, "disable_v2a_cross_attn", False)
|
| 840 |
+
|
| 841 |
+
for step_idx, t in enumerate(tqdm(
|
| 842 |
+
timesteps,
|
| 843 |
+
desc="Denoising",
|
| 844 |
+
unit="step",
|
| 845 |
+
disable=bool(getattr(args, "disable_progress", False)),
|
| 846 |
+
)):
|
| 847 |
+
audio_t = audio_timesteps[step_idx]
|
| 848 |
+
video_noisy_input = video_latents.clone()
|
| 849 |
+
|
| 850 |
+
# Per-token timestep for video (first frame = 0, rest = t)
|
| 851 |
+
t_scalar = t.item() if t.dim() == 0 else t
|
| 852 |
+
video_timestep_tokens = torch.ones(1, video_seq_len, device=device) * t_scalar
|
| 853 |
+
video_timestep_tokens[:, :first_frame_tokens] = 0
|
| 854 |
+
|
| 855 |
+
audio_timestep = audio_t.unsqueeze(0) if audio_t.dim() == 0 else audio_t
|
| 856 |
+
# video_timestep_tokens = (video_timestep_tokens / 1000) ** 0.25 * 1000
|
| 857 |
+
# Build model inputs
|
| 858 |
+
if do_classifier_free_guidance:
|
| 859 |
+
video_x_input = [video_noisy_input[0], video_noisy_input[0]]
|
| 860 |
+
video_t_input = torch.cat([video_timestep_tokens, video_timestep_tokens], dim=0)
|
| 861 |
+
video_context_input = video_prompt_embeds
|
| 862 |
+
|
| 863 |
+
audio_x_input = [audio_latents[0], audio_latents[0]]
|
| 864 |
+
audio_t_input = audio_timestep.expand(2)
|
| 865 |
+
audio_context_input = audio_prompt_embeds
|
| 866 |
+
else:
|
| 867 |
+
video_x_input = [video_noisy_input[0]]
|
| 868 |
+
video_t_input = video_timestep_tokens
|
| 869 |
+
video_context_input = video_prompt_embeds
|
| 870 |
+
|
| 871 |
+
audio_x_input = [audio_latents[0]]
|
| 872 |
+
audio_t_input = audio_timestep
|
| 873 |
+
audio_context_input = audio_prompt_embeds
|
| 874 |
+
|
| 875 |
+
video_inputs = {
|
| 876 |
+
"x": video_x_input,
|
| 877 |
+
"t": video_t_input,
|
| 878 |
+
"context": video_context_input,
|
| 879 |
+
"y": None,
|
| 880 |
+
"seq_len": video_seq_len,
|
| 881 |
+
"video_fps": float(args.fps),
|
| 882 |
+
}
|
| 883 |
+
audio_inputs = {
|
| 884 |
+
"x": audio_x_input,
|
| 885 |
+
"t": audio_t_input,
|
| 886 |
+
"context": audio_context_input,
|
| 887 |
+
"y": None,
|
| 888 |
+
"seq_len": audio_seq_len,
|
| 889 |
+
}
|
| 890 |
+
|
| 891 |
+
with torch.autocast("cuda", dtype=weight_dtype):
|
| 892 |
+
model_output = transformer(
|
| 893 |
+
video=video_inputs,
|
| 894 |
+
audio=audio_inputs,
|
| 895 |
+
dtype=weight_dtype,
|
| 896 |
+
enable_a2v=enable_a2v,
|
| 897 |
+
enable_v2a=enable_v2a,
|
| 898 |
+
)
|
| 899 |
+
|
| 900 |
+
video_pred = model_output["video"]
|
| 901 |
+
audio_pred = model_output["audio"]
|
| 902 |
+
|
| 903 |
+
# Apply CFG
|
| 904 |
+
if do_classifier_free_guidance:
|
| 905 |
+
video_pred_uncond = video_pred[0:1]
|
| 906 |
+
video_pred_text = video_pred[1:2]
|
| 907 |
+
video_noise_pred = video_pred_uncond + guidance_scale * (video_pred_text - video_pred_uncond)
|
| 908 |
+
audio_pred_uncond = audio_pred[0:1]
|
| 909 |
+
audio_pred_text = audio_pred[1:2]
|
| 910 |
+
audio_noise_pred = audio_pred_uncond + guidance_scale * (audio_pred_text - audio_pred_uncond)
|
| 911 |
+
else:
|
| 912 |
+
video_noise_pred = video_pred
|
| 913 |
+
audio_noise_pred = audio_pred
|
| 914 |
+
|
| 915 |
+
# Scheduler step
|
| 916 |
+
video_latents_denoised = video_scheduler.step(
|
| 917 |
+
video_noise_pred, t, video_latents, **video_extra_step_kwargs, return_dict=False)[0]
|
| 918 |
+
video_latents_denoised[:, :, 0:1, :, :] = first_frame_latent.to(weight_dtype)
|
| 919 |
+
video_latents = video_latents_denoised
|
| 920 |
+
|
| 921 |
+
audio_latents = audio_scheduler.step(
|
| 922 |
+
audio_noise_pred, audio_t, audio_latents, **audio_extra_step_kwargs, return_dict=False)[0]
|
| 923 |
+
|
| 924 |
+
if device.type == "cuda":
|
| 925 |
+
torch.cuda.synchronize()
|
| 926 |
+
elapsed = time.time() - start_time
|
| 927 |
+
print_info(f"Denoising completed in {elapsed:.2f}s")
|
| 928 |
+
|
| 929 |
+
if offload_flags["transformer"]:
|
| 930 |
+
transformer.to("cpu")
|
| 931 |
+
clear_cuda_cache(device)
|
| 932 |
+
|
| 933 |
+
if getattr(args, "skip_output_decode", False):
|
| 934 |
+
return None, None, num_frames
|
| 935 |
+
|
| 936 |
+
# ---- Step 9: Decode video ----
|
| 937 |
+
if offload_flags["video_vae"]:
|
| 938 |
+
video_vae.to(device)
|
| 939 |
+
|
| 940 |
+
with torch.no_grad():
|
| 941 |
+
video_decoded = video_vae.decode(video_latents.to(video_vae.dtype))[0]
|
| 942 |
+
video_decoded = (video_decoded / 2.0 + 0.5).clamp(0, 1).float().cpu()
|
| 943 |
+
|
| 944 |
+
if offload_flags["video_vae"]:
|
| 945 |
+
video_vae.to("cpu")
|
| 946 |
+
clear_cuda_cache(device)
|
| 947 |
+
|
| 948 |
+
# ---- Step 10: Decode audio (CreatorDACVAE) ----
|
| 949 |
+
if offload_flags["audio_vae"]:
|
| 950 |
+
audio_vae.to(device)
|
| 951 |
+
|
| 952 |
+
with torch.no_grad():
|
| 953 |
+
# CreatorDACVAE.decode expects [B, D, T'] and returns [B, 1, T]
|
| 954 |
+
audio_decoded = audio_vae.decode(audio_latents.float().to(audio_vae.dac.device if hasattr(audio_vae, 'dac') else device))
|
| 955 |
+
|
| 956 |
+
if offload_flags["audio_vae"]:
|
| 957 |
+
audio_vae.to("cpu")
|
| 958 |
+
clear_cuda_cache(device)
|
| 959 |
+
|
| 960 |
+
return video_decoded, audio_decoded, num_frames
|
| 961 |
+
|
| 962 |
+
|
| 963 |
+
def main():
|
| 964 |
+
args = parse_args()
|
| 965 |
+
device = init_device()
|
| 966 |
+
weight_dtype = resolve_weight_dtype(args.weight_dtype, device)
|
| 967 |
+
|
| 968 |
+
output_path = os.path.abspath(args.output)
|
| 969 |
+
if not os.path.splitext(output_path)[1]:
|
| 970 |
+
output_path += ".mp4"
|
| 971 |
+
args.output = output_path
|
| 972 |
+
output_dir = os.path.dirname(output_path) or "."
|
| 973 |
+
os.makedirs(output_dir, exist_ok=True)
|
| 974 |
+
args.output_dir = output_dir
|
| 975 |
+
|
| 976 |
+
print_info("=" * 80)
|
| 977 |
+
print_info("Creator single-image audio-video inference")
|
| 978 |
+
print_info("=" * 80)
|
| 979 |
+
print_info(f"Device: {device} | dtype: {weight_dtype}")
|
| 980 |
+
offload_flags = resolve_cpu_offload_flags(args)
|
| 981 |
+
print_info(
|
| 982 |
+
"CPU offload: "
|
| 983 |
+
f"text_encoder={'on' if offload_flags['text_encoder'] else 'off'}, "
|
| 984 |
+
f"video_vae={'on' if offload_flags['video_vae'] else 'off'}, "
|
| 985 |
+
f"audio_vae={'on' if offload_flags['audio_vae'] else 'off'}"
|
| 986 |
+
)
|
| 987 |
+
print_info(
|
| 988 |
+
"Cross attention: "
|
| 989 |
+
f"A2V={'disabled' if args.disable_a2v_cross_attn else 'enabled'}, "
|
| 990 |
+
f"V2A={'disabled' if args.disable_v2a_cross_attn else 'enabled'}"
|
| 991 |
+
)
|
| 992 |
+
|
| 993 |
+
# Load models
|
| 994 |
+
models = setup_models(args, device, weight_dtype)
|
| 995 |
+
|
| 996 |
+
item = {
|
| 997 |
+
"prompt": args.prompt,
|
| 998 |
+
"video_prompt": args.prompt,
|
| 999 |
+
"audio_prompt": args.prompt,
|
| 1000 |
+
"negative_prompt": args.negative_prompt,
|
| 1001 |
+
"audio_negative_prompt": args.negative_prompt,
|
| 1002 |
+
"duration": args.duration,
|
| 1003 |
+
"guidance_scale": args.guidance_scale,
|
| 1004 |
+
"num_inference_steps": args.num_inference_steps,
|
| 1005 |
+
"seed": args.seed,
|
| 1006 |
+
"name": "sample",
|
| 1007 |
+
}
|
| 1008 |
+
video_decoded, audio_decoded, _ = generate_joint_audio_video(
|
| 1009 |
+
args, models, device, weight_dtype, item)
|
| 1010 |
+
video_path = os.path.splitext(output_path)[0] + ".video.mp4"
|
| 1011 |
+
audio_path = os.path.splitext(output_path)[0] + ".wav"
|
| 1012 |
+
frames = (video_decoded[0].permute(1, 2, 3, 0).clamp(0, 1).numpy() * 255).astype("uint8")
|
| 1013 |
+
write_video(video_path, torch.from_numpy(frames), fps=args.fps, video_codec="h264")
|
| 1014 |
+
save_audio_wav(audio_decoded, int(models["audio_vae"].sample_rate), audio_path)
|
| 1015 |
+
mux_result = subprocess.run(
|
| 1016 |
+
["ffmpeg", "-y", "-i", video_path, "-i", audio_path,
|
| 1017 |
+
"-c:v", "copy", "-c:a", "aac", "-shortest", output_path],
|
| 1018 |
+
check=False,
|
| 1019 |
+
capture_output=True,
|
| 1020 |
+
text=True,
|
| 1021 |
+
)
|
| 1022 |
+
if mux_result.returncode != 0:
|
| 1023 |
+
raise RuntimeError(
|
| 1024 |
+
f"ffmpeg failed while muxing {output_path}:\n{mux_result.stderr.strip()}"
|
| 1025 |
+
)
|
| 1026 |
+
print_info(f"Saved: {output_path}")
|
| 1027 |
+
|
| 1028 |
+
|
| 1029 |
+
if __name__ == "__main__":
|
| 1030 |
+
main()
|
examples/case1.jpg
ADDED
|
Git LFS Details
|
examples/case2.jpg
ADDED
|
Git LFS Details
|
examples/case3.jpg
ADDED
|
Git LFS Details
|
examples/case4.jpg
ADDED
|
Git LFS Details
|
examples/case5.jpg
ADDED
|
Git LFS Details
|
examples/case6.jpg
ADDED
|
Git LFS Details
|
requirements.txt
CHANGED
|
@@ -1,18 +1,18 @@
|
|
| 1 |
-
# For ZeroGPU, gradio/spaces/huggingface_hub/torch are
|
| 2 |
-
#
|
| 3 |
-
#
|
| 4 |
-
#
|
| 5 |
-
|
| 6 |
-
diffusers==0.40.0
|
| 7 |
transformers==5.16.1
|
| 8 |
accelerate>=1.0
|
| 9 |
-
gguf
|
| 10 |
safetensors
|
|
|
|
|
|
|
| 11 |
# torch 2.6.0+cu124 predates numpy 2.5 ABI; 2.1.x is the known-good pair.
|
| 12 |
numpy==2.1.3
|
| 13 |
-
|
|
|
|
| 14 |
pillow
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
httpx
|
|
|
|
| 1 |
+
# For ZeroGPU, gradio/spaces/huggingface_hub/torch are platform-managed; the
|
| 2 |
+
# Dockerfile installs torch explicitly from the cu124 index and then this file.
|
| 3 |
+
# DreamX-Creator (Wan2.2-TI2V-5B + gated A2V/V2A) engine pins mirror the upstream
|
| 4 |
+
# Space (hugging-apps/gd-ml-dreamx-creator) adapted to our known-good base.
|
| 5 |
+
diffusers==0.37.1
|
|
|
|
| 6 |
transformers==5.16.1
|
| 7 |
accelerate>=1.0
|
|
|
|
| 8 |
safetensors
|
| 9 |
+
sentencepiece
|
| 10 |
+
# .pth checkpoints are saved pre-torch-2.6; app.py wraps torch.load(weights_only=False).
|
| 11 |
# torch 2.6.0+cu124 predates numpy 2.5 ABI; 2.1.x is the known-good pair.
|
| 12 |
numpy==2.1.3
|
| 13 |
+
omegaconf
|
| 14 |
+
einops
|
| 15 |
pillow
|
| 16 |
+
soundfile
|
| 17 |
+
tqdm
|
| 18 |
+
protobuf
|
|
|
videox_fun/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Minimal runtime package for single-image Creator inference."""
|
videox_fun/dist/__init__.py
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sequence_parallel import all_gather_sequence, all_to_all, ulysses_attention
|
| 2 |
+
|
| 3 |
+
__all__ = ["all_gather_sequence", "all_to_all", "ulysses_attention"]
|
videox_fun/dist/sequence_parallel.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Small raw-process-group helpers for Ulysses-style sequence parallelism."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
import torch.distributed as dist
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def all_gather_sequence(tensor: torch.Tensor, group=None) -> torch.Tensor:
|
| 10 |
+
"""Gather equally-sized sequence chunks along dimension 1."""
|
| 11 |
+
if not dist.is_initialized() or dist.get_world_size(group) == 1:
|
| 12 |
+
return tensor
|
| 13 |
+
tensor = tensor.contiguous()
|
| 14 |
+
chunks = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))]
|
| 15 |
+
dist.all_gather(chunks, tensor, group=group)
|
| 16 |
+
return torch.cat(chunks, dim=1).contiguous()
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def all_to_all(tensor: torch.Tensor, scatter_dim: int, gather_dim: int, group=None) -> torch.Tensor:
|
| 20 |
+
"""Scatter one dimension and gather another, matching Wan2.2's Ulysses layout."""
|
| 21 |
+
if not dist.is_initialized() or dist.get_world_size(group) == 1:
|
| 22 |
+
return tensor
|
| 23 |
+
|
| 24 |
+
world_size = dist.get_world_size(group)
|
| 25 |
+
if tensor.shape[scatter_dim] % world_size != 0:
|
| 26 |
+
raise ValueError(
|
| 27 |
+
f"Dimension {scatter_dim} ({tensor.shape[scatter_dim]}) must be divisible by "
|
| 28 |
+
f"SP size {world_size}"
|
| 29 |
+
)
|
| 30 |
+
inputs = [part.contiguous() for part in tensor.chunk(world_size, dim=scatter_dim)]
|
| 31 |
+
outputs = [torch.empty_like(inputs[0]) for _ in range(world_size)]
|
| 32 |
+
dist.all_to_all(outputs, inputs, group=group)
|
| 33 |
+
return torch.cat(outputs, dim=gather_dim).contiguous()
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def ulysses_attention(q, k, v, attention_fn, *, k_lens=None, window_size=(-1, -1), group=None):
|
| 37 |
+
"""Run attention on local sequence chunks with Ulysses head/sequence exchange."""
|
| 38 |
+
q = all_to_all(q, scatter_dim=2, gather_dim=1, group=group)
|
| 39 |
+
k = all_to_all(k, scatter_dim=2, gather_dim=1, group=group)
|
| 40 |
+
v = all_to_all(v, scatter_dim=2, gather_dim=1, group=group)
|
| 41 |
+
output = attention_fn(q, k, v, k_lens=k_lens, window_size=window_size)
|
| 42 |
+
return all_to_all(output, scatter_dim=1, gather_dim=2, group=group)
|
videox_fun/models/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Model definitions required by the release inference entrypoint."""
|
| 2 |
+
|
| 3 |
+
from .creator_dac_vae import CreatorDACVAE
|
| 4 |
+
from .wan_text_encoder import WanT5EncoderModel
|
| 5 |
+
from .wan_vae3_8 import AutoencoderKLWan3_8
|
| 6 |
+
|
| 7 |
+
__all__ = ["AutoencoderKLWan3_8", "CreatorDACVAE", "WanT5EncoderModel"]
|
videox_fun/models/attention_utils.py
ADDED
|
@@ -0,0 +1,211 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import warnings
|
| 5 |
+
|
| 6 |
+
try:
|
| 7 |
+
import flash_attn_interface
|
| 8 |
+
FLASH_ATTN_3_AVAILABLE = True
|
| 9 |
+
except ModuleNotFoundError:
|
| 10 |
+
FLASH_ATTN_3_AVAILABLE = False
|
| 11 |
+
|
| 12 |
+
try:
|
| 13 |
+
import flash_attn
|
| 14 |
+
FLASH_ATTN_2_AVAILABLE = True
|
| 15 |
+
except ModuleNotFoundError:
|
| 16 |
+
FLASH_ATTN_2_AVAILABLE = False
|
| 17 |
+
|
| 18 |
+
try:
|
| 19 |
+
major, minor = torch.cuda.get_device_capability(0)
|
| 20 |
+
if f"{major}.{minor}" == "8.0":
|
| 21 |
+
from sageattention_sm80 import sageattn
|
| 22 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 23 |
+
elif f"{major}.{minor}" == "8.6":
|
| 24 |
+
from sageattention_sm86 import sageattn
|
| 25 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 26 |
+
elif f"{major}.{minor}" == "8.9":
|
| 27 |
+
from sageattention_sm89 import sageattn
|
| 28 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 29 |
+
elif f"{major}.{minor}" == "9.0":
|
| 30 |
+
from sageattention_sm90 import sageattn
|
| 31 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 32 |
+
elif major>9:
|
| 33 |
+
from sageattention_sm120 import sageattn
|
| 34 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 35 |
+
except:
|
| 36 |
+
try:
|
| 37 |
+
from sageattention import sageattn
|
| 38 |
+
SAGE_ATTENTION_AVAILABLE = True
|
| 39 |
+
except:
|
| 40 |
+
sageattn = None
|
| 41 |
+
SAGE_ATTENTION_AVAILABLE = False
|
| 42 |
+
|
| 43 |
+
def flash_attention(
|
| 44 |
+
q,
|
| 45 |
+
k,
|
| 46 |
+
v,
|
| 47 |
+
q_lens=None,
|
| 48 |
+
k_lens=None,
|
| 49 |
+
dropout_p=0.,
|
| 50 |
+
softmax_scale=None,
|
| 51 |
+
q_scale=None,
|
| 52 |
+
causal=False,
|
| 53 |
+
window_size=(-1, -1),
|
| 54 |
+
deterministic=False,
|
| 55 |
+
dtype=torch.bfloat16,
|
| 56 |
+
version=None,
|
| 57 |
+
):
|
| 58 |
+
"""
|
| 59 |
+
q: [B, Lq, Nq, C1].
|
| 60 |
+
k: [B, Lk, Nk, C1].
|
| 61 |
+
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
|
| 62 |
+
q_lens: [B].
|
| 63 |
+
k_lens: [B].
|
| 64 |
+
dropout_p: float. Dropout probability.
|
| 65 |
+
softmax_scale: float. The scaling of QK^T before applying softmax.
|
| 66 |
+
causal: bool. Whether to apply causal attention mask.
|
| 67 |
+
window_size: (left right). If not (-1, -1), apply sliding window local attention.
|
| 68 |
+
deterministic: bool. If True, slightly slower and uses more memory.
|
| 69 |
+
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
|
| 70 |
+
"""
|
| 71 |
+
half_dtypes = (torch.float16, torch.bfloat16)
|
| 72 |
+
assert dtype in half_dtypes
|
| 73 |
+
assert q.device.type == 'cuda' and q.size(-1) <= 256
|
| 74 |
+
|
| 75 |
+
# params
|
| 76 |
+
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
| 77 |
+
|
| 78 |
+
def half(x):
|
| 79 |
+
return x if x.dtype in half_dtypes else x.to(dtype)
|
| 80 |
+
|
| 81 |
+
# preprocess query
|
| 82 |
+
if q_lens is None:
|
| 83 |
+
q = half(q.flatten(0, 1))
|
| 84 |
+
q_lens = torch.tensor(
|
| 85 |
+
[lq] * b, dtype=torch.int32).to(
|
| 86 |
+
device=q.device, non_blocking=True)
|
| 87 |
+
else:
|
| 88 |
+
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
| 89 |
+
|
| 90 |
+
# preprocess key, value
|
| 91 |
+
if k_lens is None:
|
| 92 |
+
k = half(k.flatten(0, 1))
|
| 93 |
+
v = half(v.flatten(0, 1))
|
| 94 |
+
k_lens = torch.tensor(
|
| 95 |
+
[lk] * b, dtype=torch.int32).to(
|
| 96 |
+
device=k.device, non_blocking=True)
|
| 97 |
+
else:
|
| 98 |
+
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
| 99 |
+
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
| 100 |
+
|
| 101 |
+
q = q.to(v.dtype)
|
| 102 |
+
k = k.to(v.dtype)
|
| 103 |
+
|
| 104 |
+
if q_scale is not None:
|
| 105 |
+
q = q * q_scale
|
| 106 |
+
|
| 107 |
+
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
| 108 |
+
warnings.warn(
|
| 109 |
+
'Flash attention 3 is not available, use flash attention 2 instead.'
|
| 110 |
+
)
|
| 111 |
+
|
| 112 |
+
# apply attention
|
| 113 |
+
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
| 114 |
+
# Note: dropout_p, window_size are not supported in FA3 now.
|
| 115 |
+
x = flash_attn_interface.flash_attn_varlen_func(
|
| 116 |
+
q=q,
|
| 117 |
+
k=k,
|
| 118 |
+
v=v,
|
| 119 |
+
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
| 120 |
+
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
| 121 |
+
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
| 122 |
+
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
| 123 |
+
seqused_q=None,
|
| 124 |
+
seqused_k=None,
|
| 125 |
+
max_seqlen_q=lq,
|
| 126 |
+
max_seqlen_k=lk,
|
| 127 |
+
softmax_scale=softmax_scale,
|
| 128 |
+
causal=causal,
|
| 129 |
+
deterministic=deterministic)[0].unflatten(0, (b, lq))
|
| 130 |
+
else:
|
| 131 |
+
assert FLASH_ATTN_2_AVAILABLE
|
| 132 |
+
x = flash_attn.flash_attn_varlen_func(
|
| 133 |
+
q=q,
|
| 134 |
+
k=k,
|
| 135 |
+
v=v,
|
| 136 |
+
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
| 137 |
+
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
| 138 |
+
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
| 139 |
+
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
| 140 |
+
max_seqlen_q=lq,
|
| 141 |
+
max_seqlen_k=lk,
|
| 142 |
+
dropout_p=dropout_p,
|
| 143 |
+
softmax_scale=softmax_scale,
|
| 144 |
+
causal=causal,
|
| 145 |
+
window_size=window_size,
|
| 146 |
+
deterministic=deterministic).unflatten(0, (b, lq))
|
| 147 |
+
|
| 148 |
+
# output
|
| 149 |
+
return x.type(out_dtype)
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
def attention(
|
| 153 |
+
q,
|
| 154 |
+
k,
|
| 155 |
+
v,
|
| 156 |
+
q_lens=None,
|
| 157 |
+
k_lens=None,
|
| 158 |
+
dropout_p=0.,
|
| 159 |
+
softmax_scale=None,
|
| 160 |
+
q_scale=None,
|
| 161 |
+
causal=False,
|
| 162 |
+
window_size=(-1, -1),
|
| 163 |
+
deterministic=False,
|
| 164 |
+
dtype=torch.bfloat16,
|
| 165 |
+
fa_version=None,
|
| 166 |
+
attention_type=None,
|
| 167 |
+
attn_mask=None,
|
| 168 |
+
):
|
| 169 |
+
attention_type = os.environ.get("VIDEOX_ATTENTION_TYPE", "FLASH_ATTENTION") if attention_type is None else attention_type
|
| 170 |
+
if torch.is_grad_enabled() and attention_type == "SAGE_ATTENTION":
|
| 171 |
+
attention_type = "FLASH_ATTENTION"
|
| 172 |
+
|
| 173 |
+
if attention_type == "SAGE_ATTENTION" and SAGE_ATTENTION_AVAILABLE:
|
| 174 |
+
if q_lens is not None or k_lens is not None:
|
| 175 |
+
warnings.warn(
|
| 176 |
+
'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
out = sageattn(
|
| 180 |
+
q, k, v, attn_mask=attn_mask, tensor_layout="NHD", is_causal=causal, dropout_p=dropout_p)
|
| 181 |
+
|
| 182 |
+
elif attention_type == "FLASH_ATTENTION" and (FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE):
|
| 183 |
+
return flash_attention(
|
| 184 |
+
q=q,
|
| 185 |
+
k=k,
|
| 186 |
+
v=v,
|
| 187 |
+
q_lens=q_lens,
|
| 188 |
+
k_lens=k_lens,
|
| 189 |
+
dropout_p=dropout_p,
|
| 190 |
+
softmax_scale=softmax_scale,
|
| 191 |
+
q_scale=q_scale,
|
| 192 |
+
causal=causal,
|
| 193 |
+
window_size=window_size,
|
| 194 |
+
deterministic=deterministic,
|
| 195 |
+
dtype=dtype,
|
| 196 |
+
version=fa_version,
|
| 197 |
+
)
|
| 198 |
+
else:
|
| 199 |
+
if q_lens is not None or k_lens is not None:
|
| 200 |
+
warnings.warn(
|
| 201 |
+
'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
|
| 202 |
+
)
|
| 203 |
+
q = q.transpose(1, 2)
|
| 204 |
+
k = k.transpose(1, 2)
|
| 205 |
+
v = v.transpose(1, 2)
|
| 206 |
+
|
| 207 |
+
out = torch.nn.functional.scaled_dot_product_attention(
|
| 208 |
+
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
|
| 209 |
+
|
| 210 |
+
out = out.transpose(1, 2).contiguous()
|
| 211 |
+
return out
|
videox_fun/models/creator/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""Minimal Creator audio/VAE building blocks."""
|
videox_fun/models/creator/creator_audio_dit.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Code adapted from DiffSynth-Studio's Wan DiT implementation:
|
| 3 |
+
https://github.com/modelscope/DiffSynth-Studio/blob/main/diffsynth/models/wan_video_dit.py
|
| 4 |
+
"""
|
| 5 |
+
|
| 6 |
+
import math
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn as nn
|
| 10 |
+
|
| 11 |
+
from .creator_video_dit import DiTBlock
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def sinusoidal_embedding_1d(dim, position):
|
| 15 |
+
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
|
| 16 |
+
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
|
| 17 |
+
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
| 18 |
+
return x.to(position.dtype)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def precompute_freqs_cis_1d(dim: int, end: int = 16384, theta: float = 10000.0):
|
| 22 |
+
f_freqs_cis = precompute_freqs_cis(dim, end, theta)
|
| 23 |
+
return f_freqs_cis.chunk(3, dim=-1)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def precompute_freqs_cis(dim: int, end: int = 16384, theta: float = 10000.0, s: float = 1.0):
|
| 27 |
+
# 1d rope precompute
|
| 28 |
+
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
|
| 29 |
+
[: (dim // 2)].double() / dim))
|
| 30 |
+
pos = torch.arange(end, dtype=torch.float64, device=freqs.device) * s
|
| 31 |
+
freqs = torch.outer(pos, freqs)
|
| 32 |
+
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
|
| 33 |
+
return freqs_cis
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class MLP(torch.nn.Module):
|
| 37 |
+
def __init__(self, in_dim, out_dim, has_pos_emb=False):
|
| 38 |
+
super().__init__()
|
| 39 |
+
self.proj = torch.nn.Sequential(
|
| 40 |
+
nn.LayerNorm(in_dim),
|
| 41 |
+
nn.Linear(in_dim, in_dim),
|
| 42 |
+
nn.GELU(),
|
| 43 |
+
nn.Linear(in_dim, out_dim),
|
| 44 |
+
nn.LayerNorm(out_dim)
|
| 45 |
+
)
|
| 46 |
+
self.has_pos_emb = has_pos_emb
|
| 47 |
+
if has_pos_emb:
|
| 48 |
+
self.emb_pos = torch.nn.Parameter(torch.zeros((1, 514, 1280)))
|
| 49 |
+
|
| 50 |
+
def forward(self, x):
|
| 51 |
+
if self.has_pos_emb:
|
| 52 |
+
x = x + self.emb_pos.to(dtype=x.dtype, device=x.device)
|
| 53 |
+
return self.proj(x)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
class Head(nn.Module):
|
| 57 |
+
def __init__(self, dim: int, out_dim: int, patch_size, eps: float):
|
| 58 |
+
super().__init__()
|
| 59 |
+
self.dim = dim
|
| 60 |
+
self.patch_size = patch_size
|
| 61 |
+
self.norm = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
| 62 |
+
self.head = nn.Linear(dim, out_dim * math.prod(patch_size))
|
| 63 |
+
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
| 64 |
+
|
| 65 |
+
def forward(self, x, t_mod):
|
| 66 |
+
if len(t_mod.shape) == 3:
|
| 67 |
+
shift, scale = (self.modulation.unsqueeze(0).to(dtype=t_mod.dtype, device=t_mod.device) + t_mod.unsqueeze(2)).chunk(2, dim=2)
|
| 68 |
+
x = (self.head(self.norm(x) * (1 + scale.squeeze(2)) + shift.squeeze(2)))
|
| 69 |
+
else:
|
| 70 |
+
# NOTE: `t_mod` used to be [B, C]. When B=1 broadcasting works, but for B>1
|
| 71 |
+
# it does not align with [1, 2, C]. We therefore unsqueeze at dim=1 here.
|
| 72 |
+
shift, scale = (self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod.unsqueeze(1)).chunk(2, dim=1)
|
| 73 |
+
x = (self.head(self.norm(x) * (1 + scale) + shift))
|
| 74 |
+
return x
|
videox_fun/models/creator/creator_video_dit.py
ADDED
|
@@ -0,0 +1,103 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared DiT blocks used by the Creator audio transformer."""
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
from einops import rearrange
|
| 6 |
+
from torch.nn import RMSNorm
|
| 7 |
+
|
| 8 |
+
from ..attention_utils import attention
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def sinusoidal_embedding_1d(dim, position):
|
| 12 |
+
sinusoid = torch.outer(
|
| 13 |
+
position.type(torch.float64),
|
| 14 |
+
torch.pow(10000, -torch.arange(dim // 2, dtype=torch.float64, device=position.device).div(dim // 2)),
|
| 15 |
+
)
|
| 16 |
+
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1).to(position.dtype)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
|
| 20 |
+
return (x * (1 + scale) + shift).to(shift.dtype)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def rope_apply_head_dim(x, freqs, head_dim):
|
| 24 |
+
x = rearrange(x, "b s (n d) -> b s n d", d=head_dim)
|
| 25 |
+
x_complex = torch.view_as_complex(x.to(torch.float64).reshape(*x.shape[:3], -1, 2))
|
| 26 |
+
return torch.view_as_real(x_complex * freqs).flatten(2).to(x.dtype)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class SelfAttention(nn.Module):
|
| 30 |
+
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
|
| 31 |
+
super().__init__()
|
| 32 |
+
self.num_heads = num_heads
|
| 33 |
+
self.head_dim = dim // num_heads
|
| 34 |
+
self.q = nn.Linear(dim, dim)
|
| 35 |
+
self.k = nn.Linear(dim, dim)
|
| 36 |
+
self.v = nn.Linear(dim, dim)
|
| 37 |
+
self.o = nn.Linear(dim, dim)
|
| 38 |
+
self.norm_q = RMSNorm(dim, eps=eps)
|
| 39 |
+
self.norm_k = RMSNorm(dim, eps=eps)
|
| 40 |
+
|
| 41 |
+
def forward(self, x, freqs, seq_lens=None):
|
| 42 |
+
q = rope_apply_head_dim(self.norm_q(self.q(x)), freqs, self.head_dim)
|
| 43 |
+
k = rope_apply_head_dim(self.norm_k(self.k(x)), freqs, self.head_dim)
|
| 44 |
+
v = self.v(x)
|
| 45 |
+
b, s = q.shape[:2]
|
| 46 |
+
q = q.view(b, s, self.num_heads, self.head_dim)
|
| 47 |
+
k = k.view(b, s, self.num_heads, self.head_dim)
|
| 48 |
+
v = v.view(b, s, self.num_heads, self.head_dim)
|
| 49 |
+
return self.o(attention(q, k, v, k_lens=seq_lens).flatten(2))
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class CrossAttention(nn.Module):
|
| 53 |
+
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, has_image_input: bool = False):
|
| 54 |
+
super().__init__()
|
| 55 |
+
self.num_heads = num_heads
|
| 56 |
+
self.q = nn.Linear(dim, dim)
|
| 57 |
+
self.k = nn.Linear(dim, dim)
|
| 58 |
+
self.v = nn.Linear(dim, dim)
|
| 59 |
+
self.o = nn.Linear(dim, dim)
|
| 60 |
+
self.norm_q = RMSNorm(dim, eps=eps)
|
| 61 |
+
self.norm_k = RMSNorm(dim, eps=eps)
|
| 62 |
+
|
| 63 |
+
def forward(self, x: torch.Tensor, context: torch.Tensor):
|
| 64 |
+
q = self.norm_q(self.q(x))
|
| 65 |
+
k = self.norm_k(self.k(context))
|
| 66 |
+
v = self.v(context)
|
| 67 |
+
b, s = q.shape[:2]
|
| 68 |
+
n = self.num_heads
|
| 69 |
+
d = q.shape[-1] // n
|
| 70 |
+
output = attention(
|
| 71 |
+
q.view(b, s, n, d),
|
| 72 |
+
k.view(b, k.shape[1], n, d),
|
| 73 |
+
v.view(b, v.shape[1], n, d),
|
| 74 |
+
)
|
| 75 |
+
return self.o(output.flatten(2))
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
class GateModule(nn.Module):
|
| 79 |
+
def forward(self, x, gate, residual):
|
| 80 |
+
return x + gate * residual
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class DiTBlock(nn.Module):
|
| 84 |
+
def __init__(self, has_image_input: bool, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
|
| 85 |
+
super().__init__()
|
| 86 |
+
self.self_attn = SelfAttention(dim, num_heads, eps)
|
| 87 |
+
self.cross_attn = CrossAttention(dim, num_heads, eps, has_image_input)
|
| 88 |
+
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
| 89 |
+
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
| 90 |
+
self.norm3 = nn.LayerNorm(dim, eps=eps)
|
| 91 |
+
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(approximate="tanh"), nn.Linear(ffn_dim, dim))
|
| 92 |
+
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
| 93 |
+
self.gate = GateModule()
|
| 94 |
+
|
| 95 |
+
def forward(self, x, context, t_mod, freqs, seq_lens=None):
|
| 96 |
+
chunk_dim = 2 if t_mod.ndim == 4 else 1
|
| 97 |
+
modulation = (self.modulation.to(t_mod) + t_mod).chunk(6, dim=chunk_dim)
|
| 98 |
+
if chunk_dim == 2:
|
| 99 |
+
modulation = tuple(part.squeeze(2) for part in modulation)
|
| 100 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = modulation
|
| 101 |
+
x = self.gate(x, gate_msa, self.self_attn(modulate(self.norm1(x), shift_msa, scale_msa), freqs, seq_lens))
|
| 102 |
+
x = x + self.cross_attn(self.norm3(x), context)
|
| 103 |
+
return self.gate(x, gate_mlp, self.ffn(modulate(self.norm2(x), shift_mlp, scale_mlp)))
|
videox_fun/models/creator/dac_vae.py
ADDED
|
@@ -0,0 +1,878 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Code adapted from HunyuanVideo-Foley's DAC VAE implementation:
|
| 3 |
+
https://github.com/Tencent-Hunyuan/HunyuanVideo-Foley/tree/main/hunyuanvideo_foley/models/dac_vae
|
| 4 |
+
|
| 5 |
+
We modified the original code to fit this project (e.g., integration and refactors).
|
| 6 |
+
"""
|
| 7 |
+
import math
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
from typing import List, Union
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
import torch
|
| 14 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 15 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 16 |
+
from diffusers.utils.accelerate_utils import apply_forward_hook
|
| 17 |
+
from torch import nn
|
| 18 |
+
from torch.nn.utils import weight_norm
|
| 19 |
+
|
| 20 |
+
SUPPORTED_VERSIONS = ["1.0.0"]
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
@dataclass
|
| 24 |
+
class DACFile:
|
| 25 |
+
codes: torch.Tensor
|
| 26 |
+
|
| 27 |
+
# Metadata
|
| 28 |
+
chunk_length: int
|
| 29 |
+
original_length: int
|
| 30 |
+
input_db: float
|
| 31 |
+
channels: int
|
| 32 |
+
sample_rate: int
|
| 33 |
+
padding: bool
|
| 34 |
+
dac_version: str
|
| 35 |
+
|
| 36 |
+
def save(self, path):
|
| 37 |
+
artifacts = {
|
| 38 |
+
"codes": self.codes.numpy().astype(np.uint16),
|
| 39 |
+
"metadata": {
|
| 40 |
+
"input_db": self.input_db.numpy().astype(np.float32),
|
| 41 |
+
"original_length": self.original_length,
|
| 42 |
+
"sample_rate": self.sample_rate,
|
| 43 |
+
"chunk_length": self.chunk_length,
|
| 44 |
+
"channels": self.channels,
|
| 45 |
+
"padding": self.padding,
|
| 46 |
+
"dac_version": SUPPORTED_VERSIONS[-1],
|
| 47 |
+
},
|
| 48 |
+
}
|
| 49 |
+
path = Path(path).with_suffix(".dac")
|
| 50 |
+
with open(path, "wb") as f:
|
| 51 |
+
np.save(f, artifacts)
|
| 52 |
+
return path
|
| 53 |
+
|
| 54 |
+
@classmethod
|
| 55 |
+
def load(cls, path):
|
| 56 |
+
artifacts = np.load(path, allow_pickle=True)[()]
|
| 57 |
+
codes = torch.from_numpy(artifacts["codes"].astype(int))
|
| 58 |
+
if artifacts["metadata"].get("dac_version", None) not in SUPPORTED_VERSIONS:
|
| 59 |
+
raise RuntimeError(
|
| 60 |
+
f"Given file {path} can't be loaded with this version of descript-audio-codec."
|
| 61 |
+
)
|
| 62 |
+
return cls(codes=codes, **artifacts["metadata"])
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class CodecMixin:
|
| 66 |
+
@property
|
| 67 |
+
def padding(self):
|
| 68 |
+
if not hasattr(self, "_padding"):
|
| 69 |
+
self._padding = True
|
| 70 |
+
return self._padding
|
| 71 |
+
|
| 72 |
+
@padding.setter
|
| 73 |
+
def padding(self, value):
|
| 74 |
+
assert isinstance(value, bool)
|
| 75 |
+
|
| 76 |
+
layers = [
|
| 77 |
+
l for l in self.modules() if isinstance(l, (nn.Conv1d, nn.ConvTranspose1d))
|
| 78 |
+
]
|
| 79 |
+
|
| 80 |
+
for layer in layers:
|
| 81 |
+
if value:
|
| 82 |
+
if hasattr(layer, "original_padding"):
|
| 83 |
+
layer.padding = layer.original_padding
|
| 84 |
+
else:
|
| 85 |
+
layer.original_padding = layer.padding
|
| 86 |
+
layer.padding = tuple(0 for _ in range(len(layer.padding)))
|
| 87 |
+
|
| 88 |
+
self._padding = value
|
| 89 |
+
|
| 90 |
+
def get_delay(self):
|
| 91 |
+
# Any number works here, delay is invariant to input length
|
| 92 |
+
l_out = self.get_output_length(0)
|
| 93 |
+
L = l_out
|
| 94 |
+
|
| 95 |
+
layers = []
|
| 96 |
+
for layer in self.modules():
|
| 97 |
+
if isinstance(layer, (nn.Conv1d, nn.ConvTranspose1d)):
|
| 98 |
+
layers.append(layer)
|
| 99 |
+
|
| 100 |
+
for layer in reversed(layers):
|
| 101 |
+
d = layer.dilation[0]
|
| 102 |
+
k = layer.kernel_size[0]
|
| 103 |
+
s = layer.stride[0]
|
| 104 |
+
|
| 105 |
+
if isinstance(layer, nn.ConvTranspose1d):
|
| 106 |
+
L = ((L - d * (k - 1) - 1) / s) + 1
|
| 107 |
+
elif isinstance(layer, nn.Conv1d):
|
| 108 |
+
L = (L - 1) * s + d * (k - 1) + 1
|
| 109 |
+
|
| 110 |
+
L = math.ceil(L)
|
| 111 |
+
|
| 112 |
+
l_in = L
|
| 113 |
+
|
| 114 |
+
return (l_in - l_out) // 2
|
| 115 |
+
|
| 116 |
+
def get_output_length(self, input_length):
|
| 117 |
+
L = input_length
|
| 118 |
+
# Calculate output length
|
| 119 |
+
for layer in self.modules():
|
| 120 |
+
if isinstance(layer, (nn.Conv1d, nn.ConvTranspose1d)):
|
| 121 |
+
d = layer.dilation[0]
|
| 122 |
+
k = layer.kernel_size[0]
|
| 123 |
+
s = layer.stride[0]
|
| 124 |
+
|
| 125 |
+
if isinstance(layer, nn.Conv1d):
|
| 126 |
+
L = ((L - d * (k - 1) - 1) / s) + 1
|
| 127 |
+
elif isinstance(layer, nn.ConvTranspose1d):
|
| 128 |
+
L = (L - 1) * s + d * (k - 1) + 1
|
| 129 |
+
|
| 130 |
+
L = math.floor(L)
|
| 131 |
+
return L
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def WNConv1d(*args, **kwargs):
|
| 136 |
+
return weight_norm(nn.Conv1d(*args, **kwargs))
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def WNConvTranspose1d(*args, **kwargs):
|
| 140 |
+
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
# Scripting this brings model speed up 1.4x
|
| 144 |
+
@torch.jit.script
|
| 145 |
+
def snake(x, alpha):
|
| 146 |
+
shape = x.shape
|
| 147 |
+
x = x.reshape(shape[0], shape[1], -1)
|
| 148 |
+
x = x + (alpha + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
|
| 149 |
+
x = x.reshape(shape)
|
| 150 |
+
return x
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
class Snake1d(nn.Module):
|
| 154 |
+
def __init__(self, channels):
|
| 155 |
+
super().__init__()
|
| 156 |
+
self.alpha = nn.Parameter(torch.ones(1, channels, 1))
|
| 157 |
+
|
| 158 |
+
def forward(self, x):
|
| 159 |
+
return snake(x, self.alpha)
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
import torch.nn.functional as F
|
| 163 |
+
from einops import rearrange
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
class VectorQuantize(nn.Module):
|
| 167 |
+
"""
|
| 168 |
+
Implementation of VQ similar to Karpathy's repo:
|
| 169 |
+
https://github.com/karpathy/deep-vector-quantization
|
| 170 |
+
Additionally uses following tricks from Improved VQGAN
|
| 171 |
+
(https://arxiv.org/pdf/2110.04627.pdf):
|
| 172 |
+
1. Factorized codes: Perform nearest neighbor lookup in low-dimensional space
|
| 173 |
+
for improved codebook usage
|
| 174 |
+
2. l2-normalized codes: Converts euclidean distance to cosine similarity which
|
| 175 |
+
improves training stability
|
| 176 |
+
"""
|
| 177 |
+
|
| 178 |
+
def __init__(self, input_dim: int, codebook_size: int, codebook_dim: int):
|
| 179 |
+
super().__init__()
|
| 180 |
+
self.codebook_size = codebook_size
|
| 181 |
+
self.codebook_dim = codebook_dim
|
| 182 |
+
|
| 183 |
+
self.in_proj = WNConv1d(input_dim, codebook_dim, kernel_size=1)
|
| 184 |
+
self.out_proj = WNConv1d(codebook_dim, input_dim, kernel_size=1)
|
| 185 |
+
self.codebook = nn.Embedding(codebook_size, codebook_dim)
|
| 186 |
+
|
| 187 |
+
def forward(self, z):
|
| 188 |
+
"""Quantized the input tensor using a fixed codebook and returns
|
| 189 |
+
the corresponding codebook vectors
|
| 190 |
+
|
| 191 |
+
Parameters
|
| 192 |
+
----------
|
| 193 |
+
z : Tensor[B x D x T]
|
| 194 |
+
|
| 195 |
+
Returns
|
| 196 |
+
-------
|
| 197 |
+
Tensor[B x D x T]
|
| 198 |
+
Quantized continuous representation of input
|
| 199 |
+
Tensor[1]
|
| 200 |
+
Commitment loss to train encoder to predict vectors closer to codebook
|
| 201 |
+
entries
|
| 202 |
+
Tensor[1]
|
| 203 |
+
Codebook loss to update the codebook
|
| 204 |
+
Tensor[B x T]
|
| 205 |
+
Codebook indices (quantized discrete representation of input)
|
| 206 |
+
Tensor[B x D x T]
|
| 207 |
+
Projected latents (continuous representation of input before quantization)
|
| 208 |
+
"""
|
| 209 |
+
|
| 210 |
+
# Factorized codes (ViT-VQGAN) Project input into low-dimensional space
|
| 211 |
+
z_e = self.in_proj(z) # z_e : (B x D x T)
|
| 212 |
+
z_q, indices = self.decode_latents(z_e)
|
| 213 |
+
|
| 214 |
+
commitment_loss = F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
|
| 215 |
+
codebook_loss = F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
|
| 216 |
+
|
| 217 |
+
z_q = (
|
| 218 |
+
z_e + (z_q - z_e).detach()
|
| 219 |
+
) # noop in forward pass, straight-through gradient estimator in backward pass
|
| 220 |
+
|
| 221 |
+
z_q = self.out_proj(z_q)
|
| 222 |
+
|
| 223 |
+
return z_q, commitment_loss, codebook_loss, indices, z_e
|
| 224 |
+
|
| 225 |
+
def embed_code(self, embed_id):
|
| 226 |
+
return F.embedding(embed_id, self.codebook.weight)
|
| 227 |
+
|
| 228 |
+
def decode_code(self, embed_id):
|
| 229 |
+
return self.embed_code(embed_id).transpose(1, 2)
|
| 230 |
+
|
| 231 |
+
def decode_latents(self, latents):
|
| 232 |
+
encodings = rearrange(latents, "b d t -> (b t) d")
|
| 233 |
+
codebook = self.codebook.weight # codebook: (N x D)
|
| 234 |
+
|
| 235 |
+
# L2 normalize encodings and codebook (ViT-VQGAN)
|
| 236 |
+
encodings = F.normalize(encodings)
|
| 237 |
+
codebook = F.normalize(codebook)
|
| 238 |
+
|
| 239 |
+
# Compute euclidean distance with codebook
|
| 240 |
+
dist = (
|
| 241 |
+
encodings.pow(2).sum(1, keepdim=True)
|
| 242 |
+
- 2 * encodings @ codebook.t()
|
| 243 |
+
+ codebook.pow(2).sum(1, keepdim=True).t()
|
| 244 |
+
)
|
| 245 |
+
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
|
| 246 |
+
z_q = self.decode_code(indices)
|
| 247 |
+
return z_q, indices
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
class ResidualVectorQuantize(nn.Module):
|
| 251 |
+
"""
|
| 252 |
+
Introduced in SoundStream: An end2end neural audio codec
|
| 253 |
+
https://arxiv.org/abs/2107.03312
|
| 254 |
+
"""
|
| 255 |
+
|
| 256 |
+
def __init__(
|
| 257 |
+
self,
|
| 258 |
+
input_dim: int = 512,
|
| 259 |
+
n_codebooks: int = 9,
|
| 260 |
+
codebook_size: int = 1024,
|
| 261 |
+
codebook_dim: Union[int, list] = 8,
|
| 262 |
+
quantizer_dropout: float = 0.0,
|
| 263 |
+
):
|
| 264 |
+
super().__init__()
|
| 265 |
+
if isinstance(codebook_dim, int):
|
| 266 |
+
codebook_dim = [codebook_dim for _ in range(n_codebooks)]
|
| 267 |
+
|
| 268 |
+
self.n_codebooks = n_codebooks
|
| 269 |
+
self.codebook_dim = codebook_dim
|
| 270 |
+
self.codebook_size = codebook_size
|
| 271 |
+
|
| 272 |
+
self.quantizers = nn.ModuleList(
|
| 273 |
+
[
|
| 274 |
+
VectorQuantize(input_dim, codebook_size, codebook_dim[i])
|
| 275 |
+
for i in range(n_codebooks)
|
| 276 |
+
]
|
| 277 |
+
)
|
| 278 |
+
self.quantizer_dropout = quantizer_dropout
|
| 279 |
+
|
| 280 |
+
def forward(self, z, n_quantizers: int = None):
|
| 281 |
+
"""Quantized the input tensor using a fixed set of `n` codebooks and returns
|
| 282 |
+
the corresponding codebook vectors
|
| 283 |
+
Parameters
|
| 284 |
+
----------
|
| 285 |
+
z : Tensor[B x D x T]
|
| 286 |
+
n_quantizers : int, optional
|
| 287 |
+
No. of quantizers to use
|
| 288 |
+
(n_quantizers < self.n_codebooks ex: for quantizer dropout)
|
| 289 |
+
Note: if `self.quantizer_dropout` is True, this argument is ignored
|
| 290 |
+
when in training mode, and a random number of quantizers is used.
|
| 291 |
+
Returns
|
| 292 |
+
-------
|
| 293 |
+
dict
|
| 294 |
+
A dictionary with the following keys:
|
| 295 |
+
|
| 296 |
+
"z" : Tensor[B x D x T]
|
| 297 |
+
Quantized continuous representation of input
|
| 298 |
+
"codes" : Tensor[B x N x T]
|
| 299 |
+
Codebook indices for each codebook
|
| 300 |
+
(quantized discrete representation of input)
|
| 301 |
+
"latents" : Tensor[B x N*D x T]
|
| 302 |
+
Projected latents (continuous representation of input before quantization)
|
| 303 |
+
"vq/commitment_loss" : Tensor[1]
|
| 304 |
+
Commitment loss to train encoder to predict vectors closer to codebook
|
| 305 |
+
entries
|
| 306 |
+
"vq/codebook_loss" : Tensor[1]
|
| 307 |
+
Codebook loss to update the codebook
|
| 308 |
+
"""
|
| 309 |
+
z_q = 0
|
| 310 |
+
residual = z
|
| 311 |
+
commitment_loss = 0
|
| 312 |
+
codebook_loss = 0
|
| 313 |
+
|
| 314 |
+
codebook_indices = []
|
| 315 |
+
latents = []
|
| 316 |
+
|
| 317 |
+
if n_quantizers is None:
|
| 318 |
+
n_quantizers = self.n_codebooks
|
| 319 |
+
if self.training:
|
| 320 |
+
n_quantizers = torch.ones((z.shape[0],)) * self.n_codebooks + 1
|
| 321 |
+
dropout = torch.randint(1, self.n_codebooks + 1, (z.shape[0],))
|
| 322 |
+
n_dropout = int(z.shape[0] * self.quantizer_dropout)
|
| 323 |
+
n_quantizers[:n_dropout] = dropout[:n_dropout]
|
| 324 |
+
n_quantizers = n_quantizers.to(z.device)
|
| 325 |
+
|
| 326 |
+
for i, quantizer in enumerate(self.quantizers):
|
| 327 |
+
if self.training is False and i >= n_quantizers:
|
| 328 |
+
break
|
| 329 |
+
|
| 330 |
+
z_q_i, commitment_loss_i, codebook_loss_i, indices_i, z_e_i = quantizer(
|
| 331 |
+
residual
|
| 332 |
+
)
|
| 333 |
+
|
| 334 |
+
# Create mask to apply quantizer dropout
|
| 335 |
+
mask = (
|
| 336 |
+
torch.full((z.shape[0],), fill_value=i, device=z.device) < n_quantizers
|
| 337 |
+
)
|
| 338 |
+
z_q = z_q + z_q_i * mask[:, None, None]
|
| 339 |
+
residual = residual - z_q_i
|
| 340 |
+
|
| 341 |
+
# Sum losses
|
| 342 |
+
commitment_loss += (commitment_loss_i * mask).mean()
|
| 343 |
+
codebook_loss += (codebook_loss_i * mask).mean()
|
| 344 |
+
|
| 345 |
+
codebook_indices.append(indices_i)
|
| 346 |
+
latents.append(z_e_i)
|
| 347 |
+
|
| 348 |
+
codes = torch.stack(codebook_indices, dim=1)
|
| 349 |
+
latents = torch.cat(latents, dim=1)
|
| 350 |
+
|
| 351 |
+
return z_q, codes, latents, commitment_loss, codebook_loss
|
| 352 |
+
|
| 353 |
+
def from_codes(self, codes: torch.Tensor):
|
| 354 |
+
"""Given the quantized codes, reconstruct the continuous representation
|
| 355 |
+
Parameters
|
| 356 |
+
----------
|
| 357 |
+
codes : Tensor[B x N x T]
|
| 358 |
+
Quantized discrete representation of input
|
| 359 |
+
Returns
|
| 360 |
+
-------
|
| 361 |
+
Tensor[B x D x T]
|
| 362 |
+
Quantized continuous representation of input
|
| 363 |
+
"""
|
| 364 |
+
z_q = 0.0
|
| 365 |
+
z_p = []
|
| 366 |
+
n_codebooks = codes.shape[1]
|
| 367 |
+
for i in range(n_codebooks):
|
| 368 |
+
z_p_i = self.quantizers[i].decode_code(codes[:, i, :])
|
| 369 |
+
z_p.append(z_p_i)
|
| 370 |
+
|
| 371 |
+
z_q_i = self.quantizers[i].out_proj(z_p_i)
|
| 372 |
+
z_q = z_q + z_q_i
|
| 373 |
+
return z_q, torch.cat(z_p, dim=1), codes
|
| 374 |
+
|
| 375 |
+
def from_latents(self, latents: torch.Tensor):
|
| 376 |
+
"""Given the unquantized latents, reconstruct the
|
| 377 |
+
continuous representation after quantization.
|
| 378 |
+
|
| 379 |
+
Parameters
|
| 380 |
+
----------
|
| 381 |
+
latents : Tensor[B x N x T]
|
| 382 |
+
Continuous representation of input after projection
|
| 383 |
+
|
| 384 |
+
Returns
|
| 385 |
+
-------
|
| 386 |
+
Tensor[B x D x T]
|
| 387 |
+
Quantized representation of full-projected space
|
| 388 |
+
Tensor[B x D x T]
|
| 389 |
+
Quantized representation of latent space
|
| 390 |
+
"""
|
| 391 |
+
z_q = 0
|
| 392 |
+
z_p = []
|
| 393 |
+
codes = []
|
| 394 |
+
dims = np.cumsum([0] + [q.codebook_dim for q in self.quantizers])
|
| 395 |
+
|
| 396 |
+
n_codebooks = np.where(dims <= latents.shape[1])[0].max(axis=0, keepdims=True)[
|
| 397 |
+
0
|
| 398 |
+
]
|
| 399 |
+
for i in range(n_codebooks):
|
| 400 |
+
j, k = dims[i], dims[i + 1]
|
| 401 |
+
z_p_i, codes_i = self.quantizers[i].decode_latents(latents[:, j:k, :])
|
| 402 |
+
z_p.append(z_p_i)
|
| 403 |
+
codes.append(codes_i)
|
| 404 |
+
|
| 405 |
+
z_q_i = self.quantizers[i].out_proj(z_p_i)
|
| 406 |
+
z_q = z_q + z_q_i
|
| 407 |
+
|
| 408 |
+
return z_q, torch.cat(z_p, dim=1), torch.stack(codes, dim=1)
|
| 409 |
+
|
| 410 |
+
|
| 411 |
+
class AbstractDistribution:
|
| 412 |
+
def sample(self):
|
| 413 |
+
raise NotImplementedError()
|
| 414 |
+
|
| 415 |
+
def mode(self):
|
| 416 |
+
raise NotImplementedError()
|
| 417 |
+
|
| 418 |
+
|
| 419 |
+
class DiracDistribution(AbstractDistribution):
|
| 420 |
+
def __init__(self, value):
|
| 421 |
+
self.value = value
|
| 422 |
+
|
| 423 |
+
def sample(self):
|
| 424 |
+
return self.value
|
| 425 |
+
|
| 426 |
+
def mode(self):
|
| 427 |
+
return self.value
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
class DiagonalGaussianDistribution(object):
|
| 431 |
+
def __init__(self, parameters, deterministic=False):
|
| 432 |
+
self.parameters = parameters
|
| 433 |
+
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
| 434 |
+
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
| 435 |
+
self.deterministic = deterministic
|
| 436 |
+
self.std = torch.exp(0.5 * self.logvar)
|
| 437 |
+
self.var = torch.exp(self.logvar)
|
| 438 |
+
if self.deterministic:
|
| 439 |
+
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
|
| 440 |
+
|
| 441 |
+
def sample(self):
|
| 442 |
+
x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
|
| 443 |
+
return x
|
| 444 |
+
|
| 445 |
+
def kl(self, other=None):
|
| 446 |
+
if self.deterministic:
|
| 447 |
+
return torch.Tensor([0.0])
|
| 448 |
+
else:
|
| 449 |
+
if other is None:
|
| 450 |
+
return 0.5 * torch.mean(
|
| 451 |
+
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
|
| 452 |
+
dim=[1, 2],
|
| 453 |
+
)
|
| 454 |
+
else:
|
| 455 |
+
return 0.5 * torch.mean(
|
| 456 |
+
torch.pow(self.mean - other.mean, 2) / other.var
|
| 457 |
+
+ self.var / other.var
|
| 458 |
+
- 1.0
|
| 459 |
+
- self.logvar
|
| 460 |
+
+ other.logvar,
|
| 461 |
+
dim=[1, 2],
|
| 462 |
+
)
|
| 463 |
+
|
| 464 |
+
def nll(self, sample, dims=[1, 2]):
|
| 465 |
+
if self.deterministic:
|
| 466 |
+
return torch.Tensor([0.0])
|
| 467 |
+
logtwopi = np.log(2.0 * np.pi)
|
| 468 |
+
return 0.5 * torch.sum(
|
| 469 |
+
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
| 470 |
+
dim=dims,
|
| 471 |
+
)
|
| 472 |
+
|
| 473 |
+
def mode(self):
|
| 474 |
+
return self.mean
|
| 475 |
+
|
| 476 |
+
|
| 477 |
+
def normal_kl(mean1, logvar1, mean2, logvar2):
|
| 478 |
+
"""
|
| 479 |
+
source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
|
| 480 |
+
Compute the KL divergence between two gaussians.
|
| 481 |
+
Shapes are automatically broadcasted, so batches can be compared to
|
| 482 |
+
scalars, among other use cases.
|
| 483 |
+
"""
|
| 484 |
+
tensor = None
|
| 485 |
+
for obj in (mean1, logvar1, mean2, logvar2):
|
| 486 |
+
if isinstance(obj, torch.Tensor):
|
| 487 |
+
tensor = obj
|
| 488 |
+
break
|
| 489 |
+
assert tensor is not None, "at least one argument must be a Tensor"
|
| 490 |
+
|
| 491 |
+
# Force variances to be Tensors. Broadcasting helps convert scalars to
|
| 492 |
+
# Tensors, but it does not work for torch.exp().
|
| 493 |
+
logvar1, logvar2 = [x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor) for x in (logvar1, logvar2)]
|
| 494 |
+
|
| 495 |
+
return 0.5 * (
|
| 496 |
+
-1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
|
| 497 |
+
)
|
| 498 |
+
|
| 499 |
+
|
| 500 |
+
def init_weights(m):
|
| 501 |
+
if isinstance(m, nn.Conv1d):
|
| 502 |
+
nn.init.trunc_normal_(m.weight, std=0.02)
|
| 503 |
+
nn.init.constant_(m.bias, 0)
|
| 504 |
+
|
| 505 |
+
|
| 506 |
+
class ResidualUnit(nn.Module):
|
| 507 |
+
def __init__(self, dim: int = 16, dilation: int = 1):
|
| 508 |
+
super().__init__()
|
| 509 |
+
pad = ((7 - 1) * dilation) // 2
|
| 510 |
+
self.block = nn.Sequential(
|
| 511 |
+
Snake1d(dim),
|
| 512 |
+
WNConv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad),
|
| 513 |
+
Snake1d(dim),
|
| 514 |
+
WNConv1d(dim, dim, kernel_size=1),
|
| 515 |
+
)
|
| 516 |
+
|
| 517 |
+
def forward(self, x):
|
| 518 |
+
y = self.block(x)
|
| 519 |
+
pad = (x.shape[-1] - y.shape[-1]) // 2
|
| 520 |
+
if pad > 0:
|
| 521 |
+
x = x[..., pad:-pad]
|
| 522 |
+
return x + y
|
| 523 |
+
|
| 524 |
+
|
| 525 |
+
class EncoderBlock(nn.Module):
|
| 526 |
+
def __init__(self, dim: int = 16, stride: int = 1):
|
| 527 |
+
super().__init__()
|
| 528 |
+
self.block = nn.Sequential(
|
| 529 |
+
ResidualUnit(dim // 2, dilation=1),
|
| 530 |
+
ResidualUnit(dim // 2, dilation=3),
|
| 531 |
+
ResidualUnit(dim // 2, dilation=9),
|
| 532 |
+
Snake1d(dim // 2),
|
| 533 |
+
WNConv1d(
|
| 534 |
+
dim // 2,
|
| 535 |
+
dim,
|
| 536 |
+
kernel_size=2 * stride,
|
| 537 |
+
stride=stride,
|
| 538 |
+
padding=math.ceil(stride / 2),
|
| 539 |
+
),
|
| 540 |
+
)
|
| 541 |
+
|
| 542 |
+
def forward(self, x):
|
| 543 |
+
return self.block(x)
|
| 544 |
+
|
| 545 |
+
|
| 546 |
+
class Encoder(nn.Module):
|
| 547 |
+
def __init__(
|
| 548 |
+
self,
|
| 549 |
+
d_model: int = 64,
|
| 550 |
+
strides: list = [2, 4, 8, 8],
|
| 551 |
+
d_latent: int = 64,
|
| 552 |
+
):
|
| 553 |
+
super().__init__()
|
| 554 |
+
# Create first convolution
|
| 555 |
+
self.block = [WNConv1d(1, d_model, kernel_size=7, padding=3)]
|
| 556 |
+
|
| 557 |
+
# Create EncoderBlocks that double channels as they downsample by `stride`
|
| 558 |
+
for stride in strides:
|
| 559 |
+
d_model *= 2
|
| 560 |
+
self.block += [EncoderBlock(d_model, stride=stride)]
|
| 561 |
+
|
| 562 |
+
# Create last convolution
|
| 563 |
+
self.block += [
|
| 564 |
+
Snake1d(d_model),
|
| 565 |
+
WNConv1d(d_model, d_latent, kernel_size=3, padding=1),
|
| 566 |
+
]
|
| 567 |
+
|
| 568 |
+
# Wrap black into nn.Sequential
|
| 569 |
+
self.block = nn.Sequential(*self.block)
|
| 570 |
+
self.enc_dim = d_model
|
| 571 |
+
|
| 572 |
+
def forward(self, x):
|
| 573 |
+
return self.block(x)
|
| 574 |
+
|
| 575 |
+
|
| 576 |
+
class DecoderBlock(nn.Module):
|
| 577 |
+
def __init__(self, input_dim: int = 16, output_dim: int = 8, stride: int = 1):
|
| 578 |
+
super().__init__()
|
| 579 |
+
self.block = nn.Sequential(
|
| 580 |
+
Snake1d(input_dim),
|
| 581 |
+
WNConvTranspose1d(
|
| 582 |
+
input_dim,
|
| 583 |
+
output_dim,
|
| 584 |
+
kernel_size=2 * stride,
|
| 585 |
+
stride=stride,
|
| 586 |
+
padding=math.ceil(stride / 2),
|
| 587 |
+
output_padding=stride % 2,
|
| 588 |
+
),
|
| 589 |
+
ResidualUnit(output_dim, dilation=1),
|
| 590 |
+
ResidualUnit(output_dim, dilation=3),
|
| 591 |
+
ResidualUnit(output_dim, dilation=9),
|
| 592 |
+
)
|
| 593 |
+
|
| 594 |
+
def forward(self, x):
|
| 595 |
+
return self.block(x)
|
| 596 |
+
|
| 597 |
+
|
| 598 |
+
class Decoder(nn.Module):
|
| 599 |
+
def __init__(
|
| 600 |
+
self,
|
| 601 |
+
input_channel,
|
| 602 |
+
channels,
|
| 603 |
+
rates,
|
| 604 |
+
d_out: int = 1,
|
| 605 |
+
):
|
| 606 |
+
super().__init__()
|
| 607 |
+
|
| 608 |
+
# Add first conv layer
|
| 609 |
+
layers = [WNConv1d(input_channel, channels, kernel_size=7, padding=3)]
|
| 610 |
+
|
| 611 |
+
# Add upsampling + MRF blocks
|
| 612 |
+
for i, stride in enumerate(rates):
|
| 613 |
+
input_dim = channels // 2**i
|
| 614 |
+
output_dim = channels // 2 ** (i + 1)
|
| 615 |
+
layers += [DecoderBlock(input_dim, output_dim, stride)]
|
| 616 |
+
|
| 617 |
+
# Add final conv layer
|
| 618 |
+
layers += [
|
| 619 |
+
Snake1d(output_dim),
|
| 620 |
+
WNConv1d(output_dim, d_out, kernel_size=7, padding=3),
|
| 621 |
+
nn.Tanh(),
|
| 622 |
+
]
|
| 623 |
+
|
| 624 |
+
self.model = nn.Sequential(*layers)
|
| 625 |
+
|
| 626 |
+
def forward(self, x):
|
| 627 |
+
return self.model(x)
|
| 628 |
+
|
| 629 |
+
|
| 630 |
+
class DAC(CodecMixin, ModelMixin, ConfigMixin):
|
| 631 |
+
|
| 632 |
+
@register_to_config
|
| 633 |
+
def __init__(
|
| 634 |
+
self,
|
| 635 |
+
encoder_dim: int = 128,
|
| 636 |
+
encoder_rates: List[int] = [2, 3, 4, 5, 8],
|
| 637 |
+
latent_dim: int = None, # 128
|
| 638 |
+
decoder_dim: int = 2048,
|
| 639 |
+
decoder_rates: List[int] = [8, 5, 4, 3, 2],
|
| 640 |
+
n_codebooks: int = 9,
|
| 641 |
+
codebook_size: int = 1024,
|
| 642 |
+
codebook_dim: Union[int, list] = 8,
|
| 643 |
+
quantizer_dropout: bool = False,
|
| 644 |
+
sample_rate: int = 48000,
|
| 645 |
+
continuous: bool = True,
|
| 646 |
+
use_weight_norm: bool = False,
|
| 647 |
+
):
|
| 648 |
+
super().__init__()
|
| 649 |
+
|
| 650 |
+
self.encoder_dim = encoder_dim
|
| 651 |
+
self.encoder_rates = encoder_rates
|
| 652 |
+
self.decoder_dim = decoder_dim
|
| 653 |
+
self.decoder_rates = decoder_rates
|
| 654 |
+
self.sample_rate = sample_rate
|
| 655 |
+
self.continuous = continuous
|
| 656 |
+
self.use_weight_norm = use_weight_norm
|
| 657 |
+
|
| 658 |
+
if latent_dim is None:
|
| 659 |
+
latent_dim = encoder_dim * (2 ** len(encoder_rates))
|
| 660 |
+
|
| 661 |
+
self.latent_dim = latent_dim
|
| 662 |
+
|
| 663 |
+
self.hop_length = np.prod(encoder_rates)
|
| 664 |
+
self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim)
|
| 665 |
+
|
| 666 |
+
if not continuous:
|
| 667 |
+
self.n_codebooks = n_codebooks
|
| 668 |
+
self.codebook_size = codebook_size
|
| 669 |
+
self.codebook_dim = codebook_dim
|
| 670 |
+
self.quantizer = ResidualVectorQuantize(
|
| 671 |
+
input_dim=latent_dim,
|
| 672 |
+
n_codebooks=n_codebooks,
|
| 673 |
+
codebook_size=codebook_size,
|
| 674 |
+
codebook_dim=codebook_dim,
|
| 675 |
+
quantizer_dropout=quantizer_dropout,
|
| 676 |
+
)
|
| 677 |
+
else:
|
| 678 |
+
self.quant_conv = torch.nn.Conv1d(latent_dim, 2 * latent_dim, 1)
|
| 679 |
+
self.post_quant_conv = torch.nn.Conv1d(latent_dim, latent_dim, 1)
|
| 680 |
+
|
| 681 |
+
self.decoder = Decoder(
|
| 682 |
+
latent_dim,
|
| 683 |
+
decoder_dim,
|
| 684 |
+
decoder_rates,
|
| 685 |
+
)
|
| 686 |
+
self.sample_rate = sample_rate
|
| 687 |
+
self.apply(init_weights)
|
| 688 |
+
|
| 689 |
+
if not self.use_weight_norm:
|
| 690 |
+
self.remove_weight_norm()
|
| 691 |
+
|
| 692 |
+
@property
|
| 693 |
+
def dtype(self):
|
| 694 |
+
"""Get the dtype of the model parameters."""
|
| 695 |
+
# Return the dtype of the first parameter found
|
| 696 |
+
for param in self.parameters():
|
| 697 |
+
return param.dtype
|
| 698 |
+
return torch.float32 # fallback
|
| 699 |
+
|
| 700 |
+
@property
|
| 701 |
+
def device(self):
|
| 702 |
+
"""Get the device of the model parameters."""
|
| 703 |
+
# Return the device of the first parameter found
|
| 704 |
+
for param in self.parameters():
|
| 705 |
+
return param.device
|
| 706 |
+
return torch.device('cpu') # fallback
|
| 707 |
+
|
| 708 |
+
def preprocess(self, audio_data, sample_rate):
|
| 709 |
+
if sample_rate is None:
|
| 710 |
+
sample_rate = self.sample_rate
|
| 711 |
+
assert sample_rate == self.sample_rate
|
| 712 |
+
|
| 713 |
+
length = audio_data.shape[-1]
|
| 714 |
+
right_pad = math.ceil(length / self.hop_length) * self.hop_length - length
|
| 715 |
+
audio_data = nn.functional.pad(audio_data, (0, right_pad))
|
| 716 |
+
|
| 717 |
+
return audio_data
|
| 718 |
+
|
| 719 |
+
@apply_forward_hook
|
| 720 |
+
def encode(
|
| 721 |
+
self,
|
| 722 |
+
audio_data: torch.Tensor,
|
| 723 |
+
n_quantizers: int = None,
|
| 724 |
+
):
|
| 725 |
+
"""Encode given audio data and return quantized latent codes
|
| 726 |
+
|
| 727 |
+
Parameters
|
| 728 |
+
----------
|
| 729 |
+
audio_data : Tensor[B x 1 x T]
|
| 730 |
+
Audio data to encode
|
| 731 |
+
n_quantizers : int, optional
|
| 732 |
+
Number of quantizers to use, by default None
|
| 733 |
+
If None, all quantizers are used.
|
| 734 |
+
|
| 735 |
+
Returns
|
| 736 |
+
-------
|
| 737 |
+
dict
|
| 738 |
+
A dictionary with the following keys:
|
| 739 |
+
"z" : Tensor[B x D x T]
|
| 740 |
+
Quantized continuous representation of input
|
| 741 |
+
"codes" : Tensor[B x N x T]
|
| 742 |
+
Codebook indices for each codebook
|
| 743 |
+
(quantized discrete representation of input)
|
| 744 |
+
"latents" : Tensor[B x N*D x T]
|
| 745 |
+
Projected latents (continuous representation of input before quantization)
|
| 746 |
+
"vq/commitment_loss" : Tensor[1]
|
| 747 |
+
Commitment loss to train encoder to predict vectors closer to codebook
|
| 748 |
+
entries
|
| 749 |
+
"vq/codebook_loss" : Tensor[1]
|
| 750 |
+
Codebook loss to update the codebook
|
| 751 |
+
"length" : int
|
| 752 |
+
Number of samples in input audio
|
| 753 |
+
"""
|
| 754 |
+
z = self.encoder(audio_data) # [B x D x T]
|
| 755 |
+
if not self.continuous:
|
| 756 |
+
z, codes, latents, commitment_loss, codebook_loss = self.quantizer(z, n_quantizers)
|
| 757 |
+
else:
|
| 758 |
+
z = self.quant_conv(z) # [B x 2D x T]
|
| 759 |
+
z = DiagonalGaussianDistribution(z)
|
| 760 |
+
codes, latents, commitment_loss, codebook_loss = None, None, 0, 0
|
| 761 |
+
|
| 762 |
+
return z, codes, latents, commitment_loss, codebook_loss
|
| 763 |
+
|
| 764 |
+
@apply_forward_hook
|
| 765 |
+
def decode(self, z: torch.Tensor):
|
| 766 |
+
"""Decode given latent codes and return audio data
|
| 767 |
+
|
| 768 |
+
Parameters
|
| 769 |
+
----------
|
| 770 |
+
z : Tensor[B x D x T]
|
| 771 |
+
Quantized continuous representation of input
|
| 772 |
+
length : int, optional
|
| 773 |
+
Number of samples in output audio, by default None
|
| 774 |
+
|
| 775 |
+
Returns
|
| 776 |
+
-------
|
| 777 |
+
dict
|
| 778 |
+
A dictionary with the following keys:
|
| 779 |
+
"audio" : Tensor[B x 1 x length]
|
| 780 |
+
Decoded audio data.
|
| 781 |
+
"""
|
| 782 |
+
if not self.continuous:
|
| 783 |
+
audio = self.decoder(z)
|
| 784 |
+
else:
|
| 785 |
+
z = self.post_quant_conv(z)
|
| 786 |
+
audio = self.decoder(z)
|
| 787 |
+
|
| 788 |
+
return audio
|
| 789 |
+
|
| 790 |
+
def forward(
|
| 791 |
+
self,
|
| 792 |
+
audio_data: torch.Tensor,
|
| 793 |
+
sample_rate: int = None,
|
| 794 |
+
n_quantizers: int = None,
|
| 795 |
+
):
|
| 796 |
+
"""Model forward pass
|
| 797 |
+
|
| 798 |
+
Parameters
|
| 799 |
+
----------
|
| 800 |
+
audio_data : Tensor[B x 1 x T]
|
| 801 |
+
Audio data to encode
|
| 802 |
+
sample_rate : int, optional
|
| 803 |
+
Sample rate of audio data in Hz, by default None
|
| 804 |
+
If None, defaults to `self.sample_rate`
|
| 805 |
+
n_quantizers : int, optional
|
| 806 |
+
Number of quantizers to use, by default None.
|
| 807 |
+
If None, all quantizers are used.
|
| 808 |
+
|
| 809 |
+
Returns
|
| 810 |
+
-------
|
| 811 |
+
dict
|
| 812 |
+
A dictionary with the following keys:
|
| 813 |
+
"z" : Tensor[B x D x T]
|
| 814 |
+
Quantized continuous representation of input
|
| 815 |
+
"codes" : Tensor[B x N x T]
|
| 816 |
+
Codebook indices for each codebook
|
| 817 |
+
(quantized discrete representation of input)
|
| 818 |
+
"latents" : Tensor[B x N*D x T]
|
| 819 |
+
Projected latents (continuous representation of input before quantization)
|
| 820 |
+
"vq/commitment_loss" : Tensor[1]
|
| 821 |
+
Commitment loss to train encoder to predict vectors closer to codebook
|
| 822 |
+
entries
|
| 823 |
+
"vq/codebook_loss" : Tensor[1]
|
| 824 |
+
Codebook loss to update the codebook
|
| 825 |
+
"length" : int
|
| 826 |
+
Number of samples in input audio
|
| 827 |
+
"audio" : Tensor[B x 1 x length]
|
| 828 |
+
Decoded audio data.
|
| 829 |
+
"""
|
| 830 |
+
length = audio_data.shape[-1]
|
| 831 |
+
audio_data = self.preprocess(audio_data, sample_rate)
|
| 832 |
+
if not self.continuous:
|
| 833 |
+
z, codes, latents, commitment_loss, codebook_loss = self.encode(audio_data, n_quantizers)
|
| 834 |
+
|
| 835 |
+
x = self.decode(z)
|
| 836 |
+
return {
|
| 837 |
+
"audio": x[..., :length],
|
| 838 |
+
"z": z,
|
| 839 |
+
"codes": codes,
|
| 840 |
+
"latents": latents,
|
| 841 |
+
"vq/commitment_loss": commitment_loss,
|
| 842 |
+
"vq/codebook_loss": codebook_loss,
|
| 843 |
+
}
|
| 844 |
+
else:
|
| 845 |
+
posterior, _, _, _, _ = self.encode(audio_data, n_quantizers)
|
| 846 |
+
z = posterior.sample()
|
| 847 |
+
x = self.decode(z)
|
| 848 |
+
|
| 849 |
+
kl_loss = posterior.kl()
|
| 850 |
+
kl_loss = kl_loss.mean()
|
| 851 |
+
|
| 852 |
+
return {
|
| 853 |
+
"audio": x[..., :length],
|
| 854 |
+
"z": z,
|
| 855 |
+
"kl_loss": kl_loss,
|
| 856 |
+
}
|
| 857 |
+
|
| 858 |
+
def remove_weight_norm(self):
|
| 859 |
+
"""
|
| 860 |
+
Remove weight_norm from all modules in the model.
|
| 861 |
+
This fuses the weight_g and weight_v parameters into a single weight parameter.
|
| 862 |
+
Should be called before inference for better performance.
|
| 863 |
+
Returns:
|
| 864 |
+
self: The model with weight_norm removed
|
| 865 |
+
"""
|
| 866 |
+
from torch.nn.utils import remove_weight_norm
|
| 867 |
+
for module in list(self.modules()):
|
| 868 |
+
if hasattr(module, "_forward_pre_hooks"):
|
| 869 |
+
for hook in list(module._forward_pre_hooks.values()):
|
| 870 |
+
if "WeightNorm" in str(type(hook)):
|
| 871 |
+
try:
|
| 872 |
+
remove_weight_norm(module)
|
| 873 |
+
except ValueError:
|
| 874 |
+
continue
|
| 875 |
+
if self.use_weight_norm:
|
| 876 |
+
self.use_weight_norm = False
|
| 877 |
+
self.register_to_config(use_weight_norm=False)
|
| 878 |
+
return self
|
videox_fun/models/creator_audio.py
ADDED
|
@@ -0,0 +1,353 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Creator audio diffusion transformer used by the release inference path."""
|
| 2 |
+
|
| 3 |
+
import glob
|
| 4 |
+
import json
|
| 5 |
+
import logging
|
| 6 |
+
import os
|
| 7 |
+
from typing import Optional
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
import torch.nn as nn
|
| 11 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 12 |
+
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
| 13 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 14 |
+
from einops import rearrange
|
| 15 |
+
|
| 16 |
+
from .creator.creator_video_dit import DiTBlock
|
| 17 |
+
from .creator.creator_audio_dit import (
|
| 18 |
+
Head,
|
| 19 |
+
MLP,
|
| 20 |
+
sinusoidal_embedding_1d,
|
| 21 |
+
precompute_freqs_cis_1d,
|
| 22 |
+
)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class CreatorAudioModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
| 26 |
+
"""Creator Audio DiT with wan_audio2-compatible forward interface."""
|
| 27 |
+
|
| 28 |
+
@register_to_config
|
| 29 |
+
def __init__(
|
| 30 |
+
self,
|
| 31 |
+
patch_size=(1,),
|
| 32 |
+
text_len=512,
|
| 33 |
+
in_dim=128,
|
| 34 |
+
dim=2048,
|
| 35 |
+
ffn_dim=8192,
|
| 36 |
+
freq_dim=256,
|
| 37 |
+
text_dim=4096,
|
| 38 |
+
out_dim=128,
|
| 39 |
+
num_heads=16,
|
| 40 |
+
num_layers=32,
|
| 41 |
+
eps=1e-6,
|
| 42 |
+
has_image_input=False,
|
| 43 |
+
has_image_pos_emb=False,
|
| 44 |
+
has_ref_conv=False,
|
| 45 |
+
vae_type="dac",
|
| 46 |
+
**kwargs
|
| 47 |
+
):
|
| 48 |
+
super().__init__()
|
| 49 |
+
|
| 50 |
+
self.patch_size = tuple(patch_size) if isinstance(patch_size, (list, tuple)) else (patch_size,)
|
| 51 |
+
self.text_len = text_len
|
| 52 |
+
self.in_dim = in_dim
|
| 53 |
+
self.dim = dim
|
| 54 |
+
self.ffn_dim = ffn_dim
|
| 55 |
+
self.freq_dim = freq_dim
|
| 56 |
+
self.text_dim = text_dim
|
| 57 |
+
self.out_dim = out_dim
|
| 58 |
+
self.num_heads = num_heads
|
| 59 |
+
self.num_layers = num_layers
|
| 60 |
+
self.eps = eps
|
| 61 |
+
self.has_image_input = has_image_input
|
| 62 |
+
self.vae_type = vae_type
|
| 63 |
+
|
| 64 |
+
# Patch embedding (1D conv)
|
| 65 |
+
self.patch_embedding = nn.Conv1d(
|
| 66 |
+
in_dim, dim, kernel_size=self.patch_size[0], stride=self.patch_size[0],
|
| 67 |
+
)
|
| 68 |
+
self.text_embedding = nn.Sequential(
|
| 69 |
+
nn.Linear(text_dim, dim),
|
| 70 |
+
nn.GELU(approximate="tanh"),
|
| 71 |
+
nn.Linear(dim, dim),
|
| 72 |
+
)
|
| 73 |
+
self.time_embedding = nn.Sequential(
|
| 74 |
+
nn.Linear(freq_dim, dim),
|
| 75 |
+
nn.SiLU(),
|
| 76 |
+
nn.Linear(dim, dim),
|
| 77 |
+
)
|
| 78 |
+
self.time_projection = nn.Sequential(
|
| 79 |
+
nn.SiLU(), nn.Linear(dim, dim * 6),
|
| 80 |
+
)
|
| 81 |
+
|
| 82 |
+
# DiTBlock stack
|
| 83 |
+
self.blocks = nn.ModuleList([
|
| 84 |
+
DiTBlock(has_image_input, dim, num_heads, ffn_dim, eps)
|
| 85 |
+
for _ in range(num_layers)
|
| 86 |
+
])
|
| 87 |
+
|
| 88 |
+
self.head = Head(dim, out_dim, self.patch_size, eps)
|
| 89 |
+
|
| 90 |
+
# RoPE precomputation
|
| 91 |
+
head_dim = dim // num_heads
|
| 92 |
+
self.freqs = precompute_freqs_cis_1d(head_dim)
|
| 93 |
+
|
| 94 |
+
# Optional image / reference embeddings
|
| 95 |
+
if has_image_input:
|
| 96 |
+
self.img_emb = MLP(1280, dim, has_pos_emb=has_image_pos_emb)
|
| 97 |
+
if has_ref_conv:
|
| 98 |
+
self.ref_conv = nn.Conv2d(16, dim, kernel_size=(2, 2), stride=(2, 2))
|
| 99 |
+
self.has_ref_conv = has_ref_conv
|
| 100 |
+
|
| 101 |
+
# -----------------------------------------------------------------
|
| 102 |
+
# Unpatchify
|
| 103 |
+
# -----------------------------------------------------------------
|
| 104 |
+
|
| 105 |
+
def unpatchify(self, x, grid_sizes, output_shapes):
|
| 106 |
+
"""Unpatchify tokens back to audio latent shape.
|
| 107 |
+
|
| 108 |
+
Args:
|
| 109 |
+
x: [B, seq_len, out_dim * patch_size]
|
| 110 |
+
grid_sizes: [B, 1] token lengths
|
| 111 |
+
output_shapes: list of original audio shapes [(C, T), ...]
|
| 112 |
+
|
| 113 |
+
Returns:
|
| 114 |
+
list of restored tensors matching output_shapes.
|
| 115 |
+
"""
|
| 116 |
+
output = []
|
| 117 |
+
for i, (grid_size, original_shape) in enumerate(zip(grid_sizes, output_shapes)):
|
| 118 |
+
f = int(grid_size[0].item())
|
| 119 |
+
restored = rearrange(
|
| 120 |
+
x[i, :f].unsqueeze(0),
|
| 121 |
+
"b f (p c) -> b c (f p)",
|
| 122 |
+
f=f, p=self.patch_size[0],
|
| 123 |
+
)
|
| 124 |
+
# Trim to original time length
|
| 125 |
+
orig_T = original_shape[-1]
|
| 126 |
+
restored = restored[:, :, :orig_T].squeeze(0)
|
| 127 |
+
output.append(restored)
|
| 128 |
+
return output
|
| 129 |
+
|
| 130 |
+
# -----------------------------------------------------------------
|
| 131 |
+
# RoPE frequencies
|
| 132 |
+
# -----------------------------------------------------------------
|
| 133 |
+
def _build_freqs(self, seq_len: int, device: torch.device) -> torch.Tensor:
|
| 134 |
+
"""Build RoPE complex frequencies [seq_len, 1, rope_dim]."""
|
| 135 |
+
freqs = torch.cat([
|
| 136 |
+
self.freqs[0][:seq_len].view(seq_len, -1),
|
| 137 |
+
self.freqs[1][:seq_len].view(seq_len, -1),
|
| 138 |
+
self.freqs[2][:seq_len].view(seq_len, -1),
|
| 139 |
+
], dim=-1).reshape(seq_len, 1, -1).to(device)
|
| 140 |
+
return freqs
|
| 141 |
+
|
| 142 |
+
# -----------------------------------------------------------------
|
| 143 |
+
# Weight initialisation
|
| 144 |
+
# -----------------------------------------------------------------
|
| 145 |
+
def init_weights(self):
|
| 146 |
+
for module in self.modules():
|
| 147 |
+
if isinstance(module, nn.Linear):
|
| 148 |
+
nn.init.xavier_uniform_(module.weight)
|
| 149 |
+
if module.bias is not None:
|
| 150 |
+
nn.init.zeros_(module.bias)
|
| 151 |
+
|
| 152 |
+
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
| 153 |
+
if self.patch_embedding.bias is not None:
|
| 154 |
+
nn.init.zeros_(self.patch_embedding.bias)
|
| 155 |
+
for module in self.text_embedding.modules():
|
| 156 |
+
if isinstance(module, nn.Linear):
|
| 157 |
+
nn.init.normal_(module.weight, std=0.02)
|
| 158 |
+
for module in self.time_embedding.modules():
|
| 159 |
+
if isinstance(module, nn.Linear):
|
| 160 |
+
nn.init.normal_(module.weight, std=0.02)
|
| 161 |
+
|
| 162 |
+
nn.init.zeros_(self.head.head.weight)
|
| 163 |
+
if self.head.head.bias is not None:
|
| 164 |
+
nn.init.zeros_(self.head.head.bias)
|
| 165 |
+
|
| 166 |
+
# -----------------------------------------------------------------
|
| 167 |
+
# forward (matches wan_audio2.WanAudioModel interface)
|
| 168 |
+
# -----------------------------------------------------------------
|
| 169 |
+
def forward(
|
| 170 |
+
self,
|
| 171 |
+
x,
|
| 172 |
+
t,
|
| 173 |
+
context,
|
| 174 |
+
seq_len,
|
| 175 |
+
clip_fea=None,
|
| 176 |
+
y=None,
|
| 177 |
+
dtype=torch.bfloat16,
|
| 178 |
+
**kwargs,
|
| 179 |
+
):
|
| 180 |
+
"""
|
| 181 |
+
Args:
|
| 182 |
+
x: list of [C, T] audio latent samples, or [B, C, T] tensor.
|
| 183 |
+
t: [B] or [B, T] timesteps.
|
| 184 |
+
context: list of [S, text_dim] or [B, S, text_dim] tensor.
|
| 185 |
+
seq_len: max sequence length after patching.
|
| 186 |
+
clip_fea: optional CLIP features for image conditioning.
|
| 187 |
+
y: optional conditioning latent (e.g. for i2a), same format as x.
|
| 188 |
+
"""
|
| 189 |
+
# --- Normalise inputs to lists ---
|
| 190 |
+
if isinstance(x, torch.Tensor):
|
| 191 |
+
x = [sample for sample in x]
|
| 192 |
+
if isinstance(context, torch.Tensor):
|
| 193 |
+
context = [sample for sample in context]
|
| 194 |
+
if y is not None and isinstance(y, torch.Tensor):
|
| 195 |
+
y = [sample for sample in y]
|
| 196 |
+
|
| 197 |
+
device = self.patch_embedding.weight.device
|
| 198 |
+
dtype = x[0].dtype if len(x) > 0 else self.patch_embedding.weight.dtype
|
| 199 |
+
batch_size = len(x)
|
| 200 |
+
|
| 201 |
+
# Concatenate conditioning y (e.g. for i2v/i2a)
|
| 202 |
+
if y is not None:
|
| 203 |
+
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
| 204 |
+
|
| 205 |
+
# Remember original shapes for unpatchify
|
| 206 |
+
original_audio_shapes = [tuple(sample.shape) for sample in x]
|
| 207 |
+
|
| 208 |
+
# --- Patchify per sample, then pad tokens ---
|
| 209 |
+
patchified_samples = []
|
| 210 |
+
grid_sizes_list = []
|
| 211 |
+
for sample in x:
|
| 212 |
+
tokens_i = self.patch_embedding(sample.unsqueeze(0).to(device)) # [1, dim, T']
|
| 213 |
+
tokens_i = rearrange(tokens_i, "1 c f -> f c").contiguous()
|
| 214 |
+
patchified_samples.append(tokens_i)
|
| 215 |
+
grid_sizes_list.append(tokens_i.shape[0])
|
| 216 |
+
|
| 217 |
+
grid_sizes = torch.tensor(
|
| 218 |
+
[[g] for g in grid_sizes_list], dtype=torch.long, device=device,
|
| 219 |
+
)
|
| 220 |
+
seq_lens = grid_sizes[:, 0]
|
| 221 |
+
assert seq_lens.max() <= seq_len, (
|
| 222 |
+
f"Max token length {seq_lens.max().item()} exceeds seq_len {seq_len}"
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
# Pad each sample's tokens to seq_len and stack
|
| 226 |
+
x = torch.stack([
|
| 227 |
+
torch.cat([u, u.new_zeros(seq_len - u.size(0), u.size(1))], dim=0)
|
| 228 |
+
for u in patchified_samples
|
| 229 |
+
])
|
| 230 |
+
|
| 231 |
+
# --- Time embedding ---
|
| 232 |
+
with torch.amp.autocast("cuda", dtype=torch.float32):
|
| 233 |
+
if t.dim() != 1:
|
| 234 |
+
# Per-token timesteps [B, T]
|
| 235 |
+
if t.size(1) < seq_len:
|
| 236 |
+
pad_size = seq_len - t.size(1)
|
| 237 |
+
last_elements = t[:, -1].unsqueeze(1)
|
| 238 |
+
t = torch.cat([t, last_elements.repeat(1, pad_size)], dim=1)
|
| 239 |
+
bt = t.size(0)
|
| 240 |
+
e = self.time_embedding(
|
| 241 |
+
sinusoidal_embedding_1d(self.freq_dim, t.flatten())
|
| 242 |
+
.unflatten(0, (bt, seq_len)).float()
|
| 243 |
+
)
|
| 244 |
+
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
| 245 |
+
else:
|
| 246 |
+
e = self.time_embedding(
|
| 247 |
+
sinusoidal_embedding_1d(self.freq_dim, t).float()
|
| 248 |
+
)
|
| 249 |
+
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
| 250 |
+
|
| 251 |
+
# --- Context embedding ---
|
| 252 |
+
context = self.text_embedding(
|
| 253 |
+
torch.stack([
|
| 254 |
+
torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
| 255 |
+
for u in context
|
| 256 |
+
])
|
| 257 |
+
)
|
| 258 |
+
|
| 259 |
+
# Image embedding
|
| 260 |
+
if self.has_image_input and clip_fea is not None:
|
| 261 |
+
clip_embedding = self.img_emb(clip_fea)
|
| 262 |
+
context = torch.cat([clip_embedding, context], dim=1)
|
| 263 |
+
|
| 264 |
+
# --- RoPE frequencies ---
|
| 265 |
+
freqs = self._build_freqs(seq_len, device)
|
| 266 |
+
|
| 267 |
+
for block in self.blocks:
|
| 268 |
+
x = block(x, context, e0, freqs, seq_lens=seq_lens)
|
| 269 |
+
|
| 270 |
+
x = self.head(x, e)
|
| 271 |
+
x = self.unpatchify(x, grid_sizes, original_audio_shapes)
|
| 272 |
+
return x
|
| 273 |
+
|
| 274 |
+
# -----------------------------------------------------------------
|
| 275 |
+
# from_pretrained
|
| 276 |
+
# -----------------------------------------------------------------
|
| 277 |
+
@classmethod
|
| 278 |
+
def from_pretrained(
|
| 279 |
+
cls,
|
| 280 |
+
pretrained_model_path,
|
| 281 |
+
subfolder=None,
|
| 282 |
+
transformer_additional_kwargs=None,
|
| 283 |
+
low_cpu_mem_usage=False,
|
| 284 |
+
in_dim=None,
|
| 285 |
+
out_dim=None,
|
| 286 |
+
patch_size=None,
|
| 287 |
+
torch_dtype=torch.bfloat16,
|
| 288 |
+
):
|
| 289 |
+
transformer_additional_kwargs = dict(transformer_additional_kwargs or {})
|
| 290 |
+
if subfolder is not None:
|
| 291 |
+
pretrained_model_path = os.path.join(pretrained_model_path, subfolder)
|
| 292 |
+
|
| 293 |
+
config_file = os.path.join(pretrained_model_path, "config.json")
|
| 294 |
+
if not os.path.isfile(config_file):
|
| 295 |
+
raise RuntimeError(f"{config_file} does not exist")
|
| 296 |
+
|
| 297 |
+
with open(config_file, "r") as fp:
|
| 298 |
+
config = json.load(fp)
|
| 299 |
+
|
| 300 |
+
from diffusers.utils import WEIGHTS_NAME
|
| 301 |
+
|
| 302 |
+
model_file = os.path.join(pretrained_model_path, WEIGHTS_NAME)
|
| 303 |
+
model_file_safetensors = model_file.replace(".bin", ".safetensors")
|
| 304 |
+
|
| 305 |
+
# Explicit architecture overrides
|
| 306 |
+
if in_dim is not None:
|
| 307 |
+
transformer_additional_kwargs["in_dim"] = in_dim
|
| 308 |
+
if out_dim is not None:
|
| 309 |
+
transformer_additional_kwargs["out_dim"] = out_dim
|
| 310 |
+
if patch_size is not None:
|
| 311 |
+
transformer_additional_kwargs["patch_size"] = patch_size
|
| 312 |
+
|
| 313 |
+
if "dict_mapping" in transformer_additional_kwargs:
|
| 314 |
+
for key, value in transformer_additional_kwargs["dict_mapping"].items():
|
| 315 |
+
if value not in transformer_additional_kwargs:
|
| 316 |
+
transformer_additional_kwargs[value] = config[key]
|
| 317 |
+
|
| 318 |
+
# Merge overrides into config
|
| 319 |
+
model_config = dict(config)
|
| 320 |
+
model_config.update(transformer_additional_kwargs)
|
| 321 |
+
|
| 322 |
+
model = cls.from_config(model_config, **transformer_additional_kwargs)
|
| 323 |
+
|
| 324 |
+
# --- Load state dict ---
|
| 325 |
+
if os.path.exists(model_file):
|
| 326 |
+
state_dict = torch.load(model_file, map_location="cpu")
|
| 327 |
+
elif os.path.exists(model_file_safetensors):
|
| 328 |
+
from safetensors.torch import load_file
|
| 329 |
+
state_dict = load_file(model_file_safetensors)
|
| 330 |
+
else:
|
| 331 |
+
from safetensors.torch import load_file
|
| 332 |
+
state_dict = {}
|
| 333 |
+
for shard in glob.glob(os.path.join(pretrained_model_path, "*.safetensors")):
|
| 334 |
+
state_dict.update(load_file(shard))
|
| 335 |
+
|
| 336 |
+
# Filter by shape match
|
| 337 |
+
model_sd = model.state_dict()
|
| 338 |
+
filtered_state_dict = {}
|
| 339 |
+
for key, value in state_dict.items():
|
| 340 |
+
if key in model_sd and model_sd[key].shape == value.shape:
|
| 341 |
+
filtered_state_dict[key] = value
|
| 342 |
+
else:
|
| 343 |
+
logging.info("Skipping key %s due to size mismatch or absence.", key)
|
| 344 |
+
|
| 345 |
+
missing, unexpected = model.load_state_dict(filtered_state_dict, strict=False)
|
| 346 |
+
logging.info(
|
| 347 |
+
"CreatorAudioModel missing keys: %d, unexpected keys: %d",
|
| 348 |
+
len(missing), len(unexpected),
|
| 349 |
+
)
|
| 350 |
+
return model.to(torch_dtype)
|
| 351 |
+
|
| 352 |
+
# Convenience aliases
|
| 353 |
+
CreatorAudioTransformerModel = CreatorAudioModel
|
videox_fun/models/creator_dac_vae.py
ADDED
|
@@ -0,0 +1,151 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from typing import Optional
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
from videox_fun.models.creator.dac_vae import DAC, DiagonalGaussianDistribution
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class CreatorDACVAE(torch.nn.Module):
|
| 11 |
+
"""
|
| 12 |
+
High-level wrapper around the DAC (Descript Audio Codec) VAE in continuous mode.
|
| 13 |
+
Mirrors the LTXAudioVAE interface used by the inference pipeline.
|
| 14 |
+
"""
|
| 15 |
+
|
| 16 |
+
def __init__(self, dac_model: DAC) -> None:
|
| 17 |
+
super().__init__()
|
| 18 |
+
self.dac = dac_model
|
| 19 |
+
|
| 20 |
+
@property
|
| 21 |
+
def sample_rate(self) -> int:
|
| 22 |
+
return self.dac.sample_rate
|
| 23 |
+
|
| 24 |
+
@property
|
| 25 |
+
def hop_length(self) -> int:
|
| 26 |
+
return self.dac.hop_length
|
| 27 |
+
|
| 28 |
+
@property
|
| 29 |
+
def latent_dim(self) -> int:
|
| 30 |
+
return self.dac.latent_dim
|
| 31 |
+
|
| 32 |
+
@classmethod
|
| 33 |
+
def from_pretrained(
|
| 34 |
+
cls,
|
| 35 |
+
pretrained_model_path: str | os.PathLike[str],
|
| 36 |
+
strict: bool = False,
|
| 37 |
+
) -> "CreatorDACVAE":
|
| 38 |
+
pretrained_model_path = Path(pretrained_model_path)
|
| 39 |
+
|
| 40 |
+
if pretrained_model_path.is_dir():
|
| 41 |
+
dac_model = DAC.from_pretrained(pretrained_model_path)
|
| 42 |
+
else:
|
| 43 |
+
dac_model = DAC.from_pretrained(pretrained_model_path.parent)
|
| 44 |
+
|
| 45 |
+
return cls(dac_model=dac_model)
|
| 46 |
+
|
| 47 |
+
def _preprocess_waveform(self, waveform: torch.Tensor) -> torch.Tensor:
|
| 48 |
+
"""Normalize waveform to [B, 1, T] mono and pad to hop_length boundary."""
|
| 49 |
+
"""The Creator audio VAE currently supports mono audio."""
|
| 50 |
+
if waveform.ndim == 1:
|
| 51 |
+
waveform = waveform.unsqueeze(0).unsqueeze(0)
|
| 52 |
+
elif waveform.ndim == 2:
|
| 53 |
+
if waveform.size(0) > 1:
|
| 54 |
+
waveform = waveform.mean(dim=0, keepdim=True)
|
| 55 |
+
waveform = waveform.unsqueeze(0)
|
| 56 |
+
elif waveform.ndim == 3:
|
| 57 |
+
if waveform.size(1) > 1:
|
| 58 |
+
waveform = waveform.mean(dim=1, keepdim=True)
|
| 59 |
+
|
| 60 |
+
waveform = self.dac.preprocess(waveform, self.sample_rate)
|
| 61 |
+
return waveform
|
| 62 |
+
|
| 63 |
+
@torch.inference_mode()
|
| 64 |
+
def encode_posterior(
|
| 65 |
+
self,
|
| 66 |
+
audio: torch.Tensor,
|
| 67 |
+
sampling_rate: Optional[int] = None,
|
| 68 |
+
deterministic: bool | None = None,
|
| 69 |
+
) -> DiagonalGaussianDistribution:
|
| 70 |
+
"""Encode audio waveform and return the posterior distribution.
|
| 71 |
+
|
| 72 |
+
Parameters
|
| 73 |
+
----------
|
| 74 |
+
audio : Tensor
|
| 75 |
+
Raw waveform tensor. Accepts shapes [T], [C, T], or [B, C, T].
|
| 76 |
+
sampling_rate : int, optional
|
| 77 |
+
Not used directly; kept for API compatibility with LTXAudioVAE.
|
| 78 |
+
deterministic : bool, optional
|
| 79 |
+
If True, std is zeroed so sampling returns the mean.
|
| 80 |
+
|
| 81 |
+
Returns
|
| 82 |
+
-------
|
| 83 |
+
DiagonalGaussianDistribution
|
| 84 |
+
"""
|
| 85 |
+
waveform = self._preprocess_waveform(audio)
|
| 86 |
+
posterior, _, _, _, _ = self.dac.encode(waveform)
|
| 87 |
+
if deterministic is not None:
|
| 88 |
+
posterior.deterministic = deterministic
|
| 89 |
+
if deterministic:
|
| 90 |
+
posterior.std = torch.zeros_like(posterior.mean)
|
| 91 |
+
posterior.var = torch.zeros_like(posterior.mean)
|
| 92 |
+
return posterior
|
| 93 |
+
|
| 94 |
+
@torch.inference_mode()
|
| 95 |
+
def encode(
|
| 96 |
+
self,
|
| 97 |
+
audio: torch.Tensor,
|
| 98 |
+
sampling_rate: Optional[int] = None,
|
| 99 |
+
sample: bool = False,
|
| 100 |
+
generator: Optional[torch.Generator] = None,
|
| 101 |
+
) -> torch.Tensor:
|
| 102 |
+
"""Encode audio waveform to latent representation.
|
| 103 |
+
|
| 104 |
+
Parameters
|
| 105 |
+
----------
|
| 106 |
+
audio : Tensor
|
| 107 |
+
Raw waveform tensor.
|
| 108 |
+
sampling_rate : int, optional
|
| 109 |
+
Kept for API compatibility.
|
| 110 |
+
sample : bool
|
| 111 |
+
If True, sample from the posterior; otherwise return the mean.
|
| 112 |
+
generator : torch.Generator, optional
|
| 113 |
+
RNG for reproducible sampling.
|
| 114 |
+
|
| 115 |
+
Returns
|
| 116 |
+
-------
|
| 117 |
+
Tensor [B, D, T']
|
| 118 |
+
Continuous latent codes.
|
| 119 |
+
"""
|
| 120 |
+
posterior = self.encode_posterior(audio, sampling_rate=sampling_rate)
|
| 121 |
+
if sample:
|
| 122 |
+
return posterior.sample()
|
| 123 |
+
return posterior.mode()
|
| 124 |
+
|
| 125 |
+
@torch.inference_mode()
|
| 126 |
+
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
| 127 |
+
"""Decode latent codes back to waveform.
|
| 128 |
+
|
| 129 |
+
Parameters
|
| 130 |
+
----------
|
| 131 |
+
latent : Tensor [B, D, T']
|
| 132 |
+
Continuous latent codes.
|
| 133 |
+
|
| 134 |
+
Returns
|
| 135 |
+
-------
|
| 136 |
+
Tensor [B, 1, T]
|
| 137 |
+
Reconstructed waveform.
|
| 138 |
+
"""
|
| 139 |
+
return self.dac.decode(latent)
|
| 140 |
+
|
| 141 |
+
@torch.inference_mode()
|
| 142 |
+
def reconstruct(
|
| 143 |
+
self,
|
| 144 |
+
audio: torch.Tensor,
|
| 145 |
+
sampling_rate: Optional[int] = None,
|
| 146 |
+
sample: bool = False,
|
| 147 |
+
generator: Optional[torch.Generator] = None,
|
| 148 |
+
) -> torch.Tensor:
|
| 149 |
+
"""Encode then decode (round-trip reconstruction)."""
|
| 150 |
+
latent = self.encode(audio, sampling_rate=sampling_rate, sample=sample, generator=generator)
|
| 151 |
+
return self.decode(latent)
|
videox_fun/models/creator_gating.py
ADDED
|
@@ -0,0 +1,1286 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configurable gated cross-attention joint audio-video model.
|
| 2 |
+
|
| 3 |
+
This experimental variant follows the AV cross-attention design but allows
|
| 4 |
+
audio-to-video (A2V) and video-to-audio (V2A) cross attention to be enabled
|
| 5 |
+
independently for each layer. Cross-modal attention can optionally apply a
|
| 6 |
+
fixed per-layer alpha times a per-head sigmoid gate to the attention context
|
| 7 |
+
before the output projection.
|
| 8 |
+
"""
|
| 9 |
+
|
| 10 |
+
import logging
|
| 11 |
+
import math
|
| 12 |
+
import os
|
| 13 |
+
from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple, Union
|
| 14 |
+
|
| 15 |
+
import torch
|
| 16 |
+
import torch.nn as nn
|
| 17 |
+
from einops import rearrange
|
| 18 |
+
|
| 19 |
+
from .attention_utils import attention
|
| 20 |
+
from .creator_audio import CreatorAudioModel
|
| 21 |
+
from .wan_transformer3d_prope import (
|
| 22 |
+
Wan2_2Transformer3DModel,
|
| 23 |
+
WanRMSNorm,
|
| 24 |
+
WanTransformer3DModel,
|
| 25 |
+
rope_apply_qk,
|
| 26 |
+
)
|
| 27 |
+
from .creator.creator_video_dit import sinusoidal_embedding_1d
|
| 28 |
+
from .creator.creator_video_dit import rope_apply_head_dim
|
| 29 |
+
from ..dist.sequence_parallel import all_gather_sequence, ulysses_attention
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
LayerSelection = Optional[Union[bool, str, Sequence[bool], Sequence[int], torch.Tensor]]
|
| 33 |
+
LayerAlphas = Optional[
|
| 34 |
+
Union[float, Sequence[float], Mapping[Union[int, str], float], torch.Tensor]
|
| 35 |
+
]
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
@torch.amp.autocast("cuda", enabled=False)
|
| 39 |
+
def temporal_rope_1d(
|
| 40 |
+
x: torch.Tensor,
|
| 41 |
+
temporal_positions: torch.Tensor,
|
| 42 |
+
inv_freqs_1d: torch.Tensor,
|
| 43 |
+
) -> torch.Tensor:
|
| 44 |
+
"""Apply 1D temporal RoPE to [B, L, num_heads, head_dim] tensors."""
|
| 45 |
+
dtype = x.dtype
|
| 46 |
+
batch_size, seq_len, num_heads, head_dim = x.shape
|
| 47 |
+
half_dim = head_dim // 2
|
| 48 |
+
|
| 49 |
+
freqs = torch.einsum(
|
| 50 |
+
"bl,d->bld",
|
| 51 |
+
temporal_positions.to(torch.float64),
|
| 52 |
+
inv_freqs_1d.to(x.device, torch.float64),
|
| 53 |
+
)
|
| 54 |
+
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
|
| 55 |
+
|
| 56 |
+
x_complex = torch.view_as_complex(
|
| 57 |
+
x.to(torch.float64).reshape(batch_size, seq_len, num_heads, half_dim, 2)
|
| 58 |
+
)
|
| 59 |
+
x_out = torch.view_as_real(x_complex * freqs_cis.unsqueeze(2)).flatten(3)
|
| 60 |
+
return x_out.to(dtype)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def compute_video_temporal_positions(
|
| 64 |
+
grid_sizes: torch.Tensor,
|
| 65 |
+
seq_len: int,
|
| 66 |
+
device: torch.device,
|
| 67 |
+
audio_fps: float = 48000.0 / 960.0,
|
| 68 |
+
video_fps: float = 16.0,
|
| 69 |
+
vae_temporal_stride: int = 4,
|
| 70 |
+
) -> torch.Tensor:
|
| 71 |
+
"""Compute video token temporal positions in audio-token time units."""
|
| 72 |
+
batch_size = grid_sizes.size(0)
|
| 73 |
+
positions = torch.zeros(batch_size, seq_len, device=device, dtype=torch.float64)
|
| 74 |
+
video_latent_fps = video_fps / vae_temporal_stride
|
| 75 |
+
scale = audio_fps / video_latent_fps
|
| 76 |
+
|
| 77 |
+
for sample_idx, (num_frames, height, width) in enumerate(grid_sizes.tolist()):
|
| 78 |
+
spatial_size = int(height * width)
|
| 79 |
+
num_tokens = int(num_frames * spatial_size)
|
| 80 |
+
frame_indices = torch.arange(num_tokens, device=device, dtype=torch.float64) // spatial_size
|
| 81 |
+
positions[sample_idx, :num_tokens] = frame_indices * scale
|
| 82 |
+
return positions
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def compute_audio_temporal_positions(
|
| 86 |
+
seq_lens: torch.Tensor,
|
| 87 |
+
seq_len: int,
|
| 88 |
+
device: torch.device,
|
| 89 |
+
) -> torch.Tensor:
|
| 90 |
+
"""Compute sequential temporal positions for audio tokens."""
|
| 91 |
+
batch_size = seq_lens.size(0)
|
| 92 |
+
positions = torch.zeros(batch_size, seq_len, device=device, dtype=torch.float64)
|
| 93 |
+
base_positions = torch.arange(seq_len, device=device, dtype=torch.float64)
|
| 94 |
+
for sample_idx in range(batch_size):
|
| 95 |
+
valid_len = int(seq_lens[sample_idx].item())
|
| 96 |
+
positions[sample_idx, :valid_len] = base_positions[:valid_len]
|
| 97 |
+
return positions
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def _apply_video_rope_local(
|
| 101 |
+
x: torch.Tensor,
|
| 102 |
+
grid_sizes: torch.Tensor,
|
| 103 |
+
freqs: torch.Tensor,
|
| 104 |
+
sp_rank: int,
|
| 105 |
+
sp_world_size: int,
|
| 106 |
+
) -> torch.Tensor:
|
| 107 |
+
"""Apply the Wan 3D RoPE slice belonging to this sequence-parallel rank."""
|
| 108 |
+
if sp_world_size <= 1:
|
| 109 |
+
return rope_apply_qk(x, x, grid_sizes, freqs)[0]
|
| 110 |
+
|
| 111 |
+
local_len, num_heads, complex_dim = x.size(1), x.size(2), x.size(3) // 2
|
| 112 |
+
freq_parts = freqs.split(
|
| 113 |
+
[complex_dim - 2 * (complex_dim // 3), complex_dim // 3, complex_dim // 3],
|
| 114 |
+
dim=1,
|
| 115 |
+
)
|
| 116 |
+
output = []
|
| 117 |
+
for sample_idx, (frames, height, width) in enumerate(grid_sizes.tolist()):
|
| 118 |
+
full_len = int(frames * height * width)
|
| 119 |
+
sample = x[sample_idx, :local_len].to(torch.float64)
|
| 120 |
+
sample_complex = torch.view_as_complex(
|
| 121 |
+
sample.reshape(local_len, num_heads, -1, 2)
|
| 122 |
+
)
|
| 123 |
+
full_freqs = torch.cat(
|
| 124 |
+
[
|
| 125 |
+
freq_parts[0][:frames].view(frames, 1, 1, -1).expand(frames, height, width, -1),
|
| 126 |
+
freq_parts[1][:height].view(1, height, 1, -1).expand(frames, height, width, -1),
|
| 127 |
+
freq_parts[2][:width].view(1, 1, width, -1).expand(frames, height, width, -1),
|
| 128 |
+
],
|
| 129 |
+
dim=-1,
|
| 130 |
+
).reshape(full_len, 1, -1)
|
| 131 |
+
if full_freqs.size(0) < local_len * sp_world_size:
|
| 132 |
+
full_freqs = torch.cat(
|
| 133 |
+
[
|
| 134 |
+
full_freqs,
|
| 135 |
+
torch.ones(
|
| 136 |
+
local_len * sp_world_size - full_freqs.size(0),
|
| 137 |
+
full_freqs.size(1),
|
| 138 |
+
full_freqs.size(2),
|
| 139 |
+
dtype=full_freqs.dtype,
|
| 140 |
+
device=full_freqs.device,
|
| 141 |
+
),
|
| 142 |
+
],
|
| 143 |
+
dim=0,
|
| 144 |
+
)
|
| 145 |
+
start = sp_rank * local_len
|
| 146 |
+
local_freqs = full_freqs[start : start + local_len]
|
| 147 |
+
rotated = torch.view_as_real(sample_complex * local_freqs).flatten(2)
|
| 148 |
+
if x.size(1) > local_len:
|
| 149 |
+
rotated = torch.cat([rotated, x[sample_idx, local_len:]], dim=0)
|
| 150 |
+
output.append(rotated)
|
| 151 |
+
return torch.stack(output).to(x.dtype)
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
def _to_list(value) -> List[torch.Tensor]:
|
| 155 |
+
if isinstance(value, torch.Tensor):
|
| 156 |
+
return [sample for sample in value]
|
| 157 |
+
return list(value)
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
def _build_time_embeddings(
|
| 161 |
+
time_embedding: nn.Module,
|
| 162 |
+
time_projection: nn.Module,
|
| 163 |
+
freq_dim: int,
|
| 164 |
+
dim: int,
|
| 165 |
+
timesteps: torch.Tensor,
|
| 166 |
+
seq_len: int,
|
| 167 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 168 |
+
"""Build time embeddings and projected modulations."""
|
| 169 |
+
|
| 170 |
+
with torch.amp.autocast("cuda", dtype=torch.float32):
|
| 171 |
+
if timesteps.dim() != 1:
|
| 172 |
+
if timesteps.size(1) < seq_len:
|
| 173 |
+
pad_size = seq_len - timesteps.size(1)
|
| 174 |
+
timesteps = torch.cat(
|
| 175 |
+
[timesteps, timesteps[:, -1:].repeat(1, pad_size)], dim=1
|
| 176 |
+
)
|
| 177 |
+
batch_size = timesteps.size(0)
|
| 178 |
+
embedding = time_embedding(
|
| 179 |
+
sinusoidal_embedding_1d(freq_dim, timesteps.flatten())
|
| 180 |
+
.unflatten(0, (batch_size, seq_len))
|
| 181 |
+
.float()
|
| 182 |
+
)
|
| 183 |
+
modulation = time_projection(embedding).unflatten(2, (6, dim))
|
| 184 |
+
else:
|
| 185 |
+
embedding = time_embedding(sinusoidal_embedding_1d(freq_dim, timesteps).float())
|
| 186 |
+
modulation = time_projection(embedding).unflatten(1, (6, dim))
|
| 187 |
+
return embedding, modulation
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def _embed_context(text_embedding: nn.Module, context, text_len: int) -> torch.Tensor:
|
| 191 |
+
"""Embed and right-pad text context to the model fixed text length."""
|
| 192 |
+
if isinstance(context, torch.Tensor):
|
| 193 |
+
samples = [sample for sample in context]
|
| 194 |
+
else:
|
| 195 |
+
samples = list(context)
|
| 196 |
+
return text_embedding(
|
| 197 |
+
torch.stack([
|
| 198 |
+
torch.cat([sample, sample.new_zeros(text_len - sample.size(0), sample.size(1))])
|
| 199 |
+
for sample in samples
|
| 200 |
+
])
|
| 201 |
+
)
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def _logit_from_gate_value(gate_init_value: Optional[float]) -> float:
|
| 205 |
+
"""Convert an initial gate value in [0, 1] to a sigmoid bias."""
|
| 206 |
+
if gate_init_value is None:
|
| 207 |
+
return 0.0
|
| 208 |
+
gate_value = min(max(float(gate_init_value), 1e-6), 1.0 - 1e-6)
|
| 209 |
+
return math.log(gate_value / (1.0 - gate_value))
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
def _expand_layer_selection(
|
| 213 |
+
selection: LayerSelection,
|
| 214 |
+
num_layers: int,
|
| 215 |
+
name: str,
|
| 216 |
+
) -> List[bool]:
|
| 217 |
+
"""Expand a layer-selection config to a bool mask of length ``num_layers``.
|
| 218 |
+
|
| 219 |
+
Accepted forms:
|
| 220 |
+
- ``None`` or ``True``: enable every layer.
|
| 221 |
+
- ``False``: disable every layer.
|
| 222 |
+
- bool mask with length ``num_layers``.
|
| 223 |
+
- 0/1 mask with length ``num_layers``.
|
| 224 |
+
- list/tuple/tensor of layer indices to enable.
|
| 225 |
+
- strings: ``"all"``, ``"none"``, ``"0,2,5"``.
|
| 226 |
+
"""
|
| 227 |
+
if selection is None:
|
| 228 |
+
return [False] * num_layers
|
| 229 |
+
if isinstance(selection, bool):
|
| 230 |
+
return [selection] * num_layers
|
| 231 |
+
if isinstance(selection, torch.Tensor):
|
| 232 |
+
selection = selection.cpu().tolist()
|
| 233 |
+
if isinstance(selection, str):
|
| 234 |
+
normalized = selection.strip().lower()
|
| 235 |
+
if normalized in {"", "none", "false", "off", "0"}:
|
| 236 |
+
return [False] * num_layers
|
| 237 |
+
if normalized in {"all", "true", "on", "1"}:
|
| 238 |
+
return [True] * num_layers
|
| 239 |
+
indices = [int(part.strip()) for part in selection.split(",") if part.strip()]
|
| 240 |
+
mask = [False] * num_layers
|
| 241 |
+
for layer_idx in indices:
|
| 242 |
+
if layer_idx < 0 or layer_idx >= num_layers:
|
| 243 |
+
raise ValueError(f"{name} layer index {layer_idx} out of range [0, {num_layers})")
|
| 244 |
+
mask[layer_idx] = True
|
| 245 |
+
return mask
|
| 246 |
+
|
| 247 |
+
values = list(selection)
|
| 248 |
+
if not values:
|
| 249 |
+
return [False] * num_layers
|
| 250 |
+
|
| 251 |
+
if all(isinstance(value, bool) for value in values):
|
| 252 |
+
if len(values) != num_layers:
|
| 253 |
+
raise ValueError(f"{name} bool mask must have length {num_layers}, got {len(values)}")
|
| 254 |
+
return [bool(value) for value in values]
|
| 255 |
+
|
| 256 |
+
if all(isinstance(value, int) for value in values):
|
| 257 |
+
if len(values) == num_layers and all(int(value) in {0, 1} for value in values):
|
| 258 |
+
return [bool(value) for value in values]
|
| 259 |
+
mask = [False] * num_layers
|
| 260 |
+
for layer_idx in values:
|
| 261 |
+
if layer_idx < 0 or layer_idx >= num_layers:
|
| 262 |
+
raise ValueError(f"{name} layer index {layer_idx} out of range [0, {num_layers})")
|
| 263 |
+
mask[int(layer_idx)] = True
|
| 264 |
+
return mask
|
| 265 |
+
|
| 266 |
+
raise TypeError(
|
| 267 |
+
f"{name} must be None, bool, string, bool mask, 0/1 mask, or layer-index sequence"
|
| 268 |
+
)
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
def _expand_layer_alphas(
|
| 272 |
+
values: LayerAlphas,
|
| 273 |
+
num_layers: int,
|
| 274 |
+
name: str,
|
| 275 |
+
) -> List[float]:
|
| 276 |
+
"""Expand fixed per-layer gate multipliers, defaulting each layer to 1.0.
|
| 277 |
+
|
| 278 |
+
Accepted forms are a scalar shared by all layers, a full sequence with
|
| 279 |
+
``num_layers`` entries, or a mapping of layer index to alpha. Unspecified
|
| 280 |
+
mapping entries retain the default value 1.0.
|
| 281 |
+
"""
|
| 282 |
+
if values is None:
|
| 283 |
+
return [1.0] * num_layers
|
| 284 |
+
if isinstance(values, torch.Tensor):
|
| 285 |
+
values = values.cpu().tolist()
|
| 286 |
+
|
| 287 |
+
def validate(value, layer_label: str) -> float:
|
| 288 |
+
alpha = float(value)
|
| 289 |
+
if not math.isfinite(alpha) or alpha < 0.0:
|
| 290 |
+
raise ValueError(f"{name} {layer_label} must be finite and non-negative, got {value}")
|
| 291 |
+
return alpha
|
| 292 |
+
|
| 293 |
+
if isinstance(values, (int, float)) and not isinstance(values, bool):
|
| 294 |
+
return [validate(values, "scalar")] * num_layers
|
| 295 |
+
|
| 296 |
+
if isinstance(values, Mapping):
|
| 297 |
+
alphas = [1.0] * num_layers
|
| 298 |
+
for raw_layer_idx, value in values.items():
|
| 299 |
+
try:
|
| 300 |
+
layer_idx = int(raw_layer_idx)
|
| 301 |
+
except (TypeError, ValueError) as exc:
|
| 302 |
+
raise ValueError(
|
| 303 |
+
f"{name} mapping key must be a layer index, got {raw_layer_idx!r}"
|
| 304 |
+
) from exc
|
| 305 |
+
if layer_idx < 0 or layer_idx >= num_layers:
|
| 306 |
+
raise ValueError(f"{name} layer index {layer_idx} out of range [0, {num_layers})")
|
| 307 |
+
alphas[layer_idx] = validate(value, f"layer {layer_idx}")
|
| 308 |
+
return alphas
|
| 309 |
+
|
| 310 |
+
if isinstance(values, Sequence) and not isinstance(values, (str, bytes)):
|
| 311 |
+
if len(values) != num_layers:
|
| 312 |
+
raise ValueError(f"{name} sequence must have length {num_layers}, got {len(values)}")
|
| 313 |
+
return [validate(value, f"layer {layer_idx}") for layer_idx, value in enumerate(values)]
|
| 314 |
+
|
| 315 |
+
raise TypeError(f"{name} must be None, a scalar, a mapping, or a full layer sequence")
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
class GatedCrossModalAttention(nn.Module):
|
| 319 |
+
"""Cross-modal attention with optional alpha-scaled sigmoid context gating.
|
| 320 |
+
|
| 321 |
+
For A2V, ``x`` is video hidden states and ``y`` is audio hidden states.
|
| 322 |
+
For V2A, ``x`` is audio hidden states and ``y`` is video hidden states.
|
| 323 |
+
When ``use_gating=False``, this module follows the reference cross-attn
|
| 324 |
+
behavior: query from raw ``x`` and key/value from normalized ``y``.
|
| 325 |
+
"""
|
| 326 |
+
|
| 327 |
+
def __init__(
|
| 328 |
+
self,
|
| 329 |
+
q_dim: int,
|
| 330 |
+
kv_dim: int,
|
| 331 |
+
num_heads: int,
|
| 332 |
+
eps: float = 1e-6,
|
| 333 |
+
zero_init_output: bool = False,
|
| 334 |
+
use_gating: bool = True,
|
| 335 |
+
zero_init_gating: bool = False,
|
| 336 |
+
gate_init_value: Optional[float] = None,
|
| 337 |
+
gate_alpha: float = 1.0,
|
| 338 |
+
):
|
| 339 |
+
super().__init__()
|
| 340 |
+
assert q_dim % num_heads == 0
|
| 341 |
+
self.q_dim = q_dim
|
| 342 |
+
self.kv_dim = kv_dim
|
| 343 |
+
self.num_heads = num_heads
|
| 344 |
+
self.head_dim = q_dim // num_heads
|
| 345 |
+
self.use_gating = bool(use_gating)
|
| 346 |
+
self.gate_alpha = _expand_layer_alphas(gate_alpha, 1, "gate_alpha")[0]
|
| 347 |
+
|
| 348 |
+
self.norm = nn.LayerNorm(kv_dim, eps=eps)
|
| 349 |
+
# if self.use_gating:
|
| 350 |
+
self.norm_x = nn.LayerNorm(q_dim, eps=eps)
|
| 351 |
+
|
| 352 |
+
self.q = nn.Linear(q_dim, q_dim)
|
| 353 |
+
self.k = nn.Linear(kv_dim, q_dim)
|
| 354 |
+
self.v = nn.Linear(kv_dim, q_dim)
|
| 355 |
+
self.o = nn.Linear(q_dim, q_dim)
|
| 356 |
+
|
| 357 |
+
self.norm_q = WanRMSNorm(q_dim, eps=eps)
|
| 358 |
+
self.norm_k = WanRMSNorm(q_dim, eps=eps)
|
| 359 |
+
|
| 360 |
+
if self.use_gating:
|
| 361 |
+
self.gate_hidden = nn.Linear(q_dim, num_heads, bias=False)
|
| 362 |
+
self.gate_context_norm = nn.LayerNorm(self.head_dim, eps=eps)
|
| 363 |
+
self.gate_context = nn.Linear(self.head_dim, 1, bias=False)
|
| 364 |
+
self.gate_bias = nn.Parameter(torch.empty(num_heads))
|
| 365 |
+
|
| 366 |
+
self._init_weights(
|
| 367 |
+
zero_init_output=zero_init_output,
|
| 368 |
+
zero_init_gating=zero_init_gating,
|
| 369 |
+
gate_init_value=gate_init_value,
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
def _init_weights(
|
| 373 |
+
self,
|
| 374 |
+
zero_init_output: bool,
|
| 375 |
+
zero_init_gating: bool,
|
| 376 |
+
gate_init_value: Optional[float],
|
| 377 |
+
):
|
| 378 |
+
for module in [self.q, self.k, self.v, self.o]:
|
| 379 |
+
nn.init.xavier_uniform_(module.weight)
|
| 380 |
+
if module.bias is not None:
|
| 381 |
+
nn.init.zeros_(module.bias)
|
| 382 |
+
|
| 383 |
+
if zero_init_output:
|
| 384 |
+
nn.init.zeros_(self.o.weight)
|
| 385 |
+
nn.init.zeros_(self.o.bias)
|
| 386 |
+
|
| 387 |
+
if self.use_gating:
|
| 388 |
+
if zero_init_gating:
|
| 389 |
+
nn.init.zeros_(self.gate_hidden.weight)
|
| 390 |
+
nn.init.zeros_(self.gate_context.weight)
|
| 391 |
+
else:
|
| 392 |
+
nn.init.xavier_uniform_(self.gate_hidden.weight)
|
| 393 |
+
nn.init.xavier_uniform_(self.gate_context.weight)
|
| 394 |
+
nn.init.constant_(self.gate_bias, _logit_from_gate_value(gate_init_value))
|
| 395 |
+
|
| 396 |
+
def forward(
|
| 397 |
+
self,
|
| 398 |
+
x: torch.Tensor,
|
| 399 |
+
y: torch.Tensor,
|
| 400 |
+
y_lens: Optional[torch.Tensor] = None,
|
| 401 |
+
dtype: torch.dtype = torch.bfloat16,
|
| 402 |
+
q_temporal_pos: Optional[torch.Tensor] = None,
|
| 403 |
+
k_temporal_pos: Optional[torch.Tensor] = None,
|
| 404 |
+
temporal_rope_inv_freq: Optional[torch.Tensor] = None,
|
| 405 |
+
sequence_parallel: bool = False,
|
| 406 |
+
sp_group=None,
|
| 407 |
+
) -> torch.Tensor:
|
| 408 |
+
"""
|
| 409 |
+
Args:
|
| 410 |
+
x: primary hidden states [B, Lq, q_dim]
|
| 411 |
+
y: conditioning hidden states [B, Lk, kv_dim]
|
| 412 |
+
y_lens: valid lengths of y per sample [B]
|
| 413 |
+
q_temporal_pos: [B, Lq] temporal positions for query
|
| 414 |
+
k_temporal_pos: [B, Lk] temporal positions for key
|
| 415 |
+
temporal_rope_inv_freq: [head_dim // 2] inv frequencies for temporal RoPE
|
| 416 |
+
Returns:
|
| 417 |
+
Cross-attention output [B, Lq, q_dim].
|
| 418 |
+
"""
|
| 419 |
+
batch_size = x.size(0)
|
| 420 |
+
num_heads = self.num_heads
|
| 421 |
+
head_dim = self.head_dim
|
| 422 |
+
|
| 423 |
+
# In SP mode x is a local query chunk while y is a local conditioning
|
| 424 |
+
# chunk. Cross-modal attention needs the complete conditioning sequence.
|
| 425 |
+
if sequence_parallel:
|
| 426 |
+
y = all_gather_sequence(y, group=sp_group)
|
| 427 |
+
if k_temporal_pos is not None:
|
| 428 |
+
k_temporal_pos = all_gather_sequence(
|
| 429 |
+
k_temporal_pos.unsqueeze(-1), group=sp_group
|
| 430 |
+
).squeeze(-1)
|
| 431 |
+
|
| 432 |
+
x_for_q = self.norm_x(x)
|
| 433 |
+
y_norm = self.norm(y)
|
| 434 |
+
|
| 435 |
+
query = self.norm_q(self.q(x_for_q.to(dtype))).view(batch_size, -1, num_heads, head_dim)
|
| 436 |
+
key = self.norm_k(self.k(y_norm.to(dtype))).view(batch_size, -1, num_heads, head_dim)
|
| 437 |
+
value = self.v(y_norm.to(dtype)).view(batch_size, -1, num_heads, head_dim)
|
| 438 |
+
|
| 439 |
+
if q_temporal_pos is not None and k_temporal_pos is not None and temporal_rope_inv_freq is not None:
|
| 440 |
+
query = temporal_rope_1d(query, q_temporal_pos, temporal_rope_inv_freq)
|
| 441 |
+
key = temporal_rope_1d(key, k_temporal_pos, temporal_rope_inv_freq)
|
| 442 |
+
|
| 443 |
+
context = attention(query.to(dtype), key.to(dtype), value.to(dtype), k_lens=y_lens)
|
| 444 |
+
context = context.to(dtype)
|
| 445 |
+
|
| 446 |
+
if self.use_gating:
|
| 447 |
+
hidden_gate = self.gate_hidden(x_for_q.to(dtype)).view(batch_size, -1, num_heads, 1)
|
| 448 |
+
context_gate = self.gate_context(
|
| 449 |
+
self.gate_context_norm(context)
|
| 450 |
+
)
|
| 451 |
+
gate = self.gate_alpha * torch.sigmoid(
|
| 452 |
+
hidden_gate + context_gate + self.gate_bias.view(1, 1, num_heads, 1)
|
| 453 |
+
)
|
| 454 |
+
context = gate * context
|
| 455 |
+
return self.o(context.flatten(2))
|
| 456 |
+
|
| 457 |
+
|
| 458 |
+
class GatedJointBlock(nn.Module):
|
| 459 |
+
"""One joint block with independently configurable A2V and V2A attention."""
|
| 460 |
+
|
| 461 |
+
def __init__(
|
| 462 |
+
self,
|
| 463 |
+
video_block: nn.Module,
|
| 464 |
+
audio_block: nn.Module,
|
| 465 |
+
video_dim: int,
|
| 466 |
+
audio_dim: int,
|
| 467 |
+
video_num_heads: int,
|
| 468 |
+
audio_num_heads: int,
|
| 469 |
+
enable_a2v_cross_attn: bool = True,
|
| 470 |
+
enable_v2a_cross_attn: bool = True,
|
| 471 |
+
zero_init_output: bool = False,
|
| 472 |
+
zero_init_video_cross_attn: bool | None = None,
|
| 473 |
+
zero_init_audio_cross_attn: bool | None = None,
|
| 474 |
+
use_a2v_gating: bool = True,
|
| 475 |
+
use_v2a_gating: bool = True,
|
| 476 |
+
zero_init_a2v_gating: bool = False,
|
| 477 |
+
zero_init_v2a_gating: bool = False,
|
| 478 |
+
a2v_gate_init_value: Optional[float] = None,
|
| 479 |
+
v2a_gate_init_value: Optional[float] = None,
|
| 480 |
+
a2v_gate_alpha: float = 1.0,
|
| 481 |
+
v2a_gate_alpha: float = 1.0,
|
| 482 |
+
):
|
| 483 |
+
super().__init__()
|
| 484 |
+
self.video_block = video_block
|
| 485 |
+
self.audio_block = audio_block
|
| 486 |
+
self.enable_a2v_cross_attn = bool(enable_a2v_cross_attn)
|
| 487 |
+
self.enable_v2a_cross_attn = bool(enable_v2a_cross_attn)
|
| 488 |
+
|
| 489 |
+
zero_init_video = zero_init_video_cross_attn if zero_init_video_cross_attn is not None else zero_init_output
|
| 490 |
+
zero_init_audio = zero_init_audio_cross_attn if zero_init_audio_cross_attn is not None else zero_init_output
|
| 491 |
+
|
| 492 |
+
if self.enable_a2v_cross_attn:
|
| 493 |
+
self.video_cross_attn_audio = GatedCrossModalAttention(
|
| 494 |
+
q_dim=video_dim,
|
| 495 |
+
kv_dim=audio_dim,
|
| 496 |
+
num_heads=video_num_heads,
|
| 497 |
+
zero_init_output=zero_init_video,
|
| 498 |
+
use_gating=use_a2v_gating,
|
| 499 |
+
zero_init_gating=zero_init_a2v_gating,
|
| 500 |
+
gate_init_value=a2v_gate_init_value,
|
| 501 |
+
gate_alpha=a2v_gate_alpha,
|
| 502 |
+
)
|
| 503 |
+
else:
|
| 504 |
+
self.video_cross_attn_audio = None
|
| 505 |
+
|
| 506 |
+
if self.enable_v2a_cross_attn:
|
| 507 |
+
self.audio_cross_attn_video = GatedCrossModalAttention(
|
| 508 |
+
q_dim=audio_dim,
|
| 509 |
+
kv_dim=video_dim,
|
| 510 |
+
num_heads=audio_num_heads,
|
| 511 |
+
zero_init_output=zero_init_audio,
|
| 512 |
+
use_gating=use_v2a_gating,
|
| 513 |
+
zero_init_gating=zero_init_v2a_gating,
|
| 514 |
+
gate_init_value=v2a_gate_init_value,
|
| 515 |
+
gate_alpha=v2a_gate_alpha,
|
| 516 |
+
)
|
| 517 |
+
else:
|
| 518 |
+
self.audio_cross_attn_video = None
|
| 519 |
+
|
| 520 |
+
def forward(
|
| 521 |
+
self,
|
| 522 |
+
video_x: torch.Tensor,
|
| 523 |
+
audio_x: torch.Tensor,
|
| 524 |
+
video_kwargs: Dict[str, Any],
|
| 525 |
+
audio_kwargs: Dict[str, Any],
|
| 526 |
+
dtype: torch.dtype = torch.bfloat16,
|
| 527 |
+
enable_a2v: Optional[Union[bool, torch.Tensor]] = None,
|
| 528 |
+
enable_v2a: Optional[Union[bool, torch.Tensor]] = None,
|
| 529 |
+
sequence_parallel: bool = False,
|
| 530 |
+
sp_group=None,
|
| 531 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 532 |
+
video_block = self.video_block
|
| 533 |
+
audio_block = self.audio_block
|
| 534 |
+
|
| 535 |
+
video_x = self._video_selfattn_and_text(
|
| 536 |
+
video_x, video_block, video_kwargs, dtype,
|
| 537 |
+
sequence_parallel=sequence_parallel,
|
| 538 |
+
sp_group=sp_group,
|
| 539 |
+
)
|
| 540 |
+
audio_x = self._audio_selfattn_and_text(
|
| 541 |
+
audio_x, audio_block, audio_kwargs, dtype,
|
| 542 |
+
sequence_parallel=sequence_parallel,
|
| 543 |
+
sp_group=sp_group,
|
| 544 |
+
)
|
| 545 |
+
|
| 546 |
+
temporal_rope_inv_freq = video_kwargs.get("temporal_rope_inv_freq")
|
| 547 |
+
video_temporal_pos = video_kwargs.get("temporal_positions")
|
| 548 |
+
audio_temporal_pos = audio_kwargs.get("temporal_positions")
|
| 549 |
+
|
| 550 |
+
# Runtime switches can be batch-wide booleans or per-branch CFG masks.
|
| 551 |
+
a2v_is_tensor = isinstance(enable_a2v, torch.Tensor)
|
| 552 |
+
v2a_is_tensor = isinstance(enable_v2a, torch.Tensor)
|
| 553 |
+
run_a2v = self.video_cross_attn_audio is not None and (
|
| 554 |
+
a2v_is_tensor or enable_a2v is None or enable_a2v
|
| 555 |
+
)
|
| 556 |
+
run_v2a = self.audio_cross_attn_video is not None and (
|
| 557 |
+
v2a_is_tensor or enable_v2a is None or enable_v2a
|
| 558 |
+
)
|
| 559 |
+
|
| 560 |
+
# Cache pre-cross-attention states so A2V and V2A are updated jointly:
|
| 561 |
+
# both directions must attend to the *same* pre-update snapshot, otherwise
|
| 562 |
+
# V2A would condition on the already-A2V-updated video (serial dependency).
|
| 563 |
+
video_x_pre = video_x
|
| 564 |
+
audio_x_pre = audio_x
|
| 565 |
+
|
| 566 |
+
if run_a2v:
|
| 567 |
+
a2v_result = self.video_cross_attn_audio(
|
| 568 |
+
x=video_x_pre,
|
| 569 |
+
y=audio_x_pre,
|
| 570 |
+
y_lens=audio_kwargs.get("seq_lens"),
|
| 571 |
+
dtype=dtype,
|
| 572 |
+
q_temporal_pos=video_temporal_pos,
|
| 573 |
+
k_temporal_pos=audio_temporal_pos,
|
| 574 |
+
temporal_rope_inv_freq=temporal_rope_inv_freq,
|
| 575 |
+
sequence_parallel=sequence_parallel,
|
| 576 |
+
sp_group=sp_group,
|
| 577 |
+
)
|
| 578 |
+
a2v_out = a2v_result
|
| 579 |
+
if a2v_is_tensor:
|
| 580 |
+
a2v_mask = enable_a2v.view(-1, *([1] * (a2v_out.dim() - 1))).to(
|
| 581 |
+
device=a2v_out.device, dtype=a2v_out.dtype
|
| 582 |
+
)
|
| 583 |
+
a2v_out = a2v_out * a2v_mask
|
| 584 |
+
video_x = video_x + a2v_out
|
| 585 |
+
|
| 586 |
+
if run_v2a:
|
| 587 |
+
v2a_result = self.audio_cross_attn_video(
|
| 588 |
+
x=audio_x_pre,
|
| 589 |
+
y=video_x_pre,
|
| 590 |
+
y_lens=video_kwargs.get("seq_lens"),
|
| 591 |
+
dtype=dtype,
|
| 592 |
+
q_temporal_pos=audio_temporal_pos,
|
| 593 |
+
k_temporal_pos=video_temporal_pos,
|
| 594 |
+
temporal_rope_inv_freq=temporal_rope_inv_freq,
|
| 595 |
+
sequence_parallel=sequence_parallel,
|
| 596 |
+
sp_group=sp_group,
|
| 597 |
+
)
|
| 598 |
+
v2a_out = v2a_result
|
| 599 |
+
if v2a_is_tensor:
|
| 600 |
+
v2a_mask = enable_v2a.view(-1, *([1] * (v2a_out.dim() - 1))).to(
|
| 601 |
+
device=v2a_out.device, dtype=v2a_out.dtype
|
| 602 |
+
)
|
| 603 |
+
v2a_out = v2a_out * v2a_mask
|
| 604 |
+
audio_x = audio_x + v2a_out
|
| 605 |
+
|
| 606 |
+
video_x = self._video_ffn(video_x, video_block, video_kwargs, dtype)
|
| 607 |
+
audio_x = self._audio_ffn(audio_x, audio_block, audio_kwargs, dtype)
|
| 608 |
+
return video_x, audio_x
|
| 609 |
+
|
| 610 |
+
def _video_selfattn_and_text(
|
| 611 |
+
self,
|
| 612 |
+
x: torch.Tensor,
|
| 613 |
+
block,
|
| 614 |
+
kwargs: Dict[str, Any],
|
| 615 |
+
dtype: torch.dtype,
|
| 616 |
+
sequence_parallel: bool = False,
|
| 617 |
+
sp_group=None,
|
| 618 |
+
) -> torch.Tensor:
|
| 619 |
+
e0 = kwargs["e0"]
|
| 620 |
+
seq_lens = kwargs["seq_lens"]
|
| 621 |
+
grid_sizes = kwargs["grid_sizes"]
|
| 622 |
+
freqs = kwargs["freqs"]
|
| 623 |
+
context = kwargs["context"]
|
| 624 |
+
context_lens = kwargs.get("context_lens")
|
| 625 |
+
|
| 626 |
+
if e0.dim() > 3:
|
| 627 |
+
modulation = (block.modulation.unsqueeze(0) + e0).chunk(6, dim=2)
|
| 628 |
+
modulation = [part.squeeze(2) for part in modulation]
|
| 629 |
+
else:
|
| 630 |
+
modulation = (block.modulation + e0).chunk(6, dim=1)
|
| 631 |
+
|
| 632 |
+
kwargs["_video_e"] = modulation
|
| 633 |
+
|
| 634 |
+
temp_x = block.norm1(x) * (1 + modulation[1]) + modulation[0]
|
| 635 |
+
temp_x = temp_x.to(dtype)
|
| 636 |
+
|
| 637 |
+
self_attn = block.self_attn
|
| 638 |
+
batch_size, seq_len = temp_x.shape[:2]
|
| 639 |
+
num_heads, head_dim = self_attn.num_heads, self_attn.head_dim
|
| 640 |
+
query = self_attn.norm_q(self_attn.q(temp_x)).view(batch_size, seq_len, num_heads, head_dim)
|
| 641 |
+
key = self_attn.norm_k(self_attn.k(temp_x)).view(batch_size, seq_len, num_heads, head_dim)
|
| 642 |
+
value = self_attn.v(temp_x).view(batch_size, seq_len, num_heads, head_dim)
|
| 643 |
+
if sequence_parallel:
|
| 644 |
+
sp_rank = int(kwargs["sp_rank"])
|
| 645 |
+
sp_world_size = int(kwargs["sp_world_size"])
|
| 646 |
+
query = _apply_video_rope_local(query, grid_sizes, freqs, sp_rank, sp_world_size)
|
| 647 |
+
key = _apply_video_rope_local(key, grid_sizes, freqs, sp_rank, sp_world_size)
|
| 648 |
+
attn_output = ulysses_attention(
|
| 649 |
+
query.to(dtype),
|
| 650 |
+
key.to(dtype),
|
| 651 |
+
value.to(dtype),
|
| 652 |
+
attention,
|
| 653 |
+
k_lens=seq_lens,
|
| 654 |
+
window_size=getattr(self_attn, "window_size", (-1, -1)),
|
| 655 |
+
group=sp_group,
|
| 656 |
+
)
|
| 657 |
+
else:
|
| 658 |
+
query, key = rope_apply_qk(query, key, grid_sizes, freqs)
|
| 659 |
+
attn_output = attention(
|
| 660 |
+
query.to(dtype),
|
| 661 |
+
key.to(dtype),
|
| 662 |
+
v=value.to(dtype),
|
| 663 |
+
k_lens=seq_lens,
|
| 664 |
+
window_size=getattr(self_attn, "window_size", (-1, -1)),
|
| 665 |
+
)
|
| 666 |
+
attn_output = attn_output.to(dtype).flatten(2)
|
| 667 |
+
attn_output = self_attn.o(attn_output)
|
| 668 |
+
|
| 669 |
+
x = x + attn_output * modulation[2]
|
| 670 |
+
x = x + block.cross_attn(block.norm3(x), context, context_lens, dtype)
|
| 671 |
+
return x
|
| 672 |
+
|
| 673 |
+
def _audio_selfattn_and_text(
|
| 674 |
+
self,
|
| 675 |
+
x: torch.Tensor,
|
| 676 |
+
block,
|
| 677 |
+
kwargs: Dict[str, Any],
|
| 678 |
+
dtype: torch.dtype,
|
| 679 |
+
sequence_parallel: bool = False,
|
| 680 |
+
sp_group=None,
|
| 681 |
+
) -> torch.Tensor:
|
| 682 |
+
time_mod = kwargs["e0"]
|
| 683 |
+
freqs = kwargs["freqs"]
|
| 684 |
+
context = kwargs["context"]
|
| 685 |
+
|
| 686 |
+
has_seq_mod = len(time_mod.shape) == 4
|
| 687 |
+
chunk_dim = 2 if has_seq_mod else 1
|
| 688 |
+
modulation = (
|
| 689 |
+
block.modulation.to(dtype=time_mod.dtype, device=time_mod.device) + time_mod
|
| 690 |
+
).chunk(6, dim=chunk_dim)
|
| 691 |
+
if has_seq_mod:
|
| 692 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = [
|
| 693 |
+
part.squeeze(2) for part in modulation
|
| 694 |
+
]
|
| 695 |
+
else:
|
| 696 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = modulation
|
| 697 |
+
|
| 698 |
+
kwargs["_audio_mod"] = (shift_mlp, scale_mlp, gate_mlp)
|
| 699 |
+
|
| 700 |
+
# norm1 (LayerNorm) upcasts to fp32 and the modulation terms are fp32,
|
| 701 |
+
# so input_x is fp32 here. block.self_attn calls flash-attn, which only
|
| 702 |
+
# supports fp16/bf16 -> cast down first (mirrors the video self-attn path).
|
| 703 |
+
input_x = (block.norm1(x) * (1 + scale_msa) + shift_msa).to(dtype)
|
| 704 |
+
if sequence_parallel:
|
| 705 |
+
self_attn = block.self_attn
|
| 706 |
+
batch_size, seq_len = input_x.shape[:2]
|
| 707 |
+
num_heads, head_dim = self_attn.num_heads, self_attn.head_dim
|
| 708 |
+
query = self_attn.norm_q(self_attn.q(input_x)).view(batch_size, seq_len, num_heads, head_dim)
|
| 709 |
+
key = self_attn.norm_k(self_attn.k(input_x)).view(batch_size, seq_len, num_heads, head_dim)
|
| 710 |
+
value = self_attn.v(input_x).view(batch_size, seq_len, num_heads, head_dim)
|
| 711 |
+
query = rope_apply_head_dim(
|
| 712 |
+
query.flatten(2), freqs, head_dim
|
| 713 |
+
).view(batch_size, seq_len, num_heads, head_dim)
|
| 714 |
+
key = rope_apply_head_dim(
|
| 715 |
+
key.flatten(2), freqs, head_dim
|
| 716 |
+
).view(batch_size, seq_len, num_heads, head_dim)
|
| 717 |
+
self_attn_output = ulysses_attention(
|
| 718 |
+
query,
|
| 719 |
+
key,
|
| 720 |
+
value,
|
| 721 |
+
attention,
|
| 722 |
+
k_lens=kwargs.get("seq_lens"),
|
| 723 |
+
window_size=getattr(self_attn, "window_size", (-1, -1)),
|
| 724 |
+
group=sp_group,
|
| 725 |
+
).flatten(2)
|
| 726 |
+
self_attn_output = self_attn.o(self_attn_output)
|
| 727 |
+
else:
|
| 728 |
+
self_attn_output = block.self_attn(input_x, freqs, seq_lens=kwargs.get("seq_lens"))
|
| 729 |
+
x = block.gate(x, gate_msa, self_attn_output)
|
| 730 |
+
# norm3 (LayerNorm) upcasts to fp32 and context (text embeds) may be fp32;
|
| 731 |
+
# audio cross_attn calls flash-attn (fp16/bf16 only) with no internal cast,
|
| 732 |
+
# so cast both q-source and kv-source down first.
|
| 733 |
+
x = x + block.cross_attn(block.norm3(x).to(dtype), context.to(dtype))
|
| 734 |
+
return x
|
| 735 |
+
|
| 736 |
+
def _video_ffn(
|
| 737 |
+
self, x: torch.Tensor, block, kwargs: Dict[str, Any], dtype: torch.dtype
|
| 738 |
+
) -> torch.Tensor:
|
| 739 |
+
modulation = kwargs["_video_e"]
|
| 740 |
+
temp_x = block.norm2(x) * (1 + modulation[4]) + modulation[3]
|
| 741 |
+
temp_x = temp_x.to(dtype)
|
| 742 |
+
return x + block.ffn(temp_x) * modulation[5]
|
| 743 |
+
|
| 744 |
+
def _audio_ffn(
|
| 745 |
+
self, x: torch.Tensor, block, kwargs: Dict[str, Any], dtype: torch.dtype
|
| 746 |
+
) -> torch.Tensor:
|
| 747 |
+
shift_mlp, scale_mlp, gate_mlp = kwargs["_audio_mod"]
|
| 748 |
+
input_x = block.norm2(x) * (1 + scale_mlp) + shift_mlp
|
| 749 |
+
return block.gate(x, gate_mlp, block.ffn(input_x))
|
| 750 |
+
|
| 751 |
+
|
| 752 |
+
class WanCreatorGatingAVModel(nn.Module):
|
| 753 |
+
"""Wan+Creator audio-video model with configurable gated cross attention.
|
| 754 |
+
|
| 755 |
+
A2V means the video branch attends to audio and updates video tokens.
|
| 756 |
+
V2A means the audio branch attends to video and updates audio tokens.
|
| 757 |
+
"""
|
| 758 |
+
|
| 759 |
+
def __init__(
|
| 760 |
+
self,
|
| 761 |
+
video_model: WanTransformer3DModel,
|
| 762 |
+
audio_model: CreatorAudioModel,
|
| 763 |
+
use_temporal_rope: bool = True,
|
| 764 |
+
audio_fps: float = 48000 / 960,
|
| 765 |
+
vae_temporal_stride: int = 4,
|
| 766 |
+
zero_init_cross_attn: bool = False,
|
| 767 |
+
zero_init_video_cross_attn: bool | None = None,
|
| 768 |
+
zero_init_audio_cross_attn: bool | None = None,
|
| 769 |
+
a2v_cross_attn_layers: LayerSelection = None,
|
| 770 |
+
v2a_cross_attn_layers: LayerSelection = None,
|
| 771 |
+
use_gating: bool = True,
|
| 772 |
+
use_a2v_gating: bool | None = None,
|
| 773 |
+
use_v2a_gating: bool | None = None,
|
| 774 |
+
zero_init_gating: bool = False,
|
| 775 |
+
zero_init_a2v_gating: bool | None = None,
|
| 776 |
+
zero_init_v2a_gating: bool | None = None,
|
| 777 |
+
gate_init_value: Optional[float] = None,
|
| 778 |
+
a2v_gate_init_value: Optional[float] = None,
|
| 779 |
+
v2a_gate_init_value: Optional[float] = None,
|
| 780 |
+
a2v_gate_alphas: LayerAlphas = None,
|
| 781 |
+
v2a_gate_alphas: LayerAlphas = None,
|
| 782 |
+
):
|
| 783 |
+
nn.Module.__init__(self)
|
| 784 |
+
self.video_model = video_model
|
| 785 |
+
self.audio_model = audio_model
|
| 786 |
+
|
| 787 |
+
self.video_dim = int(video_model.dim)
|
| 788 |
+
self.audio_dim = int(audio_model.dim)
|
| 789 |
+
self.video_num_heads = int(video_model.num_heads)
|
| 790 |
+
self.audio_num_heads = int(audio_model.num_heads)
|
| 791 |
+
self.num_layers = int(video_model.num_layers)
|
| 792 |
+
self.video_patch_size = tuple(int(value) for value in video_model.patch_size)
|
| 793 |
+
self.audio_patch_size = tuple(int(value) for value in audio_model.patch_size)
|
| 794 |
+
|
| 795 |
+
assert int(audio_model.num_layers) == self.num_layers, (
|
| 796 |
+
f"Video ({self.num_layers}) and audio ({audio_model.num_layers}) must have same number of layers"
|
| 797 |
+
)
|
| 798 |
+
|
| 799 |
+
a2v_enabled = _expand_layer_selection(
|
| 800 |
+
a2v_cross_attn_layers, self.num_layers, "a2v_cross_attn_layers"
|
| 801 |
+
)
|
| 802 |
+
v2a_enabled = _expand_layer_selection(
|
| 803 |
+
v2a_cross_attn_layers, self.num_layers, "v2a_cross_attn_layers"
|
| 804 |
+
)
|
| 805 |
+
self.a2v_cross_attn_layers = a2v_enabled
|
| 806 |
+
self.v2a_cross_attn_layers = v2a_enabled
|
| 807 |
+
resolved_a2v_gate_alphas = _expand_layer_alphas(
|
| 808 |
+
a2v_gate_alphas, self.num_layers, "a2v_gate_alphas"
|
| 809 |
+
)
|
| 810 |
+
resolved_v2a_gate_alphas = _expand_layer_alphas(
|
| 811 |
+
v2a_gate_alphas, self.num_layers, "v2a_gate_alphas"
|
| 812 |
+
)
|
| 813 |
+
self.a2v_gate_alphas = resolved_a2v_gate_alphas
|
| 814 |
+
self.v2a_gate_alphas = resolved_v2a_gate_alphas
|
| 815 |
+
|
| 816 |
+
resolved_use_a2v_gating = use_gating if use_a2v_gating is None else use_a2v_gating
|
| 817 |
+
resolved_use_v2a_gating = use_gating if use_v2a_gating is None else use_v2a_gating
|
| 818 |
+
resolved_zero_init_a2v_gating = (
|
| 819 |
+
zero_init_gating if zero_init_a2v_gating is None else zero_init_a2v_gating
|
| 820 |
+
)
|
| 821 |
+
resolved_zero_init_v2a_gating = (
|
| 822 |
+
zero_init_gating if zero_init_v2a_gating is None else zero_init_v2a_gating
|
| 823 |
+
)
|
| 824 |
+
resolved_a2v_gate_init_value = (
|
| 825 |
+
gate_init_value if a2v_gate_init_value is None else a2v_gate_init_value
|
| 826 |
+
)
|
| 827 |
+
resolved_v2a_gate_init_value = (
|
| 828 |
+
gate_init_value if v2a_gate_init_value is None else v2a_gate_init_value
|
| 829 |
+
)
|
| 830 |
+
video_blocks = list(video_model.blocks)
|
| 831 |
+
audio_blocks = list(audio_model.blocks)
|
| 832 |
+
self.joint_blocks = nn.ModuleList([
|
| 833 |
+
GatedJointBlock(
|
| 834 |
+
video_block=video_block,
|
| 835 |
+
audio_block=audio_block,
|
| 836 |
+
video_dim=self.video_dim,
|
| 837 |
+
audio_dim=self.audio_dim,
|
| 838 |
+
video_num_heads=self.video_num_heads,
|
| 839 |
+
audio_num_heads=self.audio_num_heads,
|
| 840 |
+
enable_a2v_cross_attn=a2v_enabled[layer_idx],
|
| 841 |
+
enable_v2a_cross_attn=v2a_enabled[layer_idx],
|
| 842 |
+
zero_init_output=zero_init_cross_attn,
|
| 843 |
+
zero_init_video_cross_attn=zero_init_video_cross_attn,
|
| 844 |
+
zero_init_audio_cross_attn=zero_init_audio_cross_attn,
|
| 845 |
+
use_a2v_gating=resolved_use_a2v_gating,
|
| 846 |
+
use_v2a_gating=resolved_use_v2a_gating,
|
| 847 |
+
zero_init_a2v_gating=resolved_zero_init_a2v_gating,
|
| 848 |
+
zero_init_v2a_gating=resolved_zero_init_v2a_gating,
|
| 849 |
+
a2v_gate_init_value=resolved_a2v_gate_init_value,
|
| 850 |
+
v2a_gate_init_value=resolved_v2a_gate_init_value,
|
| 851 |
+
a2v_gate_alpha=resolved_a2v_gate_alphas[layer_idx],
|
| 852 |
+
v2a_gate_alpha=resolved_v2a_gate_alphas[layer_idx],
|
| 853 |
+
)
|
| 854 |
+
for layer_idx, (video_block, audio_block) in enumerate(zip(video_blocks, audio_blocks))
|
| 855 |
+
])
|
| 856 |
+
|
| 857 |
+
video_model.blocks = nn.ModuleList()
|
| 858 |
+
audio_model.blocks = nn.ModuleList()
|
| 859 |
+
|
| 860 |
+
self.use_temporal_rope = use_temporal_rope
|
| 861 |
+
self.audio_fps = audio_fps
|
| 862 |
+
self.vae_temporal_stride = vae_temporal_stride
|
| 863 |
+
if use_temporal_rope:
|
| 864 |
+
head_dim = self.video_dim // self.video_num_heads
|
| 865 |
+
self.temporal_rope_inv_freq = 1.0 / (
|
| 866 |
+
10000.0 ** (torch.arange(0, head_dim, 2, dtype=torch.float64) / head_dim)
|
| 867 |
+
)
|
| 868 |
+
else:
|
| 869 |
+
self.temporal_rope_inv_freq = None
|
| 870 |
+
|
| 871 |
+
# Ulysses-style sequence-parallel inference state. All ranks keep a
|
| 872 |
+
# complete copy of the weights and exchange token/head dimensions in
|
| 873 |
+
# attention, matching the Wan2.2 inference design.
|
| 874 |
+
self.sp_world_size = 1
|
| 875 |
+
self.sp_world_rank = 0
|
| 876 |
+
self.sp_group = None
|
| 877 |
+
|
| 878 |
+
def enable_multi_gpus_inference(self, group=None) -> None:
|
| 879 |
+
"""Enable raw-process-group sequence parallelism for inference."""
|
| 880 |
+
import torch.distributed as dist
|
| 881 |
+
|
| 882 |
+
if not dist.is_initialized():
|
| 883 |
+
raise RuntimeError("Sequence-parallel inference requires an initialized process group")
|
| 884 |
+
self.sp_world_size = dist.get_world_size(group)
|
| 885 |
+
self.sp_world_rank = dist.get_rank(group)
|
| 886 |
+
self.sp_group = group
|
| 887 |
+
if self.video_num_heads % self.sp_world_size != 0:
|
| 888 |
+
raise ValueError(
|
| 889 |
+
f"Video attention heads ({self.video_num_heads}) must be divisible by "
|
| 890 |
+
f"SP size ({self.sp_world_size})"
|
| 891 |
+
)
|
| 892 |
+
if self.audio_num_heads % self.sp_world_size != 0:
|
| 893 |
+
raise ValueError(
|
| 894 |
+
f"Audio attention heads ({self.audio_num_heads}) must be divisible by "
|
| 895 |
+
f"SP size ({self.sp_world_size})"
|
| 896 |
+
)
|
| 897 |
+
|
| 898 |
+
def _apply_sequence_parallel(self, video_state: Dict[str, Any], audio_state: Dict[str, Any]):
|
| 899 |
+
"""Shard prepared video/audio token states along their sequence axes."""
|
| 900 |
+
if self.sp_world_size <= 1:
|
| 901 |
+
video_state["sp_rank"] = 0
|
| 902 |
+
video_state["sp_world_size"] = 1
|
| 903 |
+
audio_state["sp_rank"] = 0
|
| 904 |
+
audio_state["sp_world_size"] = 1
|
| 905 |
+
return video_state, audio_state
|
| 906 |
+
|
| 907 |
+
def shard(state: Dict[str, Any], *, audio: bool):
|
| 908 |
+
local_len = state["x"].size(1) // self.sp_world_size
|
| 909 |
+
rank = self.sp_world_rank
|
| 910 |
+
state["x"] = torch.chunk(state["x"], self.sp_world_size, dim=1)[rank]
|
| 911 |
+
if not audio:
|
| 912 |
+
if state["e"].dim() >= 3:
|
| 913 |
+
state["e"] = torch.chunk(state["e"], self.sp_world_size, dim=1)[rank]
|
| 914 |
+
if state["e0"].dim() >= 4:
|
| 915 |
+
state["e0"] = torch.chunk(state["e0"], self.sp_world_size, dim=1)[rank]
|
| 916 |
+
state["freqs"] = (
|
| 917 |
+
torch.chunk(state["freqs"], self.sp_world_size, dim=0)[rank]
|
| 918 |
+
if audio else state["freqs"]
|
| 919 |
+
)
|
| 920 |
+
if state.get("temporal_positions") is not None:
|
| 921 |
+
state["temporal_positions"] = torch.chunk(
|
| 922 |
+
state["temporal_positions"], self.sp_world_size, dim=1
|
| 923 |
+
)[rank]
|
| 924 |
+
state["local_seq_lens"] = (
|
| 925 |
+
state["seq_lens"] - rank * local_len
|
| 926 |
+
).clamp(min=0, max=local_len)
|
| 927 |
+
state["sp_rank"] = rank
|
| 928 |
+
state["sp_world_size"] = self.sp_world_size
|
| 929 |
+
return state
|
| 930 |
+
|
| 931 |
+
return shard(video_state, audio=False), shard(audio_state, audio=True)
|
| 932 |
+
|
| 933 |
+
def _prepare_video(self, video_inputs: Dict[str, Any], dtype: torch.dtype) -> Dict[str, Any]:
|
| 934 |
+
video_model = self.video_model
|
| 935 |
+
device = video_model.patch_embedding.weight.device
|
| 936 |
+
|
| 937 |
+
if video_model.freqs.device != device:
|
| 938 |
+
video_model.freqs = video_model.freqs.to(device)
|
| 939 |
+
|
| 940 |
+
x_list = _to_list(video_inputs["x"])
|
| 941 |
+
y = video_inputs.get("y")
|
| 942 |
+
if y is not None:
|
| 943 |
+
y_list = _to_list(y)
|
| 944 |
+
x_list = [torch.cat([sample, condition], dim=0) for sample, condition in zip(x_list, y_list)]
|
| 945 |
+
|
| 946 |
+
x_list = [video_model.patch_embedding(sample.unsqueeze(0)) for sample in x_list]
|
| 947 |
+
grid_sizes = torch.stack([
|
| 948 |
+
torch.tensor(sample.shape[2:], dtype=torch.long, device=device) for sample in x_list
|
| 949 |
+
])
|
| 950 |
+
x_list = [sample.flatten(2).transpose(1, 2) for sample in x_list]
|
| 951 |
+
seq_lens = torch.tensor([sample.size(1) for sample in x_list], dtype=torch.long, device=device)
|
| 952 |
+
|
| 953 |
+
seq_len = self._round_seq_len(int(video_inputs["seq_len"]))
|
| 954 |
+
assert int(seq_lens.max().item()) <= seq_len
|
| 955 |
+
x = torch.cat([
|
| 956 |
+
torch.cat([sample, sample.new_zeros(1, seq_len - sample.size(1), sample.size(2))], dim=1)
|
| 957 |
+
for sample in x_list
|
| 958 |
+
])
|
| 959 |
+
|
| 960 |
+
timesteps = video_inputs["t"].to(device)
|
| 961 |
+
embedding, modulation = _build_time_embeddings(
|
| 962 |
+
video_model.time_embedding,
|
| 963 |
+
video_model.time_projection,
|
| 964 |
+
int(video_model.freq_dim),
|
| 965 |
+
int(video_model.dim),
|
| 966 |
+
timesteps,
|
| 967 |
+
seq_len,
|
| 968 |
+
)
|
| 969 |
+
|
| 970 |
+
context = _embed_context(video_model.text_embedding, video_inputs["context"], int(video_model.text_len))
|
| 971 |
+
|
| 972 |
+
return {
|
| 973 |
+
"x": x,
|
| 974 |
+
"e": embedding,
|
| 975 |
+
"e0": modulation,
|
| 976 |
+
"seq_lens": seq_lens,
|
| 977 |
+
"grid_sizes": grid_sizes,
|
| 978 |
+
"freqs": video_model.freqs,
|
| 979 |
+
"context": context,
|
| 980 |
+
"context_lens": None,
|
| 981 |
+
"seq_len": seq_len,
|
| 982 |
+
}
|
| 983 |
+
|
| 984 |
+
def _prepare_audio(self, audio_inputs: Dict[str, Any], dtype: torch.dtype) -> Dict[str, Any]:
|
| 985 |
+
audio_model = self.audio_model
|
| 986 |
+
device = audio_model.patch_embedding.weight.device
|
| 987 |
+
|
| 988 |
+
x_list = _to_list(audio_inputs["x"])
|
| 989 |
+
y = audio_inputs.get("y")
|
| 990 |
+
if y is not None:
|
| 991 |
+
y_list = _to_list(y)
|
| 992 |
+
x_list = [torch.cat([sample, condition], dim=0) for sample, condition in zip(x_list, y_list)]
|
| 993 |
+
|
| 994 |
+
original_audio_shapes = [tuple(sample.shape) for sample in x_list]
|
| 995 |
+
|
| 996 |
+
patchified = []
|
| 997 |
+
grid_sizes_list = []
|
| 998 |
+
for sample in x_list:
|
| 999 |
+
tokens = audio_model.patch_embedding(sample.unsqueeze(0).to(device))
|
| 1000 |
+
tokens = rearrange(tokens, "1 c f -> f c").contiguous()
|
| 1001 |
+
patchified.append(tokens)
|
| 1002 |
+
grid_sizes_list.append(tokens.shape[0])
|
| 1003 |
+
|
| 1004 |
+
grid_sizes = torch.tensor([[grid_size] for grid_size in grid_sizes_list], dtype=torch.long, device=device)
|
| 1005 |
+
seq_lens = grid_sizes[:, 0]
|
| 1006 |
+
seq_len = self._round_seq_len(int(audio_inputs["seq_len"]))
|
| 1007 |
+
assert int(seq_lens.max().item()) <= seq_len
|
| 1008 |
+
|
| 1009 |
+
x = torch.stack([
|
| 1010 |
+
torch.cat([sample, sample.new_zeros(seq_len - sample.size(0), sample.size(1))], dim=0)
|
| 1011 |
+
for sample in patchified
|
| 1012 |
+
])
|
| 1013 |
+
|
| 1014 |
+
timesteps = audio_inputs["t"].to(device)
|
| 1015 |
+
embedding, modulation = _build_time_embeddings(
|
| 1016 |
+
audio_model.time_embedding,
|
| 1017 |
+
audio_model.time_projection,
|
| 1018 |
+
int(audio_model.freq_dim),
|
| 1019 |
+
int(audio_model.dim),
|
| 1020 |
+
timesteps,
|
| 1021 |
+
seq_len,
|
| 1022 |
+
)
|
| 1023 |
+
|
| 1024 |
+
context = _embed_context(audio_model.text_embedding, audio_inputs["context"], int(audio_model.text_len))
|
| 1025 |
+
freqs = audio_model._build_freqs(seq_len, device)
|
| 1026 |
+
|
| 1027 |
+
clip_fea = audio_inputs.get("clip_fea")
|
| 1028 |
+
if audio_model.has_image_input and clip_fea is not None:
|
| 1029 |
+
clip_embedding = audio_model.img_emb(clip_fea)
|
| 1030 |
+
context = torch.cat([clip_embedding, context], dim=1)
|
| 1031 |
+
|
| 1032 |
+
return {
|
| 1033 |
+
"x": x,
|
| 1034 |
+
"e": embedding,
|
| 1035 |
+
"e0": modulation,
|
| 1036 |
+
"seq_lens": seq_lens,
|
| 1037 |
+
"grid_sizes": grid_sizes,
|
| 1038 |
+
"freqs": freqs,
|
| 1039 |
+
"context": context,
|
| 1040 |
+
"context_lens": None,
|
| 1041 |
+
"seq_len": seq_len,
|
| 1042 |
+
"original_audio_shapes": original_audio_shapes,
|
| 1043 |
+
}
|
| 1044 |
+
|
| 1045 |
+
def forward(
|
| 1046 |
+
self,
|
| 1047 |
+
video: Dict[str, Any],
|
| 1048 |
+
audio: Dict[str, Any],
|
| 1049 |
+
dtype: torch.dtype = torch.bfloat16,
|
| 1050 |
+
return_dict: bool = True,
|
| 1051 |
+
enable_a2v: Optional[Union[bool, torch.Tensor]] = None,
|
| 1052 |
+
enable_v2a: Optional[Union[bool, torch.Tensor]] = None,
|
| 1053 |
+
):
|
| 1054 |
+
"""Forward pass of the joint audio-video model.
|
| 1055 |
+
|
| 1056 |
+
Args:
|
| 1057 |
+
video: video input dict with keys 'x', 't', 'context', 'seq_len', etc.
|
| 1058 |
+
audio: audio input dict with keys 'x', 't', 'context', 'seq_len', etc.
|
| 1059 |
+
dtype: computation dtype for attention ops (default bfloat16).
|
| 1060 |
+
return_dict: if True, return dict with 'video'/'audio' keys; else tuple.
|
| 1061 |
+
enable_a2v: gate A2V cross-attention (video attending to audio).
|
| 1062 |
+
A bool tensor can select the enabled CFG branches per sample.
|
| 1063 |
+
enable_v2a: gate V2A cross-attention (audio attending to video), same
|
| 1064 |
+
semantics as enable_a2v.
|
| 1065 |
+
|
| 1066 |
+
Returns:
|
| 1067 |
+
dict or tuple of (video_output, audio_output) tensors.
|
| 1068 |
+
"""
|
| 1069 |
+
video_state = self._prepare_video(video, dtype)
|
| 1070 |
+
audio_state = self._prepare_audio(audio, dtype)
|
| 1071 |
+
|
| 1072 |
+
device = video_state["x"].device
|
| 1073 |
+
|
| 1074 |
+
video_temporal_pos = None
|
| 1075 |
+
audio_temporal_pos = None
|
| 1076 |
+
temporal_rope_inv_freq = None
|
| 1077 |
+
if self.use_temporal_rope and self.temporal_rope_inv_freq is not None:
|
| 1078 |
+
video_fps = float(video.get("video_fps", 16.0))
|
| 1079 |
+
temporal_rope_inv_freq = self.temporal_rope_inv_freq
|
| 1080 |
+
video_temporal_pos = compute_video_temporal_positions(
|
| 1081 |
+
video_state["grid_sizes"],
|
| 1082 |
+
video_state["x"].size(1),
|
| 1083 |
+
device,
|
| 1084 |
+
audio_fps=self.audio_fps,
|
| 1085 |
+
video_fps=video_fps,
|
| 1086 |
+
vae_temporal_stride=self.vae_temporal_stride,
|
| 1087 |
+
)
|
| 1088 |
+
audio_temporal_pos = compute_audio_temporal_positions(
|
| 1089 |
+
audio_state["seq_lens"],
|
| 1090 |
+
audio_state["x"].size(1),
|
| 1091 |
+
device,
|
| 1092 |
+
)
|
| 1093 |
+
|
| 1094 |
+
video_state["temporal_positions"] = video_temporal_pos
|
| 1095 |
+
audio_state["temporal_positions"] = audio_temporal_pos
|
| 1096 |
+
video_state, audio_state = self._apply_sequence_parallel(video_state, audio_state)
|
| 1097 |
+
video_x = video_state["x"]
|
| 1098 |
+
audio_x = audio_state["x"]
|
| 1099 |
+
|
| 1100 |
+
video_kwargs = {
|
| 1101 |
+
"e0": video_state["e0"],
|
| 1102 |
+
"seq_lens": video_state["seq_lens"],
|
| 1103 |
+
"grid_sizes": video_state["grid_sizes"],
|
| 1104 |
+
"freqs": video_state["freqs"],
|
| 1105 |
+
"context": video_state["context"],
|
| 1106 |
+
"context_lens": video_state["context_lens"],
|
| 1107 |
+
"temporal_positions": video_state["temporal_positions"],
|
| 1108 |
+
"temporal_rope_inv_freq": temporal_rope_inv_freq,
|
| 1109 |
+
"sp_rank": video_state["sp_rank"],
|
| 1110 |
+
"sp_world_size": video_state["sp_world_size"],
|
| 1111 |
+
}
|
| 1112 |
+
audio_kwargs = {
|
| 1113 |
+
"e0": audio_state["e0"],
|
| 1114 |
+
"seq_lens": audio_state["seq_lens"],
|
| 1115 |
+
"freqs": audio_state["freqs"],
|
| 1116 |
+
"context": audio_state["context"],
|
| 1117 |
+
"temporal_positions": audio_state["temporal_positions"],
|
| 1118 |
+
"sp_rank": audio_state["sp_rank"],
|
| 1119 |
+
"sp_world_size": audio_state["sp_world_size"],
|
| 1120 |
+
}
|
| 1121 |
+
|
| 1122 |
+
runtime_cross_attn = {
|
| 1123 |
+
"enable_a2v": enable_a2v,
|
| 1124 |
+
"enable_v2a": enable_v2a,
|
| 1125 |
+
"sequence_parallel": self.sp_world_size > 1,
|
| 1126 |
+
"sp_group": self.sp_group,
|
| 1127 |
+
}
|
| 1128 |
+
|
| 1129 |
+
for joint_block in self.joint_blocks:
|
| 1130 |
+
video_x, audio_x = joint_block(
|
| 1131 |
+
video_x, audio_x, video_kwargs, audio_kwargs, dtype,
|
| 1132 |
+
**runtime_cross_attn,
|
| 1133 |
+
)
|
| 1134 |
+
|
| 1135 |
+
video_output = self.video_model.head(video_x, video_state["e"])
|
| 1136 |
+
audio_output = self.audio_model.head(audio_x, audio_state["e"])
|
| 1137 |
+
|
| 1138 |
+
if self.sp_world_size > 1:
|
| 1139 |
+
video_output = all_gather_sequence(video_output, group=self.sp_group)
|
| 1140 |
+
audio_output = all_gather_sequence(audio_output, group=self.sp_group)
|
| 1141 |
+
|
| 1142 |
+
video_output = torch.stack(
|
| 1143 |
+
self.video_model.unpatchify(video_output, video_state["grid_sizes"])
|
| 1144 |
+
)
|
| 1145 |
+
audio_output = torch.stack(
|
| 1146 |
+
self.audio_model.unpatchify(
|
| 1147 |
+
audio_output,
|
| 1148 |
+
audio_state["grid_sizes"],
|
| 1149 |
+
audio_state["original_audio_shapes"],
|
| 1150 |
+
)
|
| 1151 |
+
)
|
| 1152 |
+
result = {"video": video_output, "audio": audio_output}
|
| 1153 |
+
return result if return_dict else (video_output, audio_output)
|
| 1154 |
+
|
| 1155 |
+
def _round_seq_len(self, seq_len: int) -> int:
|
| 1156 |
+
if self.sp_world_size > 1:
|
| 1157 |
+
return int(math.ceil(seq_len / self.sp_world_size) * self.sp_world_size)
|
| 1158 |
+
return int(seq_len)
|
| 1159 |
+
|
| 1160 |
+
@classmethod
|
| 1161 |
+
def from_pretrained(
|
| 1162 |
+
cls,
|
| 1163 |
+
pretrained_model_path: Optional[str] = None,
|
| 1164 |
+
video_pretrained_model_path: Optional[str] = None,
|
| 1165 |
+
audio_pretrained_model_path: Optional[str] = None,
|
| 1166 |
+
video_subfolder: Optional[str] = None,
|
| 1167 |
+
audio_subfolder: Optional[str] = None,
|
| 1168 |
+
video_kwargs: Optional[Dict] = None,
|
| 1169 |
+
audio_kwargs: Optional[Dict] = None,
|
| 1170 |
+
video_model_cls=Wan2_2Transformer3DModel,
|
| 1171 |
+
audio_model_cls=CreatorAudioModel,
|
| 1172 |
+
low_cpu_mem_usage: bool = False,
|
| 1173 |
+
torch_dtype: torch.dtype = torch.bfloat16,
|
| 1174 |
+
use_temporal_rope: bool = True,
|
| 1175 |
+
audio_fps: float = 48000.0 / 960.0,
|
| 1176 |
+
vae_temporal_stride: int = 4,
|
| 1177 |
+
zero_init_cross_attn: bool = False,
|
| 1178 |
+
zero_init_video_cross_attn: bool | None = None,
|
| 1179 |
+
zero_init_audio_cross_attn: bool | None = None,
|
| 1180 |
+
a2v_cross_attn_layers: LayerSelection = None,
|
| 1181 |
+
v2a_cross_attn_layers: LayerSelection = None,
|
| 1182 |
+
use_gating: bool = True,
|
| 1183 |
+
use_a2v_gating: bool | None = None,
|
| 1184 |
+
use_v2a_gating: bool | None = None,
|
| 1185 |
+
zero_init_gating: bool = False,
|
| 1186 |
+
zero_init_a2v_gating: bool | None = None,
|
| 1187 |
+
zero_init_v2a_gating: bool | None = None,
|
| 1188 |
+
gate_init_value: Optional[float] = None,
|
| 1189 |
+
a2v_gate_init_value: Optional[float] = None,
|
| 1190 |
+
v2a_gate_init_value: Optional[float] = None,
|
| 1191 |
+
a2v_gate_alphas: LayerAlphas = None,
|
| 1192 |
+
v2a_gate_alphas: LayerAlphas = None,
|
| 1193 |
+
):
|
| 1194 |
+
video_kwargs = dict(video_kwargs or {})
|
| 1195 |
+
audio_kwargs = dict(audio_kwargs or {})
|
| 1196 |
+
|
| 1197 |
+
if video_pretrained_model_path is not None and audio_pretrained_model_path is not None:
|
| 1198 |
+
video_path = video_pretrained_model_path
|
| 1199 |
+
audio_path = audio_pretrained_model_path
|
| 1200 |
+
elif pretrained_model_path is not None:
|
| 1201 |
+
video_path = os.path.join(pretrained_model_path, "video_model")
|
| 1202 |
+
audio_path = os.path.join(pretrained_model_path, "audio_model")
|
| 1203 |
+
else:
|
| 1204 |
+
raise ValueError(
|
| 1205 |
+
"Must provide either pretrained_model_path or both video_pretrained_model_path "
|
| 1206 |
+
"and audio_pretrained_model_path"
|
| 1207 |
+
)
|
| 1208 |
+
|
| 1209 |
+
logging.info("Loading video model from: %s", video_path)
|
| 1210 |
+
video_model = video_model_cls.from_pretrained(
|
| 1211 |
+
video_path,
|
| 1212 |
+
subfolder=video_subfolder,
|
| 1213 |
+
transformer_additional_kwargs=video_kwargs,
|
| 1214 |
+
low_cpu_mem_usage=low_cpu_mem_usage,
|
| 1215 |
+
torch_dtype=torch_dtype,
|
| 1216 |
+
)
|
| 1217 |
+
|
| 1218 |
+
logging.info("Loading audio model from: %s", audio_path)
|
| 1219 |
+
audio_model = audio_model_cls.from_pretrained(
|
| 1220 |
+
audio_path,
|
| 1221 |
+
subfolder=audio_subfolder,
|
| 1222 |
+
transformer_additional_kwargs=audio_kwargs,
|
| 1223 |
+
low_cpu_mem_usage=low_cpu_mem_usage,
|
| 1224 |
+
torch_dtype=torch_dtype,
|
| 1225 |
+
)
|
| 1226 |
+
|
| 1227 |
+
model = cls(
|
| 1228 |
+
video_model=video_model,
|
| 1229 |
+
audio_model=audio_model,
|
| 1230 |
+
use_temporal_rope=use_temporal_rope,
|
| 1231 |
+
audio_fps=audio_fps,
|
| 1232 |
+
vae_temporal_stride=vae_temporal_stride,
|
| 1233 |
+
zero_init_cross_attn=zero_init_cross_attn,
|
| 1234 |
+
zero_init_video_cross_attn=zero_init_video_cross_attn,
|
| 1235 |
+
zero_init_audio_cross_attn=zero_init_audio_cross_attn,
|
| 1236 |
+
a2v_cross_attn_layers=a2v_cross_attn_layers,
|
| 1237 |
+
v2a_cross_attn_layers=v2a_cross_attn_layers,
|
| 1238 |
+
use_gating=use_gating,
|
| 1239 |
+
use_a2v_gating=use_a2v_gating,
|
| 1240 |
+
use_v2a_gating=use_v2a_gating,
|
| 1241 |
+
zero_init_gating=zero_init_gating,
|
| 1242 |
+
zero_init_a2v_gating=zero_init_a2v_gating,
|
| 1243 |
+
zero_init_v2a_gating=zero_init_v2a_gating,
|
| 1244 |
+
gate_init_value=gate_init_value,
|
| 1245 |
+
a2v_gate_init_value=a2v_gate_init_value,
|
| 1246 |
+
v2a_gate_init_value=v2a_gate_init_value,
|
| 1247 |
+
a2v_gate_alphas=a2v_gate_alphas,
|
| 1248 |
+
v2a_gate_alphas=v2a_gate_alphas,
|
| 1249 |
+
).to(torch_dtype)
|
| 1250 |
+
if pretrained_model_path is not None:
|
| 1251 |
+
cross_attn_file = os.path.join(pretrained_model_path, "cross_attn_weights.safetensors")
|
| 1252 |
+
cross_attn_file_bin = os.path.join(pretrained_model_path, "cross_attn_weights.bin")
|
| 1253 |
+
|
| 1254 |
+
if os.path.exists(cross_attn_file):
|
| 1255 |
+
from safetensors.torch import load_file
|
| 1256 |
+
|
| 1257 |
+
cross_attn_state = load_file(cross_attn_file)
|
| 1258 |
+
logging.info("Loading cross-attn weights from: %s (%d keys)", cross_attn_file, len(cross_attn_state))
|
| 1259 |
+
elif os.path.exists(cross_attn_file_bin):
|
| 1260 |
+
cross_attn_state = torch.load(cross_attn_file_bin, map_location="cpu")
|
| 1261 |
+
logging.info("Loading cross-attn weights from: %s (%d keys)", cross_attn_file_bin, len(cross_attn_state))
|
| 1262 |
+
else:
|
| 1263 |
+
cross_attn_state = None
|
| 1264 |
+
logging.warning("No cross_attn_weights found in %s, skipping.", pretrained_model_path)
|
| 1265 |
+
|
| 1266 |
+
if cross_attn_state is not None:
|
| 1267 |
+
missing, unexpected = model.load_state_dict(cross_attn_state, strict=False)
|
| 1268 |
+
logging.info("Cross-attn load: %d missing, %d unexpected keys", len(missing), len(unexpected))
|
| 1269 |
+
if unexpected:
|
| 1270 |
+
logging.warning("Unexpected keys in cross_attn_weights: %s", unexpected[:10])
|
| 1271 |
+
|
| 1272 |
+
return model
|
| 1273 |
+
|
| 1274 |
+
|
| 1275 |
+
WanCreatorCrossAttnGatingAVModel = WanCreatorGatingAVModel
|
| 1276 |
+
WanCreatorGatedCrossAttnAVModel = WanCreatorGatingAVModel
|
| 1277 |
+
|
| 1278 |
+
__all__ = [
|
| 1279 |
+
"GatedCrossModalAttention",
|
| 1280 |
+
"GatedJointBlock",
|
| 1281 |
+
"WanCreatorGatingAVModel",
|
| 1282 |
+
"WanCreatorCrossAttnGatingAVModel",
|
| 1283 |
+
"WanCreatorGatedCrossAttnAVModel",
|
| 1284 |
+
"compute_audio_temporal_positions",
|
| 1285 |
+
"compute_video_temporal_positions",
|
| 1286 |
+
]
|
videox_fun/models/wan_text_encoder.py
ADDED
|
@@ -0,0 +1,389 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from https://github.com/Wan-Video/Wan2.1/blob/main/wan/modules/t5.py
|
| 2 |
+
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
| 3 |
+
import math
|
| 4 |
+
import logging
|
| 5 |
+
from typing import Optional
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
from diffusers.configuration_utils import ConfigMixin
|
| 11 |
+
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
| 12 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def fp16_clamp(x):
|
| 16 |
+
if x.dtype == torch.float16 and torch.isinf(x).any():
|
| 17 |
+
clamp = torch.finfo(x.dtype).max - 1000
|
| 18 |
+
x = torch.clamp(x, min=-clamp, max=clamp)
|
| 19 |
+
return x
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def init_weights(m):
|
| 23 |
+
if isinstance(m, T5LayerNorm):
|
| 24 |
+
nn.init.ones_(m.weight)
|
| 25 |
+
elif isinstance(m, T5FeedForward):
|
| 26 |
+
nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)
|
| 27 |
+
nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)
|
| 28 |
+
nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)
|
| 29 |
+
elif isinstance(m, T5Attention):
|
| 30 |
+
nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)
|
| 31 |
+
nn.init.normal_(m.k.weight, std=m.dim**-0.5)
|
| 32 |
+
nn.init.normal_(m.v.weight, std=m.dim**-0.5)
|
| 33 |
+
nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)
|
| 34 |
+
elif isinstance(m, T5RelativeEmbedding):
|
| 35 |
+
nn.init.normal_(
|
| 36 |
+
m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
class GELU(nn.Module):
|
| 40 |
+
def forward(self, x):
|
| 41 |
+
return 0.5 * x * (1.0 + torch.tanh(
|
| 42 |
+
math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
class T5LayerNorm(nn.Module):
|
| 46 |
+
def __init__(self, dim, eps=1e-6):
|
| 47 |
+
super(T5LayerNorm, self).__init__()
|
| 48 |
+
self.dim = dim
|
| 49 |
+
self.eps = eps
|
| 50 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 51 |
+
|
| 52 |
+
def forward(self, x):
|
| 53 |
+
x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +
|
| 54 |
+
self.eps)
|
| 55 |
+
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
| 56 |
+
x = x.type_as(self.weight)
|
| 57 |
+
return self.weight * x
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
class T5Attention(nn.Module):
|
| 61 |
+
def __init__(self, dim, dim_attn, num_heads, dropout=0.1):
|
| 62 |
+
assert dim_attn % num_heads == 0
|
| 63 |
+
super(T5Attention, self).__init__()
|
| 64 |
+
self.dim = dim
|
| 65 |
+
self.dim_attn = dim_attn
|
| 66 |
+
self.num_heads = num_heads
|
| 67 |
+
self.head_dim = dim_attn // num_heads
|
| 68 |
+
|
| 69 |
+
# layers
|
| 70 |
+
self.q = nn.Linear(dim, dim_attn, bias=False)
|
| 71 |
+
self.k = nn.Linear(dim, dim_attn, bias=False)
|
| 72 |
+
self.v = nn.Linear(dim, dim_attn, bias=False)
|
| 73 |
+
self.o = nn.Linear(dim_attn, dim, bias=False)
|
| 74 |
+
self.dropout = nn.Dropout(dropout)
|
| 75 |
+
|
| 76 |
+
def forward(self, x, context=None, mask=None, pos_bias=None):
|
| 77 |
+
"""
|
| 78 |
+
x: [B, L1, C].
|
| 79 |
+
context: [B, L2, C] or None.
|
| 80 |
+
mask: [B, L2] or [B, L1, L2] or None.
|
| 81 |
+
"""
|
| 82 |
+
# check inputs
|
| 83 |
+
context = x if context is None else context
|
| 84 |
+
b, n, c = x.size(0), self.num_heads, self.head_dim
|
| 85 |
+
|
| 86 |
+
# compute query, key, value
|
| 87 |
+
q = self.q(x).view(b, -1, n, c)
|
| 88 |
+
k = self.k(context).view(b, -1, n, c)
|
| 89 |
+
v = self.v(context).view(b, -1, n, c)
|
| 90 |
+
|
| 91 |
+
# attention bias
|
| 92 |
+
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
|
| 93 |
+
if pos_bias is not None:
|
| 94 |
+
attn_bias += pos_bias
|
| 95 |
+
if mask is not None:
|
| 96 |
+
assert mask.ndim in [2, 3]
|
| 97 |
+
mask = mask.view(b, 1, 1,
|
| 98 |
+
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
| 99 |
+
attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
|
| 100 |
+
|
| 101 |
+
# compute attention (T5 does not use scaling)
|
| 102 |
+
attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias
|
| 103 |
+
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
| 104 |
+
x = torch.einsum('bnij,bjnc->binc', attn, v)
|
| 105 |
+
|
| 106 |
+
# output
|
| 107 |
+
x = x.reshape(b, -1, n * c)
|
| 108 |
+
x = self.o(x)
|
| 109 |
+
x = self.dropout(x)
|
| 110 |
+
return x
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
class T5FeedForward(nn.Module):
|
| 114 |
+
|
| 115 |
+
def __init__(self, dim, dim_ffn, dropout=0.1):
|
| 116 |
+
super(T5FeedForward, self).__init__()
|
| 117 |
+
self.dim = dim
|
| 118 |
+
self.dim_ffn = dim_ffn
|
| 119 |
+
|
| 120 |
+
# layers
|
| 121 |
+
self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
|
| 122 |
+
self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
|
| 123 |
+
self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
|
| 124 |
+
self.dropout = nn.Dropout(dropout)
|
| 125 |
+
|
| 126 |
+
def forward(self, x):
|
| 127 |
+
x = self.fc1(x) * self.gate(x)
|
| 128 |
+
x = self.dropout(x)
|
| 129 |
+
x = self.fc2(x)
|
| 130 |
+
x = self.dropout(x)
|
| 131 |
+
return x
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
class T5SelfAttention(nn.Module):
|
| 135 |
+
def __init__(self,
|
| 136 |
+
dim,
|
| 137 |
+
dim_attn,
|
| 138 |
+
dim_ffn,
|
| 139 |
+
num_heads,
|
| 140 |
+
num_buckets,
|
| 141 |
+
shared_pos=True,
|
| 142 |
+
dropout=0.1):
|
| 143 |
+
super(T5SelfAttention, self).__init__()
|
| 144 |
+
self.dim = dim
|
| 145 |
+
self.dim_attn = dim_attn
|
| 146 |
+
self.dim_ffn = dim_ffn
|
| 147 |
+
self.num_heads = num_heads
|
| 148 |
+
self.num_buckets = num_buckets
|
| 149 |
+
self.shared_pos = shared_pos
|
| 150 |
+
|
| 151 |
+
# layers
|
| 152 |
+
self.norm1 = T5LayerNorm(dim)
|
| 153 |
+
self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
| 154 |
+
self.norm2 = T5LayerNorm(dim)
|
| 155 |
+
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
|
| 156 |
+
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
|
| 157 |
+
num_buckets, num_heads, bidirectional=True)
|
| 158 |
+
|
| 159 |
+
def forward(self, x, mask=None, pos_bias=None):
|
| 160 |
+
e = pos_bias if self.shared_pos else self.pos_embedding(
|
| 161 |
+
x.size(1), x.size(1))
|
| 162 |
+
x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
|
| 163 |
+
x = fp16_clamp(x + self.ffn(self.norm2(x)))
|
| 164 |
+
return x
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class T5CrossAttention(nn.Module):
|
| 168 |
+
def __init__(self,
|
| 169 |
+
dim,
|
| 170 |
+
dim_attn,
|
| 171 |
+
dim_ffn,
|
| 172 |
+
num_heads,
|
| 173 |
+
num_buckets,
|
| 174 |
+
shared_pos=True,
|
| 175 |
+
dropout=0.1):
|
| 176 |
+
super(T5CrossAttention, self).__init__()
|
| 177 |
+
self.dim = dim
|
| 178 |
+
self.dim_attn = dim_attn
|
| 179 |
+
self.dim_ffn = dim_ffn
|
| 180 |
+
self.num_heads = num_heads
|
| 181 |
+
self.num_buckets = num_buckets
|
| 182 |
+
self.shared_pos = shared_pos
|
| 183 |
+
|
| 184 |
+
# layers
|
| 185 |
+
self.norm1 = T5LayerNorm(dim)
|
| 186 |
+
self.self_attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
| 187 |
+
self.norm2 = T5LayerNorm(dim)
|
| 188 |
+
self.cross_attn = T5Attention(dim, dim_attn, num_heads, dropout)
|
| 189 |
+
self.norm3 = T5LayerNorm(dim)
|
| 190 |
+
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
|
| 191 |
+
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
|
| 192 |
+
num_buckets, num_heads, bidirectional=False)
|
| 193 |
+
|
| 194 |
+
def forward(self,
|
| 195 |
+
x,
|
| 196 |
+
mask=None,
|
| 197 |
+
encoder_states=None,
|
| 198 |
+
encoder_mask=None,
|
| 199 |
+
pos_bias=None):
|
| 200 |
+
e = pos_bias if self.shared_pos else self.pos_embedding(
|
| 201 |
+
x.size(1), x.size(1))
|
| 202 |
+
x = fp16_clamp(x + self.self_attn(self.norm1(x), mask=mask, pos_bias=e))
|
| 203 |
+
x = fp16_clamp(x + self.cross_attn(
|
| 204 |
+
self.norm2(x), context=encoder_states, mask=encoder_mask))
|
| 205 |
+
x = fp16_clamp(x + self.ffn(self.norm3(x)))
|
| 206 |
+
return x
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
class T5RelativeEmbedding(nn.Module):
|
| 210 |
+
def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
|
| 211 |
+
super(T5RelativeEmbedding, self).__init__()
|
| 212 |
+
self.num_buckets = num_buckets
|
| 213 |
+
self.num_heads = num_heads
|
| 214 |
+
self.bidirectional = bidirectional
|
| 215 |
+
self.max_dist = max_dist
|
| 216 |
+
|
| 217 |
+
# layers
|
| 218 |
+
self.embedding = nn.Embedding(num_buckets, num_heads)
|
| 219 |
+
|
| 220 |
+
def forward(self, lq, lk):
|
| 221 |
+
device = self.embedding.weight.device
|
| 222 |
+
# rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
|
| 223 |
+
# torch.arange(lq).unsqueeze(1).to(device)
|
| 224 |
+
if torch.device(type="meta") != device:
|
| 225 |
+
rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
|
| 226 |
+
torch.arange(lq, device=device).unsqueeze(1)
|
| 227 |
+
else:
|
| 228 |
+
rel_pos = torch.arange(lk).unsqueeze(0) - \
|
| 229 |
+
torch.arange(lq).unsqueeze(1)
|
| 230 |
+
rel_pos = self._relative_position_bucket(rel_pos)
|
| 231 |
+
rel_pos_embeds = self.embedding(rel_pos)
|
| 232 |
+
rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
|
| 233 |
+
0) # [1, N, Lq, Lk]
|
| 234 |
+
return rel_pos_embeds.contiguous()
|
| 235 |
+
|
| 236 |
+
def _relative_position_bucket(self, rel_pos):
|
| 237 |
+
# preprocess
|
| 238 |
+
if self.bidirectional:
|
| 239 |
+
num_buckets = self.num_buckets // 2
|
| 240 |
+
rel_buckets = (rel_pos > 0).long() * num_buckets
|
| 241 |
+
rel_pos = torch.abs(rel_pos)
|
| 242 |
+
else:
|
| 243 |
+
num_buckets = self.num_buckets
|
| 244 |
+
rel_buckets = 0
|
| 245 |
+
rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
|
| 246 |
+
|
| 247 |
+
# embeddings for small and large positions
|
| 248 |
+
max_exact = num_buckets // 2
|
| 249 |
+
rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
|
| 250 |
+
math.log(self.max_dist / max_exact) *
|
| 251 |
+
(num_buckets - max_exact)).long()
|
| 252 |
+
rel_pos_large = torch.min(
|
| 253 |
+
rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
|
| 254 |
+
rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
|
| 255 |
+
return rel_buckets
|
| 256 |
+
|
| 257 |
+
class WanT5EncoderModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
| 258 |
+
def __init__(self,
|
| 259 |
+
vocab,
|
| 260 |
+
dim,
|
| 261 |
+
dim_attn,
|
| 262 |
+
dim_ffn,
|
| 263 |
+
num_heads,
|
| 264 |
+
num_layers,
|
| 265 |
+
num_buckets,
|
| 266 |
+
shared_pos=True,
|
| 267 |
+
dropout=0.1):
|
| 268 |
+
super(WanT5EncoderModel, self).__init__()
|
| 269 |
+
self.dim = dim
|
| 270 |
+
self.dim_attn = dim_attn
|
| 271 |
+
self.dim_ffn = dim_ffn
|
| 272 |
+
self.num_heads = num_heads
|
| 273 |
+
self.num_layers = num_layers
|
| 274 |
+
self.num_buckets = num_buckets
|
| 275 |
+
self.shared_pos = shared_pos
|
| 276 |
+
|
| 277 |
+
# layers
|
| 278 |
+
self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
|
| 279 |
+
else nn.Embedding(vocab, dim)
|
| 280 |
+
self.pos_embedding = T5RelativeEmbedding(
|
| 281 |
+
num_buckets, num_heads, bidirectional=True) if shared_pos else None
|
| 282 |
+
self.dropout = nn.Dropout(dropout)
|
| 283 |
+
self.blocks = nn.ModuleList([
|
| 284 |
+
T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
|
| 285 |
+
shared_pos, dropout) for _ in range(num_layers)
|
| 286 |
+
])
|
| 287 |
+
self.norm = T5LayerNorm(dim)
|
| 288 |
+
|
| 289 |
+
# initialize weights
|
| 290 |
+
self.apply(init_weights)
|
| 291 |
+
|
| 292 |
+
def forward(
|
| 293 |
+
self,
|
| 294 |
+
input_ids: Optional[torch.LongTensor] = None,
|
| 295 |
+
attention_mask: Optional[torch.FloatTensor] = None,
|
| 296 |
+
):
|
| 297 |
+
x = self.token_embedding(input_ids)
|
| 298 |
+
x = self.dropout(x)
|
| 299 |
+
e = self.pos_embedding(x.size(1),
|
| 300 |
+
x.size(1)) if self.shared_pos else None
|
| 301 |
+
for block in self.blocks:
|
| 302 |
+
x = block(x, attention_mask, pos_bias=e)
|
| 303 |
+
x = self.norm(x)
|
| 304 |
+
x = self.dropout(x)
|
| 305 |
+
return (x, )
|
| 306 |
+
|
| 307 |
+
@classmethod
|
| 308 |
+
def from_pretrained(cls, pretrained_model_path, additional_kwargs={}, low_cpu_mem_usage=False, torch_dtype=torch.bfloat16):
|
| 309 |
+
def filter_kwargs(cls, kwargs):
|
| 310 |
+
import inspect
|
| 311 |
+
sig = inspect.signature(cls.__init__)
|
| 312 |
+
valid_params = set(sig.parameters.keys()) - {'self', 'cls'}
|
| 313 |
+
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
| 314 |
+
return filtered_kwargs
|
| 315 |
+
|
| 316 |
+
if low_cpu_mem_usage:
|
| 317 |
+
try:
|
| 318 |
+
import re
|
| 319 |
+
|
| 320 |
+
from diffusers import __version__ as diffusers_version
|
| 321 |
+
if diffusers_version >= "0.33.0":
|
| 322 |
+
from diffusers.models.model_loading_utils import \
|
| 323 |
+
load_model_dict_into_meta
|
| 324 |
+
else:
|
| 325 |
+
from diffusers.models.modeling_utils import \
|
| 326 |
+
load_model_dict_into_meta
|
| 327 |
+
from diffusers.utils import is_accelerate_available
|
| 328 |
+
if is_accelerate_available():
|
| 329 |
+
import accelerate
|
| 330 |
+
|
| 331 |
+
# Instantiate model with empty weights
|
| 332 |
+
with accelerate.init_empty_weights():
|
| 333 |
+
model = cls(**filter_kwargs(cls, additional_kwargs))
|
| 334 |
+
|
| 335 |
+
param_device = "cpu"
|
| 336 |
+
if pretrained_model_path.endswith(".safetensors"):
|
| 337 |
+
from safetensors.torch import load_file
|
| 338 |
+
state_dict = load_file(pretrained_model_path)
|
| 339 |
+
else:
|
| 340 |
+
state_dict = torch.load(pretrained_model_path, map_location="cpu")
|
| 341 |
+
|
| 342 |
+
if diffusers_version >= "0.33.0":
|
| 343 |
+
# Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit:
|
| 344 |
+
# https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785.
|
| 345 |
+
load_model_dict_into_meta(
|
| 346 |
+
model,
|
| 347 |
+
state_dict,
|
| 348 |
+
dtype=torch_dtype,
|
| 349 |
+
model_name_or_path=pretrained_model_path,
|
| 350 |
+
)
|
| 351 |
+
else:
|
| 352 |
+
# move the params from meta device to cpu
|
| 353 |
+
missing_keys = set(model.state_dict().keys()) - set(state_dict.keys())
|
| 354 |
+
if len(missing_keys) > 0:
|
| 355 |
+
raise ValueError(
|
| 356 |
+
f"Cannot load {cls} from {pretrained_model_path} because the following keys are"
|
| 357 |
+
f" missing: \n {', '.join(missing_keys)}. \n Please make sure to pass"
|
| 358 |
+
" `low_cpu_mem_usage=False` and `device_map=None` if you want to randomly initialize"
|
| 359 |
+
" those weights or else make sure your checkpoint file is correct."
|
| 360 |
+
)
|
| 361 |
+
|
| 362 |
+
unexpected_keys = load_model_dict_into_meta(
|
| 363 |
+
model,
|
| 364 |
+
state_dict,
|
| 365 |
+
device=param_device,
|
| 366 |
+
dtype=torch_dtype,
|
| 367 |
+
model_name_or_path=pretrained_model_path,
|
| 368 |
+
)
|
| 369 |
+
|
| 370 |
+
if cls._keys_to_ignore_on_load_unexpected is not None:
|
| 371 |
+
for pat in cls._keys_to_ignore_on_load_unexpected:
|
| 372 |
+
unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None]
|
| 373 |
+
|
| 374 |
+
if len(unexpected_keys) > 0:
|
| 375 |
+
logging.warning("Unused text encoder keys: %d", len(unexpected_keys))
|
| 376 |
+
|
| 377 |
+
return model
|
| 378 |
+
except Exception:
|
| 379 |
+
logging.warning("Falling back to regular text encoder loading")
|
| 380 |
+
|
| 381 |
+
model = cls(**filter_kwargs(cls, additional_kwargs))
|
| 382 |
+
if pretrained_model_path.endswith(".safetensors"):
|
| 383 |
+
from safetensors.torch import load_file, safe_open
|
| 384 |
+
state_dict = load_file(pretrained_model_path)
|
| 385 |
+
else:
|
| 386 |
+
state_dict = torch.load(pretrained_model_path, map_location="cpu")
|
| 387 |
+
m, u = model.load_state_dict(state_dict, strict=False)
|
| 388 |
+
model = model.to(torch_dtype)
|
| 389 |
+
return model
|
videox_fun/models/wan_transformer3d_prope.py
ADDED
|
@@ -0,0 +1,1035 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from https://github.com/Wan-Video/Wan2.1/blob/main/wan/modules/model.py
|
| 2 |
+
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
| 3 |
+
|
| 4 |
+
import glob
|
| 5 |
+
import json
|
| 6 |
+
import math
|
| 7 |
+
import os
|
| 8 |
+
import logging
|
| 9 |
+
import torch
|
| 10 |
+
import torch.cuda.amp as amp
|
| 11 |
+
import torch.nn as nn
|
| 12 |
+
|
| 13 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 14 |
+
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
| 15 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 16 |
+
from .attention_utils import attention
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def sinusoidal_embedding_1d(dim, position):
|
| 20 |
+
# preprocess
|
| 21 |
+
assert dim % 2 == 0
|
| 22 |
+
half = dim // 2
|
| 23 |
+
position = position.type(torch.float64)
|
| 24 |
+
|
| 25 |
+
# calculation
|
| 26 |
+
sinusoid = torch.outer(
|
| 27 |
+
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
|
| 28 |
+
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
| 29 |
+
return x
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
@amp.autocast(enabled=False)
|
| 33 |
+
def rope_params(max_seq_len, dim, theta=10000):
|
| 34 |
+
assert dim % 2 == 0
|
| 35 |
+
freqs = torch.outer(
|
| 36 |
+
torch.arange(max_seq_len),
|
| 37 |
+
1.0 / torch.pow(theta,
|
| 38 |
+
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
|
| 39 |
+
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
| 40 |
+
return freqs
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
# Similar to diffusers.pipelines.hunyuandit.pipeline_hunyuandit.get_resize_crop_region_for_grid
|
| 44 |
+
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
|
| 45 |
+
tw = tgt_width
|
| 46 |
+
th = tgt_height
|
| 47 |
+
h, w = src
|
| 48 |
+
r = h / w
|
| 49 |
+
if r > (th / tw):
|
| 50 |
+
resize_height = th
|
| 51 |
+
resize_width = int(round(th / h * w))
|
| 52 |
+
else:
|
| 53 |
+
resize_width = tw
|
| 54 |
+
resize_height = int(round(tw / w * h))
|
| 55 |
+
|
| 56 |
+
crop_top = int(round((th - resize_height) / 2.0))
|
| 57 |
+
crop_left = int(round((tw - resize_width) / 2.0))
|
| 58 |
+
|
| 59 |
+
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
@amp.autocast(enabled=False)
|
| 63 |
+
def rope_apply(x, grid_sizes, freqs):
|
| 64 |
+
n, c = x.size(2), x.size(3) // 2
|
| 65 |
+
|
| 66 |
+
# split freqs
|
| 67 |
+
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
| 68 |
+
|
| 69 |
+
# loop over samples
|
| 70 |
+
output = []
|
| 71 |
+
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
| 72 |
+
seq_len = f * h * w
|
| 73 |
+
|
| 74 |
+
# precompute multipliers
|
| 75 |
+
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float32).reshape(
|
| 76 |
+
seq_len, n, -1, 2))
|
| 77 |
+
freqs_i = torch.cat([
|
| 78 |
+
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
| 79 |
+
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
| 80 |
+
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
| 81 |
+
],
|
| 82 |
+
dim=-1).reshape(seq_len, 1, -1)
|
| 83 |
+
|
| 84 |
+
# apply rotary embedding
|
| 85 |
+
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
| 86 |
+
x_i = torch.cat([x_i, x[i, seq_len:]])
|
| 87 |
+
|
| 88 |
+
# append to collection
|
| 89 |
+
output.append(x_i)
|
| 90 |
+
return torch.stack(output).to(x.dtype)
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def rope_apply_qk(q, k, grid_sizes, freqs):
|
| 94 |
+
q = rope_apply(q, grid_sizes, freqs)
|
| 95 |
+
k = rope_apply(k, grid_sizes, freqs)
|
| 96 |
+
return q, k
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
class WanRMSNorm(nn.Module):
|
| 100 |
+
|
| 101 |
+
def __init__(self, dim, eps=1e-5):
|
| 102 |
+
super().__init__()
|
| 103 |
+
self.dim = dim
|
| 104 |
+
self.eps = eps
|
| 105 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 106 |
+
|
| 107 |
+
def forward(self, x):
|
| 108 |
+
r"""
|
| 109 |
+
Args:
|
| 110 |
+
x(Tensor): Shape [B, L, C]
|
| 111 |
+
"""
|
| 112 |
+
return self._norm(x.float()).type_as(x) * self.weight
|
| 113 |
+
|
| 114 |
+
def _norm(self, x):
|
| 115 |
+
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
class WanLayerNorm(nn.LayerNorm):
|
| 119 |
+
|
| 120 |
+
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
| 121 |
+
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
| 122 |
+
|
| 123 |
+
def forward(self, x):
|
| 124 |
+
r"""
|
| 125 |
+
Args:
|
| 126 |
+
x(Tensor): Shape [B, L, C]
|
| 127 |
+
"""
|
| 128 |
+
return super().forward(x.float()).type_as(x)
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
class WanSelfAttention(nn.Module):
|
| 132 |
+
|
| 133 |
+
def __init__(self,
|
| 134 |
+
dim,
|
| 135 |
+
num_heads,
|
| 136 |
+
window_size=(-1, -1),
|
| 137 |
+
qk_norm=True,
|
| 138 |
+
eps=1e-6):
|
| 139 |
+
assert dim % num_heads == 0
|
| 140 |
+
super().__init__()
|
| 141 |
+
self.dim = dim
|
| 142 |
+
self.num_heads = num_heads
|
| 143 |
+
self.head_dim = dim // num_heads
|
| 144 |
+
self.window_size = window_size
|
| 145 |
+
self.qk_norm = qk_norm
|
| 146 |
+
self.eps = eps
|
| 147 |
+
|
| 148 |
+
# layers
|
| 149 |
+
self.q = nn.Linear(dim, dim)
|
| 150 |
+
self.k = nn.Linear(dim, dim)
|
| 151 |
+
self.v = nn.Linear(dim, dim)
|
| 152 |
+
self.o = nn.Linear(dim, dim)
|
| 153 |
+
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
| 154 |
+
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
| 155 |
+
|
| 156 |
+
def forward(self, x, seq_lens, grid_sizes, freqs, dtype=torch.bfloat16, t=0, **kwargs):
|
| 157 |
+
r"""
|
| 158 |
+
Args:
|
| 159 |
+
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
| 160 |
+
seq_lens(Tensor): Shape [B]
|
| 161 |
+
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
| 162 |
+
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
| 163 |
+
"""
|
| 164 |
+
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
| 165 |
+
|
| 166 |
+
# query, key, value function
|
| 167 |
+
def qkv_fn(x):
|
| 168 |
+
q = self.norm_q(self.q(x.to(dtype))).view(b, s, n, d)
|
| 169 |
+
k = self.norm_k(self.k(x.to(dtype))).view(b, s, n, d)
|
| 170 |
+
v = self.v(x.to(dtype)).view(b, s, n, d)
|
| 171 |
+
return q, k, v
|
| 172 |
+
|
| 173 |
+
q, k, v = qkv_fn(x)
|
| 174 |
+
|
| 175 |
+
q, k = rope_apply_qk(q, k, grid_sizes, freqs)
|
| 176 |
+
|
| 177 |
+
x = attention(
|
| 178 |
+
q.to(dtype),
|
| 179 |
+
k.to(dtype),
|
| 180 |
+
v=v.to(dtype),
|
| 181 |
+
k_lens=seq_lens,
|
| 182 |
+
window_size=self.window_size)
|
| 183 |
+
x = x.to(dtype)
|
| 184 |
+
|
| 185 |
+
# output
|
| 186 |
+
x = x.flatten(2)
|
| 187 |
+
x = self.o(x)
|
| 188 |
+
return x
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
class WanT2VCrossAttention(WanSelfAttention):
|
| 192 |
+
|
| 193 |
+
def forward(self, x, context, context_lens, dtype=torch.bfloat16, t=0):
|
| 194 |
+
r"""
|
| 195 |
+
Args:
|
| 196 |
+
x(Tensor): Shape [B, L1, C]
|
| 197 |
+
context(Tensor): Shape [B, L2, C]
|
| 198 |
+
context_lens(Tensor): Shape [B]
|
| 199 |
+
"""
|
| 200 |
+
b, n, d = x.size(0), self.num_heads, self.head_dim
|
| 201 |
+
|
| 202 |
+
# compute query, key, value
|
| 203 |
+
q = self.norm_q(self.q(x.to(dtype))).view(b, -1, n, d)
|
| 204 |
+
k = self.norm_k(self.k(context.to(dtype))).view(b, -1, n, d)
|
| 205 |
+
v = self.v(context.to(dtype)).view(b, -1, n, d)
|
| 206 |
+
|
| 207 |
+
# compute attention
|
| 208 |
+
x = attention(
|
| 209 |
+
q.to(dtype),
|
| 210 |
+
k.to(dtype),
|
| 211 |
+
v.to(dtype),
|
| 212 |
+
k_lens=context_lens
|
| 213 |
+
)
|
| 214 |
+
x = x.to(dtype)
|
| 215 |
+
|
| 216 |
+
# output
|
| 217 |
+
x = x.flatten(2)
|
| 218 |
+
x = self.o(x)
|
| 219 |
+
return x
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
class WanI2VCrossAttention(WanSelfAttention):
|
| 223 |
+
|
| 224 |
+
def __init__(self,
|
| 225 |
+
dim,
|
| 226 |
+
num_heads,
|
| 227 |
+
window_size=(-1, -1),
|
| 228 |
+
qk_norm=True,
|
| 229 |
+
eps=1e-6):
|
| 230 |
+
super().__init__(dim, num_heads, window_size, qk_norm, eps)
|
| 231 |
+
|
| 232 |
+
self.k_img = nn.Linear(dim, dim)
|
| 233 |
+
self.v_img = nn.Linear(dim, dim)
|
| 234 |
+
# self.alpha = nn.Parameter(torch.zeros((1, )))
|
| 235 |
+
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
| 236 |
+
|
| 237 |
+
def forward(self, x, context, context_lens, dtype=torch.bfloat16, t=0):
|
| 238 |
+
r"""
|
| 239 |
+
Args:
|
| 240 |
+
x(Tensor): Shape [B, L1, C]
|
| 241 |
+
context(Tensor): Shape [B, L2, C]
|
| 242 |
+
context_lens(Tensor): Shape [B]
|
| 243 |
+
"""
|
| 244 |
+
context_img = context[:, :257]
|
| 245 |
+
context = context[:, 257:]
|
| 246 |
+
b, n, d = x.size(0), self.num_heads, self.head_dim
|
| 247 |
+
|
| 248 |
+
# compute query, key, value
|
| 249 |
+
q = self.norm_q(self.q(x.to(dtype))).view(b, -1, n, d)
|
| 250 |
+
k = self.norm_k(self.k(context.to(dtype))).view(b, -1, n, d)
|
| 251 |
+
v = self.v(context.to(dtype)).view(b, -1, n, d)
|
| 252 |
+
k_img = self.norm_k_img(self.k_img(context_img.to(dtype))).view(b, -1, n, d)
|
| 253 |
+
v_img = self.v_img(context_img.to(dtype)).view(b, -1, n, d)
|
| 254 |
+
|
| 255 |
+
img_x = attention(
|
| 256 |
+
q.to(dtype),
|
| 257 |
+
k_img.to(dtype),
|
| 258 |
+
v_img.to(dtype),
|
| 259 |
+
k_lens=None
|
| 260 |
+
)
|
| 261 |
+
img_x = img_x.to(dtype)
|
| 262 |
+
# compute attention
|
| 263 |
+
x = attention(
|
| 264 |
+
q.to(dtype),
|
| 265 |
+
k.to(dtype),
|
| 266 |
+
v.to(dtype),
|
| 267 |
+
k_lens=context_lens
|
| 268 |
+
)
|
| 269 |
+
x = x.to(dtype)
|
| 270 |
+
|
| 271 |
+
# output
|
| 272 |
+
x = x.flatten(2)
|
| 273 |
+
img_x = img_x.flatten(2)
|
| 274 |
+
x = x + img_x
|
| 275 |
+
x = self.o(x)
|
| 276 |
+
return x
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
class WanCrossAttention(WanSelfAttention):
|
| 280 |
+
def forward(self, x, context, context_lens, dtype=torch.bfloat16, t=0):
|
| 281 |
+
r"""
|
| 282 |
+
Args:
|
| 283 |
+
x(Tensor): Shape [B, L1, C]
|
| 284 |
+
context(Tensor): Shape [B, L2, C]
|
| 285 |
+
context_lens(Tensor): Shape [B]
|
| 286 |
+
"""
|
| 287 |
+
b, n, d = x.size(0), self.num_heads, self.head_dim
|
| 288 |
+
# compute query, key, value
|
| 289 |
+
q = self.norm_q(self.q(x.to(dtype))).view(b, -1, n, d)
|
| 290 |
+
k = self.norm_k(self.k(context.to(dtype))).view(b, -1, n, d)
|
| 291 |
+
v = self.v(context.to(dtype)).view(b, -1, n, d)
|
| 292 |
+
# compute attention
|
| 293 |
+
x = attention(q.to(dtype), k.to(dtype), v.to(dtype), k_lens=context_lens)
|
| 294 |
+
# output
|
| 295 |
+
x = x.flatten(2)
|
| 296 |
+
x = self.o(x.to(dtype))
|
| 297 |
+
return x
|
| 298 |
+
|
| 299 |
+
|
| 300 |
+
WAN_CROSSATTENTION_CLASSES = {
|
| 301 |
+
't2v_cross_attn': WanT2VCrossAttention,
|
| 302 |
+
'i2v_cross_attn': WanI2VCrossAttention,
|
| 303 |
+
'cross_attn': WanCrossAttention,
|
| 304 |
+
}
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
class WanAttentionBlock(nn.Module):
|
| 308 |
+
|
| 309 |
+
def __init__(self,
|
| 310 |
+
cross_attn_type,
|
| 311 |
+
dim,
|
| 312 |
+
ffn_dim,
|
| 313 |
+
num_heads,
|
| 314 |
+
window_size=(-1, -1),
|
| 315 |
+
qk_norm=True,
|
| 316 |
+
cross_attn_norm=False,
|
| 317 |
+
eps=1e-6,
|
| 318 |
+
**kwargs):
|
| 319 |
+
super().__init__()
|
| 320 |
+
self.dim = dim
|
| 321 |
+
self.ffn_dim = ffn_dim
|
| 322 |
+
self.num_heads = num_heads
|
| 323 |
+
self.window_size = window_size
|
| 324 |
+
self.qk_norm = qk_norm
|
| 325 |
+
self.cross_attn_norm = cross_attn_norm
|
| 326 |
+
self.eps = eps
|
| 327 |
+
|
| 328 |
+
# layers
|
| 329 |
+
self.norm1 = WanLayerNorm(dim, eps)
|
| 330 |
+
self.self_attn = WanSelfAttention(dim, num_heads, window_size, qk_norm,
|
| 331 |
+
eps)
|
| 332 |
+
self.norm3 = WanLayerNorm(
|
| 333 |
+
dim, eps,
|
| 334 |
+
elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
| 335 |
+
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](dim,
|
| 336 |
+
num_heads,
|
| 337 |
+
(-1, -1),
|
| 338 |
+
qk_norm,
|
| 339 |
+
eps)
|
| 340 |
+
self.norm2 = WanLayerNorm(dim, eps)
|
| 341 |
+
self.ffn = nn.Sequential(
|
| 342 |
+
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
| 343 |
+
nn.Linear(ffn_dim, dim))
|
| 344 |
+
|
| 345 |
+
# modulation
|
| 346 |
+
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
| 347 |
+
|
| 348 |
+
def forward(
|
| 349 |
+
self,
|
| 350 |
+
x,
|
| 351 |
+
e,
|
| 352 |
+
seq_lens,
|
| 353 |
+
grid_sizes,
|
| 354 |
+
freqs,
|
| 355 |
+
context,
|
| 356 |
+
context_lens,
|
| 357 |
+
dtype=torch.bfloat16,
|
| 358 |
+
t=0,
|
| 359 |
+
group=None,
|
| 360 |
+
):
|
| 361 |
+
r"""
|
| 362 |
+
Args:
|
| 363 |
+
x(Tensor): Shape [B, L, C]
|
| 364 |
+
e(Tensor): Shape [B, 6, C]
|
| 365 |
+
seq_lens(Tensor): Shape [B], length of each sequence in batch
|
| 366 |
+
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
| 367 |
+
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
| 368 |
+
"""
|
| 369 |
+
if e.dim() > 3:
|
| 370 |
+
e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
|
| 371 |
+
e = [e.squeeze(2) for e in e]
|
| 372 |
+
else:
|
| 373 |
+
e = (self.modulation + e).chunk(6, dim=1)
|
| 374 |
+
|
| 375 |
+
# self-attention (RoPE branch)
|
| 376 |
+
temp_x = self.norm1(x) * (1 + e[1]) + e[0]
|
| 377 |
+
temp_x = temp_x.to(dtype)
|
| 378 |
+
|
| 379 |
+
y = self.self_attn(temp_x, seq_lens, grid_sizes, freqs, dtype, group=group)
|
| 380 |
+
|
| 381 |
+
x = x + y * e[2]
|
| 382 |
+
|
| 383 |
+
# cross-attention & ffn function
|
| 384 |
+
def cross_attn_ffn(x, context, context_lens, e):
|
| 385 |
+
# cross-attention
|
| 386 |
+
x = x + self.cross_attn(self.norm3(x), context, context_lens, dtype)
|
| 387 |
+
|
| 388 |
+
# ffn function
|
| 389 |
+
temp_x = self.norm2(x) * (1 + e[4]) + e[3]
|
| 390 |
+
temp_x = temp_x.to(dtype)
|
| 391 |
+
|
| 392 |
+
y = self.ffn(temp_x)
|
| 393 |
+
x = x + y * e[5]
|
| 394 |
+
return x
|
| 395 |
+
|
| 396 |
+
x = cross_attn_ffn(x, context, context_lens, e)
|
| 397 |
+
return x
|
| 398 |
+
|
| 399 |
+
|
| 400 |
+
class Head(nn.Module):
|
| 401 |
+
|
| 402 |
+
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
|
| 403 |
+
super().__init__()
|
| 404 |
+
self.dim = dim
|
| 405 |
+
self.out_dim = out_dim
|
| 406 |
+
self.patch_size = patch_size
|
| 407 |
+
self.eps = eps
|
| 408 |
+
|
| 409 |
+
# layers
|
| 410 |
+
out_dim = math.prod(patch_size) * out_dim
|
| 411 |
+
self.norm = WanLayerNorm(dim, eps)
|
| 412 |
+
self.head = nn.Linear(dim, out_dim)
|
| 413 |
+
|
| 414 |
+
# modulation
|
| 415 |
+
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
| 416 |
+
|
| 417 |
+
def forward(self, x, e):
|
| 418 |
+
r"""
|
| 419 |
+
Args:
|
| 420 |
+
x(Tensor): Shape [B, L1, C]
|
| 421 |
+
e(Tensor): Shape [B, C]
|
| 422 |
+
"""
|
| 423 |
+
if e.dim() > 2:
|
| 424 |
+
e = (self.modulation.unsqueeze(0) + e.unsqueeze(2)).chunk(2, dim=2)
|
| 425 |
+
e = [e.squeeze(2) for e in e]
|
| 426 |
+
else:
|
| 427 |
+
e = (self.modulation + e.unsqueeze(1)).chunk(2, dim=1)
|
| 428 |
+
|
| 429 |
+
x = (self.head(self.norm(x) * (1 + e[1]) + e[0]))
|
| 430 |
+
return x
|
| 431 |
+
|
| 432 |
+
|
| 433 |
+
class MLPProj(torch.nn.Module):
|
| 434 |
+
|
| 435 |
+
def __init__(self, in_dim, out_dim):
|
| 436 |
+
super().__init__()
|
| 437 |
+
|
| 438 |
+
self.proj = torch.nn.Sequential(
|
| 439 |
+
torch.nn.LayerNorm(in_dim), torch.nn.Linear(in_dim, in_dim),
|
| 440 |
+
torch.nn.GELU(), torch.nn.Linear(in_dim, out_dim),
|
| 441 |
+
torch.nn.LayerNorm(out_dim))
|
| 442 |
+
|
| 443 |
+
def forward(self, image_embeds):
|
| 444 |
+
clip_extra_context_tokens = self.proj(image_embeds)
|
| 445 |
+
return clip_extra_context_tokens
|
| 446 |
+
|
| 447 |
+
|
| 448 |
+
|
| 449 |
+
class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
| 450 |
+
r"""
|
| 451 |
+
Wan diffusion backbone supporting both text-to-video and image-to-video.
|
| 452 |
+
"""
|
| 453 |
+
|
| 454 |
+
# ignore_for_config = [
|
| 455 |
+
# 'patch_size', 'cross_attn_norm', 'qk_norm', 'text_dim', 'window_size'
|
| 456 |
+
# ]
|
| 457 |
+
# _no_split_modules = ['WanAttentionBlock']
|
| 458 |
+
@register_to_config
|
| 459 |
+
def __init__(
|
| 460 |
+
self,
|
| 461 |
+
model_type='t2v',
|
| 462 |
+
patch_size=(1, 2, 2),
|
| 463 |
+
text_len=512,
|
| 464 |
+
in_dim=16,
|
| 465 |
+
dim=2048,
|
| 466 |
+
ffn_dim=8192,
|
| 467 |
+
freq_dim=256,
|
| 468 |
+
text_dim=4096,
|
| 469 |
+
out_dim=16,
|
| 470 |
+
num_heads=16,
|
| 471 |
+
num_layers=32,
|
| 472 |
+
window_size=(-1, -1),
|
| 473 |
+
qk_norm=True,
|
| 474 |
+
cross_attn_norm=True,
|
| 475 |
+
eps=1e-6,
|
| 476 |
+
in_channels=16,
|
| 477 |
+
hidden_size=2048,
|
| 478 |
+
add_control_adapter=True,
|
| 479 |
+
control_adapter_type="baseline",
|
| 480 |
+
in_dim_control_adapter=24,
|
| 481 |
+
downscale_factor_control_adapter=8,
|
| 482 |
+
add_ref_conv=False,
|
| 483 |
+
in_dim_ref_conv=16,
|
| 484 |
+
cross_attn_type=None,
|
| 485 |
+
cam_method='prope',
|
| 486 |
+
attn_compress=2,
|
| 487 |
+
cam_self_attn_layers=None,
|
| 488 |
+
traning=False,
|
| 489 |
+
):
|
| 490 |
+
r"""
|
| 491 |
+
Initialize the diffusion model backbone.
|
| 492 |
+
|
| 493 |
+
Args:
|
| 494 |
+
model_type (`str`, *optional*, defaults to 't2v'):
|
| 495 |
+
Model variant - 't2v' (text-to-video) or 'i2v' (image-to-video)
|
| 496 |
+
patch_size (`tuple`, *optional*, defaults to (1, 2, 2)):
|
| 497 |
+
3D patch dimensions for video embedding (t_patch, h_patch, w_patch)
|
| 498 |
+
text_len (`int`, *optional*, defaults to 512):
|
| 499 |
+
Fixed length for text embeddings
|
| 500 |
+
in_dim (`int`, *optional*, defaults to 16):
|
| 501 |
+
Input video channels (C_in)
|
| 502 |
+
dim (`int`, *optional*, defaults to 2048):
|
| 503 |
+
Hidden dimension of the transformer
|
| 504 |
+
ffn_dim (`int`, *optional*, defaults to 8192):
|
| 505 |
+
Intermediate dimension in feed-forward network
|
| 506 |
+
freq_dim (`int`, *optional*, defaults to 256):
|
| 507 |
+
Dimension for sinusoidal time embeddings
|
| 508 |
+
text_dim (`int`, *optional*, defaults to 4096):
|
| 509 |
+
Input dimension for text embeddings
|
| 510 |
+
out_dim (`int`, *optional*, defaults to 16):
|
| 511 |
+
Output video channels (C_out)
|
| 512 |
+
num_heads (`int`, *optional*, defaults to 16):
|
| 513 |
+
Number of attention heads
|
| 514 |
+
num_layers (`int`, *optional*, defaults to 32):
|
| 515 |
+
Number of transformer blocks
|
| 516 |
+
window_size (`tuple`, *optional*, defaults to (-1, -1)):
|
| 517 |
+
Window size for local attention (-1 indicates global attention)
|
| 518 |
+
qk_norm (`bool`, *optional*, defaults to True):
|
| 519 |
+
Enable query/key normalization
|
| 520 |
+
cross_attn_norm (`bool`, *optional*, defaults to False):
|
| 521 |
+
Enable cross-attention normalization
|
| 522 |
+
eps (`float`, *optional*, defaults to 1e-6):
|
| 523 |
+
Epsilon value for normalization layers
|
| 524 |
+
"""
|
| 525 |
+
|
| 526 |
+
super().__init__()
|
| 527 |
+
|
| 528 |
+
# assert model_type in ['t2v', 'i2v', 'ti2v']
|
| 529 |
+
self.model_type = model_type
|
| 530 |
+
|
| 531 |
+
self.patch_size = patch_size
|
| 532 |
+
self.text_len = text_len
|
| 533 |
+
self.in_dim = in_dim
|
| 534 |
+
self.dim = dim
|
| 535 |
+
self.ffn_dim = ffn_dim
|
| 536 |
+
self.freq_dim = freq_dim
|
| 537 |
+
self.text_dim = text_dim
|
| 538 |
+
self.out_dim = out_dim
|
| 539 |
+
self.num_heads = num_heads
|
| 540 |
+
self.num_layers = num_layers
|
| 541 |
+
self.window_size = window_size
|
| 542 |
+
self.qk_norm = qk_norm
|
| 543 |
+
self.cross_attn_norm = cross_attn_norm
|
| 544 |
+
self.eps = eps
|
| 545 |
+
|
| 546 |
+
# embeddings
|
| 547 |
+
self.patch_embedding = nn.Conv3d(
|
| 548 |
+
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
| 549 |
+
self.text_embedding = nn.Sequential(
|
| 550 |
+
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
| 551 |
+
nn.Linear(dim, dim))
|
| 552 |
+
|
| 553 |
+
self.time_embedding = nn.Sequential(
|
| 554 |
+
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
| 555 |
+
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
| 556 |
+
|
| 557 |
+
# blocks
|
| 558 |
+
if cross_attn_type is None:
|
| 559 |
+
if model_type == "t2v":
|
| 560 |
+
cross_attn_type = 't2v_cross_attn'
|
| 561 |
+
else:
|
| 562 |
+
cross_attn_type = 'i2v_cross_attn'
|
| 563 |
+
# else:
|
| 564 |
+
# cross_attn_type = "cross_attn"
|
| 565 |
+
|
| 566 |
+
attn_class = WanAttentionBlock # by default use WanAttentionBlock
|
| 567 |
+
self.control_adapter = None
|
| 568 |
+
|
| 569 |
+
self.blocks = nn.ModuleList([
|
| 570 |
+
attn_class(cross_attn_type, dim, ffn_dim, num_heads,
|
| 571 |
+
window_size, qk_norm, cross_attn_norm, eps)
|
| 572 |
+
|
| 573 |
+
for _ in range(num_layers)
|
| 574 |
+
])
|
| 575 |
+
|
| 576 |
+
# head
|
| 577 |
+
self.head = Head(dim, out_dim, patch_size, eps)
|
| 578 |
+
|
| 579 |
+
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
| 580 |
+
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
| 581 |
+
d = dim // num_heads
|
| 582 |
+
self.d = d
|
| 583 |
+
self.dim = dim
|
| 584 |
+
self.freqs = torch.cat(
|
| 585 |
+
[
|
| 586 |
+
rope_params(1024, d - 4 * (d // 6)),
|
| 587 |
+
rope_params(1024, 2 * (d // 6)),
|
| 588 |
+
rope_params(1024, 2 * (d // 6))
|
| 589 |
+
],
|
| 590 |
+
dim=1
|
| 591 |
+
)
|
| 592 |
+
|
| 593 |
+
if model_type == 'i2v':
|
| 594 |
+
self.img_emb = MLPProj(1280, dim)
|
| 595 |
+
|
| 596 |
+
if add_ref_conv:
|
| 597 |
+
self.ref_conv = nn.Conv2d(in_dim_ref_conv, dim, kernel_size=patch_size[1:], stride=patch_size[1:])
|
| 598 |
+
else:
|
| 599 |
+
self.ref_conv = None
|
| 600 |
+
|
| 601 |
+
self.customize_initialize_weights()
|
| 602 |
+
|
| 603 |
+
def customize_initialize_weights(self):
|
| 604 |
+
self.init_weights()
|
| 605 |
+
|
| 606 |
+
def forward(
|
| 607 |
+
self,
|
| 608 |
+
x,
|
| 609 |
+
t,
|
| 610 |
+
context,
|
| 611 |
+
seq_len,
|
| 612 |
+
clip_fea=None,
|
| 613 |
+
y=None,
|
| 614 |
+
y_camera=None,
|
| 615 |
+
full_ref=None,
|
| 616 |
+
subject_ref=None,
|
| 617 |
+
cond_flag=True,
|
| 618 |
+
dtype=torch.bfloat16,
|
| 619 |
+
):
|
| 620 |
+
r"""
|
| 621 |
+
Forward pass through the diffusion model
|
| 622 |
+
|
| 623 |
+
Args:
|
| 624 |
+
x (List[Tensor]):
|
| 625 |
+
List of input video tensors, each with shape [C_in, F, H, W]
|
| 626 |
+
t (Tensor):
|
| 627 |
+
Diffusion timesteps tensor of shape [B]
|
| 628 |
+
context (List[Tensor]):
|
| 629 |
+
List of text embeddings each with shape [L, C]
|
| 630 |
+
seq_len (`int`):
|
| 631 |
+
Maximum sequence length for positional encoding
|
| 632 |
+
clip_fea (Tensor, *optional*):
|
| 633 |
+
CLIP image features for image-to-video mode
|
| 634 |
+
y (List[Tensor], *optional*):
|
| 635 |
+
Conditional video inputs for image-to-video mode, same shape as x
|
| 636 |
+
cond_flag (`bool`, *optional*, defaults to True):
|
| 637 |
+
Flag to indicate whether to forward the condition input
|
| 638 |
+
|
| 639 |
+
Returns:
|
| 640 |
+
List[Tensor]:
|
| 641 |
+
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
| 642 |
+
"""
|
| 643 |
+
|
| 644 |
+
device = self.patch_embedding.weight.device
|
| 645 |
+
dtype = x[0].dtype
|
| 646 |
+
if self.freqs.device != device and torch.device(type="meta") != device:
|
| 647 |
+
self.freqs = self.freqs.to(device)
|
| 648 |
+
|
| 649 |
+
if y is not None:
|
| 650 |
+
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
| 651 |
+
|
| 652 |
+
# embeddings
|
| 653 |
+
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
| 654 |
+
grid_sizes = torch.stack(
|
| 655 |
+
[torch.tensor(u.shape[2:], dtype=torch.long, device=device) for u in x])
|
| 656 |
+
|
| 657 |
+
x = [u.flatten(2).transpose(1, 2) for u in x]
|
| 658 |
+
|
| 659 |
+
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long, device=device)
|
| 660 |
+
assert seq_lens.max() <= seq_len
|
| 661 |
+
x = torch.cat([
|
| 662 |
+
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
| 663 |
+
dim=1) for u in x
|
| 664 |
+
])
|
| 665 |
+
|
| 666 |
+
# time embeddings
|
| 667 |
+
with amp.autocast(dtype=torch.float32):
|
| 668 |
+
if t.dim() != 1:
|
| 669 |
+
if t.size(1) < seq_len:
|
| 670 |
+
pad_size = seq_len - t.size(1)
|
| 671 |
+
last_elements = t[:, -1].unsqueeze(1)
|
| 672 |
+
padding = last_elements.repeat(1, pad_size)
|
| 673 |
+
t = torch.cat([t, padding], dim=1)
|
| 674 |
+
bt = t.size(0)
|
| 675 |
+
ft = t.flatten()
|
| 676 |
+
e = self.time_embedding(
|
| 677 |
+
sinusoidal_embedding_1d(self.freq_dim,
|
| 678 |
+
ft).unflatten(0, (bt, seq_len)).float())
|
| 679 |
+
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
| 680 |
+
else:
|
| 681 |
+
e = self.time_embedding(
|
| 682 |
+
sinusoidal_embedding_1d(self.freq_dim, t).float())
|
| 683 |
+
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
| 684 |
+
|
| 685 |
+
# context
|
| 686 |
+
context_lens = None
|
| 687 |
+
context = self.text_embedding(
|
| 688 |
+
torch.stack([
|
| 689 |
+
torch.cat(
|
| 690 |
+
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
| 691 |
+
for u in context
|
| 692 |
+
]))
|
| 693 |
+
|
| 694 |
+
# if clip_fea is not None:
|
| 695 |
+
# context_clip = self.img_emb(clip_fea) # bs x 257 x dim
|
| 696 |
+
# context = torch.concat([context_clip, context], dim=1)
|
| 697 |
+
|
| 698 |
+
for block in self.blocks:
|
| 699 |
+
x = block(
|
| 700 |
+
x, e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes,
|
| 701 |
+
freqs=self.freqs, context=context, context_lens=context_lens,
|
| 702 |
+
dtype=dtype, t=t, group=None,
|
| 703 |
+
)
|
| 704 |
+
|
| 705 |
+
x = self.head(x, e)
|
| 706 |
+
|
| 707 |
+
return torch.stack(self.unpatchify(x, grid_sizes))
|
| 708 |
+
|
| 709 |
+
|
| 710 |
+
def unpatchify(self, x, grid_sizes):
|
| 711 |
+
r"""
|
| 712 |
+
Reconstruct video tensors from patch embeddings.
|
| 713 |
+
|
| 714 |
+
Args:
|
| 715 |
+
x (List[Tensor]):
|
| 716 |
+
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
| 717 |
+
grid_sizes (Tensor):
|
| 718 |
+
Original spatial-temporal grid dimensions before patching,
|
| 719 |
+
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
|
| 720 |
+
|
| 721 |
+
Returns:
|
| 722 |
+
List[Tensor]:
|
| 723 |
+
Reconstructed video tensors with shape [C_out, F, H / 8, W / 8]
|
| 724 |
+
"""
|
| 725 |
+
|
| 726 |
+
c = self.out_dim
|
| 727 |
+
out = []
|
| 728 |
+
for u, v in zip(x, grid_sizes.tolist()):
|
| 729 |
+
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
| 730 |
+
u = torch.einsum('fhwpqrc->cfphqwr', u)
|
| 731 |
+
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
| 732 |
+
out.append(u)
|
| 733 |
+
return out
|
| 734 |
+
|
| 735 |
+
def init_weights(self):
|
| 736 |
+
r"""
|
| 737 |
+
Initialize model parameters using Xavier initialization.
|
| 738 |
+
"""
|
| 739 |
+
|
| 740 |
+
# basic init
|
| 741 |
+
for m in self.modules():
|
| 742 |
+
if isinstance(m, nn.Linear):
|
| 743 |
+
nn.init.xavier_uniform_(m.weight)
|
| 744 |
+
if m.bias is not None:
|
| 745 |
+
nn.init.zeros_(m.bias)
|
| 746 |
+
|
| 747 |
+
# init embeddings
|
| 748 |
+
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
| 749 |
+
for m in self.text_embedding.modules():
|
| 750 |
+
if isinstance(m, nn.Linear):
|
| 751 |
+
nn.init.normal_(m.weight, std=.02)
|
| 752 |
+
for m in self.time_embedding.modules():
|
| 753 |
+
if isinstance(m, nn.Linear):
|
| 754 |
+
nn.init.normal_(m.weight, std=.02)
|
| 755 |
+
|
| 756 |
+
# init output layer
|
| 757 |
+
nn.init.zeros_(self.head.head.weight)
|
| 758 |
+
|
| 759 |
+
@classmethod
|
| 760 |
+
def from_pretrained(
|
| 761 |
+
cls, pretrained_model_path, subfolder=None, transformer_additional_kwargs={},
|
| 762 |
+
low_cpu_mem_usage=False, torch_dtype=torch.bfloat16
|
| 763 |
+
):
|
| 764 |
+
if subfolder is not None:
|
| 765 |
+
pretrained_model_path = os.path.join(pretrained_model_path, subfolder)
|
| 766 |
+
config_file = os.path.join(pretrained_model_path, 'config.json')
|
| 767 |
+
if not os.path.isfile(config_file):
|
| 768 |
+
raise RuntimeError(f"{config_file} does not exist")
|
| 769 |
+
with open(config_file, "r") as f:
|
| 770 |
+
config = json.load(f)
|
| 771 |
+
|
| 772 |
+
from diffusers.utils import WEIGHTS_NAME
|
| 773 |
+
model_file = os.path.join(pretrained_model_path, WEIGHTS_NAME)
|
| 774 |
+
model_file_safetensors = model_file.replace(".bin", ".safetensors")
|
| 775 |
+
|
| 776 |
+
if "dict_mapping" in transformer_additional_kwargs.keys():
|
| 777 |
+
for key in transformer_additional_kwargs["dict_mapping"]:
|
| 778 |
+
transformer_additional_kwargs[transformer_additional_kwargs["dict_mapping"][key]] = config[key]
|
| 779 |
+
|
| 780 |
+
if low_cpu_mem_usage:
|
| 781 |
+
try:
|
| 782 |
+
import re
|
| 783 |
+
|
| 784 |
+
from diffusers import __version__ as diffusers_version
|
| 785 |
+
if diffusers_version >= "0.33.0":
|
| 786 |
+
from diffusers.models.model_loading_utils import \
|
| 787 |
+
load_model_dict_into_meta
|
| 788 |
+
else:
|
| 789 |
+
from diffusers.models.modeling_utils import \
|
| 790 |
+
load_model_dict_into_meta
|
| 791 |
+
from diffusers.utils import is_accelerate_available
|
| 792 |
+
if is_accelerate_available():
|
| 793 |
+
import accelerate
|
| 794 |
+
|
| 795 |
+
# Instantiate model with empty weights
|
| 796 |
+
with accelerate.init_empty_weights():
|
| 797 |
+
model = cls.from_config(config, **transformer_additional_kwargs)
|
| 798 |
+
|
| 799 |
+
param_device = "cpu"
|
| 800 |
+
if os.path.exists(model_file):
|
| 801 |
+
state_dict = torch.load(model_file, map_location="cpu")
|
| 802 |
+
elif os.path.exists(model_file_safetensors):
|
| 803 |
+
from safetensors.torch import load_file, safe_open
|
| 804 |
+
state_dict = load_file(model_file_safetensors)
|
| 805 |
+
else:
|
| 806 |
+
from safetensors.torch import load_file, safe_open
|
| 807 |
+
model_files_safetensors = glob.glob(os.path.join(pretrained_model_path, "*.safetensors"))
|
| 808 |
+
state_dict = {}
|
| 809 |
+
for _model_file_safetensors in model_files_safetensors:
|
| 810 |
+
_state_dict = load_file(_model_file_safetensors)
|
| 811 |
+
for key in _state_dict:
|
| 812 |
+
state_dict[key] = _state_dict[key]
|
| 813 |
+
|
| 814 |
+
if model.state_dict()['patch_embedding.weight'].size() != state_dict['patch_embedding.weight'].size():
|
| 815 |
+
model.state_dict()['patch_embedding.weight'][:, :state_dict['patch_embedding.weight'].size()[1], :, :] = state_dict['patch_embedding.weight'][:, :model.state_dict()['patch_embedding.weight'].size()[1], :, :]
|
| 816 |
+
model.state_dict()['patch_embedding.weight'][:, state_dict['patch_embedding.weight'].size()[1]:, :, :] = 0
|
| 817 |
+
state_dict['patch_embedding.weight'] = model.state_dict()['patch_embedding.weight']
|
| 818 |
+
|
| 819 |
+
filtered_state_dict = {}
|
| 820 |
+
for key in state_dict:
|
| 821 |
+
if key in model.state_dict() and model.state_dict()[key].size() == state_dict[key].size():
|
| 822 |
+
filtered_state_dict[key] = state_dict[key]
|
| 823 |
+
|
| 824 |
+
model_keys = set(model.state_dict().keys())
|
| 825 |
+
loaded_keys = set(filtered_state_dict.keys())
|
| 826 |
+
missing_keys = model_keys - loaded_keys
|
| 827 |
+
|
| 828 |
+
def initialize_missing_parameters(missing_keys, model_state_dict, torch_dtype=None):
|
| 829 |
+
initialized_dict = {}
|
| 830 |
+
|
| 831 |
+
with torch.no_grad():
|
| 832 |
+
for key in missing_keys:
|
| 833 |
+
param_shape = model_state_dict[key].shape
|
| 834 |
+
param_dtype = torch_dtype if torch_dtype is not None else model_state_dict[key].dtype
|
| 835 |
+
if 'weight' in key:
|
| 836 |
+
if any(norm_type in key for norm_type in ['norm', 'ln_', 'layer_norm', 'group_norm', 'batch_norm']):
|
| 837 |
+
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
|
| 838 |
+
elif 'embedding' in key or 'embed' in key:
|
| 839 |
+
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
|
| 840 |
+
elif 'head' in key or 'output' in key or 'proj_out' in key or 'project_out' in key or "out" in key or 'out_proj' in key:
|
| 841 |
+
logging.info(f"Initialize {key} with zeros")
|
| 842 |
+
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
|
| 843 |
+
elif len(param_shape) >= 2:
|
| 844 |
+
initialized_dict[key] = torch.empty(param_shape, dtype=param_dtype)
|
| 845 |
+
nn.init.xavier_uniform_(initialized_dict[key])
|
| 846 |
+
else:
|
| 847 |
+
initialized_dict[key] = torch.randn(param_shape, dtype=param_dtype) * 0.02
|
| 848 |
+
elif 'bias' in key:
|
| 849 |
+
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
|
| 850 |
+
elif 'running_mean' in key:
|
| 851 |
+
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
|
| 852 |
+
elif 'running_var' in key:
|
| 853 |
+
initialized_dict[key] = torch.ones(param_shape, dtype=param_dtype)
|
| 854 |
+
elif 'num_batches_tracked' in key:
|
| 855 |
+
initialized_dict[key] = torch.zeros(param_shape, dtype=torch.long)
|
| 856 |
+
else:
|
| 857 |
+
initialized_dict[key] = torch.zeros(param_shape, dtype=param_dtype)
|
| 858 |
+
|
| 859 |
+
return initialized_dict
|
| 860 |
+
|
| 861 |
+
if missing_keys:
|
| 862 |
+
initialized_params = initialize_missing_parameters(
|
| 863 |
+
missing_keys,
|
| 864 |
+
model.state_dict(),
|
| 865 |
+
torch_dtype
|
| 866 |
+
)
|
| 867 |
+
filtered_state_dict.update(initialized_params)
|
| 868 |
+
# logging.info(f"Weights of Camera_adapter: {initialized_params}")
|
| 869 |
+
|
| 870 |
+
if diffusers_version >= "0.33.0":
|
| 871 |
+
# Diffusers has refactored `load_model_dict_into_meta` since version 0.33.0 in this commit:
|
| 872 |
+
# https://github.com/huggingface/diffusers/commit/f5929e03060d56063ff34b25a8308833bec7c785.
|
| 873 |
+
load_model_dict_into_meta(
|
| 874 |
+
model,
|
| 875 |
+
filtered_state_dict,
|
| 876 |
+
dtype=torch_dtype,
|
| 877 |
+
model_name_or_path=pretrained_model_path,
|
| 878 |
+
)
|
| 879 |
+
else:
|
| 880 |
+
model._convert_deprecated_attention_blocks(filtered_state_dict)
|
| 881 |
+
unexpected_keys = load_model_dict_into_meta(
|
| 882 |
+
model,
|
| 883 |
+
filtered_state_dict,
|
| 884 |
+
device=param_device,
|
| 885 |
+
dtype=torch_dtype,
|
| 886 |
+
model_name_or_path=pretrained_model_path,
|
| 887 |
+
)
|
| 888 |
+
|
| 889 |
+
if cls._keys_to_ignore_on_load_unexpected is not None:
|
| 890 |
+
for pat in cls._keys_to_ignore_on_load_unexpected:
|
| 891 |
+
unexpected_keys = [k for k in unexpected_keys if re.search(pat, k) is None]
|
| 892 |
+
|
| 893 |
+
if len(unexpected_keys) > 0:
|
| 894 |
+
logging.warning("Unused checkpoint keys: %d", len(unexpected_keys))
|
| 895 |
+
|
| 896 |
+
return model
|
| 897 |
+
except Exception:
|
| 898 |
+
logging.warning("Falling back to regular transformer loading")
|
| 899 |
+
|
| 900 |
+
model = cls.from_config(config, **transformer_additional_kwargs)
|
| 901 |
+
if os.path.exists(model_file):
|
| 902 |
+
state_dict = torch.load(model_file, map_location="cpu")
|
| 903 |
+
elif os.path.exists(model_file_safetensors):
|
| 904 |
+
from safetensors.torch import load_file, safe_open
|
| 905 |
+
state_dict = load_file(model_file_safetensors)
|
| 906 |
+
else:
|
| 907 |
+
from safetensors.torch import load_file, safe_open
|
| 908 |
+
model_files_safetensors = glob.glob(os.path.join(pretrained_model_path, "*.safetensors"))
|
| 909 |
+
state_dict = {}
|
| 910 |
+
for _model_file_safetensors in model_files_safetensors:
|
| 911 |
+
_state_dict = load_file(_model_file_safetensors)
|
| 912 |
+
for key in _state_dict:
|
| 913 |
+
state_dict[key] = _state_dict[key]
|
| 914 |
+
|
| 915 |
+
if model.state_dict()['patch_embedding.weight'].size() != state_dict['patch_embedding.weight'].size():
|
| 916 |
+
model.state_dict()['patch_embedding.weight'][:, :state_dict['patch_embedding.weight'].size()[1], :, :] = state_dict['patch_embedding.weight'][:, :model.state_dict()['patch_embedding.weight'].size()[1], :, :]
|
| 917 |
+
model.state_dict()['patch_embedding.weight'][:, state_dict['patch_embedding.weight'].size()[1]:, :, :] = 0
|
| 918 |
+
state_dict['patch_embedding.weight'] = model.state_dict()['patch_embedding.weight']
|
| 919 |
+
|
| 920 |
+
tmp_state_dict = {}
|
| 921 |
+
for key in state_dict:
|
| 922 |
+
if key in model.state_dict().keys() and model.state_dict()[key].size() == state_dict[key].size():
|
| 923 |
+
tmp_state_dict[key] = state_dict[key]
|
| 924 |
+
|
| 925 |
+
state_dict = tmp_state_dict
|
| 926 |
+
|
| 927 |
+
m, u = model.load_state_dict(state_dict, strict=False)
|
| 928 |
+
model = model.to(torch_dtype)
|
| 929 |
+
return model
|
| 930 |
+
|
| 931 |
+
|
| 932 |
+
class Wan2_2Transformer3DModel(WanTransformer3DModel):
|
| 933 |
+
r"""
|
| 934 |
+
Wan diffusion backbone supporting both text-to-video and image-to-video.
|
| 935 |
+
"""
|
| 936 |
+
|
| 937 |
+
# ignore_for_config = [
|
| 938 |
+
# 'patch_size', 'cross_attn_norm', 'qk_norm', 'text_dim', 'window_size'
|
| 939 |
+
# ]
|
| 940 |
+
# _no_split_modules = ['WanAttentionBlock']
|
| 941 |
+
def __init__(
|
| 942 |
+
self,
|
| 943 |
+
model_type='t2v',
|
| 944 |
+
patch_size=(1, 2, 2),
|
| 945 |
+
text_len=512,
|
| 946 |
+
in_dim=16,
|
| 947 |
+
dim=2048,
|
| 948 |
+
ffn_dim=8192,
|
| 949 |
+
freq_dim=256,
|
| 950 |
+
text_dim=4096,
|
| 951 |
+
out_dim=16,
|
| 952 |
+
num_heads=16,
|
| 953 |
+
num_layers=32,
|
| 954 |
+
window_size=(-1, -1),
|
| 955 |
+
qk_norm=True,
|
| 956 |
+
cross_attn_norm=True,
|
| 957 |
+
eps=1e-6,
|
| 958 |
+
in_channels=16,
|
| 959 |
+
hidden_size=2048,
|
| 960 |
+
add_control_adapter=False,
|
| 961 |
+
in_dim_control_adapter=24,
|
| 962 |
+
downscale_factor_control_adapter=8,
|
| 963 |
+
add_ref_conv=False,
|
| 964 |
+
control_adapter_type="baseline",
|
| 965 |
+
in_dim_ref_conv=16,
|
| 966 |
+
cam_method='prope',
|
| 967 |
+
attn_compress=2,
|
| 968 |
+
cam_self_attn_layers=None,
|
| 969 |
+
):
|
| 970 |
+
r"""
|
| 971 |
+
Initialize the diffusion model backbone.
|
| 972 |
+
Args:
|
| 973 |
+
model_type (`str`, *optional*, defaults to 't2v'):
|
| 974 |
+
Model variant - 't2v' (text-to-video) or 'i2v' (image-to-video)
|
| 975 |
+
patch_size (`tuple`, *optional*, defaults to (1, 2, 2)):
|
| 976 |
+
3D patch dimensions for video embedding (t_patch, h_patch, w_patch)
|
| 977 |
+
text_len (`int`, *optional*, defaults to 512):
|
| 978 |
+
Fixed length for text embeddings
|
| 979 |
+
in_dim (`int`, *optional*, defaults to 16):
|
| 980 |
+
Input video channels (C_in)
|
| 981 |
+
dim (`int`, *optional*, defaults to 2048):
|
| 982 |
+
Hidden dimension of the transformer
|
| 983 |
+
ffn_dim (`int`, *optional*, defaults to 8192):
|
| 984 |
+
Intermediate dimension in feed-forward network
|
| 985 |
+
freq_dim (`int`, *optional*, defaults to 256):
|
| 986 |
+
Dimension for sinusoidal time embeddings
|
| 987 |
+
text_dim (`int`, *optional*, defaults to 4096):
|
| 988 |
+
Input dimension for text embeddings
|
| 989 |
+
out_dim (`int`, *optional*, defaults to 16):
|
| 990 |
+
Output video channels (C_out)
|
| 991 |
+
num_heads (`int`, *optional*, defaults to 16):
|
| 992 |
+
Number of attention heads
|
| 993 |
+
num_layers (`int`, *optional*, defaults to 32):
|
| 994 |
+
Number of transformer blocks
|
| 995 |
+
window_size (`tuple`, *optional*, defaults to (-1, -1)):
|
| 996 |
+
Window size for local attention (-1 indicates global attention)
|
| 997 |
+
qk_norm (`bool`, *optional*, defaults to True):
|
| 998 |
+
Enable query/key normalization
|
| 999 |
+
cross_attn_norm (`bool`, *optional*, defaults to False):
|
| 1000 |
+
Enable cross-attention normalization
|
| 1001 |
+
eps (`float`, *optional*, defaults to 1e-6):
|
| 1002 |
+
Epsilon value for normalization layers
|
| 1003 |
+
"""
|
| 1004 |
+
super().__init__(
|
| 1005 |
+
model_type=model_type,
|
| 1006 |
+
patch_size=patch_size,
|
| 1007 |
+
text_len=text_len,
|
| 1008 |
+
in_dim=in_dim,
|
| 1009 |
+
dim=dim,
|
| 1010 |
+
ffn_dim=ffn_dim,
|
| 1011 |
+
freq_dim=freq_dim,
|
| 1012 |
+
text_dim=text_dim,
|
| 1013 |
+
out_dim=out_dim,
|
| 1014 |
+
num_heads=num_heads,
|
| 1015 |
+
num_layers=num_layers,
|
| 1016 |
+
window_size=window_size,
|
| 1017 |
+
qk_norm=qk_norm,
|
| 1018 |
+
cross_attn_norm=cross_attn_norm,
|
| 1019 |
+
eps=eps,
|
| 1020 |
+
in_channels=in_channels,
|
| 1021 |
+
hidden_size=hidden_size,
|
| 1022 |
+
add_control_adapter=add_control_adapter,
|
| 1023 |
+
control_adapter_type=control_adapter_type,
|
| 1024 |
+
in_dim_control_adapter=in_dim_control_adapter,
|
| 1025 |
+
downscale_factor_control_adapter=downscale_factor_control_adapter,
|
| 1026 |
+
add_ref_conv=add_ref_conv,
|
| 1027 |
+
in_dim_ref_conv=in_dim_ref_conv,
|
| 1028 |
+
cross_attn_type="cross_attn",
|
| 1029 |
+
cam_method=cam_method,
|
| 1030 |
+
attn_compress=attn_compress,
|
| 1031 |
+
cam_self_attn_layers=cam_self_attn_layers,
|
| 1032 |
+
)
|
| 1033 |
+
|
| 1034 |
+
if hasattr(self, "img_emb"):
|
| 1035 |
+
del self.img_emb
|
videox_fun/models/wan_vae3_8.py
ADDED
|
@@ -0,0 +1,1248 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
| 2 |
+
from typing import Tuple, Union
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.cuda.amp as amp
|
| 6 |
+
import torch.nn as nn
|
| 7 |
+
import torch.nn.functional as F
|
| 8 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 9 |
+
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
| 10 |
+
from diffusers.models.autoencoders.vae import (DecoderOutput,
|
| 11 |
+
DiagonalGaussianDistribution)
|
| 12 |
+
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
| 13 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 14 |
+
from diffusers.utils.accelerate_utils import apply_forward_hook
|
| 15 |
+
from einops import rearrange
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
CACHE_T = 2
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class CausalConv3d(nn.Conv3d):
|
| 22 |
+
"""
|
| 23 |
+
Causal 3d convolusion.
|
| 24 |
+
"""
|
| 25 |
+
|
| 26 |
+
def __init__(self, *args, **kwargs):
|
| 27 |
+
super().__init__(*args, **kwargs)
|
| 28 |
+
self._padding = (
|
| 29 |
+
self.padding[2],
|
| 30 |
+
self.padding[2],
|
| 31 |
+
self.padding[1],
|
| 32 |
+
self.padding[1],
|
| 33 |
+
2 * self.padding[0],
|
| 34 |
+
0,
|
| 35 |
+
)
|
| 36 |
+
self.padding = (0, 0, 0)
|
| 37 |
+
|
| 38 |
+
def forward(self, x, cache_x=None):
|
| 39 |
+
padding = list(self._padding)
|
| 40 |
+
if cache_x is not None and self._padding[4] > 0:
|
| 41 |
+
cache_x = cache_x.to(x.device)
|
| 42 |
+
x = torch.cat([cache_x, x], dim=2)
|
| 43 |
+
padding[4] -= cache_x.shape[2]
|
| 44 |
+
x = F.pad(x, padding)
|
| 45 |
+
|
| 46 |
+
return super().forward(x)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class RMS_norm(nn.Module):
|
| 50 |
+
|
| 51 |
+
def __init__(self, dim, channel_first=True, images=True, bias=False):
|
| 52 |
+
super().__init__()
|
| 53 |
+
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
| 54 |
+
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
| 55 |
+
|
| 56 |
+
self.channel_first = channel_first
|
| 57 |
+
self.scale = dim**0.5
|
| 58 |
+
self.gamma = nn.Parameter(torch.ones(shape))
|
| 59 |
+
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
|
| 60 |
+
|
| 61 |
+
def forward(self, x):
|
| 62 |
+
return (F.normalize(x, dim=(1 if self.channel_first else -1)) *
|
| 63 |
+
self.scale * self.gamma + self.bias)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
class Upsample(nn.Upsample):
|
| 67 |
+
|
| 68 |
+
def forward(self, x):
|
| 69 |
+
"""
|
| 70 |
+
Fix bfloat16 support for nearest neighbor interpolation.
|
| 71 |
+
"""
|
| 72 |
+
return super().forward(x.float()).type_as(x)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class Resample(nn.Module):
|
| 76 |
+
|
| 77 |
+
def __init__(self, dim, mode):
|
| 78 |
+
assert mode in (
|
| 79 |
+
"none",
|
| 80 |
+
"upsample2d",
|
| 81 |
+
"upsample3d",
|
| 82 |
+
"downsample2d",
|
| 83 |
+
"downsample3d",
|
| 84 |
+
)
|
| 85 |
+
super().__init__()
|
| 86 |
+
self.dim = dim
|
| 87 |
+
self.mode = mode
|
| 88 |
+
|
| 89 |
+
# layers
|
| 90 |
+
if mode == "upsample2d":
|
| 91 |
+
self.resample = nn.Sequential(
|
| 92 |
+
Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
|
| 93 |
+
nn.Conv2d(dim, dim, 3, padding=1),
|
| 94 |
+
)
|
| 95 |
+
elif mode == "upsample3d":
|
| 96 |
+
self.resample = nn.Sequential(
|
| 97 |
+
Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
|
| 98 |
+
nn.Conv2d(dim, dim, 3, padding=1),
|
| 99 |
+
# nn.Conv2d(dim, dim//2, 3, padding=1)
|
| 100 |
+
)
|
| 101 |
+
self.time_conv = CausalConv3d(
|
| 102 |
+
dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
|
| 103 |
+
elif mode == "downsample2d":
|
| 104 |
+
self.resample = nn.Sequential(
|
| 105 |
+
nn.ZeroPad2d((0, 1, 0, 1)),
|
| 106 |
+
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
| 107 |
+
elif mode == "downsample3d":
|
| 108 |
+
self.resample = nn.Sequential(
|
| 109 |
+
nn.ZeroPad2d((0, 1, 0, 1)),
|
| 110 |
+
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
| 111 |
+
self.time_conv = CausalConv3d(
|
| 112 |
+
dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
|
| 113 |
+
else:
|
| 114 |
+
self.resample = nn.Identity()
|
| 115 |
+
|
| 116 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 117 |
+
b, c, t, h, w = x.size()
|
| 118 |
+
if self.mode == "upsample3d":
|
| 119 |
+
if feat_cache is not None:
|
| 120 |
+
idx = feat_idx[0]
|
| 121 |
+
if feat_cache[idx] is None:
|
| 122 |
+
feat_cache[idx] = "Rep"
|
| 123 |
+
feat_idx[0] += 1
|
| 124 |
+
else:
|
| 125 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 126 |
+
if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
|
| 127 |
+
feat_cache[idx] != "Rep"):
|
| 128 |
+
# cache last frame of last two chunk
|
| 129 |
+
cache_x = torch.cat(
|
| 130 |
+
[
|
| 131 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 132 |
+
cache_x.device),
|
| 133 |
+
cache_x,
|
| 134 |
+
],
|
| 135 |
+
dim=2,
|
| 136 |
+
)
|
| 137 |
+
if (cache_x.shape[2] < 2 and feat_cache[idx] is not None and
|
| 138 |
+
feat_cache[idx] == "Rep"):
|
| 139 |
+
cache_x = torch.cat(
|
| 140 |
+
[
|
| 141 |
+
torch.zeros_like(cache_x).to(cache_x.device),
|
| 142 |
+
cache_x
|
| 143 |
+
],
|
| 144 |
+
dim=2,
|
| 145 |
+
)
|
| 146 |
+
if feat_cache[idx] == "Rep":
|
| 147 |
+
x = self.time_conv(x)
|
| 148 |
+
else:
|
| 149 |
+
x = self.time_conv(x, feat_cache[idx])
|
| 150 |
+
feat_cache[idx] = cache_x
|
| 151 |
+
feat_idx[0] += 1
|
| 152 |
+
x = x.reshape(b, 2, c, t, h, w)
|
| 153 |
+
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
|
| 154 |
+
3)
|
| 155 |
+
x = x.reshape(b, c, t * 2, h, w)
|
| 156 |
+
t = x.shape[2]
|
| 157 |
+
x = rearrange(x, "b c t h w -> (b t) c h w")
|
| 158 |
+
x = self.resample(x)
|
| 159 |
+
x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
|
| 160 |
+
|
| 161 |
+
if self.mode == "downsample3d":
|
| 162 |
+
if feat_cache is not None:
|
| 163 |
+
idx = feat_idx[0]
|
| 164 |
+
if feat_cache[idx] is None:
|
| 165 |
+
feat_cache[idx] = x.clone()
|
| 166 |
+
feat_idx[0] += 1
|
| 167 |
+
else:
|
| 168 |
+
cache_x = x[:, :, -1:, :, :].clone()
|
| 169 |
+
x = self.time_conv(
|
| 170 |
+
torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
|
| 171 |
+
feat_cache[idx] = cache_x
|
| 172 |
+
feat_idx[0] += 1
|
| 173 |
+
return x
|
| 174 |
+
|
| 175 |
+
def init_weight(self, conv):
|
| 176 |
+
conv_weight = conv.weight.detach().clone()
|
| 177 |
+
nn.init.zeros_(conv_weight)
|
| 178 |
+
c1, c2, t, h, w = conv_weight.size()
|
| 179 |
+
one_matrix = torch.eye(c1, c2)
|
| 180 |
+
init_matrix = one_matrix
|
| 181 |
+
nn.init.zeros_(conv_weight)
|
| 182 |
+
conv_weight.data[:, :, 1, 0, 0] = init_matrix # * 0.5
|
| 183 |
+
conv.weight = nn.Parameter(conv_weight)
|
| 184 |
+
nn.init.zeros_(conv.bias.data)
|
| 185 |
+
|
| 186 |
+
def init_weight2(self, conv):
|
| 187 |
+
conv_weight = conv.weight.data.detach().clone()
|
| 188 |
+
nn.init.zeros_(conv_weight)
|
| 189 |
+
c1, c2, t, h, w = conv_weight.size()
|
| 190 |
+
init_matrix = torch.eye(c1 // 2, c2)
|
| 191 |
+
conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
|
| 192 |
+
conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
|
| 193 |
+
conv.weight = nn.Parameter(conv_weight)
|
| 194 |
+
nn.init.zeros_(conv.bias.data)
|
| 195 |
+
|
| 196 |
+
|
| 197 |
+
class ResidualBlock(nn.Module):
|
| 198 |
+
|
| 199 |
+
def __init__(self, in_dim, out_dim, dropout=0.0):
|
| 200 |
+
super().__init__()
|
| 201 |
+
self.in_dim = in_dim
|
| 202 |
+
self.out_dim = out_dim
|
| 203 |
+
|
| 204 |
+
# layers
|
| 205 |
+
self.residual = nn.Sequential(
|
| 206 |
+
RMS_norm(in_dim, images=False),
|
| 207 |
+
nn.SiLU(),
|
| 208 |
+
CausalConv3d(in_dim, out_dim, 3, padding=1),
|
| 209 |
+
RMS_norm(out_dim, images=False),
|
| 210 |
+
nn.SiLU(),
|
| 211 |
+
nn.Dropout(dropout),
|
| 212 |
+
CausalConv3d(out_dim, out_dim, 3, padding=1),
|
| 213 |
+
)
|
| 214 |
+
self.shortcut = (
|
| 215 |
+
CausalConv3d(in_dim, out_dim, 1)
|
| 216 |
+
if in_dim != out_dim else nn.Identity())
|
| 217 |
+
|
| 218 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 219 |
+
h = self.shortcut(x)
|
| 220 |
+
for layer in self.residual:
|
| 221 |
+
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
| 222 |
+
idx = feat_idx[0]
|
| 223 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 224 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 225 |
+
# cache last frame of last two chunk
|
| 226 |
+
cache_x = torch.cat(
|
| 227 |
+
[
|
| 228 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 229 |
+
cache_x.device),
|
| 230 |
+
cache_x,
|
| 231 |
+
],
|
| 232 |
+
dim=2,
|
| 233 |
+
)
|
| 234 |
+
x = layer(x, feat_cache[idx])
|
| 235 |
+
feat_cache[idx] = cache_x
|
| 236 |
+
feat_idx[0] += 1
|
| 237 |
+
else:
|
| 238 |
+
x = layer(x)
|
| 239 |
+
return x + h
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
class AttentionBlock(nn.Module):
|
| 243 |
+
"""
|
| 244 |
+
Causal self-attention with a single head.
|
| 245 |
+
"""
|
| 246 |
+
|
| 247 |
+
def __init__(self, dim):
|
| 248 |
+
super().__init__()
|
| 249 |
+
self.dim = dim
|
| 250 |
+
|
| 251 |
+
# layers
|
| 252 |
+
self.norm = RMS_norm(dim)
|
| 253 |
+
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
| 254 |
+
self.proj = nn.Conv2d(dim, dim, 1)
|
| 255 |
+
|
| 256 |
+
# zero out the last layer params
|
| 257 |
+
nn.init.zeros_(self.proj.weight)
|
| 258 |
+
|
| 259 |
+
def forward(self, x):
|
| 260 |
+
identity = x
|
| 261 |
+
b, c, t, h, w = x.size()
|
| 262 |
+
x = rearrange(x, "b c t h w -> (b t) c h w")
|
| 263 |
+
x = self.norm(x)
|
| 264 |
+
# compute query, key, value
|
| 265 |
+
q, k, v = (
|
| 266 |
+
self.to_qkv(x).reshape(b * t, 1, c * 3,
|
| 267 |
+
-1).permute(0, 1, 3,
|
| 268 |
+
2).contiguous().chunk(3, dim=-1))
|
| 269 |
+
|
| 270 |
+
# apply attention
|
| 271 |
+
x = F.scaled_dot_product_attention(
|
| 272 |
+
q,
|
| 273 |
+
k,
|
| 274 |
+
v,
|
| 275 |
+
)
|
| 276 |
+
x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
|
| 277 |
+
|
| 278 |
+
# output
|
| 279 |
+
x = self.proj(x)
|
| 280 |
+
x = rearrange(x, "(b t) c h w-> b c t h w", t=t)
|
| 281 |
+
return x + identity
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def patchify(x, patch_size):
|
| 285 |
+
if patch_size == 1:
|
| 286 |
+
return x
|
| 287 |
+
if x.dim() == 4:
|
| 288 |
+
x = rearrange(
|
| 289 |
+
x, "b c (h q) (w r) -> b (c r q) h w", q=patch_size, r=patch_size)
|
| 290 |
+
elif x.dim() == 5:
|
| 291 |
+
x = rearrange(
|
| 292 |
+
x,
|
| 293 |
+
"b c f (h q) (w r) -> b (c r q) f h w",
|
| 294 |
+
q=patch_size,
|
| 295 |
+
r=patch_size,
|
| 296 |
+
)
|
| 297 |
+
else:
|
| 298 |
+
raise ValueError(f"Invalid input shape: {x.shape}")
|
| 299 |
+
|
| 300 |
+
return x
|
| 301 |
+
|
| 302 |
+
|
| 303 |
+
def unpatchify(x, patch_size):
|
| 304 |
+
if patch_size == 1:
|
| 305 |
+
return x
|
| 306 |
+
|
| 307 |
+
if x.dim() == 4:
|
| 308 |
+
x = rearrange(
|
| 309 |
+
x, "b (c r q) h w -> b c (h q) (w r)", q=patch_size, r=patch_size)
|
| 310 |
+
elif x.dim() == 5:
|
| 311 |
+
x = rearrange(
|
| 312 |
+
x,
|
| 313 |
+
"b (c r q) f h w -> b c f (h q) (w r)",
|
| 314 |
+
q=patch_size,
|
| 315 |
+
r=patch_size,
|
| 316 |
+
)
|
| 317 |
+
return x
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
class AvgDown3D(nn.Module):
|
| 321 |
+
|
| 322 |
+
def __init__(
|
| 323 |
+
self,
|
| 324 |
+
in_channels,
|
| 325 |
+
out_channels,
|
| 326 |
+
factor_t,
|
| 327 |
+
factor_s=1,
|
| 328 |
+
):
|
| 329 |
+
super().__init__()
|
| 330 |
+
self.in_channels = in_channels
|
| 331 |
+
self.out_channels = out_channels
|
| 332 |
+
self.factor_t = factor_t
|
| 333 |
+
self.factor_s = factor_s
|
| 334 |
+
self.factor = self.factor_t * self.factor_s * self.factor_s
|
| 335 |
+
|
| 336 |
+
assert in_channels * self.factor % out_channels == 0
|
| 337 |
+
self.group_size = in_channels * self.factor // out_channels
|
| 338 |
+
|
| 339 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 340 |
+
pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
|
| 341 |
+
pad = (0, 0, 0, 0, pad_t, 0)
|
| 342 |
+
x = F.pad(x, pad)
|
| 343 |
+
B, C, T, H, W = x.shape
|
| 344 |
+
x = x.view(
|
| 345 |
+
B,
|
| 346 |
+
C,
|
| 347 |
+
T // self.factor_t,
|
| 348 |
+
self.factor_t,
|
| 349 |
+
H // self.factor_s,
|
| 350 |
+
self.factor_s,
|
| 351 |
+
W // self.factor_s,
|
| 352 |
+
self.factor_s,
|
| 353 |
+
)
|
| 354 |
+
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
|
| 355 |
+
x = x.view(
|
| 356 |
+
B,
|
| 357 |
+
C * self.factor,
|
| 358 |
+
T // self.factor_t,
|
| 359 |
+
H // self.factor_s,
|
| 360 |
+
W // self.factor_s,
|
| 361 |
+
)
|
| 362 |
+
x = x.view(
|
| 363 |
+
B,
|
| 364 |
+
self.out_channels,
|
| 365 |
+
self.group_size,
|
| 366 |
+
T // self.factor_t,
|
| 367 |
+
H // self.factor_s,
|
| 368 |
+
W // self.factor_s,
|
| 369 |
+
)
|
| 370 |
+
x = x.mean(dim=2)
|
| 371 |
+
return x
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
class DupUp3D(nn.Module):
|
| 375 |
+
|
| 376 |
+
def __init__(
|
| 377 |
+
self,
|
| 378 |
+
in_channels: int,
|
| 379 |
+
out_channels: int,
|
| 380 |
+
factor_t,
|
| 381 |
+
factor_s=1,
|
| 382 |
+
):
|
| 383 |
+
super().__init__()
|
| 384 |
+
self.in_channels = in_channels
|
| 385 |
+
self.out_channels = out_channels
|
| 386 |
+
|
| 387 |
+
self.factor_t = factor_t
|
| 388 |
+
self.factor_s = factor_s
|
| 389 |
+
self.factor = self.factor_t * self.factor_s * self.factor_s
|
| 390 |
+
|
| 391 |
+
assert out_channels * self.factor % in_channels == 0
|
| 392 |
+
self.repeats = out_channels * self.factor // in_channels
|
| 393 |
+
|
| 394 |
+
def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
|
| 395 |
+
x = x.repeat_interleave(self.repeats, dim=1)
|
| 396 |
+
x = x.view(
|
| 397 |
+
x.size(0),
|
| 398 |
+
self.out_channels,
|
| 399 |
+
self.factor_t,
|
| 400 |
+
self.factor_s,
|
| 401 |
+
self.factor_s,
|
| 402 |
+
x.size(2),
|
| 403 |
+
x.size(3),
|
| 404 |
+
x.size(4),
|
| 405 |
+
)
|
| 406 |
+
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
|
| 407 |
+
x = x.view(
|
| 408 |
+
x.size(0),
|
| 409 |
+
self.out_channels,
|
| 410 |
+
x.size(2) * self.factor_t,
|
| 411 |
+
x.size(4) * self.factor_s,
|
| 412 |
+
x.size(6) * self.factor_s,
|
| 413 |
+
)
|
| 414 |
+
if first_chunk:
|
| 415 |
+
x = x[:, :, self.factor_t - 1:, :, :]
|
| 416 |
+
return x
|
| 417 |
+
|
| 418 |
+
|
| 419 |
+
class Down_ResidualBlock(nn.Module):
|
| 420 |
+
|
| 421 |
+
def __init__(self,
|
| 422 |
+
in_dim,
|
| 423 |
+
out_dim,
|
| 424 |
+
dropout,
|
| 425 |
+
mult,
|
| 426 |
+
temperal_downsample=False,
|
| 427 |
+
down_flag=False):
|
| 428 |
+
super().__init__()
|
| 429 |
+
|
| 430 |
+
# Shortcut path with downsample
|
| 431 |
+
self.avg_shortcut = AvgDown3D(
|
| 432 |
+
in_dim,
|
| 433 |
+
out_dim,
|
| 434 |
+
factor_t=2 if temperal_downsample else 1,
|
| 435 |
+
factor_s=2 if down_flag else 1,
|
| 436 |
+
)
|
| 437 |
+
|
| 438 |
+
# Main path with residual blocks and downsample
|
| 439 |
+
downsamples = []
|
| 440 |
+
for _ in range(mult):
|
| 441 |
+
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
| 442 |
+
in_dim = out_dim
|
| 443 |
+
|
| 444 |
+
# Add the final downsample block
|
| 445 |
+
if down_flag:
|
| 446 |
+
mode = "downsample3d" if temperal_downsample else "downsample2d"
|
| 447 |
+
downsamples.append(Resample(out_dim, mode=mode))
|
| 448 |
+
|
| 449 |
+
self.downsamples = nn.Sequential(*downsamples)
|
| 450 |
+
|
| 451 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 452 |
+
x_copy = x.clone()
|
| 453 |
+
for module in self.downsamples:
|
| 454 |
+
x = module(x, feat_cache, feat_idx)
|
| 455 |
+
|
| 456 |
+
return x + self.avg_shortcut(x_copy)
|
| 457 |
+
|
| 458 |
+
|
| 459 |
+
class Up_ResidualBlock(nn.Module):
|
| 460 |
+
|
| 461 |
+
def __init__(self,
|
| 462 |
+
in_dim,
|
| 463 |
+
out_dim,
|
| 464 |
+
dropout,
|
| 465 |
+
mult,
|
| 466 |
+
temperal_upsample=False,
|
| 467 |
+
up_flag=False):
|
| 468 |
+
super().__init__()
|
| 469 |
+
# Shortcut path with upsample
|
| 470 |
+
if up_flag:
|
| 471 |
+
self.avg_shortcut = DupUp3D(
|
| 472 |
+
in_dim,
|
| 473 |
+
out_dim,
|
| 474 |
+
factor_t=2 if temperal_upsample else 1,
|
| 475 |
+
factor_s=2 if up_flag else 1,
|
| 476 |
+
)
|
| 477 |
+
else:
|
| 478 |
+
self.avg_shortcut = None
|
| 479 |
+
|
| 480 |
+
# Main path with residual blocks and upsample
|
| 481 |
+
upsamples = []
|
| 482 |
+
for _ in range(mult):
|
| 483 |
+
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
| 484 |
+
in_dim = out_dim
|
| 485 |
+
|
| 486 |
+
# Add the final upsample block
|
| 487 |
+
if up_flag:
|
| 488 |
+
mode = "upsample3d" if temperal_upsample else "upsample2d"
|
| 489 |
+
upsamples.append(Resample(out_dim, mode=mode))
|
| 490 |
+
|
| 491 |
+
self.upsamples = nn.Sequential(*upsamples)
|
| 492 |
+
|
| 493 |
+
def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
|
| 494 |
+
x_main = x.clone()
|
| 495 |
+
for module in self.upsamples:
|
| 496 |
+
x_main = module(x_main, feat_cache, feat_idx)
|
| 497 |
+
if self.avg_shortcut is not None:
|
| 498 |
+
x_shortcut = self.avg_shortcut(x, first_chunk)
|
| 499 |
+
return x_main + x_shortcut
|
| 500 |
+
else:
|
| 501 |
+
return x_main
|
| 502 |
+
|
| 503 |
+
|
| 504 |
+
class Encoder3d(nn.Module):
|
| 505 |
+
|
| 506 |
+
def __init__(
|
| 507 |
+
self,
|
| 508 |
+
dim=128,
|
| 509 |
+
z_dim=4,
|
| 510 |
+
dim_mult=[1, 2, 4, 4],
|
| 511 |
+
num_res_blocks=2,
|
| 512 |
+
attn_scales=[],
|
| 513 |
+
temperal_downsample=[True, True, False],
|
| 514 |
+
dropout=0.0,
|
| 515 |
+
):
|
| 516 |
+
super().__init__()
|
| 517 |
+
self.dim = dim
|
| 518 |
+
self.z_dim = z_dim
|
| 519 |
+
self.dim_mult = dim_mult
|
| 520 |
+
self.num_res_blocks = num_res_blocks
|
| 521 |
+
self.attn_scales = attn_scales
|
| 522 |
+
self.temperal_downsample = temperal_downsample
|
| 523 |
+
|
| 524 |
+
# dimensions
|
| 525 |
+
dims = [dim * u for u in [1] + dim_mult]
|
| 526 |
+
scale = 1.0
|
| 527 |
+
|
| 528 |
+
# init block
|
| 529 |
+
self.conv1 = CausalConv3d(12, dims[0], 3, padding=1)
|
| 530 |
+
|
| 531 |
+
# downsample blocks
|
| 532 |
+
downsamples = []
|
| 533 |
+
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
| 534 |
+
t_down_flag = (
|
| 535 |
+
temperal_downsample[i]
|
| 536 |
+
if i < len(temperal_downsample) else False)
|
| 537 |
+
downsamples.append(
|
| 538 |
+
Down_ResidualBlock(
|
| 539 |
+
in_dim=in_dim,
|
| 540 |
+
out_dim=out_dim,
|
| 541 |
+
dropout=dropout,
|
| 542 |
+
mult=num_res_blocks,
|
| 543 |
+
temperal_downsample=t_down_flag,
|
| 544 |
+
down_flag=i != len(dim_mult) - 1,
|
| 545 |
+
))
|
| 546 |
+
scale /= 2.0
|
| 547 |
+
self.downsamples = nn.Sequential(*downsamples)
|
| 548 |
+
|
| 549 |
+
# middle blocks
|
| 550 |
+
self.middle = nn.Sequential(
|
| 551 |
+
ResidualBlock(out_dim, out_dim, dropout),
|
| 552 |
+
AttentionBlock(out_dim),
|
| 553 |
+
ResidualBlock(out_dim, out_dim, dropout),
|
| 554 |
+
)
|
| 555 |
+
|
| 556 |
+
# # output blocks
|
| 557 |
+
self.head = nn.Sequential(
|
| 558 |
+
RMS_norm(out_dim, images=False),
|
| 559 |
+
nn.SiLU(),
|
| 560 |
+
CausalConv3d(out_dim, z_dim, 3, padding=1),
|
| 561 |
+
)
|
| 562 |
+
|
| 563 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 564 |
+
|
| 565 |
+
if feat_cache is not None:
|
| 566 |
+
idx = feat_idx[0]
|
| 567 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 568 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 569 |
+
cache_x = torch.cat(
|
| 570 |
+
[
|
| 571 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 572 |
+
cache_x.device),
|
| 573 |
+
cache_x,
|
| 574 |
+
],
|
| 575 |
+
dim=2,
|
| 576 |
+
)
|
| 577 |
+
x = self.conv1(x, feat_cache[idx])
|
| 578 |
+
feat_cache[idx] = cache_x
|
| 579 |
+
feat_idx[0] += 1
|
| 580 |
+
else:
|
| 581 |
+
x = self.conv1(x)
|
| 582 |
+
|
| 583 |
+
## downsamples
|
| 584 |
+
for layer in self.downsamples:
|
| 585 |
+
if feat_cache is not None:
|
| 586 |
+
x = layer(x, feat_cache, feat_idx)
|
| 587 |
+
else:
|
| 588 |
+
x = layer(x)
|
| 589 |
+
|
| 590 |
+
## middle
|
| 591 |
+
for layer in self.middle:
|
| 592 |
+
if isinstance(layer, ResidualBlock) and feat_cache is not None:
|
| 593 |
+
x = layer(x, feat_cache, feat_idx)
|
| 594 |
+
else:
|
| 595 |
+
x = layer(x)
|
| 596 |
+
|
| 597 |
+
## head
|
| 598 |
+
for layer in self.head:
|
| 599 |
+
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
| 600 |
+
idx = feat_idx[0]
|
| 601 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 602 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 603 |
+
cache_x = torch.cat(
|
| 604 |
+
[
|
| 605 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 606 |
+
cache_x.device),
|
| 607 |
+
cache_x,
|
| 608 |
+
],
|
| 609 |
+
dim=2,
|
| 610 |
+
)
|
| 611 |
+
x = layer(x, feat_cache[idx])
|
| 612 |
+
feat_cache[idx] = cache_x
|
| 613 |
+
feat_idx[0] += 1
|
| 614 |
+
else:
|
| 615 |
+
x = layer(x)
|
| 616 |
+
|
| 617 |
+
return x
|
| 618 |
+
|
| 619 |
+
|
| 620 |
+
class Decoder3d(nn.Module):
|
| 621 |
+
|
| 622 |
+
def __init__(
|
| 623 |
+
self,
|
| 624 |
+
dim=128,
|
| 625 |
+
z_dim=4,
|
| 626 |
+
dim_mult=[1, 2, 4, 4],
|
| 627 |
+
num_res_blocks=2,
|
| 628 |
+
attn_scales=[],
|
| 629 |
+
temperal_upsample=[False, True, True],
|
| 630 |
+
dropout=0.0,
|
| 631 |
+
):
|
| 632 |
+
super().__init__()
|
| 633 |
+
self.dim = dim
|
| 634 |
+
self.z_dim = z_dim
|
| 635 |
+
self.dim_mult = dim_mult
|
| 636 |
+
self.num_res_blocks = num_res_blocks
|
| 637 |
+
self.attn_scales = attn_scales
|
| 638 |
+
self.temperal_upsample = temperal_upsample
|
| 639 |
+
|
| 640 |
+
# dimensions
|
| 641 |
+
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
| 642 |
+
scale = 1.0 / 2**(len(dim_mult) - 2)
|
| 643 |
+
# init block
|
| 644 |
+
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
|
| 645 |
+
|
| 646 |
+
# middle blocks
|
| 647 |
+
self.middle = nn.Sequential(
|
| 648 |
+
ResidualBlock(dims[0], dims[0], dropout),
|
| 649 |
+
AttentionBlock(dims[0]),
|
| 650 |
+
ResidualBlock(dims[0], dims[0], dropout),
|
| 651 |
+
)
|
| 652 |
+
|
| 653 |
+
# upsample blocks
|
| 654 |
+
upsamples = []
|
| 655 |
+
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
| 656 |
+
t_up_flag = temperal_upsample[i] if i < len(
|
| 657 |
+
temperal_upsample) else False
|
| 658 |
+
upsamples.append(
|
| 659 |
+
Up_ResidualBlock(
|
| 660 |
+
in_dim=in_dim,
|
| 661 |
+
out_dim=out_dim,
|
| 662 |
+
dropout=dropout,
|
| 663 |
+
mult=num_res_blocks + 1,
|
| 664 |
+
temperal_upsample=t_up_flag,
|
| 665 |
+
up_flag=i != len(dim_mult) - 1,
|
| 666 |
+
))
|
| 667 |
+
self.upsamples = nn.Sequential(*upsamples)
|
| 668 |
+
|
| 669 |
+
# output blocks
|
| 670 |
+
self.head = nn.Sequential(
|
| 671 |
+
RMS_norm(out_dim, images=False),
|
| 672 |
+
nn.SiLU(),
|
| 673 |
+
CausalConv3d(out_dim, 12, 3, padding=1),
|
| 674 |
+
)
|
| 675 |
+
|
| 676 |
+
def forward(self, x, feat_cache=None, feat_idx=[0], first_chunk=False):
|
| 677 |
+
if feat_cache is not None:
|
| 678 |
+
idx = feat_idx[0]
|
| 679 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 680 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 681 |
+
cache_x = torch.cat(
|
| 682 |
+
[
|
| 683 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 684 |
+
cache_x.device),
|
| 685 |
+
cache_x,
|
| 686 |
+
],
|
| 687 |
+
dim=2,
|
| 688 |
+
)
|
| 689 |
+
x = self.conv1(x, feat_cache[idx])
|
| 690 |
+
feat_cache[idx] = cache_x
|
| 691 |
+
feat_idx[0] += 1
|
| 692 |
+
else:
|
| 693 |
+
x = self.conv1(x)
|
| 694 |
+
|
| 695 |
+
for layer in self.middle:
|
| 696 |
+
if isinstance(layer, ResidualBlock) and feat_cache is not None:
|
| 697 |
+
x = layer(x, feat_cache, feat_idx)
|
| 698 |
+
else:
|
| 699 |
+
x = layer(x)
|
| 700 |
+
|
| 701 |
+
## upsamples
|
| 702 |
+
for layer in self.upsamples:
|
| 703 |
+
if feat_cache is not None:
|
| 704 |
+
x = layer(x, feat_cache, feat_idx, first_chunk)
|
| 705 |
+
else:
|
| 706 |
+
x = layer(x)
|
| 707 |
+
|
| 708 |
+
## head
|
| 709 |
+
for layer in self.head:
|
| 710 |
+
if isinstance(layer, CausalConv3d) and feat_cache is not None:
|
| 711 |
+
idx = feat_idx[0]
|
| 712 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 713 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 714 |
+
cache_x = torch.cat(
|
| 715 |
+
[
|
| 716 |
+
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
| 717 |
+
cache_x.device),
|
| 718 |
+
cache_x,
|
| 719 |
+
],
|
| 720 |
+
dim=2,
|
| 721 |
+
)
|
| 722 |
+
x = layer(x, feat_cache[idx])
|
| 723 |
+
feat_cache[idx] = cache_x
|
| 724 |
+
feat_idx[0] += 1
|
| 725 |
+
else:
|
| 726 |
+
x = layer(x)
|
| 727 |
+
return x
|
| 728 |
+
|
| 729 |
+
|
| 730 |
+
def count_conv3d(model):
|
| 731 |
+
count = 0
|
| 732 |
+
for m in model.modules():
|
| 733 |
+
if isinstance(m, CausalConv3d):
|
| 734 |
+
count += 1
|
| 735 |
+
return count
|
| 736 |
+
|
| 737 |
+
|
| 738 |
+
class AutoencoderKLWan2_2_(nn.Module):
|
| 739 |
+
|
| 740 |
+
def __init__(
|
| 741 |
+
self,
|
| 742 |
+
dim=160,
|
| 743 |
+
dec_dim=256,
|
| 744 |
+
z_dim=16,
|
| 745 |
+
dim_mult=[1, 2, 4, 4],
|
| 746 |
+
num_res_blocks=2,
|
| 747 |
+
attn_scales=[],
|
| 748 |
+
temperal_downsample=[True, True, False],
|
| 749 |
+
dropout=0.0,
|
| 750 |
+
):
|
| 751 |
+
super().__init__()
|
| 752 |
+
self.dim = dim
|
| 753 |
+
self.z_dim = z_dim
|
| 754 |
+
self.dim_mult = dim_mult
|
| 755 |
+
self.num_res_blocks = num_res_blocks
|
| 756 |
+
self.attn_scales = attn_scales
|
| 757 |
+
self.temperal_downsample = temperal_downsample
|
| 758 |
+
self.temperal_upsample = temperal_downsample[::-1]
|
| 759 |
+
|
| 760 |
+
# modules
|
| 761 |
+
self.encoder = Encoder3d(
|
| 762 |
+
dim,
|
| 763 |
+
z_dim * 2,
|
| 764 |
+
dim_mult,
|
| 765 |
+
num_res_blocks,
|
| 766 |
+
attn_scales,
|
| 767 |
+
self.temperal_downsample,
|
| 768 |
+
dropout,
|
| 769 |
+
)
|
| 770 |
+
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
|
| 771 |
+
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
|
| 772 |
+
self.decoder = Decoder3d(
|
| 773 |
+
dec_dim,
|
| 774 |
+
z_dim,
|
| 775 |
+
dim_mult,
|
| 776 |
+
num_res_blocks,
|
| 777 |
+
attn_scales,
|
| 778 |
+
self.temperal_upsample,
|
| 779 |
+
dropout,
|
| 780 |
+
)
|
| 781 |
+
|
| 782 |
+
def forward(self, x, scale=[0, 1]):
|
| 783 |
+
mu = self.encode(x, scale)
|
| 784 |
+
x_recon = self.decode(mu, scale)
|
| 785 |
+
return x_recon, mu
|
| 786 |
+
|
| 787 |
+
def encode(self, x, scale):
|
| 788 |
+
self.clear_cache()
|
| 789 |
+
# z: [b,c,t,h,w]
|
| 790 |
+
scale = [item.to(x.device, x.dtype) for item in scale]
|
| 791 |
+
x = patchify(x, patch_size=2)
|
| 792 |
+
t = x.shape[2]
|
| 793 |
+
iter_ = 1 + (t - 1) // 4
|
| 794 |
+
for i in range(iter_):
|
| 795 |
+
self._enc_conv_idx = [0]
|
| 796 |
+
if i == 0:
|
| 797 |
+
out = self.encoder(
|
| 798 |
+
x[:, :, :1, :, :],
|
| 799 |
+
feat_cache=self._enc_feat_map,
|
| 800 |
+
feat_idx=self._enc_conv_idx,
|
| 801 |
+
)
|
| 802 |
+
else:
|
| 803 |
+
out_ = self.encoder(
|
| 804 |
+
x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
|
| 805 |
+
feat_cache=self._enc_feat_map,
|
| 806 |
+
feat_idx=self._enc_conv_idx,
|
| 807 |
+
)
|
| 808 |
+
out = torch.cat([out, out_], 2)
|
| 809 |
+
mu, log_var = self.conv1(out).chunk(2, dim=1)
|
| 810 |
+
if isinstance(scale[0], torch.Tensor):
|
| 811 |
+
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
|
| 812 |
+
1, self.z_dim, 1, 1, 1)
|
| 813 |
+
else:
|
| 814 |
+
mu = (mu - scale[0]) * scale[1]
|
| 815 |
+
x = torch.cat([mu, log_var], dim = 1)
|
| 816 |
+
self.clear_cache()
|
| 817 |
+
return x
|
| 818 |
+
|
| 819 |
+
def decode(self, z, scale):
|
| 820 |
+
self.clear_cache()
|
| 821 |
+
# z: [b,c,t,h,w]
|
| 822 |
+
scale = [item.to(z.device, z.dtype) for item in scale]
|
| 823 |
+
if isinstance(scale[0], torch.Tensor):
|
| 824 |
+
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
|
| 825 |
+
1, self.z_dim, 1, 1, 1)
|
| 826 |
+
else:
|
| 827 |
+
z = z / scale[1] + scale[0]
|
| 828 |
+
iter_ = z.shape[2]
|
| 829 |
+
x = self.conv2(z)
|
| 830 |
+
for i in range(iter_):
|
| 831 |
+
self._conv_idx = [0]
|
| 832 |
+
if i == 0:
|
| 833 |
+
out = self.decoder(
|
| 834 |
+
x[:, :, i:i + 1, :, :],
|
| 835 |
+
feat_cache=self._feat_map,
|
| 836 |
+
feat_idx=self._conv_idx,
|
| 837 |
+
first_chunk=True,
|
| 838 |
+
)
|
| 839 |
+
else:
|
| 840 |
+
out_ = self.decoder(
|
| 841 |
+
x[:, :, i:i + 1, :, :],
|
| 842 |
+
feat_cache=self._feat_map,
|
| 843 |
+
feat_idx=self._conv_idx,
|
| 844 |
+
)
|
| 845 |
+
out = torch.cat([out, out_], 2)
|
| 846 |
+
out = unpatchify(out, patch_size=2)
|
| 847 |
+
self.clear_cache()
|
| 848 |
+
return out
|
| 849 |
+
|
| 850 |
+
def reparameterize(self, mu, log_var):
|
| 851 |
+
std = torch.exp(0.5 * log_var)
|
| 852 |
+
eps = torch.randn_like(std)
|
| 853 |
+
return eps * std + mu
|
| 854 |
+
|
| 855 |
+
def sample(self, imgs, deterministic=False):
|
| 856 |
+
mu, log_var = self.encode(imgs)
|
| 857 |
+
if deterministic:
|
| 858 |
+
return mu
|
| 859 |
+
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
|
| 860 |
+
return mu + std * torch.randn_like(std)
|
| 861 |
+
|
| 862 |
+
def clear_cache(self):
|
| 863 |
+
self._conv_num = count_conv3d(self.decoder)
|
| 864 |
+
self._conv_idx = [0]
|
| 865 |
+
self._feat_map = [None] * self._conv_num
|
| 866 |
+
# cache encode
|
| 867 |
+
self._enc_conv_num = count_conv3d(self.encoder)
|
| 868 |
+
self._enc_conv_idx = [0]
|
| 869 |
+
self._enc_feat_map = [None] * self._enc_conv_num
|
| 870 |
+
|
| 871 |
+
|
| 872 |
+
def _video_vae(pretrained_path=None, z_dim=16, dim=160, device="cpu", **kwargs):
|
| 873 |
+
# params
|
| 874 |
+
cfg = dict(
|
| 875 |
+
dim=dim,
|
| 876 |
+
z_dim=z_dim,
|
| 877 |
+
dim_mult=[1, 2, 4, 4],
|
| 878 |
+
num_res_blocks=2,
|
| 879 |
+
attn_scales=[],
|
| 880 |
+
temperal_downsample=[True, True, True],
|
| 881 |
+
dropout=0.0,
|
| 882 |
+
)
|
| 883 |
+
cfg.update(**kwargs)
|
| 884 |
+
|
| 885 |
+
# init model
|
| 886 |
+
model = AutoencoderKLWan2_2_(**cfg)
|
| 887 |
+
|
| 888 |
+
return model
|
| 889 |
+
|
| 890 |
+
|
| 891 |
+
class AutoencoderKLWan3_8(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
| 892 |
+
@register_to_config
|
| 893 |
+
def __init__(
|
| 894 |
+
self,
|
| 895 |
+
latent_channels=48,
|
| 896 |
+
c_dim=160,
|
| 897 |
+
vae_pth=None,
|
| 898 |
+
dim_mult=[1, 2, 4, 4],
|
| 899 |
+
temperal_downsample=[False, True, True],
|
| 900 |
+
temporal_compression_ratio=4,
|
| 901 |
+
spatial_compression_ratio=8
|
| 902 |
+
):
|
| 903 |
+
super().__init__()
|
| 904 |
+
mean = torch.tensor(
|
| 905 |
+
[
|
| 906 |
+
-0.2289,
|
| 907 |
+
-0.0052,
|
| 908 |
+
-0.1323,
|
| 909 |
+
-0.2339,
|
| 910 |
+
-0.2799,
|
| 911 |
+
0.0174,
|
| 912 |
+
0.1838,
|
| 913 |
+
0.1557,
|
| 914 |
+
-0.1382,
|
| 915 |
+
0.0542,
|
| 916 |
+
0.2813,
|
| 917 |
+
0.0891,
|
| 918 |
+
0.1570,
|
| 919 |
+
-0.0098,
|
| 920 |
+
0.0375,
|
| 921 |
+
-0.1825,
|
| 922 |
+
-0.2246,
|
| 923 |
+
-0.1207,
|
| 924 |
+
-0.0698,
|
| 925 |
+
0.5109,
|
| 926 |
+
0.2665,
|
| 927 |
+
-0.2108,
|
| 928 |
+
-0.2158,
|
| 929 |
+
0.2502,
|
| 930 |
+
-0.2055,
|
| 931 |
+
-0.0322,
|
| 932 |
+
0.1109,
|
| 933 |
+
0.1567,
|
| 934 |
+
-0.0729,
|
| 935 |
+
0.0899,
|
| 936 |
+
-0.2799,
|
| 937 |
+
-0.1230,
|
| 938 |
+
-0.0313,
|
| 939 |
+
-0.1649,
|
| 940 |
+
0.0117,
|
| 941 |
+
0.0723,
|
| 942 |
+
-0.2839,
|
| 943 |
+
-0.2083,
|
| 944 |
+
-0.0520,
|
| 945 |
+
0.3748,
|
| 946 |
+
0.0152,
|
| 947 |
+
0.1957,
|
| 948 |
+
0.1433,
|
| 949 |
+
-0.2944,
|
| 950 |
+
0.3573,
|
| 951 |
+
-0.0548,
|
| 952 |
+
-0.1681,
|
| 953 |
+
-0.0667,
|
| 954 |
+
], dtype=torch.float32
|
| 955 |
+
)
|
| 956 |
+
std = torch.tensor(
|
| 957 |
+
[
|
| 958 |
+
0.4765,
|
| 959 |
+
1.0364,
|
| 960 |
+
0.4514,
|
| 961 |
+
1.1677,
|
| 962 |
+
0.5313,
|
| 963 |
+
0.4990,
|
| 964 |
+
0.4818,
|
| 965 |
+
0.5013,
|
| 966 |
+
0.8158,
|
| 967 |
+
1.0344,
|
| 968 |
+
0.5894,
|
| 969 |
+
1.0901,
|
| 970 |
+
0.6885,
|
| 971 |
+
0.6165,
|
| 972 |
+
0.8454,
|
| 973 |
+
0.4978,
|
| 974 |
+
0.5759,
|
| 975 |
+
0.3523,
|
| 976 |
+
0.7135,
|
| 977 |
+
0.6804,
|
| 978 |
+
0.5833,
|
| 979 |
+
1.4146,
|
| 980 |
+
0.8986,
|
| 981 |
+
0.5659,
|
| 982 |
+
0.7069,
|
| 983 |
+
0.5338,
|
| 984 |
+
0.4889,
|
| 985 |
+
0.4917,
|
| 986 |
+
0.4069,
|
| 987 |
+
0.4999,
|
| 988 |
+
0.6866,
|
| 989 |
+
0.4093,
|
| 990 |
+
0.5709,
|
| 991 |
+
0.6065,
|
| 992 |
+
0.6415,
|
| 993 |
+
0.4944,
|
| 994 |
+
0.5726,
|
| 995 |
+
1.2042,
|
| 996 |
+
0.5458,
|
| 997 |
+
1.6887,
|
| 998 |
+
0.3971,
|
| 999 |
+
1.0600,
|
| 1000 |
+
0.3943,
|
| 1001 |
+
0.5537,
|
| 1002 |
+
0.5444,
|
| 1003 |
+
0.4089,
|
| 1004 |
+
0.7468,
|
| 1005 |
+
0.7744,
|
| 1006 |
+
], dtype=torch.float32
|
| 1007 |
+
)
|
| 1008 |
+
self.scale = [mean, 1.0 / std]
|
| 1009 |
+
|
| 1010 |
+
# init model
|
| 1011 |
+
self.model = _video_vae(
|
| 1012 |
+
pretrained_path=vae_pth,
|
| 1013 |
+
z_dim=latent_channels,
|
| 1014 |
+
dim=c_dim,
|
| 1015 |
+
dim_mult=dim_mult,
|
| 1016 |
+
temperal_downsample=temperal_downsample,
|
| 1017 |
+
).eval().requires_grad_(False)
|
| 1018 |
+
|
| 1019 |
+
self.use_tiling = False
|
| 1020 |
+
self.tile_size = (34, 34)
|
| 1021 |
+
self.tile_stride = (18, 16)
|
| 1022 |
+
|
| 1023 |
+
def enable_tiling(self, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 1024 |
+
"""Enable spatial tiling for encode/decode to reduce peak GPU memory."""
|
| 1025 |
+
self.use_tiling = True
|
| 1026 |
+
self.tile_size = tile_size
|
| 1027 |
+
self.tile_stride = tile_stride
|
| 1028 |
+
|
| 1029 |
+
def disable_tiling(self):
|
| 1030 |
+
"""Disable spatial tiling."""
|
| 1031 |
+
self.use_tiling = False
|
| 1032 |
+
|
| 1033 |
+
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
| 1034 |
+
x = [
|
| 1035 |
+
self.model.encode(u.unsqueeze(0), self.scale).squeeze(0)
|
| 1036 |
+
for u in x
|
| 1037 |
+
]
|
| 1038 |
+
x = torch.stack(x)
|
| 1039 |
+
return x
|
| 1040 |
+
|
| 1041 |
+
@apply_forward_hook
|
| 1042 |
+
def encode(
|
| 1043 |
+
self, x: torch.Tensor, return_dict: bool = True
|
| 1044 |
+
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
|
| 1045 |
+
if self.use_tiling:
|
| 1046 |
+
return self.tiled_encode(x, tile_size=self.tile_size, tile_stride=self.tile_stride)
|
| 1047 |
+
|
| 1048 |
+
h = self._encode(x)
|
| 1049 |
+
|
| 1050 |
+
posterior = DiagonalGaussianDistribution(h)
|
| 1051 |
+
|
| 1052 |
+
if not return_dict:
|
| 1053 |
+
return (posterior,)
|
| 1054 |
+
return AutoencoderKLOutput(latent_dist=posterior)
|
| 1055 |
+
|
| 1056 |
+
def _decode(self, zs):
|
| 1057 |
+
dec = [
|
| 1058 |
+
self.model.decode(u.unsqueeze(0), self.scale).clamp_(-1, 1).squeeze(0)
|
| 1059 |
+
for u in zs
|
| 1060 |
+
]
|
| 1061 |
+
dec = torch.stack(dec)
|
| 1062 |
+
|
| 1063 |
+
return DecoderOutput(sample=dec)
|
| 1064 |
+
|
| 1065 |
+
@apply_forward_hook
|
| 1066 |
+
def decode(self, z: torch.Tensor, return_dict: bool = True) -> Union[DecoderOutput, torch.Tensor]:
|
| 1067 |
+
if self.use_tiling:
|
| 1068 |
+
return self.tiled_decode(z, tile_size=self.tile_size, tile_stride=self.tile_stride)
|
| 1069 |
+
|
| 1070 |
+
decoded = self._decode(z).sample
|
| 1071 |
+
|
| 1072 |
+
if not return_dict:
|
| 1073 |
+
return (decoded,)
|
| 1074 |
+
return DecoderOutput(sample=decoded)
|
| 1075 |
+
|
| 1076 |
+
@staticmethod
|
| 1077 |
+
def _build_1d_mask(length, left_bound, right_bound, border_width):
|
| 1078 |
+
"""Build a 1D blending mask with linear ramp at non-boundary edges."""
|
| 1079 |
+
mask = torch.ones((length,))
|
| 1080 |
+
if not left_bound:
|
| 1081 |
+
mask[:border_width] = (torch.arange(border_width) + 1) / border_width
|
| 1082 |
+
if not right_bound:
|
| 1083 |
+
mask[-border_width:] = torch.flip(
|
| 1084 |
+
(torch.arange(border_width) + 1) / border_width, dims=(0,)
|
| 1085 |
+
)
|
| 1086 |
+
return mask
|
| 1087 |
+
|
| 1088 |
+
def _build_spatial_mask(self, data, is_bound, border_width):
|
| 1089 |
+
"""Build a 2D spatial blending mask from H and W 1D masks.
|
| 1090 |
+
|
| 1091 |
+
Args:
|
| 1092 |
+
data: tensor of shape [B, C, T, H, W] (only H, W are used).
|
| 1093 |
+
is_bound: (top, bottom, left, right) booleans.
|
| 1094 |
+
border_width: (border_h, border_w) in pixel space.
|
| 1095 |
+
"""
|
| 1096 |
+
_, _, _, height, width = data.shape
|
| 1097 |
+
h_mask = self._build_1d_mask(height, is_bound[0], is_bound[1], border_width[0])
|
| 1098 |
+
w_mask = self._build_1d_mask(width, is_bound[2], is_bound[3], border_width[1])
|
| 1099 |
+
h_mask = h_mask.unsqueeze(1).expand(height, width)
|
| 1100 |
+
w_mask = w_mask.unsqueeze(0).expand(height, width)
|
| 1101 |
+
mask = torch.stack([h_mask, w_mask]).min(dim=0).values
|
| 1102 |
+
return mask.reshape(1, 1, 1, height, width)
|
| 1103 |
+
|
| 1104 |
+
def tiled_decode(self, z: torch.Tensor, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 1105 |
+
"""Decode latent with spatial tiling to reduce peak GPU memory.
|
| 1106 |
+
|
| 1107 |
+
Args:
|
| 1108 |
+
z: Latent tensor [B, C, T, H_latent, W_latent].
|
| 1109 |
+
tile_size: (tile_h, tile_w) in latent spatial units.
|
| 1110 |
+
tile_stride: (stride_h, stride_w) in latent spatial units.
|
| 1111 |
+
|
| 1112 |
+
Returns:
|
| 1113 |
+
DecoderOutput with decoded video [B, 3, T_out, H_out, W_out] clamped to [-1, 1].
|
| 1114 |
+
"""
|
| 1115 |
+
upsampling_factor = self.config.spatial_compression_ratio
|
| 1116 |
+
temporal_ratio = self.config.temporal_compression_ratio
|
| 1117 |
+
|
| 1118 |
+
_, _, latent_t, latent_h, latent_w = z.shape
|
| 1119 |
+
size_h, size_w = tile_size
|
| 1120 |
+
stride_h, stride_w = tile_stride
|
| 1121 |
+
|
| 1122 |
+
# Build tile task list (skip redundant trailing tiles)
|
| 1123 |
+
tasks = []
|
| 1124 |
+
for h in range(0, latent_h, stride_h):
|
| 1125 |
+
if h - stride_h >= 0 and h - stride_h + size_h >= latent_h:
|
| 1126 |
+
continue
|
| 1127 |
+
for w in range(0, latent_w, stride_w):
|
| 1128 |
+
if w - stride_w >= 0 and w - stride_w + size_w >= latent_w:
|
| 1129 |
+
continue
|
| 1130 |
+
tasks.append((h, h + size_h, w, w + size_w))
|
| 1131 |
+
|
| 1132 |
+
data_device = "cpu"
|
| 1133 |
+
computation_device = z.device
|
| 1134 |
+
|
| 1135 |
+
out_t = latent_t * temporal_ratio - (temporal_ratio - 1)
|
| 1136 |
+
out_h = latent_h * upsampling_factor
|
| 1137 |
+
out_w = latent_w * upsampling_factor
|
| 1138 |
+
|
| 1139 |
+
weight = torch.zeros((1, 1, out_t, out_h, out_w), dtype=z.dtype, device=data_device)
|
| 1140 |
+
values = torch.zeros((1, 3, out_t, out_h, out_w), dtype=z.dtype, device=data_device)
|
| 1141 |
+
|
| 1142 |
+
border_h = (size_h - stride_h) * upsampling_factor
|
| 1143 |
+
border_w = (size_w - stride_w) * upsampling_factor
|
| 1144 |
+
|
| 1145 |
+
for h_start, h_end, w_start, w_end in tasks:
|
| 1146 |
+
tile_latent = z[:, :, :, h_start:h_end, w_start:w_end].to(computation_device)
|
| 1147 |
+
tile_decoded = self.model.decode(tile_latent, self.scale).float().clamp_(-1, 1).to(data_device)
|
| 1148 |
+
|
| 1149 |
+
is_bound = (h_start == 0, h_end >= latent_h, w_start == 0, w_end >= latent_w)
|
| 1150 |
+
mask = self._build_spatial_mask(
|
| 1151 |
+
tile_decoded, is_bound=is_bound, border_width=(border_h, border_w)
|
| 1152 |
+
).to(dtype=z.dtype, device=data_device)
|
| 1153 |
+
|
| 1154 |
+
target_h = h_start * upsampling_factor
|
| 1155 |
+
target_w = w_start * upsampling_factor
|
| 1156 |
+
th = tile_decoded.shape[3]
|
| 1157 |
+
tw = tile_decoded.shape[4]
|
| 1158 |
+
values[:, :, :, target_h:target_h + th, target_w:target_w + tw] += tile_decoded * mask
|
| 1159 |
+
weight[:, :, :, target_h:target_h + th, target_w:target_w + tw] += mask
|
| 1160 |
+
|
| 1161 |
+
decoded = (values / weight).clamp_(-1, 1)
|
| 1162 |
+
return DecoderOutput(sample=decoded.to(z.device))
|
| 1163 |
+
|
| 1164 |
+
def tiled_encode(self, x: torch.Tensor, tile_size=(34, 34), tile_stride=(18, 16)):
|
| 1165 |
+
"""Encode video with spatial tiling to reduce peak GPU memory.
|
| 1166 |
+
|
| 1167 |
+
Args:
|
| 1168 |
+
x: Pixel video tensor [B, C, T, H, W].
|
| 1169 |
+
tile_size: (tile_h, tile_w) in latent spatial units.
|
| 1170 |
+
tile_stride: (stride_h, stride_w) in latent spatial units.
|
| 1171 |
+
|
| 1172 |
+
Returns:
|
| 1173 |
+
AutoencoderKLOutput with latent distribution.
|
| 1174 |
+
"""
|
| 1175 |
+
upsampling_factor = self.config.spatial_compression_ratio
|
| 1176 |
+
temporal_ratio = self.config.temporal_compression_ratio
|
| 1177 |
+
latent_channels = self.config.latent_channels
|
| 1178 |
+
|
| 1179 |
+
_, _, pixel_t, pixel_h, pixel_w = x.shape
|
| 1180 |
+
size_h = tile_size[0] * upsampling_factor
|
| 1181 |
+
size_w = tile_size[1] * upsampling_factor
|
| 1182 |
+
stride_h = tile_stride[0] * upsampling_factor
|
| 1183 |
+
stride_w = tile_stride[1] * upsampling_factor
|
| 1184 |
+
|
| 1185 |
+
tasks = []
|
| 1186 |
+
for h in range(0, pixel_h, stride_h):
|
| 1187 |
+
if h - stride_h >= 0 and h - stride_h + size_h >= pixel_h:
|
| 1188 |
+
continue
|
| 1189 |
+
for w in range(0, pixel_w, stride_w):
|
| 1190 |
+
if w - stride_w >= 0 and w - stride_w + size_w >= pixel_w:
|
| 1191 |
+
continue
|
| 1192 |
+
tasks.append((h, h + size_h, w, w + size_w))
|
| 1193 |
+
|
| 1194 |
+
data_device = "cpu"
|
| 1195 |
+
computation_device = x.device
|
| 1196 |
+
|
| 1197 |
+
latent_h = pixel_h // upsampling_factor
|
| 1198 |
+
latent_w = pixel_w // upsampling_factor
|
| 1199 |
+
latent_t = (pixel_t + temporal_ratio - 1) // temporal_ratio
|
| 1200 |
+
|
| 1201 |
+
weight = torch.zeros((1, 1, latent_t, latent_h, latent_w), dtype=x.dtype, device=data_device)
|
| 1202 |
+
values = torch.zeros((1, latent_channels * 2, latent_t, latent_h, latent_w), dtype=x.dtype, device=data_device)
|
| 1203 |
+
|
| 1204 |
+
border_h = tile_size[0] - tile_stride[0]
|
| 1205 |
+
border_w = tile_size[1] - tile_stride[1]
|
| 1206 |
+
|
| 1207 |
+
for h_start, h_end, w_start, w_end in tasks:
|
| 1208 |
+
tile_pixel = x[:, :, :, h_start:h_end, w_start:w_end].to(computation_device)
|
| 1209 |
+
tile_encoded = self.model.encode(tile_pixel, self.scale).float().to(data_device)
|
| 1210 |
+
|
| 1211 |
+
is_bound = (h_start == 0, h_end >= pixel_h, w_start == 0, w_end >= pixel_w)
|
| 1212 |
+
mask = self._build_spatial_mask(
|
| 1213 |
+
tile_encoded, is_bound=is_bound, border_width=(border_h, border_w)
|
| 1214 |
+
).to(dtype=x.dtype, device=data_device)
|
| 1215 |
+
|
| 1216 |
+
target_h = h_start // upsampling_factor
|
| 1217 |
+
target_w = w_start // upsampling_factor
|
| 1218 |
+
th = tile_encoded.shape[3]
|
| 1219 |
+
tw = tile_encoded.shape[4]
|
| 1220 |
+
values[:, :, :, target_h:target_h + th, target_w:target_w + tw] += tile_encoded * mask
|
| 1221 |
+
weight[:, :, :, target_h:target_h + th, target_w:target_w + tw] += mask
|
| 1222 |
+
|
| 1223 |
+
latent = (values / weight).to(x.device)
|
| 1224 |
+
posterior = DiagonalGaussianDistribution(latent)
|
| 1225 |
+
return AutoencoderKLOutput(latent_dist=posterior)
|
| 1226 |
+
|
| 1227 |
+
@classmethod
|
| 1228 |
+
def from_pretrained(cls, pretrained_model_path, additional_kwargs={}):
|
| 1229 |
+
def filter_kwargs(cls, kwargs):
|
| 1230 |
+
import inspect
|
| 1231 |
+
sig = inspect.signature(cls.__init__)
|
| 1232 |
+
valid_params = set(sig.parameters.keys()) - {'self', 'cls'}
|
| 1233 |
+
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
| 1234 |
+
return filtered_kwargs
|
| 1235 |
+
|
| 1236 |
+
model = cls(**filter_kwargs(cls, additional_kwargs))
|
| 1237 |
+
if pretrained_model_path.endswith(".safetensors"):
|
| 1238 |
+
from safetensors.torch import load_file, safe_open
|
| 1239 |
+
state_dict = load_file(pretrained_model_path)
|
| 1240 |
+
else:
|
| 1241 |
+
state_dict = torch.load(pretrained_model_path, map_location="cpu")
|
| 1242 |
+
tmp_state_dict = {}
|
| 1243 |
+
for key in state_dict:
|
| 1244 |
+
tmp_state_dict["model." + key] = state_dict[key]
|
| 1245 |
+
state_dict = tmp_state_dict
|
| 1246 |
+
m, u = model.load_state_dict(state_dict, strict=False)
|
| 1247 |
+
|
| 1248 |
+
return model
|