Abdullahcoder54 commited on
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 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: Video Creator
3
- emoji: 🔥
4
  colorFrom: green
5
  colorTo: blue
6
  sdk: docker
7
  app_file: app.py
8
  pinned: false
 
 
9
  ---
10
 
11
- Docker-hosted LTX-2.5 Space for the AI Shorts Factory backend. Four Gradio
12
- endpoints via `/call/{fn}` (text_to_video, text_to_video_av, image_to_video,
13
- image_to_video_av). Requires a GPU tier — the free CPU tier (2 vCPU/16GB)
14
- boots but generation OOMs.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- """LTX-2.5 Space — ZeroGPU inference API for the AI Shorts Factory backend.
2
 
3
- Runs the distilled LTX-2.5 GGUF (Q4_K_M) on zero GPU as a Gradio demo. Four
4
- endpoints are exposed through Gradio's `/call/{fn}` protocol and consumed by
5
- the backend provider `app.ai.visuals.ltx`:
6
 
7
- text_to_video(prompt, width, height, num_frames, enhance_prompt)
8
- text_to_video_av(...) # video + synchronized 48kHz narration audio
9
- image_to_video(prompt, image_url, width, height, num_frames, enhance_prompt)
10
- image_to_video_av(...) # video + synchronized narration audio
 
11
 
12
- Every function returns ``(video_path, first_frame_path)``. Scene 2..N of the
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
- The backend always sends an already-composed prompt (quoted dialogue for the
17
- exact script words) and passes ``enhance_prompt=False``; the prompt enhancer
18
- is a hard no-op here so no rewrite happens.
19
- """
 
 
 
20
 
21
- from __future__ import annotations
 
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
- # Distilled model: guidance is baked into the weights, few-step schedule.
56
- INFERENCE_STEPS = int(os.environ.get("LTX_STEPS", "8"))
57
- GUIDANCE_SCALE = float(os.environ.get("LTX_GUIDANCE", "1.0"))
58
- NEGATIVE_PROMPT = "low quality, blurry, distorted, watermark"
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
- _lock = threading.Lock()
66
- _model = None
67
 
 
 
 
 
 
 
 
 
68
 
69
- # ----------------------------------------------------------------- model load
 
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
- reader = GGUFReader(path)
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
- def _build_transformer(gguf_path: str):
145
- try:
146
- # diffusers 0.40 (pinned): the LTX-2 video transformer is loadable from
147
- # a single GGUF file via FromOriginalModelMixin.
148
- from diffusers import LTX2VideoTransformer3DModel as TransformerCls
149
- except ImportError: # older versions only expose the generic AutoModel
150
- from diffusers import AutoModel as TransformerCls
151
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
152
  try:
153
- # GGUF config lives in quantizers since 0.33; newer releases keep it there.
154
- from diffusers.quantizers.quantization_config import GGUFQuantizationConfig
155
- except ImportError: # legacy export path
156
- from diffusers.utils import GGUFQuantizationConfig
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
- print("[ltx] downloading GGUF transformer…")
190
- gguf_path = _ensure_gguf()
 
 
 
 
 
191
 
192
- print("[ltx] building quantized transformer…")
193
- transformer = _build_transformer(gguf_path)
 
 
 
194
 
195
- from diffusers import LTX2ImageToVideoPipeline, LTX2Pipeline
196
 
197
- built = {
198
- "transformer": transformer,
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
- for pipe in (built["t2v"], built["i2v"]):
219
- pipe.enable_model_cpu_offload()
220
- _model = built
221
- print("[ltx] model ready")
222
- return _model
223
 
224
 
225
- # ZeroGPU rule #2: load the model at module scope so weights are packed once
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
- # ----------------------------------------------------------------- generation
 
234
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
235
 
236
- def _try_call(pipe, **kwargs):
237
- """Call the pipeline, disabling the (Gemma-4 based) prompt enhancer."""
238
- try:
239
- return pipe(enhance_prompt=False, **kwargs)
240
- except TypeError:
241
- return pipe(**kwargs)
242
-
243
-
244
- def _extract_frames(output) -> list:
245
- frames = output.frames
246
- if isinstance(frames[0], (list, tuple)):
247
- return list(frames[0])
248
- return list(frames)
249
-
250
-
251
- def _extract_audio(output):
252
- audio = getattr(output, "audio", None)
253
- if audio is None:
254
- return None
255
- if isinstance(audio, (list, tuple)):
256
- arr, sr = audio[0], audio[1] if len(audio) > 1 else 48000
257
- else:
258
- arr, sr = audio, 48000
259
- if torch.is_tensor(arr):
260
- arr = arr.detach().float().cpu().numpy()
261
- arr = np.asarray(arr)
262
- if arr.ndim == 3: # (batch, channels, samples)
263
- arr = arr[0]
264
- if arr.ndim == 2: # collapse to mono; channel axis is the smaller dimension
265
- arr = arr.mean(axis=int(arr.shape[0] < arr.shape[1]))
266
- return np.ascontiguousarray(arr.astype(np.float32)), int(sr)
267
-
268
-
269
- def _mux_audio(video_path, frames, fps, audio) -> str:
270
- from diffusers.utils import export_to_video
271
-
272
- export_to_video(frames, video_path, fps=fps)
273
- if audio is None:
274
- return video_path
275
-
276
- import imageio_ffmpeg
277
- from scipy.io import wavfile
278
-
279
- arr, sr = audio
280
- vp = Path(video_path)
281
- wav_path = vp.with_suffix(".wav")
282
- wavfile.write(str(wav_path), sr, arr)
283
- ffmpeg = imageio_ffmpeg.get_ffmpeg_exe()
284
- muxed = str(vp.with_name(vp.stem + "_mux.mp4"))
285
- subprocess.run(
286
- [
287
- ffmpeg, "-y",
288
- "-i", str(video_path),
289
- "-i", str(wav_path),
290
- "-map", "0:v", "-map", "1:a",
291
- "-c:v", "copy", "-c:a", "aac", "-shortest",
292
- muxed,
293
- ],
294
- check=True,
295
- capture_output=True,
 
 
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
- num_frames = max(9, int(num_frames))
325
- num_frames = 1 + (num_frames - 1 - ((num_frames - 1) % 8)) # 1 + 8n
326
- width = int(width) // 32 * 32
327
- height = int(height) // 32 * 32
328
-
329
- model = _load_model()
330
- generator = torch.Generator(device="cuda").manual_seed(SEED)
331
-
332
- image = _load_image(image_url)
333
- common = dict(
334
- prompt=prompt,
335
- negative_prompt=NEGATIVE_PROMPT,
336
- width=width,
337
- height=height,
338
- num_frames=num_frames,
339
- num_inference_steps=INFERENCE_STEPS,
340
- guidance_scale=GUIDANCE_SCALE,
341
- generator=generator,
342
  )
343
- if image is not None:
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
- def text_to_video(prompt, width, height, num_frames, enhance_prompt=False):
363
- return _generate(prompt, "", width, height, num_frames, enhance_prompt, with_audio=False)
 
 
 
 
 
364
 
 
365
 
366
- def text_to_video_av(prompt, width, height, num_frames, enhance_prompt=False):
367
- return _generate(prompt, "", width, height, num_frames, enhance_prompt, with_audio=True)
368
 
369
-
370
- def image_to_video(prompt, image_url, width, height, num_frames, enhance_prompt=False):
371
- return _generate(prompt, image_url, width, height, num_frames, enhance_prompt, with_audio=False)
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
- width = gr.Slider(256, 768, value=544, step=32, label="Width (÷32)")
392
- height = gr.Slider(256, 1408, value=960, step=32, label="Height (÷32)")
393
- num_frames = gr.Slider(9, 801, value=97, step=8, label="Frames (8n+1)")
394
- enhance = gr.Checkbox(value=False, label="Enhance prompt (kept off)")
395
- video_out = gr.Video(label="Clip")
396
- frame_out = gr.Image(label="First frame (for chaining)")
397
-
398
- with gr.Tab("Text → Video"):
399
- t2v_prompt = gr.Textbox(label="Prompt", lines=4)
400
- t2v_btn = gr.Button("Generate (video only)")
401
- t2v_btn.click(
402
- t2v_fns[0],
403
- inputs=[t2v_prompt, width, height, num_frames, enhance],
404
- outputs=[video_out, frame_out],
405
  )
406
 
407
- with gr.Tab("Text → Video + Audio"):
408
- t2va_prompt = gr.Textbox(label="Prompt", lines=4)
409
- t2va_btn = gr.Button("Generate (with narration)")
410
- t2va_btn.click(
411
- t2v_fns[1],
412
- inputs=[t2va_prompt, width, height, num_frames, enhance],
413
- outputs=[video_out, frame_out],
414
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
415
 
416
- with gr.Tab("Image → Video"):
417
- i2v_prompt = gr.Textbox(label="Prompt", lines=4)
418
- i2v_image = gr.Textbox(
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
- with gr.Tab("Image → Video + Audio"):
430
- i2va_prompt = gr.Textbox(label="Prompt", lines=4)
431
- i2va_image = gr.Textbox(
432
- label="First-frame image URL (last frame from the previous clip)",
433
- value="",
434
- )
435
- i2va_btn = gr.Button("Generate (with narration)")
436
- i2va_btn.click(
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.queue(default_concurrency_limit=1).launch(server_name="0.0.0.0")
 
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

  • SHA256: 9398bf37f28f2d54a6d1222ceeba53897d7bca4f3be2b94f33bbdc37eedec479
  • Pointer size: 130 Bytes
  • Size of remote file: 92.8 kB
examples/case2.jpg ADDED

Git LFS Details

  • SHA256: 717580395f94795f59d3f753fc7dceb841efbe5cae4783e2a8255226c576fe92
  • Pointer size: 131 Bytes
  • Size of remote file: 146 kB
examples/case3.jpg ADDED

Git LFS Details

  • SHA256: 7cd93e6e9a1a8a5eed7233e7623d2b65905d01163549333be0255fd566043385
  • Pointer size: 130 Bytes
  • Size of remote file: 29.9 kB
examples/case4.jpg ADDED

Git LFS Details

  • SHA256: cc1e04f26401b08c4c317b12dd257fc965865dd6e8020c813b918369d2dc06eb
  • Pointer size: 130 Bytes
  • Size of remote file: 46.9 kB
examples/case5.jpg ADDED

Git LFS Details

  • SHA256: 6eac5e4a3fb79810f6324ccb6b7189ff28f5b9bf189a0a6162de3f0212bafa85
  • Pointer size: 130 Bytes
  • Size of remote file: 94.2 kB
examples/case6.jpg ADDED

Git LFS Details

  • SHA256: fdb9fd60480cac43e3ad7d5b175ff8ff67a021164f82fea66d184e4e10cb2cb6
  • Pointer size: 130 Bytes
  • Size of remote file: 81.1 kB
requirements.txt CHANGED
@@ -1,18 +1,18 @@
1
- # For ZeroGPU, gradio/spaces/huggingface_hub/torch are preinstalled and
2
- # platform-managed; the Dockerfile installs them explicitly.
3
- # Pinned to 0.40: LTX2 video transformer + GGUF single-file loading both work
4
- # here (LTX2VideoTransformer3DModel + GGUFQuantizer). Earlier versions lack
5
- # AutoModel/from_single_file support for LTX-2 GGUF.
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
- scipy
 
14
  pillow
15
- sentencepiece
16
- protobuf
17
- imageio-ffmpeg
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