umer1995 commited on
Commit
f0a4e91
·
verified ·
1 Parent(s): 732d269

Fun CN: fp8 stream DiT + local bnb4 TE + xlarge (no bf16 host dump / no remote TE)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +6 -0
  2. README.md +19 -9
  3. app.py +335 -67
  4. requirements.txt +1 -0
  5. vendor/VideoX-Fun/LICENSE +201 -0
  6. vendor/VideoX-Fun/README.md +733 -0
  7. vendor/VideoX-Fun/config/flux2/flux2_control.yaml +5 -0
  8. vendor/VideoX-Fun/config/qwenimage/qwenimage_control.yaml +5 -0
  9. vendor/VideoX-Fun/config/wan2.1/wan_civitai.yaml +39 -0
  10. vendor/VideoX-Fun/config/wan2.2/wan_civitai_5b.yaml +41 -0
  11. vendor/VideoX-Fun/config/wan2.2/wan_civitai_animate.yaml +41 -0
  12. vendor/VideoX-Fun/config/wan2.2/wan_civitai_i2v.yaml +43 -0
  13. vendor/VideoX-Fun/config/wan2.2/wan_civitai_s2v.yaml +44 -0
  14. vendor/VideoX-Fun/config/wan2.2/wan_civitai_t2v.yaml +43 -0
  15. vendor/VideoX-Fun/config/z_image/z_image_control.yaml +5 -0
  16. vendor/VideoX-Fun/config/z_image/z_image_control_2.0.yaml +8 -0
  17. vendor/VideoX-Fun/config/z_image/z_image_control_2.1.yaml +8 -0
  18. vendor/VideoX-Fun/config/z_image/z_image_control_2.1_lite.yaml +8 -0
  19. vendor/VideoX-Fun/config/zero_stage2_config.json +16 -0
  20. vendor/VideoX-Fun/config/zero_stage3_config.json +28 -0
  21. vendor/VideoX-Fun/config/zero_stage3_config_cpu_offload.json +28 -0
  22. vendor/VideoX-Fun/examples/cogvideox_fun/app.py +73 -0
  23. vendor/VideoX-Fun/examples/cogvideox_fun/launch_api.py +90 -0
  24. vendor/VideoX-Fun/examples/cogvideox_fun/post_infer.py +150 -0
  25. vendor/VideoX-Fun/examples/cogvideox_fun/post_infer_queue.py +145 -0
  26. vendor/VideoX-Fun/examples/cogvideox_fun/predict_i2v.py +328 -0
  27. vendor/VideoX-Fun/examples/cogvideox_fun/predict_t2v.py +268 -0
  28. vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v.py +263 -0
  29. vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v_control.py +248 -0
  30. vendor/VideoX-Fun/examples/ernie_image/predict_t2i.py +210 -0
  31. vendor/VideoX-Fun/examples/fantasytalking/predict_s2v.py +335 -0
  32. vendor/VideoX-Fun/examples/flashhead/predict_s2v.py +262 -0
  33. vendor/VideoX-Fun/examples/flux/predict_t2i.py +224 -0
  34. vendor/VideoX-Fun/examples/flux2/predict_t2i.py +218 -0
  35. vendor/VideoX-Fun/examples/flux2_fun/predict_i2i_inpaint.py +258 -0
  36. vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control.py +258 -0
  37. vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control_ref.py +258 -0
  38. vendor/VideoX-Fun/examples/hunyuanvideo/predict_i2v.py +270 -0
  39. vendor/VideoX-Fun/examples/hunyuanvideo/predict_t2v.py +255 -0
  40. vendor/VideoX-Fun/examples/infinitetalk/predict_s2v.py +319 -0
  41. vendor/VideoX-Fun/examples/lens/predict_t2i.py +226 -0
  42. vendor/VideoX-Fun/examples/longcatvideo/predict_i2v.py +247 -0
  43. vendor/VideoX-Fun/examples/longcatvideo/predict_s2v_avatar.py +293 -0
  44. vendor/VideoX-Fun/examples/longcatvideo/predict_t2v.py +239 -0
  45. vendor/VideoX-Fun/examples/ltx2.3/predict_i2v.py +305 -0
  46. vendor/VideoX-Fun/examples/ltx2.3/predict_t2v.py +300 -0
  47. vendor/VideoX-Fun/examples/ltx2/predict_i2v.py +281 -0
  48. vendor/VideoX-Fun/examples/ltx2/predict_i2v_upsample.py +326 -0
  49. vendor/VideoX-Fun/examples/ltx2/predict_t2v.py +276 -0
  50. vendor/VideoX-Fun/examples/mova/predict_i2v.py +380 -0
.gitattributes CHANGED
@@ -33,3 +33,9 @@ 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
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/before_vcut/--C66yU3LjM_2.mp4 filter=lfs diff=lfs merge=lfs -text
37
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/train/--C66yU3LjM_2-Scene-001.mp4 filter=lfs diff=lfs merge=lfs -text
38
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/train/--C66yU3LjM_2-Scene-002.mp4 filter=lfs diff=lfs merge=lfs -text
39
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/train/--C66yU3LjM_2-Scene-003.mp4 filter=lfs diff=lfs merge=lfs -text
40
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/train/--C66yU3LjM_2-Scene-004.mp4 filter=lfs diff=lfs merge=lfs -text
41
+ vendor/VideoX-Fun/videox_fun/video_caption/datasets/panda_70m/train/--C66yU3LjM_2-Scene-005.mp4 filter=lfs diff=lfs merge=lfs -text
README.md CHANGED
@@ -8,7 +8,7 @@ sdk_version: "5.49.1"
8
  app_file: app.py
9
  pinned: false
10
  license: other
11
- short_description: Real Fun depth ControlNet (NOT soft image=depth)
12
  ---
13
 
14
  # Track R — FLUX.2 Fun depth ControlNet (ZeroGPU)
@@ -16,22 +16,32 @@ short_description: Real Fun depth ControlNet (NOT soft image=depth)
16
  **Hard rule:** This Space runs **real** ALIMAMA / VideoX-Fun Fun ControlNet Union depth.
17
  Soft Flux2 `image=depth` is **banned forever** on this path.
18
 
19
- ## Stack (VRAM choice)
20
 
21
  | Piece | Value |
22
  |---|---|
23
- | Painter + CN | VideoX-Fun `Flux2ControlPipeline` + `FLUX.2-dev-Fun-Controlnet-Union-2602` |
24
- | Base | `black-forest-labs/FLUX.2-dev` |
25
- | GPU | `@spaces.GPU(duration=300, size="large")` = **48GB** @ **1×** Pro minutes |
26
- | Memory mode | VideoX-Fun **`model_cpu_offload_and_qfloat8`** (official low-VRAM Fun CN path) |
27
- | Weights | **HF Mount volumes** (not ephemeral download): `/data/FLUX.2-dev` + `/data/Fun-CN` |
28
- | Escalation | If OOM → redeploy with `size="xlarge"` + `model_cpu_offload` ( quota) |
 
29
 
30
  Requires Space secret **`HF_TOKEN`** (gated FLUX.2-dev license accepted on the account).
31
 
 
 
 
 
 
 
 
 
 
32
  ## API
33
 
34
- Same client contract as Plan A soft Space: `/generate_still`
35
  (`positive`, `negative`, `depth_image`, `seed`, `width`, `height`, `steps`, `guidance`, `cn_strength`)
36
 
37
  Depth is **required**. Soft depth is never used.
 
8
  app_file: app.py
9
  pinned: false
10
  license: other
11
+ short_description: Real Fun depth ControlNet fp8+xlarge (NOT soft image=depth)
12
  ---
13
 
14
  # Track R — FLUX.2 Fun depth ControlNet (ZeroGPU)
 
16
  **Hard rule:** This Space runs **real** ALIMAMA / VideoX-Fun Fun ControlNet Union depth.
17
  Soft Flux2 `image=depth` is **banned forever** on this path.
18
 
19
+ ## Stack (2026-07-21 host-RAM fix)
20
 
21
  | Piece | Value |
22
  |---|---|
23
+ | Painter | FLUX.2-dev DiT **streamed to float8** (never full bf16 host materialize) |
24
+ | CN | VideoX-Fun `Flux2ControlPipeline` + `FLUX.2-dev-Fun-Controlnet-Union-2602` |
25
+ | TE | **Local** quantized `diffusers/FLUX.2-dev-bnb-4bit` (NO HF remote TE) |
26
+ | GPU | `@spaces.GPU(duration=300, size="xlarge")` = **96GB** @ **2×** Pro minutes |
27
+ | Memory mode | `model_cpu_offload_and_qfloat8` (weights already fp8 at load) |
28
+ | Weights | HF Mount `/data/FLUX.2-dev` + `/data/Fun-CN` (+ TE download/cache) |
29
+ | Abort | `NFA_FUN_CN_LOAD_DEADLINE_SEC` (default 480) if `Fun CN loaded` not reached |
30
 
31
  Requires Space secret **`HF_TOKEN`** (gated FLUX.2-dev license accepted on the account).
32
 
33
+ ## Env
34
+
35
+ | Var | Default |
36
+ |---|---|
37
+ | `NFA_FUN_CN_GPU_SIZE` | `xlarge` |
38
+ | `NFA_FUN_CN_MEM_MODE` | `model_cpu_offload_and_qfloat8` |
39
+ | `NFA_FLUX2_TE_MODEL_ID` | `diffusers/FLUX.2-dev-bnb-4bit` |
40
+ | `NFA_FUN_CN_LOAD_DEADLINE_SEC` | `480` |
41
+
42
  ## API
43
 
44
+ Same client contract: `/generate_still`
45
  (`positive`, `negative`, `depth_image`, `seed`, `width`, `height`, `steps`, `guidance`, `cn_strength`)
46
 
47
  Depth is **required**. Soft depth is never used.
app.py CHANGED
@@ -2,19 +2,23 @@
2
 
3
  REAL Fun ControlNet Union depth — NOT soft Flux2 image=depth (banned forever).
4
 
5
- VRAM choice (documented):
6
- size=\"large\" (48GB, Pro) + VideoX-Fun model_cpu_offload_and_qfloat8.
7
- Weights via HF Mount volumes (not 178GB ephemeral download).
8
- CPU-preload outside @spaces.GPU so ZeroGPU minutes are not burned on load.
9
- Escalation if OOM: size=\"xlarge\" + model_cpu_offload (2× quota).
10
  """
11
 
12
  from __future__ import annotations
13
 
 
 
 
14
  import os
15
  import shutil
16
  import subprocess
17
  import sys
 
18
  import traceback
19
  from pathlib import Path
20
  from typing import Optional
@@ -30,6 +34,9 @@ APP_DIR = Path(__file__).resolve().parent
30
  _VX = APP_DIR / "vendor" / "VideoX-Fun"
31
  _VX_CACHE = Path.home() / "VideoX-Fun"
32
 
 
 
 
33
 
34
  def _patch_videox_inits(root: Path) -> None:
35
  models_init = root / "videox_fun" / "models" / "__init__.py"
@@ -114,29 +121,32 @@ if HF_TOKEN:
114
  BASE_MODEL = os.environ.get(
115
  "NFA_FLUX2_MODEL_ID", "black-forest-labs/FLUX.2-dev"
116
  ).strip()
 
 
 
 
117
  CN_REPO = os.environ.get(
118
  "NFA_FUN_CN_REPO", "alibaba-pai/FLUX.2-dev-Fun-Controlnet-Union"
119
  ).strip()
120
  CN_FILE = os.environ.get(
121
  "NFA_FUN_CN_FILE", "FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors"
122
  ).strip()
123
- GPU_SIZE = (os.environ.get("NFA_FUN_CN_GPU_SIZE") or "large").strip().lower()
 
124
  if GPU_SIZE not in ("large", "xlarge"):
125
- GPU_SIZE = "large"
126
  GPU_DURATION = int(os.environ.get("NFA_FUN_CN_GPU_DURATION") or "300")
127
  WEIGHT_DTYPE = torch.bfloat16
128
  MEM_MODE = (
129
- os.environ.get("NFA_FUN_CN_MEM_MODE")
130
- or (
131
- "model_cpu_offload"
132
- if GPU_SIZE == "xlarge"
133
- else "model_cpu_offload_and_qfloat8"
134
- )
135
  ).strip()
 
 
136
 
137
  CONFIG_PATH = APP_DIR / "config" / "flux2_control.yaml"
138
  MODEL_DIR = Path(os.environ.get("NFA_FLUX2_MOUNT") or "/data/FLUX.2-dev")
139
  CN_MOUNT_DIR = Path(os.environ.get("NFA_FUN_CN_MOUNT") or "/data/Fun-CN")
 
140
  CACHE_ROOT = Path(
141
  os.environ.get("NFA_FUN_CN_CACHE") or (Path.home() / ".cache" / "nfa_fun_cn")
142
  )
@@ -145,6 +155,30 @@ _PIPE = None
145
  _PIPE_OFFLOAD_READY = False
146
  _CN_FILE_PATH: Path | None = None
147
  _GET_IMAGE_LATENT = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
148
 
149
 
150
  def _resolve_cn_path() -> Path:
@@ -188,7 +222,6 @@ def _ensure_weights() -> None:
188
  "transformer/*",
189
  "vae/*",
190
  "tokenizer/*",
191
- "text_encoder/*",
192
  "scheduler/*",
193
  ],
194
  )
@@ -199,6 +232,37 @@ def _ensure_weights() -> None:
199
  _CN_FILE_PATH = Path(path)
200
 
201
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
202
  def _prep_depth(depth_image: Image.Image, width: int, height: int) -> Image.Image:
203
  img = depth_image.convert("RGB")
204
  if img.size != (width, height):
@@ -210,18 +274,236 @@ def _compose_prompt(positive: str, negative: str) -> tuple[str, str]:
210
  return (positive or "").strip(), ((negative or "").strip() or " ")
211
 
212
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
  def get_pipe(*, prepare_gpu_offload: bool = False):
214
- """CPU-load Fun CN pipeline; arm GPU offload only inside @spaces.GPU."""
215
- global _PIPE, _PIPE_OFFLOAD_READY, _GET_IMAGE_LATENT
216
  if _PIPE is None:
 
 
217
  _ensure_weights()
 
218
  _ensure_videox_on_path()
219
  from diffusers import FlowMatchEulerDiscreteScheduler
220
- from safetensors.torch import load_file
221
- from transformers import Mistral3ForConditionalGeneration, PixtralProcessor
222
- from videox_fun.models.flux2_transformer2d_control import (
223
- Flux2ControlTransformer2DModel,
224
- )
225
  from videox_fun.models.flux2_vae import AutoencoderKLFlux2
226
  from videox_fun.pipeline.pipeline_flux2_control import Flux2ControlPipeline
227
  from videox_fun.utils.utils import get_image_latent
@@ -229,37 +511,16 @@ def get_pipe(*, prepare_gpu_offload: bool = False):
229
  _GET_IMAGE_LATENT = get_image_latent
230
  model_name = str(MODEL_DIR)
231
  cn_file = str(_resolve_cn_path())
232
- config = OmegaConf.load(str(CONFIG_PATH))
233
- print(
234
- f"[nfa-fun-cn] CPU-load Flux2Control + Fun CN mem={MEM_MODE}",
235
- flush=True,
236
- )
237
- transformer = Flux2ControlTransformer2DModel.from_pretrained(
238
- model_name,
239
- subfolder="transformer",
240
- low_cpu_mem_usage=True,
241
- torch_dtype=WEIGHT_DTYPE,
242
- transformer_additional_kwargs=OmegaConf.to_container(
243
- config["transformer_additional_kwargs"]
244
- ),
245
- ).to(WEIGHT_DTYPE)
246
- state_dict = load_file(cn_file)
247
- state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
248
- missing, unexpected = transformer.load_state_dict(state_dict, strict=False)
249
- print(
250
- f"[nfa-fun-cn] Fun CN loaded missing={len(missing)} unexpected={len(unexpected)}",
251
- flush=True,
252
- )
253
  vae = AutoencoderKLFlux2.from_pretrained(model_name, subfolder="vae").to(
254
  WEIGHT_DTYPE
255
  )
 
256
  tokenizer = PixtralProcessor.from_pretrained(model_name, subfolder="tokenizer")
257
- text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
258
- model_name,
259
- subfolder="text_encoder",
260
- torch_dtype=WEIGHT_DTYPE,
261
- low_cpu_mem_usage=True,
262
- )
263
  scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
264
  model_name, subfolder="scheduler"
265
  )
@@ -270,32 +531,33 @@ def get_pipe(*, prepare_gpu_offload: bool = False):
270
  transformer=transformer,
271
  scheduler=scheduler,
272
  )
273
- print("[nfa-fun-cn] Flux2ControlPipeline CPU-ready (REAL Fun CN)", flush=True)
 
 
 
 
 
 
274
 
275
  if prepare_gpu_offload and not _PIPE_OFFLOAD_READY and torch.cuda.is_available():
276
- from videox_fun.utils.fp8_optimization import (
277
- convert_model_weight_to_float8,
278
- convert_weight_dtype_wrapper,
279
- )
280
 
281
  device = "cuda"
282
  transformer = _PIPE.transformer
283
- if MEM_MODE == "model_cpu_offload_and_qfloat8":
284
- convert_model_weight_to_float8(
285
- transformer,
286
- exclude_module_name=["img_in", "txt_in", "timestep"],
287
- device=device,
288
- )
289
  convert_weight_dtype_wrapper(transformer, WEIGHT_DTYPE)
290
- _PIPE.enable_model_cpu_offload(device=device)
291
- elif MEM_MODE == "sequential_cpu_offload":
292
  _PIPE.enable_sequential_cpu_offload(device=device)
293
- elif MEM_MODE == "model_cpu_offload":
294
  _PIPE.enable_model_cpu_offload(device=device)
295
  else:
296
  _PIPE.to(device=device)
297
  _PIPE_OFFLOAD_READY = True
298
- print(f"[nfa-fun-cn] GPU offload armed mem={MEM_MODE}", flush=True)
299
  return _PIPE
300
 
301
 
@@ -336,7 +598,8 @@ def _generate_still_gpu(
336
  generator = torch.Generator(device=device).manual_seed(int(seed))
337
  print(
338
  f"[nfa-fun-cn] REAL Fun CN generate seed={seed} {w}x{h} steps={steps} "
339
- f"cn={strength} path=videox_fun_flux2_control",
 
340
  flush=True,
341
  )
342
  with torch.no_grad():
@@ -376,6 +639,10 @@ def generate_still(
376
  _ensure_weights()
377
  # Wall-clock CPU load — does not burn ZeroGPU minutes.
378
  get_pipe(prepare_gpu_offload=False)
 
 
 
 
379
  return _generate_still_gpu(
380
  positive,
381
  negative or "",
@@ -399,10 +666,11 @@ with gr.Blocks(title="NFA Track R FLUX.2 Fun CN ZeroGPU") as demo:
399
  gr.Markdown(
400
  "## NFA Track R — **Real Fun depth ControlNet** (ZeroGPU)\n"
401
  f"- Stack: VideoX-Fun `Flux2ControlPipeline` + `{CN_FILE}`\n"
402
- f"- Base: `{BASE_MODEL}` (HF Mount `/data/FLUX.2-dev`)\n"
 
403
  f"- GPU: `size={GPU_SIZE}` duration={GPU_DURATION}s mem=`{MEM_MODE}`\n"
404
- "- Soft `image=depth` is **banned** on this Space.\n"
405
- "- First call CPU-loads weights (slow once), then GPU Fun CN infer."
406
  )
407
  with gr.Row():
408
  with gr.Column():
 
2
 
3
  REAL Fun ControlNet Union depth — NOT soft Flux2 image=depth (banned forever).
4
 
5
+ Host-RAM fix (2026-07-21):
6
+ Prior hang: bf16 from_pretrained DiT+TE (~112GB) then post-hoc qfloat8.
7
+ Now: stream DiT shards straight into float8 (never full bf16 materialize),
8
+ local quantized TE (bnb-4bit / fp8-class NO HF remote TE),
9
+ ZeroGPU size=xlarge (96GB). Abort early if Fun CN not loaded in time.
10
  """
11
 
12
  from __future__ import annotations
13
 
14
+ import gc
15
+ import glob
16
+ import json
17
  import os
18
  import shutil
19
  import subprocess
20
  import sys
21
+ import time
22
  import traceback
23
  from pathlib import Path
24
  from typing import Optional
 
34
  _VX = APP_DIR / "vendor" / "VideoX-Fun"
35
  _VX_CACHE = Path.home() / "VideoX-Fun"
36
 
37
+ # Exclude from float8 (VideoX Fun CN + embedding / timestep stability).
38
+ _FP8_EXCLUDE = ("img_in", "txt_in", "timestep", "control_img_in", "embed")
39
+
40
 
41
  def _patch_videox_inits(root: Path) -> None:
42
  models_init = root / "videox_fun" / "models" / "__init__.py"
 
121
  BASE_MODEL = os.environ.get(
122
  "NFA_FLUX2_MODEL_ID", "black-forest-labs/FLUX.2-dev"
123
  ).strip()
124
+ # Pre-quantized local TE (bnb-4bit). Remote TE is banned (HF endpoint broken).
125
+ TE_MODEL = os.environ.get(
126
+ "NFA_FLUX2_TE_MODEL_ID", "diffusers/FLUX.2-dev-bnb-4bit"
127
+ ).strip()
128
  CN_REPO = os.environ.get(
129
  "NFA_FUN_CN_REPO", "alibaba-pai/FLUX.2-dev-Fun-Controlnet-Union"
130
  ).strip()
131
  CN_FILE = os.environ.get(
132
  "NFA_FUN_CN_FILE", "FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors"
133
  ).strip()
134
+ # Director lock: 96GB xlarge + Q8/fp8 painter (not soft; not bf16 host dump).
135
+ GPU_SIZE = (os.environ.get("NFA_FUN_CN_GPU_SIZE") or "xlarge").strip().lower()
136
  if GPU_SIZE not in ("large", "xlarge"):
137
+ GPU_SIZE = "xlarge"
138
  GPU_DURATION = int(os.environ.get("NFA_FUN_CN_GPU_DURATION") or "300")
139
  WEIGHT_DTYPE = torch.bfloat16
140
  MEM_MODE = (
141
+ os.environ.get("NFA_FUN_CN_MEM_MODE") or "model_cpu_offload_and_qfloat8"
 
 
 
 
 
142
  ).strip()
143
+ # Abort CPU load if Fun CN not ready — do not thrash 35+ min again.
144
+ LOAD_DEADLINE_SEC = int(os.environ.get("NFA_FUN_CN_LOAD_DEADLINE_SEC") or "480")
145
 
146
  CONFIG_PATH = APP_DIR / "config" / "flux2_control.yaml"
147
  MODEL_DIR = Path(os.environ.get("NFA_FLUX2_MOUNT") or "/data/FLUX.2-dev")
148
  CN_MOUNT_DIR = Path(os.environ.get("NFA_FUN_CN_MOUNT") or "/data/Fun-CN")
149
+ TE_MOUNT_DIR = Path(os.environ.get("NFA_FLUX2_TE_MOUNT") or "/data/FLUX.2-TE-bnb4")
150
  CACHE_ROOT = Path(
151
  os.environ.get("NFA_FUN_CN_CACHE") or (Path.home() / ".cache" / "nfa_fun_cn")
152
  )
 
155
  _PIPE_OFFLOAD_READY = False
156
  _CN_FILE_PATH: Path | None = None
157
  _GET_IMAGE_LATENT = None
158
+ _FUN_CN_LOADED = False
159
+ _LOAD_T0: float | None = None
160
+
161
+
162
+ def _rss_gb() -> float:
163
+ try:
164
+ import resource
165
+
166
+ # Linux: ru_maxrss is KB
167
+ return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / (1024 * 1024)
168
+ except Exception: # noqa: BLE001
169
+ return -1.0
170
+
171
+
172
+ def _check_load_deadline(stage: str) -> None:
173
+ if _LOAD_T0 is None:
174
+ return
175
+ elapsed = time.time() - _LOAD_T0
176
+ if elapsed > LOAD_DEADLINE_SEC and not _FUN_CN_LOADED:
177
+ raise RuntimeError(
178
+ f"FUN_CN_LOAD_ABORT: stage={stage} elapsed={elapsed:.0f}s "
179
+ f"> deadline={LOAD_DEADLINE_SEC}s (never reached Fun CN loaded). "
180
+ "Refusing to thrash ZeroGPU host RAM like the prior bf16 hang."
181
+ )
182
 
183
 
184
  def _resolve_cn_path() -> Path:
 
222
  "transformer/*",
223
  "vae/*",
224
  "tokenizer/*",
 
225
  "scheduler/*",
226
  ],
227
  )
 
232
  _CN_FILE_PATH = Path(path)
233
 
234
 
235
+ def _ensure_te_weights() -> Path:
236
+ """Local quantized TE only — never remote TE."""
237
+ if TE_MOUNT_DIR.is_dir() and (
238
+ (TE_MOUNT_DIR / "text_encoder").is_dir()
239
+ or (TE_MOUNT_DIR / "config.json").is_file()
240
+ ):
241
+ print(f"[nfa-fun-cn] TE mount={TE_MOUNT_DIR}", flush=True)
242
+ return TE_MOUNT_DIR
243
+ cache_te = CACHE_ROOT / "FLUX.2-TE-bnb4"
244
+ if (cache_te / "text_encoder").is_dir() or (cache_te / "model_index.json").is_file():
245
+ print(f"[nfa-fun-cn] TE cache={cache_te}", flush=True)
246
+ return cache_te
247
+ print(
248
+ f"[nfa-fun-cn] downloading local quantized TE from {TE_MODEL} "
249
+ "(NO remote TE)",
250
+ flush=True,
251
+ )
252
+ CACHE_ROOT.mkdir(parents=True, exist_ok=True)
253
+ snapshot_download(
254
+ repo_id=TE_MODEL,
255
+ local_dir=str(cache_te),
256
+ token=HF_TOKEN or None,
257
+ allow_patterns=[
258
+ "model_index.json",
259
+ "text_encoder/*",
260
+ "tokenizer/*",
261
+ ],
262
+ )
263
+ return cache_te
264
+
265
+
266
  def _prep_depth(depth_image: Image.Image, width: int, height: int) -> Image.Image:
267
  img = depth_image.convert("RGB")
268
  if img.size != (width, height):
 
274
  return (positive or "").strip(), ((negative or "").strip() or " ")
275
 
276
 
277
+ def _fp8_dtype_for_key(key: str) -> torch.dtype:
278
+ for ex in _FP8_EXCLUDE:
279
+ if ex in key:
280
+ return WEIGHT_DTYPE
281
+ return torch.float8_e4m3fn
282
+
283
+
284
+ def _set_tensor(model: torch.nn.Module, key: str, tensor: torch.Tensor) -> None:
285
+ from accelerate.utils import set_module_tensor_to_device
286
+
287
+ target = _fp8_dtype_for_key(key)
288
+ # float8 cast must go through a float dtype first on some builds
289
+ if target == torch.float8_e4m3fn and tensor.dtype not in (
290
+ torch.float8_e4m3fn,
291
+ torch.float8_e5m2,
292
+ ):
293
+ value = tensor.detach().to(dtype=torch.bfloat16).to(dtype=target)
294
+ else:
295
+ value = tensor.detach().to(dtype=target)
296
+ set_module_tensor_to_device(model, key, device="cpu", value=value, dtype=target)
297
+
298
+
299
+ def _stream_shards_into_model(
300
+ model: torch.nn.Module,
301
+ shard_paths: list[str],
302
+ *,
303
+ label: str,
304
+ ) -> None:
305
+ """Load one safetensors shard at a time → fp8; never accumulate full bf16."""
306
+ from safetensors import safe_open
307
+
308
+ model_sd = model.state_dict()
309
+ loaded = 0
310
+ skipped = 0
311
+ for i, path in enumerate(shard_paths):
312
+ _check_load_deadline(f"{label}_shard_{i}")
313
+ print(
314
+ f"[nfa-fun-cn] {label} shard {i+1}/{len(shard_paths)} "
315
+ f"rss_max≈{_rss_gb():.1f}GB path={Path(path).name}",
316
+ flush=True,
317
+ )
318
+ with safe_open(path, framework="pt", device="cpu") as f:
319
+ for key in f.keys():
320
+ if key not in model_sd:
321
+ skipped += 1
322
+ continue
323
+ tensor = f.get_tensor(key)
324
+ if tuple(tensor.shape) != tuple(model_sd[key].shape):
325
+ print(
326
+ f"[nfa-fun-cn] skip size mismatch {key} "
327
+ f"{tuple(tensor.shape)} vs {tuple(model_sd[key].shape)}",
328
+ flush=True,
329
+ )
330
+ skipped += 1
331
+ continue
332
+ _set_tensor(model, key, tensor)
333
+ loaded += 1
334
+ del tensor
335
+ gc.collect()
336
+ print(
337
+ f"[nfa-fun-cn] {label} stream done loaded={loaded} skipped={skipped} "
338
+ f"rss_max≈{_rss_gb():.1f}GB",
339
+ flush=True,
340
+ )
341
+
342
+
343
+ def _init_missing_control_params(model: torch.nn.Module) -> None:
344
+ """Mirror VideoX missing-key init for control blocks (zeros / clones)."""
345
+ from accelerate.utils import set_module_tensor_to_device
346
+
347
+ sd = {k: v for k, v in model.named_parameters()}
348
+ # named_parameters may still be meta; use state_dict meta shapes
349
+ meta_sd = model.state_dict()
350
+ missing = [
351
+ k
352
+ for k, v in meta_sd.items()
353
+ if getattr(v, "is_meta", False) or (hasattr(v, "device") and v.device.type == "meta")
354
+ ]
355
+ if not missing:
356
+ # Also detect uninitialized via device
357
+ missing = []
358
+ for name, param in model.named_parameters():
359
+ if param.device.type == "meta":
360
+ missing.append(name)
361
+ if not missing:
362
+ print("[nfa-fun-cn] no meta params left before Fun CN overlay", flush=True)
363
+ return
364
+
365
+ print(f"[nfa-fun-cn] init {len(missing)} missing/meta params", flush=True)
366
+ with torch.no_grad():
367
+ for key in missing:
368
+ shape = tuple(meta_sd[key].shape)
369
+ dtype = _fp8_dtype_for_key(key)
370
+ # Prefer clone from non-control twin when VideoX does
371
+ twin = key.replace("control_", "")
372
+ if "control" in key and twin in sd and sd[twin].device.type != "meta":
373
+ value = sd[twin].detach().to(dtype=torch.bfloat16).to(dtype=dtype)
374
+ elif "after_proj" in key or "before_proj" in key:
375
+ value = torch.zeros(shape, dtype=dtype)
376
+ elif "bias" in key:
377
+ value = torch.zeros(shape, dtype=dtype)
378
+ else:
379
+ value = torch.zeros(shape, dtype=dtype)
380
+ set_module_tensor_to_device(model, key, device="cpu", value=value, dtype=dtype)
381
+
382
+
383
+ def _overlay_fun_cn(model: torch.nn.Module, cn_path: Path) -> tuple[int, int]:
384
+ """Apply Fun CN Union weights without loading a second full DiT."""
385
+ from safetensors import safe_open
386
+
387
+ model_sd = model.state_dict()
388
+ loaded = 0
389
+ skipped = 0
390
+ print(
391
+ f"[nfa-fun-cn] Fun CN overlay {cn_path.name} "
392
+ f"rss_max≈{_rss_gb():.1f}GB",
393
+ flush=True,
394
+ )
395
+ with safe_open(str(cn_path), framework="pt", device="cpu") as f:
396
+ keys = list(f.keys())
397
+ if keys == ["state_dict"]:
398
+ # rare wrapper — fall back to full load of inner dict only
399
+ from safetensors.torch import load_file
400
+
401
+ wrapped = load_file(str(cn_path))
402
+ inner = wrapped.get("state_dict", wrapped)
403
+ for key, tensor in inner.items():
404
+ if key not in model_sd or tuple(tensor.shape) != tuple(model_sd[key].shape):
405
+ skipped += 1
406
+ continue
407
+ _set_tensor(model, key, tensor)
408
+ loaded += 1
409
+ del wrapped, inner
410
+ gc.collect()
411
+ else:
412
+ for key in keys:
413
+ if key not in model_sd:
414
+ skipped += 1
415
+ continue
416
+ tensor = f.get_tensor(key)
417
+ if tuple(tensor.shape) != tuple(model_sd[key].shape):
418
+ skipped += 1
419
+ continue
420
+ _set_tensor(model, key, tensor)
421
+ loaded += 1
422
+ del tensor
423
+ gc.collect()
424
+ return loaded, skipped
425
+
426
+
427
+ def _load_control_transformer_fp8(model_name: str, cn_file: str):
428
+ """Q8-class painter: empty meta → stream bf16 shards as float8 → Fun CN."""
429
+ global _FUN_CN_LOADED
430
+ import accelerate
431
+ from videox_fun.models.flux2_transformer2d_control import (
432
+ Flux2ControlTransformer2DModel,
433
+ )
434
+
435
+ config_file = os.path.join(model_name, "transformer", "config.json")
436
+ if not os.path.isfile(config_file):
437
+ raise FileNotFoundError(config_file)
438
+ with open(config_file, "r", encoding="utf-8") as fh:
439
+ config = json.load(fh)
440
+ extra = OmegaConf.to_container(OmegaConf.load(str(CONFIG_PATH))[
441
+ "transformer_additional_kwargs"
442
+ ])
443
+
444
+ print(
445
+ "[nfa-fun-cn] CPU-load Flux2Control as float8 stream "
446
+ f"(NOT full bf16) mem={MEM_MODE} gpu_size={GPU_SIZE}",
447
+ flush=True,
448
+ )
449
+ with accelerate.init_empty_weights():
450
+ transformer = Flux2ControlTransformer2DModel.from_config(config, **extra)
451
+
452
+ shard_dir = os.path.join(model_name, "transformer")
453
+ shards = sorted(glob.glob(os.path.join(shard_dir, "*.safetensors")))
454
+ if not shards:
455
+ raise FileNotFoundError(f"No transformer shards under {shard_dir}")
456
+ _stream_shards_into_model(transformer, shards, label="DiT-fp8")
457
+ _check_load_deadline("after_dit_stream")
458
+ _init_missing_control_params(transformer)
459
+ _check_load_deadline("after_control_init")
460
+ loaded, skipped = _overlay_fun_cn(transformer, Path(cn_file))
461
+ _FUN_CN_LOADED = True
462
+ print(
463
+ f"[nfa-fun-cn] Fun CN loaded overlay_ok={loaded} skipped={skipped} "
464
+ f"rss_max≈{_rss_gb():.1f}GB (REAL Fun CN, fp8 painter)",
465
+ flush=True,
466
+ )
467
+ return transformer
468
+
469
+
470
+ def _load_local_quantized_te(te_root: Path):
471
+ """Local quantized Mistral TE — never HF remote TE."""
472
+ from transformers import Mistral3ForConditionalGeneration
473
+
474
+ te_path = te_root / "text_encoder"
475
+ if not te_path.is_dir():
476
+ te_path = te_root
477
+ print(
478
+ f"[nfa-fun-cn] loading LOCAL quantized TE from {te_path} "
479
+ "(bnb-4bit / no remote TE)",
480
+ flush=True,
481
+ )
482
+ _check_load_deadline("te_start")
483
+ text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
484
+ str(te_path),
485
+ torch_dtype=WEIGHT_DTYPE,
486
+ low_cpu_mem_usage=True,
487
+ device_map="cpu",
488
+ )
489
+ print(
490
+ f"[nfa-fun-cn] local TE ready rss_max≈{_rss_gb():.1f}GB",
491
+ flush=True,
492
+ )
493
+ return text_encoder
494
+
495
+
496
  def get_pipe(*, prepare_gpu_offload: bool = False):
497
+ """CPU-load Fun CN pipeline (fp8 DiT stream); arm GPU offload inside @spaces.GPU."""
498
+ global _PIPE, _PIPE_OFFLOAD_READY, _GET_IMAGE_LATENT, _LOAD_T0, _FUN_CN_LOADED
499
  if _PIPE is None:
500
+ _LOAD_T0 = time.time()
501
+ _FUN_CN_LOADED = False
502
  _ensure_weights()
503
+ te_root = _ensure_te_weights()
504
  _ensure_videox_on_path()
505
  from diffusers import FlowMatchEulerDiscreteScheduler
506
+ from transformers import PixtralProcessor
 
 
 
 
507
  from videox_fun.models.flux2_vae import AutoencoderKLFlux2
508
  from videox_fun.pipeline.pipeline_flux2_control import Flux2ControlPipeline
509
  from videox_fun.utils.utils import get_image_latent
 
511
  _GET_IMAGE_LATENT = get_image_latent
512
  model_name = str(MODEL_DIR)
513
  cn_file = str(_resolve_cn_path())
514
+
515
+ transformer = _load_control_transformer_fp8(model_name, cn_file)
516
+ _check_load_deadline("post_fun_cn")
517
+
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
518
  vae = AutoencoderKLFlux2.from_pretrained(model_name, subfolder="vae").to(
519
  WEIGHT_DTYPE
520
  )
521
+ # Tokenizer from base FLUX.2 mount (same vocab as TE).
522
  tokenizer = PixtralProcessor.from_pretrained(model_name, subfolder="tokenizer")
523
+ text_encoder = _load_local_quantized_te(te_root)
 
 
 
 
 
524
  scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
525
  model_name, subfolder="scheduler"
526
  )
 
531
  transformer=transformer,
532
  scheduler=scheduler,
533
  )
534
+ elapsed = time.time() - _LOAD_T0
535
+ print(
536
+ f"[nfa-fun-cn] Flux2ControlPipeline CPU-ready "
537
+ f"(REAL Fun CN + fp8 DiT + local TE) in {elapsed:.0f}s "
538
+ f"rss_max≈{_rss_gb():.1f}GB",
539
+ flush=True,
540
+ )
541
 
542
  if prepare_gpu_offload and not _PIPE_OFFLOAD_READY and torch.cuda.is_available():
543
+ from videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper
 
 
 
544
 
545
  device = "cuda"
546
  transformer = _PIPE.transformer
547
+ # Weights already float8 from stream load — only wrap compute dtype.
548
+ if MEM_MODE in (
549
+ "model_cpu_offload_and_qfloat8",
550
+ "model_full_load_and_qfloat8",
551
+ ):
 
552
  convert_weight_dtype_wrapper(transformer, WEIGHT_DTYPE)
553
+ if MEM_MODE == "sequential_cpu_offload":
 
554
  _PIPE.enable_sequential_cpu_offload(device=device)
555
+ elif MEM_MODE in ("model_cpu_offload", "model_cpu_offload_and_qfloat8"):
556
  _PIPE.enable_model_cpu_offload(device=device)
557
  else:
558
  _PIPE.to(device=device)
559
  _PIPE_OFFLOAD_READY = True
560
+ print(f"[nfa-fun-cn] GPU offload armed mem={MEM_MODE} size={GPU_SIZE}", flush=True)
561
  return _PIPE
562
 
563
 
 
598
  generator = torch.Generator(device=device).manual_seed(int(seed))
599
  print(
600
  f"[nfa-fun-cn] REAL Fun CN generate seed={seed} {w}x{h} steps={steps} "
601
+ f"cn={strength} path=videox_fun_flux2_control painter=fp8 te=local_bnb4 "
602
+ f"size={GPU_SIZE}",
603
  flush=True,
604
  )
605
  with torch.no_grad():
 
639
  _ensure_weights()
640
  # Wall-clock CPU load — does not burn ZeroGPU minutes.
641
  get_pipe(prepare_gpu_offload=False)
642
+ if not _FUN_CN_LOADED:
643
+ raise RuntimeError(
644
+ "FUN_CN_LOAD_ABORT: pipeline built but Fun CN flag false"
645
+ )
646
  return _generate_still_gpu(
647
  positive,
648
  negative or "",
 
666
  gr.Markdown(
667
  "## NFA Track R — **Real Fun depth ControlNet** (ZeroGPU)\n"
668
  f"- Stack: VideoX-Fun `Flux2ControlPipeline` + `{CN_FILE}`\n"
669
+ f"- Painter: **float8 stream** from `{BASE_MODEL}` (no full bf16 host dump)\n"
670
+ f"- TE: **local quantized** `{TE_MODEL}` (NO remote TE)\n"
671
  f"- GPU: `size={GPU_SIZE}` duration={GPU_DURATION}s mem=`{MEM_MODE}`\n"
672
+ f"- Load deadline: {LOAD_DEADLINE_SEC}s to reach `Fun CN loaded`\n"
673
+ "- Soft `image=depth` is **banned** on this Space."
674
  )
675
  with gr.Row():
676
  with gr.Column():
requirements.txt CHANGED
@@ -15,6 +15,7 @@ huggingface_hub
15
  requests
16
  diffusers>=0.36.0
17
  transformers>=4.46.2
 
18
  # VideoX-Fun runtime (code cloned at Space start; pip package misses submodules)
19
  timm
20
  tomesd
 
15
  requests
16
  diffusers>=0.36.0
17
  transformers>=4.46.2
18
+ bitsandbytes>=0.45.0
19
  # VideoX-Fun runtime (code cloned at Space start; pip package misses submodules)
20
  timm
21
  tomesd
vendor/VideoX-Fun/LICENSE ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction,
10
+ and distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by
13
+ the copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all
16
+ other entities that control, are controlled by, or are under common
17
+ control with that entity. For the purposes of this definition,
18
+ "control" means (i) the power, direct or indirect, to cause the
19
+ direction or management of such entity, whether by contract or
20
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
21
+ outstanding shares, or (iii) beneficial ownership of such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity
24
+ exercising permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation
28
+ source, and configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical
31
+ transformation or translation of a Source form, including but
32
+ not limited to compiled object code, generated documentation,
33
+ and conversions to other media types.
34
+
35
+ "Work" shall mean the work of authorship, whether in Source or
36
+ Object form, made available under the License, as indicated by a
37
+ copyright notice that is included in or attached to the work
38
+ (an example is provided in the Appendix below).
39
+
40
+ "Derivative Works" shall mean any work, whether in Source or Object
41
+ form, that is based on (or derived from) the Work and for which the
42
+ editorial revisions, annotations, elaborations, or other modifications
43
+ represent, as a whole, an original work of authorship. For the purposes
44
+ of this License, Derivative Works shall not include works that remain
45
+ separable from, or merely link (or bind by name) to the interfaces of,
46
+ the Work and Derivative Works thereof.
47
+
48
+ "Contribution" shall mean any work of authorship, including
49
+ the original version of the Work and any modifications or additions
50
+ to that Work or Derivative Works thereof, that is intentionally
51
+ submitted to Licensor for inclusion in the Work by the copyright owner
52
+ or by an individual or Legal Entity authorized to submit on behalf of
53
+ the copyright owner. For the purposes of this definition, "submitted"
54
+ means any form of electronic, verbal, or written communication sent
55
+ to the Licensor or its representatives, including but not limited to
56
+ communication on electronic mailing lists, source code control systems,
57
+ and issue tracking systems that are managed by, or on behalf of, the
58
+ Licensor for the purpose of discussing and improving the Work, but
59
+ excluding communication that is conspicuously marked or otherwise
60
+ designated in writing by the copyright owner as "Not a Contribution."
61
+
62
+ "Contributor" shall mean Licensor and any individual or Legal Entity
63
+ on behalf of whom a Contribution has been received by Licensor and
64
+ subsequently incorporated within the Work.
65
+
66
+ 2. Grant of Copyright License. Subject to the terms and conditions of
67
+ this License, each Contributor hereby grants to You a perpetual,
68
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
69
+ copyright license to reproduce, prepare Derivative Works of,
70
+ publicly display, publicly perform, sublicense, and distribute the
71
+ Work and such Derivative Works in Source or Object form.
72
+
73
+ 3. Grant of Patent License. Subject to the terms and conditions of
74
+ this License, each Contributor hereby grants to You a perpetual,
75
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
76
+ (except as stated in this section) patent license to make, have made,
77
+ use, offer to sell, sell, import, and otherwise transfer the Work,
78
+ where such license applies only to those patent claims licensable
79
+ by such Contributor that are necessarily infringed by their
80
+ Contribution(s) alone or by combination of their Contribution(s)
81
+ with the Work to which such Contribution(s) was submitted. If You
82
+ institute patent litigation against any entity (including a
83
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
84
+ or a Contribution incorporated within the Work constitutes direct
85
+ or contributory patent infringement, then any patent licenses
86
+ granted to You under this License for that Work shall terminate
87
+ as of the date such litigation is filed.
88
+
89
+ 4. Redistribution. You may reproduce and distribute copies of the
90
+ Work or Derivative Works thereof in any medium, with or without
91
+ modifications, and in Source or Object form, provided that You
92
+ meet the following conditions:
93
+
94
+ (a) You must give any other recipients of the Work or
95
+ Derivative Works a copy of this License; and
96
+
97
+ (b) You must cause any modified files to carry prominent notices
98
+ stating that You changed the files; and
99
+
100
+ (c) You must retain, in the Source form of any Derivative Works
101
+ that You distribute, all copyright, patent, trademark, and
102
+ attribution notices from the Source form of the Work,
103
+ excluding those notices that do not pertain to any part of
104
+ the Derivative Works; and
105
+
106
+ (d) If the Work includes a "NOTICE" text file as part of its
107
+ distribution, then any Derivative Works that You distribute must
108
+ include a readable copy of the attribution notices contained
109
+ within such NOTICE file, excluding those notices that do not
110
+ pertain to any part of the Derivative Works, in at least one
111
+ of the following places: within a NOTICE text file distributed
112
+ as part of the Derivative Works; within the Source form or
113
+ documentation, if provided along with the Derivative Works; or,
114
+ within a display generated by the Derivative Works, if and
115
+ wherever such third-party notices normally appear. The contents
116
+ of the NOTICE file are for informational purposes only and
117
+ do not modify the License. You may add Your own attribution
118
+ notices within Derivative Works that You distribute, alongside
119
+ or as an addendum to the NOTICE text from the Work, provided
120
+ that such additional attribution notices cannot be construed
121
+ as modifying the License.
122
+
123
+ You may add Your own copyright statement to Your modifications and
124
+ may provide additional or different license terms and conditions
125
+ for use, reproduction, or distribution of Your modifications, or
126
+ for any such Derivative Works as a whole, provided Your use,
127
+ reproduction, and distribution of the Work otherwise complies with
128
+ the conditions stated in this License.
129
+
130
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
131
+ any Contribution intentionally submitted for inclusion in the Work
132
+ by You to the Licensor shall be under the terms and conditions of
133
+ this License, without any additional terms or conditions.
134
+ Notwithstanding the above, nothing herein shall supersede or modify
135
+ the terms of any separate license agreement you may have executed
136
+ with Licensor regarding such Contributions.
137
+
138
+ 6. Trademarks. This License does not grant permission to use the trade
139
+ names, trademarks, service marks, or product names of the Licensor,
140
+ except as required for reasonable and customary use in describing the
141
+ origin of the Work and reproducing the content of the NOTICE file.
142
+
143
+ 7. Disclaimer of Warranty. Unless required by applicable law or
144
+ agreed to in writing, Licensor provides the Work (and each
145
+ Contributor provides its Contributions) on an "AS IS" BASIS,
146
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
147
+ implied, including, without limitation, any warranties or conditions
148
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
149
+ PARTICULAR PURPOSE. You are solely responsible for determining the
150
+ appropriateness of using or redistributing the Work and assume any
151
+ risks associated with Your exercise of permissions under this License.
152
+
153
+ 8. Limitation of Liability. In no event and under no legal theory,
154
+ whether in tort (including negligence), contract, or otherwise,
155
+ unless required by applicable law (such as deliberate and grossly
156
+ negligent acts) or agreed to in writing, shall any Contributor be
157
+ liable to You for damages, including any direct, indirect, special,
158
+ incidental, or consequential damages of any character arising as a
159
+ result of this License or out of the use or inability to use the
160
+ Work (including but not limited to damages for loss of goodwill,
161
+ work stoppage, computer failure or malfunction, or any and all
162
+ other commercial damages or losses), even if such Contributor
163
+ has been advised of the possibility of such damages.
164
+
165
+ 9. Accepting Warranty or Additional Liability. While redistributing
166
+ the Work or Derivative Works thereof, You may choose to offer,
167
+ and charge a fee for, acceptance of support, warranty, indemnity,
168
+ or other liability obligations and/or rights consistent with this
169
+ License. However, in accepting such obligations, You may act only
170
+ on Your own behalf and on Your sole responsibility, not on behalf
171
+ of any other Contributor, and only if You agree to indemnify,
172
+ defend, and hold each Contributor harmless for any liability
173
+ incurred by, or claims asserted against, such Contributor by reason
174
+ of your accepting any such warranty or additional liability.
175
+
176
+ END OF TERMS AND CONDITIONS
177
+
178
+ APPENDIX: How to apply the Apache License to your work.
179
+
180
+ To apply the Apache License to your work, attach the following
181
+ boilerplate notice, with the fields enclosed by brackets "[]"
182
+ replaced with your own identifying information. (Don't include
183
+ the brackets!) The text should be enclosed in the appropriate
184
+ comment syntax for the file format. We also recommend that a
185
+ file or class name and description of purpose be included on the
186
+ same "printed page" as the copyright notice for easier
187
+ identification within third-party archives.
188
+
189
+ Copyright [yyyy] [name of copyright owner]
190
+
191
+ Licensed under the Apache License, Version 2.0 (the "License");
192
+ you may not use this file except in compliance with the License.
193
+ You may obtain a copy of the License at
194
+
195
+ http://www.apache.org/licenses/LICENSE-2.0
196
+
197
+ Unless required by applicable law or agreed to in writing, software
198
+ distributed under the License is distributed on an "AS IS" BASIS,
199
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
200
+ See the License for the specific language governing permissions and
201
+ limitations under the License.
vendor/VideoX-Fun/README.md ADDED
@@ -0,0 +1,733 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # VideoX-Fun
2
+
3
+ 😊 Welcome!
4
+
5
+ CogVideoX-Fun:
6
+ [![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/CogVideoX-Fun-5b)
7
+
8
+ Wan-Fun:
9
+ [![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-yellow)](https://huggingface.co/spaces/alibaba-pai/Wan2.1-Fun-1.3B-InP)
10
+
11
+ English | [简体中文](./README_zh-CN.md) | [日本語](./README_ja-JP.md)
12
+
13
+ # Table of Contents
14
+ - [Introduction](#introduction)
15
+ - [Quick Start](#quick-start)
16
+ - [Video Result](#video-result)
17
+ - [How to Use](#how-to-use)
18
+ - [Model zoo](#model-zoo)
19
+ - [Reference](#reference)
20
+ - [Citation](#citation)
21
+ - [Limitations and Risks](#limitations-and-risks)
22
+ - [License](#license)
23
+
24
+ # Introduction
25
+ VideoX-Fun is a video generation pipeline that can be used to generate AI images and videos, as well as to train baseline and Lora models for Diffusion Transformer. We support direct prediction from pre-trained baseline models to generate videos with different resolutions, durations, and FPS. Additionally, we also support users in training their own baseline and Lora models to perform specific style transformations.
26
+
27
+ We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start).
28
+
29
+ What's New:
30
+ - Added support for Wan 2.2 series models, Wan-VACE control model, Fantasy Talking digital human model, Qwen-Image, Flux image generation models, and more. [2025.10.16]
31
+ - Update Wan2.1-Fun-V1.1: Support for 14B and 1.3B model Control + Reference Image models, support for camera control, and the Inpaint model has been retrained for improved performance. [2025.04.25]
32
+ - Update Wan2.1-Fun-V1.0: Support I2V and Control models for 14B and 1.3B models, with support for start and end frame prediction. [2025.03.26]
33
+ - Update CogVideoX-Fun-V1.5: Upload I2V model and related training/prediction code. [2024.12.16]
34
+ - Reward Lora Support: Train Lora using reward backpropagation techniques to optimize generated videos, making them better aligned with human preferences. [More Information](scripts/README_TRAIN_REWARD.md). New version of the control model supports various control conditions such as Canny, Depth, Pose, MLSD, etc. [2024.11.21]
35
+ - Diffusers Support: CogVideoX-Fun Control is now supported in diffusers. Thanks to [a-r-r-o-w](https://github.com/a-r-r-o-w) for contributing support in this [PR](https://github.com/huggingface/diffusers/pull/9671). Check out the [documentation](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox) for more details. [2024.10.16]
36
+ - Update CogVideoX-Fun-V1.1: Retrain i2v model, add Noise to increase the motion amplitude of the video. Upload control model training code and Control model. [2024.09.29]
37
+ - Update CogVideoX-Fun-V1.0: Initial code release! Now supports Windows and Linux. Supports video generation at arbitrary resolutions from 256x256x49 to 1024x1024x49 for 2B and 5B models. [2024.09.18]
38
+
39
+ Function:
40
+ - [Data Preprocessing](#data-preprocess)
41
+ - [Train DiT](#dit-train)
42
+ - [Video Generation](#video-gen)
43
+
44
+ Our UI interface is as follows:
45
+ ![ui](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/ui.jpg)
46
+
47
+ # Quick Start
48
+ ### 1. Cloud usage: AliyunDSW/Docker
49
+ #### a. From AliyunDSW
50
+ DSW has free GPU time, which can be applied once by a user and is valid for 3 months after applying.
51
+
52
+ Aliyun provide free GPU time in [Freetier](https://free.aliyun.com/?product=9602825&crowd=enterprise&spm=5176.28055625.J_5831864660.1.e939154aRgha4e&scm=20140722.M_9974135.P_110.MO_1806-ID_9974135-MID_9974135-CID_30683-ST_8512-V_1), get it and use in Aliyun PAI-DSW to start CogVideoX-Fun within 5min!
53
+
54
+ [![DSW Notebook](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/asset/dsw.png)](https://gallery.pai-ml.com/#/preview/deepLearning/cv/cogvideox_fun)
55
+
56
+ #### b. From ComfyUI
57
+ Our ComfyUI is as follows, please refer to [ComfyUI README](comfyui/README.md) for details.
58
+ ![workflow graph](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/cogvideoxfunv1_workflow_i2v.jpg)
59
+
60
+ #### c. From docker
61
+ If you are using docker, please make sure that the graphics card driver and CUDA environment have been installed correctly in your machine.
62
+
63
+ Then execute the following commands in this way:
64
+
65
+ ```
66
+ # pull image
67
+ docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
68
+
69
+ # enter image
70
+ docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
71
+
72
+ # clone code
73
+ git clone https://github.com/aigc-apps/VideoX-Fun.git
74
+
75
+ # enter VideoX-Fun's dir
76
+ cd VideoX-Fun
77
+
78
+ # download weights
79
+ mkdir models/Diffusion_Transformer
80
+ mkdir models/Personalized_Model
81
+
82
+ # Please use the hugginface link or modelscope link to download the model.
83
+ # CogVideoX-Fun
84
+ # https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-InP
85
+ # https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-InP
86
+
87
+ # Wan
88
+ # https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-InP
89
+ # https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-InP
90
+ ```
91
+
92
+ ### 2. Local install: Environment Check/Downloading/Installation
93
+ #### a. Environment Check
94
+ We have verified this repo execution on the following environment:
95
+
96
+ The detailed of Windows:
97
+ - OS: Windows 10
98
+ - python: python3.10 & python3.11
99
+ - pytorch: torch2.2.0
100
+ - CUDA: 11.8 & 12.1
101
+ - CUDNN: 8+
102
+ - GPU: Nvidia-3060 12G & Nvidia-3090 24G
103
+
104
+ The detailed of Linux:
105
+ - OS: Ubuntu 20.04, CentOS
106
+ - python: python3.10 & python3.11
107
+ - pytorch: torch2.2.0
108
+ - CUDA: 11.8 & 12.1
109
+ - CUDNN: 8+
110
+ - GPU:Nvidia-V100 16G & Nvidia-A10 24G & Nvidia-A100 40G & Nvidia-A100 80G
111
+
112
+ We need about 60GB available on disk (for saving weights), please check!
113
+
114
+ #### b. Weights
115
+ We'd better place the [weights](#model-zoo) along the specified path:
116
+
117
+ **Via ComfyUI**:
118
+ Put the models into the ComfyUI weights folder `ComfyUI/models/Fun_Models/`:
119
+ ```
120
+ 📦 ComfyUI/
121
+ ├── 📂 models/
122
+ │ └── 📂 Fun_Models/
123
+ │ ├── 📂 CogVideoX-Fun-V1.1-2b-InP/
124
+ │ ├── 📂 CogVideoX-Fun-V1.1-5b-InP/
125
+ │ ├── 📂 Wan2.1-Fun-14B-InP
126
+ │ └── 📂 Wan2.1-Fun-1.3B-InP/
127
+ ```
128
+
129
+ **Run its own python file or UI interface**:
130
+ ```
131
+ 📦 models/
132
+ ├── 📂 Diffusion_Transformer/
133
+ │ ├── 📂 CogVideoX-Fun-V1.1-2b-InP/
134
+ │ ├── 📂 CogVideoX-Fun-V1.1-5b-InP/
135
+ │ ├── 📂 Wan2.1-Fun-14B-InP
136
+ │ └── 📂 Wan2.1-Fun-1.3B-InP/
137
+ ├── 📂 Personalized_Model/
138
+ │ └── your trained trainformer model / your trained lora model (for UI load)
139
+ ```
140
+
141
+ # Video Result
142
+
143
+ ### Wan2.1-Fun-V1.1-14B-InP && Wan2.1-Fun-V1.1-1.3B-InP
144
+
145
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
146
+ <tr>
147
+ <td>
148
+ <video src="https://github.com/user-attachments/assets/d6a46051-8fe6-4174-be12-95ee52c96298" width="100%" controls preload loop></video>
149
+ </td>
150
+ <td>
151
+ <video src="https://github.com/user-attachments/assets/8572c656-8548-4b1f-9ec8-8107c6236cb1" width="100%" controls preload loop></video>
152
+ </td>
153
+ <td>
154
+ <video src="https://github.com/user-attachments/assets/d3411c95-483d-4e30-bc72-483c2b288918" width="100%" controls preload loop></video>
155
+ </td>
156
+ <td>
157
+ <video src="https://github.com/user-attachments/assets/b2f5addc-06bd-49d9-b925-973090a32800" width="100%" controls preload loop></video>
158
+ </td>
159
+ </tr>
160
+ </table>
161
+
162
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
163
+ <tr>
164
+ <td>
165
+ <video src="https://github.com/user-attachments/assets/747b6ab8-9617-4ba2-84a0-b51c0efbd4f8" width="100%" controls preload loop></video>
166
+ </td>
167
+ <td>
168
+ <video src="https://github.com/user-attachments/assets/ae94dcda-9d5e-4bae-a86f-882c4282a367" width="100%" controls preload loop></video>
169
+ </td>
170
+ <td>
171
+ <video src="https://github.com/user-attachments/assets/a4aa1a82-e162-4ab5-8f05-72f79568a191" width="100%" controls preload loop></video>
172
+ </td>
173
+ <td>
174
+ <video src="https://github.com/user-attachments/assets/83c005b8-ccbc-44a0-a845-c0472763119c" width="100%" controls preload loop></video>
175
+ </td>
176
+ </tr>
177
+ </table>
178
+
179
+ ### Wan2.1-Fun-V1.1-14B-Control && Wan2.1-Fun-V1.1-1.3B-Control
180
+
181
+ Generic Control Video + Reference Image:
182
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
183
+ <tr>
184
+ <td>
185
+ Reference Image
186
+ </td>
187
+ <td>
188
+ Control Video
189
+ </td>
190
+ <td>
191
+ Wan2.1-Fun-V1.1-14B-Control
192
+ </td>
193
+ <td>
194
+ Wan2.1-Fun-V1.1-1.3B-Control
195
+ </td>
196
+ <tr>
197
+ <td>
198
+ <image src="https://github.com/user-attachments/assets/221f2879-3b1b-4fbd-84f9-c3e0b0b3533e" width="100%" controls preload loop></image>
199
+ </td>
200
+ <td>
201
+ <video src="https://github.com/user-attachments/assets/f361af34-b3b3-4be4-9d03-cd478cb3dfc5" width="100%" controls preload loop></video>
202
+ </td>
203
+ <td>
204
+ <video src="https://github.com/user-attachments/assets/85e2f00b-6ef0-4922-90ab-4364afb2c93d" width="100%" controls preload loop></video>
205
+ </td>
206
+ <td>
207
+ <video src="https://github.com/user-attachments/assets/1f3fe763-2754-4215-bc9a-ae804950d4b3" width="100%" controls preload loop></video>
208
+ </td>
209
+ <tr>
210
+ </table>
211
+
212
+
213
+ Generic Control Video (Canny, Pose, Depth, etc.) and Trajectory Control:
214
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
215
+ <tr>
216
+ <td>
217
+ <video src="https://github.com/user-attachments/assets/f35602c4-9f0a-4105-9762-1e3a88abbac6" width="100%" controls preload loop></video>
218
+ </td>
219
+ <td>
220
+ <video src="https://github.com/user-attachments/assets/8b0f0e87-f1be-4915-bb35-2d53c852333e" width="100%" controls preload loop></video>
221
+ </td>
222
+ <td>
223
+ <video src="https://github.com/user-attachments/assets/972012c1-772b-427a-bce6-ba8b39edcfad" width="100%" controls preload loop></video>
224
+ </td>
225
+ <tr>
226
+ </table>
227
+
228
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
229
+ <tr>
230
+ <td>
231
+ <video src="https://github.com/user-attachments/assets/ce62d0bd-82c0-4d7b-9c49-7e0e4b605745" width="100%" controls preload loop></video>
232
+ </td>
233
+ <td>
234
+ <video src="https://github.com/user-attachments/assets/89dfbffb-c4a6-4821-bcef-8b1489a3ca00" width="100%" controls preload loop></video>
235
+ </td>
236
+ <td>
237
+ <video src="https://github.com/user-attachments/assets/72a43e33-854f-4349-861b-c959510d1a84" width="100%" controls preload loop></video>
238
+ </td>
239
+ <tr>
240
+ <td>
241
+ <video src="https://github.com/user-attachments/assets/bb0ce13d-dee0-4049-9eec-c92f3ebc1358" width="100%" controls preload loop></video>
242
+ </td>
243
+ <td>
244
+ <video src="https://github.com/user-attachments/assets/7840c333-7bec-4582-ba63-20a39e1139c4" width="100%" controls preload loop></video>
245
+ </td>
246
+ <td>
247
+ <video src="https://github.com/user-attachments/assets/85147d30-ae09-4f36-a077-2167f7a578c0" width="100%" controls preload loop></video>
248
+ </td>
249
+ </tr>
250
+ </table>
251
+
252
+ ### Wan2.1-Fun-V1.1-14B-Control-Camera && Wan2.1-Fun-V1.1-1.3B-Control-Camera
253
+
254
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
255
+ <tr>
256
+ <td>
257
+ Pan Up
258
+ </td>
259
+ <td>
260
+ Pan Left
261
+ </td>
262
+ <td>
263
+ Pan Right
264
+ </td>
265
+ <tr>
266
+ <td>
267
+ <video src="https://github.com/user-attachments/assets/869fe2ef-502a-484e-8656-fe9e626b9f63" width="100%" controls preload loop></video>
268
+ </td>
269
+ <td>
270
+ <video src="https://github.com/user-attachments/assets/2d4185c8-d6ec-4831-83b4-b1dbfc3616fa" width="100%" controls preload loop></video>
271
+ </td>
272
+ <td>
273
+ <video src="https://github.com/user-attachments/assets/7dfb7cad-ed24-4acc-9377-832445a07ec7" width="100%" controls preload loop></video>
274
+ </td>
275
+ <tr>
276
+ <td>
277
+ Pan Down
278
+ </td>
279
+ <td>
280
+ Pan Up + Pan Left
281
+ </td>
282
+ <td>
283
+ Pan Up + Pan Right
284
+ </td>
285
+ <tr>
286
+ <td>
287
+ <video src="https://github.com/user-attachments/assets/3ea3a08d-f2df-43a2-976e-bf2659345373" width="100%" controls preload loop></video>
288
+ </td>
289
+ <td>
290
+ <video src="https://github.com/user-attachments/assets/4a85b028-4120-4293-886b-b8afe2d01713" width="100%" controls preload loop></video>
291
+ </td>
292
+ <td>
293
+ <video src="https://github.com/user-attachments/assets/ad0d58c1-13ef-450c-b658-4fed7ff5ed36" width="100%" controls preload loop></video>
294
+ </td>
295
+ </tr>
296
+ </table>
297
+
298
+ ### CogVideoX-Fun-V1.1-5B
299
+
300
+ Resolution-1024
301
+
302
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
303
+ <tr>
304
+ <td>
305
+ <video src="https://github.com/user-attachments/assets/34e7ec8f-293e-4655-bb14-5e1ee476f788" width="100%" controls preload loop></video>
306
+ </td>
307
+ <td>
308
+ <video src="https://github.com/user-attachments/assets/7809c64f-eb8c-48a9-8bdc-ca9261fd5434" width="100%" controls preload loop></video>
309
+ </td>
310
+ <td>
311
+ <video src="https://github.com/user-attachments/assets/8e76aaa4-c602-44ac-bcb4-8b24b72c386c" width="100%" controls preload loop></video>
312
+ </td>
313
+ <td>
314
+ <video src="https://github.com/user-attachments/assets/19dba894-7c35-4f25-b15c-384167ab3b03" width="100%" controls preload loop></video>
315
+ </td>
316
+ </tr>
317
+ </table>
318
+
319
+
320
+ Resolution-768
321
+
322
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
323
+ <tr>
324
+ <td>
325
+ <video src="https://github.com/user-attachments/assets/0bc339b9-455b-44fd-8917-80272d702737" width="100%" controls preload loop></video>
326
+ </td>
327
+ <td>
328
+ <video src="https://github.com/user-attachments/assets/70a043b9-6721-4bd9-be47-78b7ec5c27e9" width="100%" controls preload loop></video>
329
+ </td>
330
+ <td>
331
+ <video src="https://github.com/user-attachments/assets/d5dd6c09-14f3-40f8-8b6d-91e26519b8ac" width="100%" controls preload loop></video>
332
+ </td>
333
+ <td>
334
+ <video src="https://github.com/user-attachments/assets/9327e8bc-4f17-46b0-b50d-38c250a9483a" width="100%" controls preload loop></video>
335
+ </td>
336
+ </tr>
337
+ </table>
338
+
339
+ Resolution-512
340
+
341
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
342
+ <tr>
343
+ <td>
344
+ <video src="https://github.com/user-attachments/assets/ef407030-8062-454d-aba3-131c21e6b58c" width="100%" controls preload loop></video>
345
+ </td>
346
+ <td>
347
+ <video src="https://github.com/user-attachments/assets/7610f49e-38b6-4214-aa48-723ae4d1b07e" width="100%" controls preload loop></video>
348
+ </td>
349
+ <td>
350
+ <video src="https://github.com/user-attachments/assets/1fff0567-1e15-415c-941e-53ee8ae2c841" width="100%" controls preload loop></video>
351
+ </td>
352
+ <td>
353
+ <video src="https://github.com/user-attachments/assets/bcec48da-b91b-43a0-9d50-cf026e00fa4f" width="100%" controls preload loop></video>
354
+ </td>
355
+ </tr>
356
+ </table>
357
+
358
+ ### CogVideoX-Fun-V1.1-5B-Control
359
+
360
+ <table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
361
+ <tr>
362
+ <td>
363
+ <video src="https://github.com/user-attachments/assets/53002ce2-dd18-4d4f-8135-b6f68364cabd" width="100%" controls preload loop></video>
364
+ </td>
365
+ <td>
366
+ <video src="https://github.com/user-attachments/assets/a1a07cf8-d86d-4cd2-831f-18a6c1ceee1d" width="100%" controls preload loop></video>
367
+ </td>
368
+ <td>
369
+ <video src="https://github.com/user-attachments/assets/3224804f-342d-4947-918d-d9fec8e3d273" width="100%" controls preload loop></video>
370
+ </td>
371
+ <tr>
372
+ <td>
373
+ A young woman with beautiful clear eyes and blonde hair, wearing white clothes and twisting her body, with the camera focused on her face. High quality, masterpiece, best quality, high resolution, ultra-fine, dreamlike.
374
+ </td>
375
+ <td>
376
+ A young woman with beautiful clear eyes and blonde hair, wearing white clothes and twisting her body, with the camera focused on her face. High quality, masterpiece, best quality, high resolution, ultra-fine, dreamlike.
377
+ </td>
378
+ <td>
379
+ A young bear.
380
+ </td>
381
+ </tr>
382
+ <tr>
383
+ <td>
384
+ <video src="https://github.com/user-attachments/assets/ea908454-684b-4d60-b562-3db229a250a9" width="100%" controls preload loop></video>
385
+ </td>
386
+ <td>
387
+ <video src="https://github.com/user-attachments/assets/ffb7c6fc-8b69-453b-8aad-70dfae3899b9" width="100%" controls preload loop></video>
388
+ </td>
389
+ <td>
390
+ <video src="https://github.com/user-attachments/assets/d3f757a3-3551-4dcb-9372-7a61469813f5" width="100%" controls preload loop></video>
391
+ </td>
392
+ </tr>
393
+ </table>
394
+
395
+ # How to Use
396
+
397
+ <h3 id="video-gen">1. Generation</h3>
398
+
399
+ #### a. GPU Memory Optimization
400
+ Since Wan2.1 has a very large number of parameters, we need to consider memory optimization strategies to adapt to consumer-grade GPUs. We provide `GPU_memory_mode` for each prediction file, allowing you to choose between `model_cpu_offload`, `model_cpu_offload_and_qfloat8`, and `sequential_cpu_offload`. This solution is also applicable to CogVideoX-Fun generation.
401
+
402
+ - `model_cpu_offload`: The entire model is moved to the CPU after use, saving some GPU memory.
403
+ - `model_cpu_offload_and_qfloat8`: The entire model is moved to the CPU after use, and the transformer model is quantized to float8, saving more GPU memory.
404
+ - `sequential_cpu_offload`: Each layer of the model is moved to the CPU after use. It is slower but saves a significant amount of GPU memory.
405
+
406
+ `qfloat8` may slightly reduce model performance but saves more GPU memory. If you have sufficient GPU memory, it is recommended to use `model_cpu_offload`.
407
+
408
+ #### b. Using ComfyUI
409
+ For details, refer to [ComfyUI README](comfyui/README.md).
410
+
411
+ #### c. Running Python Files
412
+
413
+ ##### i. Single-GPU Inference:
414
+
415
+ - **Step 1**: Download the corresponding [weights](#model-zoo) and place them in the `models` folder.
416
+ - **Step 2**: Use different files for prediction based on the weights and prediction goals. This library currently supports CogVideoX-Fun, Wan2.1, and Wan2.1-Fun. Different models are distinguished by folder names under the `examples` folder, and their supported features vary. Use them accordingly. Below is an example using CogVideoX-Fun:
417
+ - **Text-to-Video**:
418
+ - Modify `prompt`, `neg_prompt`, `guidance_scale`, and `seed` in the file `examples/cogvideox_fun/predict_t2v.py`.
419
+ - Run the file `examples/cogvideox_fun/predict_t2v.py` and wait for the results. The generated videos will be saved in the folder `samples/cogvideox-fun-videos`.
420
+ - **Image-to-Video**:
421
+ - Modify `validation_image_start`, `validation_image_end`, `prompt`, `neg_prompt`, `guidance_scale`, and `seed` in the file `examples/cogvideox_fun/predict_i2v.py`.
422
+ - `validation_image_start` is the starting image of the video, and `validation_image_end` is the ending image of the video.
423
+ - Run the file `examples/cogvideox_fun/predict_i2v.py` and wait for the results. The generated videos will be saved in the folder `samples/cogvideox-fun-videos_i2v`.
424
+ - **Video-to-Video**:
425
+ - Modify `validation_video`, `validation_image_end`, `prompt`, `neg_prompt`, `guidance_scale`, and `seed` in the file `examples/cogvideox_fun/predict_v2v.py`.
426
+ - `validation_video` is the reference video for video-to-video generation. You can use the following demo video: [Demo Video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1/play_guitar.mp4).
427
+ - Run the file `examples/cogvideox_fun/predict_v2v.py` and wait for the results. The generated videos will be saved in the folder `samples/cogvideox-fun-videos_v2v`.
428
+ - **Controlled Video Generation (Canny, Pose, Depth, etc.)**:
429
+ - Modify `control_video`, `validation_image_end`, `prompt`, `neg_prompt`, `guidance_scale`, and `seed` in the file `examples/cogvideox_fun/predict_v2v_control.py`.
430
+ - `control_video` is the control video extracted using operators such as Canny, Pose, or Depth. You can use the following demo video: [Demo Video](https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/cogvideox_fun/asset/v1.1/pose.mp4).
431
+ - Run the file `examples/cogvideox_fun/predict_v2v_control.py` and wait for the results. The generated videos will be saved in the folder `samples/cogvideox-fun-videos_v2v_control`.
432
+ - **Step 3**: If you want to integrate other backbones or Loras trained by yourself, modify `lora_path` and relevant paths in `examples/{model_name}/predict_t2v.py` or `examples/{model_name}/predict_i2v.py` as needed.
433
+
434
+ ##### ii. Multi-GPU Inference:
435
+ When using multi-GPU inference, please make sure to install the xfuser. We recommend installing xfuser==0.4.2 and yunchang==0.6.2.
436
+ ```
437
+ pip install xfuser==0.4.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
438
+ pip install yunchang==0.6.2 --progress-bar off -i https://mirrors.aliyun.com/pypi/simple/
439
+ ```
440
+
441
+ Please ensure that the product of `ulysses_degree` and `ring_degree` equals the number of GPUs being used. For example, if you are using 8 GPUs, you can set `ulysses_degree=2` and `ring_degree=4`, or alternatively `ulysses_degree=4` and `ring_degree=2`.
442
+
443
+ - `ulysses_degree` performs parallelization after splitting across the heads.
444
+ - `ring_degree` performs parallelization after splitting across the sequence.
445
+
446
+ Compared to `ulysses_degree`, `ring_degree` incurs higher communication costs. Therefore, when setting these parameters, you should take into account both the sequence length and the number of heads in the model.
447
+
448
+ Let’s take 8-GPU parallel inference as an example:
449
+
450
+ - **For Wan2.1-Fun-V1.1-14B-InP**, which has 40 heads, `ulysses_degree` should be set to a divisor of 40 (e.g., 2, 4, 8, etc.). Thus, when using 8 GPUs for parallel inference, you can set `ulysses_degree=8` and `ring_degree=1`.
451
+
452
+ - **For Wan2.1-Fun-V1.1-1.3B-InP**, which has 12 heads, `ulysses_degree` should be set to a divisor of 12 (e.g., 2, 4, etc.). Thus, when using 8 GPUs for parallel inference, you can set `ulysses_degree=4` and `ring_degree=2`.
453
+
454
+ After setting the parameters, run the following command for parallel inference:
455
+
456
+ ```sh
457
+ torchrun --nproc-per-node=8 examples/wan2.1_fun/predict_t2v.py
458
+ ```
459
+
460
+ #### d. Using the Web UI
461
+ The web UI supports text-to-video, image-to-video, video-to-video, and controlled video generation (Canny, Pose, Depth, etc.). This library currently supports CogVideoX-Fun, Wan2.1, and Wan2.1-Fun. Different models are distinguished by folder names under the `examples` folder, and their supported features vary. Use them accordingly. Below is an example using CogVideoX-Fun:
462
+
463
+ - **Step 1**: Download the corresponding [weights](#model-zoo) and place them in the `models` folder.
464
+ - **Step 2**: Run the file `examples/cogvideox_fun/app.py` to access the Gradio interface.
465
+ - **Step 3**: Select the generation model on the page, fill in `prompt`, `neg_prompt`, `guidance_scale`, and `seed`, click "Generate," and wait for the results. The generated videos will be saved in the `sample` folder.
466
+
467
+ ### 2. Model Training
468
+ A complete model training pipeline should include data preprocessing and Video DiT training. The training process for different models is similar, and the data formats are also similar:
469
+
470
+ <h4 id="data-preprocess">a. data preprocessing</h4>
471
+
472
+ We have provided a simple demo of training the Lora model through image data, which can be found in the [wiki](https://github.com/aigc-apps/CogVideoX-Fun/wiki/Training-Lora) for details.
473
+
474
+ A complete data preprocessing link for long video segmentation, cleaning, and description can refer to [README](cogvideox/video_caption/README.md) in the video captions section.
475
+
476
+ If you want to train a text to image and video generation model. You need to arrange the dataset in this format.
477
+
478
+ ```
479
+ 📦 project/
480
+ ├── 📂 datasets/
481
+ │ ├── 📂 internal_datasets/
482
+ │ ├── 📂 train/
483
+ │ │ ├── 📄 00000001.mp4
484
+ │ │ ├── 📄 00000002.jpg
485
+ │ │ └── 📄 .....
486
+ │ └── 📄 json_of_internal_datasets.json
487
+ ```
488
+
489
+ The json_of_internal_datasets.json is a standard JSON file. The file_path in the json can to be set as relative path, as shown in below:
490
+ ```json
491
+ [
492
+ {
493
+ "file_path": "train/00000001.mp4",
494
+ "text": "A group of young men in suits and sunglasses are walking down a city street.",
495
+ "type": "video"
496
+ },
497
+ {
498
+ "file_path": "train/00000002.jpg",
499
+ "text": "A group of young men in suits and sunglasses are walking down a city street.",
500
+ "type": "image"
501
+ },
502
+ .....
503
+ ]
504
+ ```
505
+
506
+ You can also set the path as absolute path as follow:
507
+ ```json
508
+ [
509
+ {
510
+ "file_path": "/mnt/data/videos/00000001.mp4",
511
+ "text": "A group of young men in suits and sunglasses are walking down a city street.",
512
+ "type": "video"
513
+ },
514
+ {
515
+ "file_path": "/mnt/data/train/00000001.jpg",
516
+ "text": "A group of young men in suits and sunglasses are walking down a city street.",
517
+ "type": "image"
518
+ },
519
+ .....
520
+ ]
521
+ ```
522
+
523
+ <h4 id="dit-train">b. Video DiT training </h4>
524
+
525
+ If the data format is relative path during data preprocessing, please set ```scripts/{model_name}/train.sh``` as follow.
526
+ ```
527
+ export DATASET_NAME="datasets/internal_datasets/"
528
+ export DATASET_META_NAME="datasets/internal_datasets/json_of_internal_datasets.json"
529
+ ```
530
+
531
+ If the data format is absolute path during data preprocessing, please set ```scripts/train.sh``` as follow.
532
+ ```
533
+ export DATASET_NAME=""
534
+ export DATASET_META_NAME="/mnt/data/json_of_internal_datasets.json"
535
+ ```
536
+
537
+ Then, we run scripts/train.sh.
538
+ ```sh
539
+ sh scripts/train.sh
540
+ ```
541
+
542
+ For details on some parameter settings:
543
+ Wan2.1-Fun can be found in [Readme Train](scripts/wan2.1_fun/README_TRAIN.md) and [Readme Lora](scripts/wan2.1_fun/README_TRAIN_LORA.md).
544
+ Wan2.1 can be found in [Readme Train](scripts/wan2.1/README_TRAIN.md) and [Readme Lora](scripts/wan2.1/README_TRAIN_LORA.md).
545
+ CogVideoX-Fun can be found in [Readme Train](scripts/cogvideox_fun/README_TRAIN.md) and [Readme Lora](scripts/cogvideox_fun/README_TRAIN_LORA.md).
546
+
547
+
548
+ # Model zoo
549
+ ## 1. Wan2.2-Fun
550
+
551
+ | Name | Storage Size | Hugging Face | Model Scope | Description |
552
+ |--|--|--|--|--|
553
+ | Wan2.2-Fun-A14B-InP | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-InP) | Wan2.2-Fun-14B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. |
554
+ | Wan2.2-Fun-A14B-Control | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control)| Wan2.2-Fun-14B video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc., and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction at 81 frames, trained at 16 frames per second, with multilingual prediction support. |
555
+ | Wan2.2-Fun-A14B-Control-Camera | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-A14B-Control-Camera)| Wan2.2-Fun-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
556
+ | Wan2.2-VACE-Fun-A14B | 64.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-VACE-Fun-A14B) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-VACE-Fun-A14B) | Control weights for Wan2.2 trained using the VACE scheme (based on the base model Wan2.2-T2V-A14B), supporting various control conditions such as Canny, Depth, Pose, MLSD, trajectory control, etc. It supports video generation by specifying the subject. It supports multi-resolution (512, 768, 1024) video prediction, and is trained with 81 frames at 16 FPS. It also supports multi-language prediction. |
557
+ | Wan2.2-Fun-5B-InP | 23.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-5B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-5B-InP) | Wan2.2-Fun-5B text-to-video weights trained at 121 frames, 24 FPS, supporting first/last frame prediction. |
558
+ | Wan2.2-Fun-5B-Control | 23.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-5B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-5B-Control)| Wan2.2-Fun-5B video control weights, supporting control conditions like Canny, Depth, Pose, MLSD, and trajectory control. Trained at 121 frames, 24 FPS, with multilingual prediction support. |
559
+ | Wan2.2-Fun-5B-Control-Camera | 23.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-5B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.2-Fun-5B-Control-Camera)| Wan2.2-Fun-5B camera lens control weights. Trained at 121 frames, 24 FPS, with multilingual prediction support. |
560
+
561
+
562
+ ## 2. Wan2.2
563
+
564
+ | Name | Hugging Face | Model Scope | Description |
565
+ |--|--|--|--|
566
+ | Wan2.2-TI2V-5B | [🤗Link](https://huggingface.co/Wan-AI/Wan2.2-TI2V-5B) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.2-TI2V-5B) | Wan2.2-5B Text-to-Video Weights |
567
+ | Wan2.2-T2V-14B | [🤗Link](https://huggingface.co/Wan-AI/Wan2.2-T2V-A14B) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.2-T2V-A14B) | Wan2.2-14B Text-to-Video Weights |
568
+ | Wan2.2-I2V-A14B | [🤗Link](https://huggingface.co/Wan-AI/Wan2.2-I2V-A14B) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.2-I2V-A14B) | Wan2.2-I2V-A14B Image-to-Video Weights |
569
+
570
+ ## 3. Wan2.1-Fun
571
+
572
+ V1.1:
573
+ | Name | Storage Size | Hugging Face | Model Scope | Description |
574
+ |------|--------------|--------------|-------------|-------------|
575
+ | Wan2.1-Fun-V1.1-1.3B-InP | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-InP) | Wan2.1-Fun-V1.1-1.3B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. |
576
+ | Wan2.1-Fun-V1.1-14B-InP | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-InP) | Wan2.1-Fun-V1.1-14B text-to-video generation weights, trained at multiple resolutions, supports start-end image prediction. |
577
+ | Wan2.1-Fun-V1.1-1.3B-Control | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control) | Wan2.1-Fun-V1.1-1.3B video control weights support various control conditions such as Canny, Depth, Pose, MLSD, etc., supports reference image + control condition-based control, and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
578
+ | Wan2.1-Fun-V1.1-14B-Control | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control) | Wan2.1-Fun-V1.1-14B video control weights support various control conditions such as Canny, Depth, Pose, MLSD, etc., supports reference image + control condition-based control, and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
579
+ | Wan2.1-Fun-V1.1-1.3B-Control-Camera | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-1.3B-Control-Camera) | Wan2.1-Fun-V1.1-1.3B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
580
+ | Wan2.1-Fun-V1.1-14B-Control-Camera | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-V1.1-14B-Control-Camera) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-V1.1-14B-Control-Camera) | Wan2.1-Fun-V1.1-14B camera lens control weights. Supports multi-resolution (512, 768, 1024) video prediction, trained with 81 frames at 16 FPS, supports multilingual prediction. |
581
+
582
+ V1.0:
583
+ | Name | Storage Space | Hugging Face | Model Scope | Description |
584
+ |--|--|--|--|--|
585
+ | Wan2.1-Fun-1.3B-InP | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-1.3B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-1.3B-InP) | Wan2.1-Fun-1.3B text-to-video weights, trained at multiple resolutions, supporting start and end frame prediction. |
586
+ | Wan2.1-Fun-14B-InP | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-14B-InP) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-14B-InP) | Wan2.1-Fun-14B text-to-video weights, trained at multiple resolutions, supporting start and end frame prediction. |
587
+ | Wan2.1-Fun-1.3B-Control | 19.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-1.3B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-1.3B-Control) | Wan2.1-Fun-1.3B video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc., and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction at 81 frames, trained at 16 frames per second, with multilingual prediction support. |
588
+ | Wan2.1-Fun-14B-Control | 47.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-14B-Control) | [😄Link](https://modelscope.cn/models/PAI/Wan2.1-Fun-14B-Control) | Wan2.1-Fun-14B video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc., and trajectory control. Supports multi-resolution (512, 768, 1024) video prediction at 81 frames, trained at 16 frames per second, with multilingual prediction support. |
589
+
590
+ ## 4. Wan2.1
591
+
592
+ | Name | Hugging Face | Model Scope | Description |
593
+ |--|--|--|--|
594
+ | Wan2.1-T2V-1.3B | [🤗Link](https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-1.3B) | Wanxiang 2.1-1.3B text-to-video weights |
595
+ | Wan2.1-T2V-14B | [🤗Link](https://huggingface.co/Wan-AI/Wan2.1-T2V-14B) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-14B) | Wanxiang 2.1-14B text-to-video weights |
596
+ | Wan2.1-I2V-14B-480P | [🤗Link](https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-480P) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-480P) | Wanxiang 2.1-14B-480P image-to-video weights |
597
+ | Wan2.1-I2V-14B-720P| [🤗Link](https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-720P) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-720P) | Wanxiang 2.1-14B-720P image-to-video weights |
598
+
599
+ ## 5. FantasyTalking
600
+
601
+ | Name | Storage | Hugging Face | Model Scope | Description |
602
+ |--|--|--|--|--|
603
+ | Wan2.1-I2V-14B-720P | - | [🤗Link](https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-720P) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-720P) | Wan 2.1-14B-720P image-to-video model weights |
604
+ | Wav2Vec | - | [🤗Link](https://huggingface.co/facebook/wav2vec2-base-960h) | [😄Link](https://modelscope.cn/models/AI-ModelScope/wav2vec2-base-960h) | Wav2Vec model; place inside the Wan2.1-I2V-14B-720P folder and rename to `audio_encoder` |
605
+ | FantasyTalking model | - | [🤗Link](https://huggingface.co/acvlab/FantasyTalking/) | [😄Link](https://www.modelscope.cn/models/amap_cvlab/FantasyTalking/) | Official audio-conditioned weights |
606
+
607
+ ## 6. Qwen-Image
608
+
609
+ | Name | Storage | Hugging Face | Model Scope | Description |
610
+ |--|--|--|--|--|
611
+ | Qwen-Image | [🤗Link](https://huggingface.co/Qwen/Qwen-Image) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image) | Official Qwen-Image weights |
612
+ | Qwen-Image-Edit | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-Edit) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-Edit) | Official Qwen-Image-Edit weights |
613
+ | Qwen-Image-Edit-2509 | [🤗Link](https://huggingface.co/Qwen/Qwen-Image-Edit-2509) | [😄Link](https://modelscope.cn/models/Qwen/Qwen-Image-Edit-2509) | Official Qwen-Image-Edit-2509 weights |
614
+
615
+ ## 7. Qwen-Image-Fun
616
+
617
+ | Name | Storage | Hugging Face | Model Scope | Description |
618
+ |--|--|--|--|--|
619
+ | Qwen-Image-2512-Fun-Controlnet-Union | - | [🤗Link](https://huggingface.co/alibaba-pai/Qwen-Image-2512-Fun-Controlnet-Union) | [😄Link](https://modelscope.cn/models/PAI/Qwen-Image-2512-Fun-Controlnet-Union) | ControlNet weights for Qwen-Image-2512, supporting multiple control conditions such as Canny, Depth, Pose, MLSD, Scribble, etc. |
620
+
621
+ ## 8. Z-Image
622
+
623
+ | Name | Storage | Hugging Face | Model Scope | Description |
624
+ |--|--|--|--|--|
625
+ | Z-Image | [🤗Link](https://huggingface.co/Tongyi-MAI/Z-Image) | [😄Link](https://www.modelscope.cn/models/Tongyi-MAI/Z-Image) | Official weights for Z-Image |
626
+ | Z-Image-Turbo | [🤗Link](https://huggingface.co/Tongyi-MAI/Z-Image-Turbo) | [😄Link](https://www.modelscope.cn/models/Tongyi-MAI/Z-Image-Turbo) | Official weights for Z-Image-Turbo |
627
+
628
+ ## 9. Z-Image-Fun
629
+
630
+ | Name | Storage | Hugging Face | Model Scope | Description |
631
+ |--|--|--|--|--|
632
+ | Z-Image-Fun-Controlnet-Union-2.1 | - | [🤗Link](https://huggingface.co/alibaba-pai/Z-Image-Fun-Controlnet-Union-2.1) | [😄Link](https://modelscope.cn/models/PAI/Z-Image-Fun-Controlnet-Union-2.1) | ControlNet weights for Z-Image. Compared to the first version, it adds to more layers and has been trained for a longer period. It supports multiple control conditions including Canny, Depth, Pose, MLSD, Scribble and Gray. |
633
+ | Z-Image-Fun-Lora-Distill | - | [🤗Link](https://huggingface.co/alibaba-pai/Z-Image-Fun-Lora-Distill) | [😄Link](https://modelscope.cn/models/PAI/Z-Image-Fun-Lora-Distill) | This is a Distill LoRA for Z-Image that distills both steps and CFG. This model does not require CFG and uses 8 steps for inference. |
634
+ | Z-Image-Turbo-Fun-Controlnet-Union | - | [🤗Link](https://huggingface.co/alibaba-pai/Z-Image-Turbo-Fun-Controlnet-Union) | [😄Link](https://modelscope.cn/models/PAI/Z-Image-Turbo-Fun-Controlnet-Union) | ControlNet weights for Z-Image-Turbo, supporting multiple control conditions such as Canny, Depth, Pose, MLSD, etc. |
635
+ | Z-Image-Turbo-Fun-Controlnet-Union-2.1 | - | [🤗Link](https://huggingface.co/alibaba-pai/Z-Image-Turbo-Fun-Controlnet-Union-2.1) | [😄Link](https://modelscope.cn/models/PAI/Z-Image-Turbo-Fun-Controlnet-Union-2.1) | ControlNet weights for Z-Image-Turbo. Compared to the first version, it adds to more layers and has been trained for a longer period. It supports multiple control conditions including Canny, Depth, Pose, MLSD, and more. |
636
+
637
+ ## 10. Flux
638
+
639
+ | Name | Storage | Hugging Face | Model Scope | Description |
640
+ |--|--|--|--|--|
641
+ | FLUX.1-dev | [🤗Link](https://huggingface.co/black-forest-labs/FLUX.1-dev) | [😄Link](https://www.modelscope.cn/models/black-forest-labs/FLUX.1-dev) | Official FLUX.1-dev weights |
642
+ | FLUX.2-dev | [🤗Link](https://huggingface.co/black-forest-labs/FLUX.2-dev) | [😄Link](https://www.modelscope.cn/models/black-forest-labs/FLUX.2-dev) | Official FLUX.2-dev weights |
643
+
644
+ ## 11. Flux-Fun
645
+
646
+ | Name | Storage | Hugging Face | Model Scope | Description |
647
+ |--|--|--|--|--|
648
+ | Flux.2-dev-Fun-Controlnet-Union | - | [🤗Link](https://huggingface.co/alibaba-pai/FLUX.2-dev-Fun-Controlnet-Union) | [😄Link](https://modelscope.cn/models/PAI/FLUX.2-dev-Fun-Controlnet-Union) | Flux.2-dev control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc. |
649
+
650
+ ## 12. HunyuanVideo
651
+
652
+ | Name | Storage | Hugging Face | Model Scope | Description |
653
+ |--|--|--|--|--|
654
+ | HunyuanVideo | [🤗Link](https://huggingface.co/hunyuanvideo-community/HunyuanVideo) | - | HunyuanVideo-diffusers weights |
655
+ | HunyuanVideo-I2V | [🤗Link](https://huggingface.co/hunyuanvideo-community/HunyuanVideo-I2V) | - | HunyuanVideo-I2V-diffusers weights |
656
+
657
+ ## 13. CogVideoX-Fun
658
+
659
+ V1.5:
660
+
661
+ | Name | Storage Space | Hugging Face | Model Scope | Description |
662
+ |--|--|--|--|--|
663
+ | CogVideoX-Fun-V1.5-5b-InP | 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.5-5b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.5-5b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024) and has been trained on 85 frames at a rate of 8 frames per second. |
664
+ | CogVideoX-Fun-V1.5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.5-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by CogVideoX-Fun-V1.5 to better match human preferences. |
665
+
666
+ V1.1:
667
+
668
+ | Name | Storage Space | Hugging Face | Model Scope | Description |
669
+ |--|--|--|--|--|
670
+ | CogVideoX-Fun-V1.1-2b-InP | 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. |
671
+ | CogVideoX-Fun-V1.1-5b-InP | 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Noise has been added to the reference image, and the amplitude of motion is greater compared to V1.0. |
672
+ | CogVideoX-Fun-V1.1-2b-Pose | 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.|
673
+ | CogVideoX-Fun-V1.1-2b-Control | 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-2b-Control) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-2b-Control) | Our official control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Supporting various control conditions such as Canny, Depth, Pose, MLSD, etc.|
674
+ | CogVideoX-Fun-V1.1-5b-Pose | 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Pose) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-Pose) | Our official pose-control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second.|
675
+ | CogVideoX-Fun-V1.1-5b-Control | 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Control) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-5b-Control) | Our official control video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. Supporting various control conditions such as Canny, Depth, Pose, MLSD, etc.|
676
+ | CogVideoX-Fun-V1.1-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-V1.1-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by CogVideoX-Fun-V1.1 to better match human preferences. |
677
+
678
+ <details>
679
+ <summary>(Obsolete) V1.0:</summary>
680
+
681
+ | Name | Storage Space | Hugging Face | Model Scope | Description |
682
+ |--|--|--|--|--|
683
+ | CogVideoX-Fun-2b-InP | 13.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-2b-InP) | [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-2b-InP) | Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. |
684
+ | CogVideoX-Fun-5b-InP | 20.0 GB | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-5b-InP)| [😄Link](https://modelscope.cn/models/PAI/CogVideoX-Fun-5b-InP)| Our official graph-generated video model is capable of predicting videos at multiple resolutions (512, 768, 1024, 1280) and has been trained on 49 frames at a rate of 8 frames per second. |
685
+ </details>
686
+
687
+ # Reference
688
+ - CogVideo: https://github.com/THUDM/CogVideo/
689
+ - EasyAnimate: https://github.com/aigc-apps/EasyAnimate
690
+ - Wan2.1: https://github.com/Wan-Video/Wan2.1/
691
+ - Wan2.2: https://github.com/Wan-Video/Wan2.2/
692
+ - Diffusers: https://github.com/huggingface/diffusers
693
+ - Qwen-Image: https://github.com/QwenLM/Qwen-Image
694
+ - Self-Forcing: https://github.com/guandeh17/Self-Forcing
695
+ - Flux: https://github.com/black-forest-labs/flux
696
+ - Flux2: https://github.com/black-forest-labs/flux2
697
+ - HunyuanVideo: https://github.com/Tencent-Hunyuan/HunyuanVideo
698
+ - ComfyUI-KJNodes: https://github.com/kijai/ComfyUI-KJNodes
699
+ - ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper
700
+ - ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper
701
+ - CameraCtrl: https://github.com/hehao13/CameraCtrl
702
+
703
+ # Citation
704
+
705
+ If you use VideoX-Fun in your research or project, please cite it as follows:
706
+
707
+ ```bibtex
708
+ @misc{aigc_apps_VideoX_Fun_2026,
709
+ author = {aigc-apps},
710
+ title = {VideoX-Fun: A Video Generation Pipeline for Diffusion Transformer},
711
+ year = {2026},
712
+ publisher = {GitHub},
713
+ url = {https://github.com/aigc-apps/VideoX-Fun}
714
+ }
715
+ ```
716
+
717
+ # Limitations and Risks
718
+
719
+ - Generated videos may have artifacts or quality issues, especially in complex scenes.
720
+ - The model may struggle with fine details, text rendering, or specific artistic styles.
721
+ - Performance varies with input prompt quality, resolution, and other parameters.
722
+ - The technology could be misused to create misleading content (e.g., deepfakes). Users are responsible for ethical use.
723
+ - The model may reflect biases present in the training data.
724
+ - Users should respect privacy and copyright when using real people's images or videos.
725
+
726
+ We encourage responsible use and recommend implementing safeguards in production environments.
727
+
728
+ # License
729
+ This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
730
+
731
+ The CogVideoX-2B model (including its corresponding Transformers module and VAE module) is released under the [Apache 2.0 License](LICENSE).
732
+
733
+ The CogVideoX-5B model (Transformers module) is released under the [CogVideoX LICENSE](https://huggingface.co/THUDM/CogVideoX-5b/blob/main/LICENSE).
vendor/VideoX-Fun/config/flux2/flux2_control.yaml ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: flux2
3
+ transformer_additional_kwargs:
4
+ control_layers: [0, 2, 4, 6]
5
+ control_in_dim: 260
vendor/VideoX-Fun/config/qwenimage/qwenimage_control.yaml ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: qwenimage
3
+ transformer_additional_kwargs:
4
+ control_layers: [0, 12, 24, 36, 48]
5
+ control_in_dim: 132
vendor/VideoX-Fun/config/wan2.1/wan_civitai.yaml ADDED
@@ -0,0 +1,39 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_subpath: ./
5
+ dict_mapping:
6
+ in_dim: in_channels
7
+ dim: hidden_size
8
+
9
+ vae_kwargs:
10
+ vae_subpath: Wan2.1_VAE.pth
11
+ temporal_compression_ratio: 4
12
+ spatial_compression_ratio: 8
13
+
14
+ text_encoder_kwargs:
15
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
16
+ tokenizer_subpath: google/umt5-xxl
17
+ text_length: 512
18
+ vocab: 256384
19
+ dim: 4096
20
+ dim_attn: 4096
21
+ dim_ffn: 10240
22
+ num_heads: 64
23
+ num_layers: 24
24
+ num_buckets: 32
25
+ shared_pos: False
26
+ dropout: 0.0
27
+
28
+ scheduler_kwargs:
29
+ scheduler_subpath: null
30
+ num_train_timesteps: 1000
31
+ shift: 5.0
32
+ use_dynamic_shifting: false
33
+ base_shift: 0.5
34
+ max_shift: 1.15
35
+ base_image_seq_len: 256
36
+ max_image_seq_len: 4096
37
+
38
+ image_encoder_kwargs:
39
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/wan2.2/wan_civitai_5b.yaml ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_low_noise_model_subpath: ./
5
+ transformer_combination_type: "single"
6
+ dict_mapping:
7
+ in_dim: in_channels
8
+ dim: hidden_size
9
+
10
+ vae_kwargs:
11
+ vae_type: "AutoencoderKLWan3_8"
12
+ vae_subpath: Wan2.2_VAE.pth
13
+ temporal_compression_ratio: 4
14
+ spatial_compression_ratio: 16
15
+
16
+ text_encoder_kwargs:
17
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
18
+ tokenizer_subpath: google/umt5-xxl
19
+ text_length: 512
20
+ vocab: 256384
21
+ dim: 4096
22
+ dim_attn: 4096
23
+ dim_ffn: 10240
24
+ num_heads: 64
25
+ num_layers: 24
26
+ num_buckets: 32
27
+ shared_pos: False
28
+ dropout: 0.0
29
+
30
+ scheduler_kwargs:
31
+ scheduler_subpath: null
32
+ num_train_timesteps: 1000
33
+ shift: 5.0
34
+ use_dynamic_shifting: false
35
+ base_shift: 0.5
36
+ max_shift: 1.15
37
+ base_image_seq_len: 256
38
+ max_image_seq_len: 4096
39
+
40
+ image_encoder_kwargs:
41
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/wan2.2/wan_civitai_animate.yaml ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_low_noise_model_subpath: ./
5
+ transformer_combination_type: "single"
6
+ dict_mapping:
7
+ in_dim: in_channels
8
+ dim: hidden_size
9
+
10
+ vae_kwargs:
11
+ vae_type: "AutoencoderKLWan"
12
+ vae_subpath: Wan2.1_VAE.pth
13
+ temporal_compression_ratio: 4
14
+ spatial_compression_ratio: 8
15
+
16
+ text_encoder_kwargs:
17
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
18
+ tokenizer_subpath: google/umt5-xxl
19
+ text_length: 512
20
+ vocab: 256384
21
+ dim: 4096
22
+ dim_attn: 4096
23
+ dim_ffn: 10240
24
+ num_heads: 64
25
+ num_layers: 24
26
+ num_buckets: 32
27
+ shared_pos: False
28
+ dropout: 0.0
29
+
30
+ scheduler_kwargs:
31
+ scheduler_subpath: null
32
+ num_train_timesteps: 1000
33
+ shift: 5.0
34
+ use_dynamic_shifting: false
35
+ base_shift: 0.5
36
+ max_shift: 1.15
37
+ base_image_seq_len: 256
38
+ max_image_seq_len: 4096
39
+
40
+ image_encoder_kwargs:
41
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/wan2.2/wan_civitai_i2v.yaml ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_low_noise_model_subpath: ./low_noise_model
5
+ transformer_high_noise_model_subpath: ./high_noise_model
6
+ transformer_combination_type: "moe"
7
+ boundary: 0.900
8
+ dict_mapping:
9
+ in_dim: in_channels
10
+ dim: hidden_size
11
+
12
+ vae_kwargs:
13
+ vae_type: "AutoencoderKLWan"
14
+ vae_subpath: Wan2.1_VAE.pth
15
+ temporal_compression_ratio: 4
16
+ spatial_compression_ratio: 8
17
+
18
+ text_encoder_kwargs:
19
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
20
+ tokenizer_subpath: google/umt5-xxl
21
+ text_length: 512
22
+ vocab: 256384
23
+ dim: 4096
24
+ dim_attn: 4096
25
+ dim_ffn: 10240
26
+ num_heads: 64
27
+ num_layers: 24
28
+ num_buckets: 32
29
+ shared_pos: False
30
+ dropout: 0.0
31
+
32
+ scheduler_kwargs:
33
+ scheduler_subpath: null
34
+ num_train_timesteps: 1000
35
+ shift: 5.0
36
+ use_dynamic_shifting: false
37
+ base_shift: 0.5
38
+ max_shift: 1.15
39
+ base_image_seq_len: 256
40
+ max_image_seq_len: 4096
41
+
42
+ image_encoder_kwargs:
43
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/wan2.2/wan_civitai_s2v.yaml ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_low_noise_model_subpath: ./
5
+ transformer_combination_type: "single"
6
+ dict_mapping:
7
+ in_dim: in_channels
8
+ dim: hidden_size
9
+
10
+ vae_kwargs:
11
+ vae_type: "AutoencoderKLWan"
12
+ vae_subpath: Wan2.1_VAE.pth
13
+ temporal_compression_ratio: 4
14
+ spatial_compression_ratio: 8
15
+
16
+ text_encoder_kwargs:
17
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
18
+ tokenizer_subpath: google/umt5-xxl
19
+ text_length: 512
20
+ vocab: 256384
21
+ dim: 4096
22
+ dim_attn: 4096
23
+ dim_ffn: 10240
24
+ num_heads: 64
25
+ num_layers: 24
26
+ num_buckets: 32
27
+ shared_pos: False
28
+ dropout: 0.0
29
+
30
+ audio_encoder_kwargs:
31
+ audio_encoder_subpath: wav2vec2-large-xlsr-53-english
32
+
33
+ scheduler_kwargs:
34
+ scheduler_subpath: null
35
+ num_train_timesteps: 1000
36
+ shift: 3.0
37
+ use_dynamic_shifting: false
38
+ base_shift: 0.5
39
+ max_shift: 1.15
40
+ base_image_seq_len: 256
41
+ max_image_seq_len: 4096
42
+
43
+ image_encoder_kwargs:
44
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/wan2.2/wan_civitai_t2v.yaml ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ format: civitai
2
+ pipeline: Wan
3
+ transformer_additional_kwargs:
4
+ transformer_low_noise_model_subpath: ./low_noise_model
5
+ transformer_high_noise_model_subpath: ./high_noise_model
6
+ transformer_combination_type: "moe"
7
+ boundary: 0.875
8
+ dict_mapping:
9
+ in_dim: in_channels
10
+ dim: hidden_size
11
+
12
+ vae_kwargs:
13
+ vae_type: "AutoencoderKLWan"
14
+ vae_subpath: Wan2.1_VAE.pth
15
+ temporal_compression_ratio: 4
16
+ spatial_compression_ratio: 8
17
+
18
+ text_encoder_kwargs:
19
+ text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
20
+ tokenizer_subpath: google/umt5-xxl
21
+ text_length: 512
22
+ vocab: 256384
23
+ dim: 4096
24
+ dim_attn: 4096
25
+ dim_ffn: 10240
26
+ num_heads: 64
27
+ num_layers: 24
28
+ num_buckets: 32
29
+ shared_pos: False
30
+ dropout: 0.0
31
+
32
+ scheduler_kwargs:
33
+ scheduler_subpath: null
34
+ num_train_timesteps: 1000
35
+ shift: 12.0
36
+ use_dynamic_shifting: false
37
+ base_shift: 0.5
38
+ max_shift: 1.15
39
+ base_image_seq_len: 256
40
+ max_image_seq_len: 4096
41
+
42
+ image_encoder_kwargs:
43
+ image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth
vendor/VideoX-Fun/config/z_image/z_image_control.yaml ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: z_image
3
+ transformer_additional_kwargs:
4
+ control_layers_places: [0, 5, 10, 15, 20, 25]
5
+ control_in_dim: 16
vendor/VideoX-Fun/config/z_image/z_image_control_2.0.yaml ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: z_image
3
+ transformer_additional_kwargs:
4
+ control_layers_places: [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28]
5
+ control_refiner_layers_places: [0, 1]
6
+ add_control_noise_refiner: true
7
+ add_control_noise_refiner_correctly: false
8
+ control_in_dim: 33
vendor/VideoX-Fun/config/z_image/z_image_control_2.1.yaml ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: z_image
3
+ transformer_additional_kwargs:
4
+ control_layers_places: [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28]
5
+ control_refiner_layers_places: [0, 1]
6
+ add_control_noise_refiner: true
7
+ add_control_noise_refiner_correctly: true
8
+ control_in_dim: 33
vendor/VideoX-Fun/config/z_image/z_image_control_2.1_lite.yaml ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ format: diffusers
2
+ pipeline: z_image
3
+ transformer_additional_kwargs:
4
+ control_layers_places: [0, 10, 20]
5
+ control_refiner_layers_places: [0, 1]
6
+ add_control_noise_refiner: true
7
+ add_control_noise_refiner_correctly: true
8
+ control_in_dim: 33
vendor/VideoX-Fun/config/zero_stage2_config.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bf16": {
3
+ "enabled": true
4
+ },
5
+ "train_micro_batch_size_per_gpu": 1,
6
+ "train_batch_size": "auto",
7
+ "gradient_accumulation_steps": "auto",
8
+ "dump_state": true,
9
+ "zero_optimization": {
10
+ "stage": 2,
11
+ "overlap_comm": true,
12
+ "contiguous_gradients": true,
13
+ "sub_group_size": 1e9,
14
+ "reduce_bucket_size": 5e8
15
+ }
16
+ }
vendor/VideoX-Fun/config/zero_stage3_config.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bf16": {
3
+ "enabled": true
4
+ },
5
+ "train_micro_batch_size_per_gpu": 1,
6
+ "train_batch_size": "auto",
7
+ "gradient_accumulation_steps": "auto",
8
+ "gradient_clipping": "auto",
9
+ "steps_per_print": 2000,
10
+ "wall_clock_breakdown": false,
11
+ "zero_optimization": {
12
+ "stage": 3,
13
+ "overlap_comm": true,
14
+ "contiguous_gradients": true,
15
+ "reduce_bucket_size": 5e8,
16
+ "sub_group_size": 1e9,
17
+ "stage3_max_live_parameters": 1e9,
18
+ "stage3_max_reuse_distance": 1e9,
19
+ "stage3_gather_16bit_weights_on_model_save": "auto",
20
+ "offload_optimizer": {
21
+ "device": "none"
22
+ },
23
+ "offload_param": {
24
+ "device": "none"
25
+ }
26
+ }
27
+ }
28
+
vendor/VideoX-Fun/config/zero_stage3_config_cpu_offload.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bf16": {
3
+ "enabled": true
4
+ },
5
+ "train_micro_batch_size_per_gpu": 1,
6
+ "train_batch_size": "auto",
7
+ "gradient_accumulation_steps": "auto",
8
+ "gradient_clipping": "auto",
9
+ "steps_per_print": 2000,
10
+ "wall_clock_breakdown": false,
11
+ "zero_optimization": {
12
+ "stage": 3,
13
+ "overlap_comm": true,
14
+ "contiguous_gradients": true,
15
+ "reduce_bucket_size": 5e8,
16
+ "sub_group_size": 1e9,
17
+ "stage3_max_live_parameters": 1e9,
18
+ "stage3_max_reuse_distance": 1e9,
19
+ "stage3_gather_16bit_weights_on_model_save": "auto",
20
+ "offload_optimizer": {
21
+ "device": "cpu"
22
+ },
23
+ "offload_param": {
24
+ "device": "cpu"
25
+ }
26
+ }
27
+ }
28
+
vendor/VideoX-Fun/examples/cogvideox_fun/app.py ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import time
4
+
5
+ import torch
6
+
7
+ current_file_path = os.path.abspath(__file__)
8
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
9
+ for project_root in project_roots:
10
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
11
+
12
+ from videox_fun.api.api import (infer_forward_api,
13
+ update_diffusion_transformer_api)
14
+ from videox_fun.ui.controller import ddpm_scheduler_dict
15
+ from videox_fun.ui.cogvideox_fun_ui import ui, ui_client, ui_host
16
+
17
+ if __name__ == "__main__":
18
+ # Choose the ui mode
19
+ # "normal" refers to the standard UI, which allows users to click to switch models, change model types, and more.
20
+ # "host" represents the hosting mode, where the model is loaded directly at startup and can be accessed via
21
+ # the API to return generation results.
22
+ # "client" represents the client mode, offering a simple UI that sends requests to a remote API for generation.
23
+ ui_mode = "normal"
24
+
25
+ # GPU memory mode, which can be chosen in [model_full_load, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
26
+ # model_full_load means that the entire model will be moved to the GPU.
27
+ #
28
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
29
+ #
30
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
31
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
32
+ #
33
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
34
+ # resulting in slower speeds but saving a large amount of GPU memory.
35
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
36
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
37
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
38
+ compile_dit = False
39
+
40
+ # Use torch.float16 if GPU does not support torch.bfloat16
41
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
42
+ weight_dtype = torch.bfloat16
43
+
44
+ # Server ip
45
+ server_name = "0.0.0.0"
46
+ server_port = 7860
47
+
48
+ # Params below is used when ui_mode = "host"
49
+ model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
50
+ # "Inpaint" or "Control"
51
+ model_type = "Inpaint"
52
+
53
+ if ui_mode == "host":
54
+ demo, controller = ui_host(GPU_memory_mode, ddpm_scheduler_dict, model_name, model_type, compile_dit, weight_dtype)
55
+ elif ui_mode == "client":
56
+ demo, controller = ui_client(ddpm_scheduler_dict, model_name)
57
+ else:
58
+ demo, controller = ui(GPU_memory_mode, ddpm_scheduler_dict, compile_dit, weight_dtype)
59
+
60
+ # launch gradio
61
+ app, _, _ = demo.queue(status_update_rate=1).launch(
62
+ server_name=server_name,
63
+ server_port=server_port,
64
+ prevent_thread_lock=True
65
+ )
66
+
67
+ # launch api
68
+ infer_forward_api(None, app, controller)
69
+ update_diffusion_transformer_api(None, app, controller)
70
+
71
+ # not close the python
72
+ while True:
73
+ time.sleep(5)
vendor/VideoX-Fun/examples/cogvideox_fun/launch_api.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ import sys
4
+ import time
5
+
6
+ import gradio as gr
7
+ import ray
8
+ import torch
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.api.api_multi_nodes import (MultiNodesEngine,
16
+ multi_nodes_infer_forward_api)
17
+ from videox_fun.ui.controller import flow_scheduler_dict
18
+ from videox_fun.ui.cogvideox_fun_ui import CogVideoXFunController
19
+
20
+ def main():
21
+ parser = argparse.ArgumentParser(description='xDiT HTTP Service')
22
+ parser.add_argument('--world_size', type=int, default=8, help='Number of parallel workers')
23
+ parser.add_argument(
24
+ '--gpu_memory_mode', type=str, default="model_cpu_offload", help='''
25
+ GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8].
26
+ model_full_load means that the entire model will be moved to the GPU.
27
+
28
+ model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
29
+ and the transformer model has been quantized to float8, which can save more GPU memory.
30
+
31
+ model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
32
+
33
+ model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
34
+ and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ '''
36
+ )
37
+ parser.add_argument(
38
+ '--compile_dit', action='store_true', help='''
39
+ Enable compile dit.
40
+ Compile will give a speedup in fixed resolution and need a little GPU memory.
41
+ The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
42
+ '''
43
+ )
44
+ parser.add_argument('--fsdp_dit', action='store_true', help="Use DIT FSDP to save more GPU memory in multi gpus.")
45
+ parser.add_argument('--fsdp_text_encoder', action='store_true', help="Use Text Encoder FSDP to save more GPU memory in multi gpus.")
46
+ parser.add_argument('--ulysses_degree', type=int, default=4, help='Degree of Ulysses configuration')
47
+ parser.add_argument('--ring_degree', type=int, default=2, help='Degree of Ring configuration')
48
+ parser.add_argument('--weight_dtype', type=str, default='bf16', help='Weight data type')
49
+ parser.add_argument('--server_name', type=str, default="0.0.0.0", help='Server IP address')
50
+ parser.add_argument('--server_port', type=int, default=7860, help='Server Port')
51
+ parser.add_argument('--model_name', type=str, default="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP", help='Model path')
52
+ parser.add_argument('--model_type', type=str, default="Inpaint", help='Model type (Inpaint/Control)')
53
+ parser.add_argument('--savedir_sample', type=str, default=None, help='The save directory for samples')
54
+ args = parser.parse_args()
55
+
56
+ weight_dtype = torch.float32
57
+ if args.weight_dtype == "bf16":
58
+ weight_dtype = torch.bfloat16
59
+ elif args.weight_dtype == "fp16":
60
+ weight_dtype = torch.float16
61
+
62
+ engine = MultiNodesEngine(
63
+ world_size=args.world_size, Controller=CogVideoXFunController,
64
+ GPU_memory_mode=args.gpu_memory_mode, scheduler_dict=flow_scheduler_dict, model_name=args.model_name, model_type=args.model_type, config_path=None,
65
+ ulysses_degree=args.ulysses_degree, ring_degree=args.ring_degree,
66
+ fsdp_dit=args.fsdp_dit, fsdp_text_encoder=args.fsdp_text_encoder, compile_dit=args.compile_dit,
67
+ weight_dtype=weight_dtype, savedir_sample=args.savedir_sample,
68
+ )
69
+
70
+ def gr_launch():
71
+ # launch gradio
72
+ with gr.Blocks() as demo:
73
+ gr.Markdown("")
74
+ app, _, _ = demo.queue(status_update_rate=1).launch(
75
+ server_name=args.server_name,
76
+ server_port=args.server_port,
77
+ prevent_thread_lock=True
78
+ )
79
+
80
+ # launch api
81
+ multi_nodes_infer_forward_api(None, app, engine)
82
+
83
+ gr_launch()
84
+
85
+ # not close the python
86
+ while True:
87
+ time.sleep(5)
88
+
89
+ if __name__ == "__main__":
90
+ main()
vendor/VideoX-Fun/examples/cogvideox_fun/post_infer.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import json
3
+ import time
4
+ from datetime import datetime
5
+
6
+ import requests
7
+ import base64
8
+
9
+
10
+ def post_diffusion_transformer(diffusion_transformer_path, url='http://127.0.0.1:7860'):
11
+ datas = json.dumps({
12
+ "diffusion_transformer_path": diffusion_transformer_path
13
+ })
14
+ r = requests.post(f'{url}/cogvideox_fun/update_diffusion_transformer', data=datas, timeout=1500)
15
+ data = r.content.decode('utf-8')
16
+ return data
17
+
18
+ def post_update_edition(edition, url='http://0.0.0.0:7860'):
19
+ datas = json.dumps({
20
+ "edition": edition
21
+ })
22
+ r = requests.post(f'{url}/cogvideox_fun/update_edition', data=datas, timeout=1500)
23
+ data = r.content.decode('utf-8')
24
+ return data
25
+
26
+
27
+ def post_infer(
28
+ generation_method,
29
+ length_slider,
30
+ url='http://127.0.0.1:7860',
31
+ POST_TOKEN="",
32
+ timeout=5000,
33
+ base_model_path="none",
34
+ lora_model_path="none",
35
+ lora_alpha_slider=0.55,
36
+ prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.",
37
+ negative_prompt_textbox="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion.",
38
+ sampler_dropdown="Flow",
39
+ sample_step_slider=50,
40
+ width_slider=672,
41
+ height_slider=384,
42
+ cfg_scale_slider=6,
43
+ seed_textbox=43
44
+ ):
45
+ # Prepare the data payload
46
+ datas = json.dumps({
47
+ "base_model_path": base_model_path,
48
+ "lora_model_path": lora_model_path,
49
+ "lora_alpha_slider": lora_alpha_slider,
50
+ "prompt_textbox": prompt_textbox,
51
+ "negative_prompt_textbox": negative_prompt_textbox,
52
+ "sampler_dropdown": sampler_dropdown,
53
+ "sample_step_slider": sample_step_slider,
54
+ "width_slider": width_slider,
55
+ "height_slider": height_slider,
56
+ "generation_method": generation_method,
57
+ "length_slider": length_slider,
58
+ "cfg_scale_slider": cfg_scale_slider,
59
+ "seed_textbox": seed_textbox,
60
+ })
61
+
62
+ # Initialize session and set headers
63
+ session = requests.session()
64
+ session.headers.update({"Authorization": POST_TOKEN})
65
+
66
+ # Send POST request
67
+ if url[-1] == "/":
68
+ url = url[:-1]
69
+ post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout)
70
+
71
+ data = post_r.content.decode('utf-8')
72
+ return data
73
+
74
+ if __name__ == '__main__':
75
+ # initiate time
76
+ time_start = time.time()
77
+
78
+ # The Url you want to post
79
+ POST_URL = 'http://0.0.0.0:7860'
80
+ # Used in EAS. If you don't need Authorization, please set it to empty string.
81
+ TOKEN = ''
82
+
83
+ # -------------------------- #
84
+ # Step 1: update edition
85
+ # -------------------------- #
86
+ # diffusion_transformer_path = "models/Diffusion_Transformer/CogVideoX-Fun-2b-InP"
87
+ # outputs = post_diffusion_transformer(diffusion_transformer_path)
88
+ # print('Output update edition: ', outputs)
89
+
90
+ # -------------------------- #
91
+ # Step 2: infer
92
+ # -------------------------- #
93
+ # "Video Generation" and "Image Generation"
94
+ generation_method = "Video Generation"
95
+ # Video length
96
+ length_slider = 49
97
+ # Used in Lora models
98
+ lora_model_path = "none"
99
+ lora_alpha_slider = 0.55
100
+ # Prompts
101
+ prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
102
+ negative_prompt_textbox = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion."
103
+ # Sampler name
104
+ sampler_dropdown = "Euler"
105
+ # Sampler steps
106
+ sample_step_slider = 50
107
+ # height and width
108
+ width_slider = 672
109
+ height_slider = 384
110
+ # cfg scale
111
+ cfg_scale_slider = 6
112
+ seed_textbox = 43
113
+
114
+ outputs = post_infer(
115
+ generation_method,
116
+ length_slider,
117
+ lora_model_path=lora_model_path,
118
+ lora_alpha_slider=lora_alpha_slider,
119
+ prompt_textbox=prompt_textbox,
120
+ negative_prompt_textbox=negative_prompt_textbox,
121
+ sampler_dropdown=sampler_dropdown,
122
+ sample_step_slider=sample_step_slider,
123
+ width_slider=width_slider,
124
+ height_slider=height_slider,
125
+ cfg_scale_slider=cfg_scale_slider,
126
+ seed_textbox=seed_textbox,
127
+ url=POST_URL,
128
+ POST_TOKEN=TOKEN
129
+ )
130
+
131
+ # Get decoded data
132
+ outputs = json.loads(outputs)
133
+ base64_encoding = outputs["base64_encoding"]
134
+ decoded_data = base64.b64decode(base64_encoding)
135
+
136
+ is_image = True if generation_method == "Image Generation" else False
137
+ if is_image or length_slider == 1:
138
+ file_path = "1.png"
139
+ else:
140
+ file_path = "1.mp4"
141
+ with open(file_path, "wb") as file:
142
+ file.write(decoded_data)
143
+
144
+ # End of record time
145
+ # The calculated time difference is the execution time of the program, expressed in seconds / s
146
+ time_end = time.time()
147
+ time_sum = (time_end - time_start)
148
+ print('# --------------------------------------------------------- #')
149
+ print(f'# Total expenditure: {time_sum}s')
150
+ print('# --------------------------------------------------------- #')
vendor/VideoX-Fun/examples/cogvideox_fun/post_infer_queue.py ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import json
3
+ import time
4
+ import urllib.parse
5
+ import requests
6
+
7
+
8
+ def post_infer(
9
+ generation_method,
10
+ length_slider,
11
+ url='http://127.0.0.1:7860',
12
+ POST_TOKEN="",
13
+ timeout=5,
14
+ base_model_path="none",
15
+ lora_model_path="none",
16
+ lora_alpha_slider=0.55,
17
+ prompt_textbox="A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic.",
18
+ negative_prompt_textbox="The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion.",
19
+ sampler_dropdown="Flow",
20
+ sample_step_slider=50,
21
+ width_slider=672,
22
+ height_slider=384,
23
+ cfg_scale_slider=6,
24
+ seed_textbox=43
25
+ ):
26
+ # Prepare the data payload
27
+ datas = json.dumps({
28
+ "base_model_path": base_model_path,
29
+ "lora_model_path": lora_model_path,
30
+ "lora_alpha_slider": lora_alpha_slider,
31
+ "prompt_textbox": prompt_textbox,
32
+ "negative_prompt_textbox": negative_prompt_textbox,
33
+ "sampler_dropdown": sampler_dropdown,
34
+ "sample_step_slider": sample_step_slider,
35
+ "width_slider": width_slider,
36
+ "height_slider": height_slider,
37
+ "generation_method": generation_method,
38
+ "length_slider": length_slider,
39
+ "cfg_scale_slider": cfg_scale_slider,
40
+ "seed_textbox": seed_textbox,
41
+ })
42
+
43
+ # Initialize session and set headers
44
+ session = requests.session()
45
+ session.headers.update({"Authorization": POST_TOKEN})
46
+
47
+ # Send POST request
48
+ if url[-1] == "/":
49
+ url = url[:-1]
50
+ post_r = session.post(f'{url}/videox_fun/infer_forward', data=datas, timeout=timeout)
51
+
52
+ # Extract request ID from POST response headers
53
+ request_id = post_r.headers.get("X-Eas-Queueservice-Request-Id")
54
+
55
+ # Prepare query parameters for GET request
56
+ query = {
57
+ '_index_': '0',
58
+ '_length_': '1',
59
+ '_timeout_': str(timeout),
60
+ '_raw_': 'false',
61
+ '_auto_delete_': 'true',
62
+ }
63
+ if request_id:
64
+ query['requestId'] = request_id
65
+
66
+ query_str = urllib.parse.urlencode(query)
67
+
68
+ # Polling GET request until status code is not 204
69
+ status_code = 204
70
+ while status_code == 204:
71
+ if query_str:
72
+ get_r = session.get(f'{url}/sink?{query_str}', timeout=timeout)
73
+ else:
74
+ get_r = session.get(f'{url}/sink', timeout=timeout)
75
+ status_code = get_r.status_code
76
+ # Decode and return the response content
77
+ data = get_r.content.decode('utf-8')
78
+ return data
79
+
80
+ if __name__ == '__main__':
81
+ # initiate time
82
+ time_start = time.time()
83
+
84
+ # EAS队列配置
85
+ EAS_URL = 'http://17xxxxxxxxx.pai-eas.aliyuncs.com/api/predict/xxxxxxxx'
86
+ # Use in EAS Queue
87
+ TOKEN = 'xxxxxxxx'
88
+
89
+ # "Video Generation" and "Image Generation"
90
+ generation_method = "Video Generation"
91
+ # Video length
92
+ length_slider = 49
93
+ # Used in Lora models
94
+ lora_model_path = "none"
95
+ lora_alpha_slider = 0.55
96
+ # Prompts
97
+ prompt_textbox = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
98
+ negative_prompt_textbox = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion."
99
+ # Sampler name
100
+ sampler_dropdown = "Euler"
101
+ # Sampler steps
102
+ sample_step_slider = 50
103
+ # height and width
104
+ width_slider = 672
105
+ height_slider = 384
106
+ # cfg scale
107
+ cfg_scale_slider = 6
108
+ seed_textbox = 43
109
+
110
+ outputs = post_infer(
111
+ generation_method,
112
+ length_slider,
113
+ lora_model_path=lora_model_path,
114
+ lora_alpha_slider=lora_alpha_slider,
115
+ prompt_textbox=prompt_textbox,
116
+ negative_prompt_textbox=negative_prompt_textbox,
117
+ sampler_dropdown=sampler_dropdown,
118
+ sample_step_slider=sample_step_slider,
119
+ width_slider=width_slider,
120
+ height_slider=height_slider,
121
+ cfg_scale_slider=cfg_scale_slider,
122
+ seed_textbox=seed_textbox,
123
+ url=EAS_URL,
124
+ POST_TOKEN=TOKEN
125
+ )
126
+ # Get decoded data
127
+ outputs = json.loads(base64.b64decode(json.loads(outputs)[0]['data']))
128
+ base64_encoding = outputs["base64_encoding"]
129
+ decoded_data = base64.b64decode(base64_encoding)
130
+
131
+ is_image = True if generation_method == "Image Generation" else False
132
+ if is_image or length_slider == 1:
133
+ file_path = "1.png"
134
+ else:
135
+ file_path = "1.mp4"
136
+ with open(file_path, "wb") as file:
137
+ file.write(decoded_data)
138
+
139
+ # End of record time
140
+ # The calculated time difference is the execution time of the program, expressed in seconds / s
141
+ time_end = time.time()
142
+ time_sum = (time_end - time_start)
143
+ print('# --------------------------------------------------------- #')
144
+ print(f'# Total expenditure: {time_sum}s')
145
+ print('# --------------------------------------------------------- #')
vendor/VideoX-Fun/examples/cogvideox_fun/predict_i2v.py ADDED
@@ -0,0 +1,328 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import (CogVideoXDDIMScheduler, DDIMScheduler,
7
+ DPMSolverMultistepScheduler,
8
+ EulerAncestralDiscreteScheduler, EulerDiscreteScheduler,
9
+ PNDMScheduler)
10
+ from PIL import Image
11
+
12
+ current_file_path = os.path.abspath(__file__)
13
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
14
+ for project_root in project_roots:
15
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
16
+
17
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
18
+ from videox_fun.models import (AutoencoderKLCogVideoX,
19
+ CogVideoXTransformer3DModel, T5EncoderModel,
20
+ T5Tokenizer)
21
+ from videox_fun.pipeline import (CogVideoXFunInpaintPipeline,
22
+ CogVideoXFunPipeline)
23
+ from videox_fun.utils import (register_auto_device_hook,
24
+ safe_enable_group_offload)
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
29
+
30
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
31
+ # model_full_load means that the entire model will be moved to the GPU.
32
+ #
33
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
37
+ #
38
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
42
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
43
+ #
44
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
45
+ # resulting in slower speeds but saving a large amount of GPU memory.
46
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
47
+ # Multi GPUs config
48
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
49
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
50
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
51
+ ulysses_degree = 1
52
+ ring_degree = 1
53
+ # Use FSDP to save more GPU memory in multi gpus.
54
+ fsdp_dit = False
55
+ fsdp_text_encoder = True
56
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
57
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
58
+ compile_dit = False
59
+
60
+ # Config and model path
61
+ model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
62
+
63
+ # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" "DDIM_Cog" and "DDIM_Origin"
64
+ sampler_name = "DDIM_Origin"
65
+
66
+ # Load pretrained model if need
67
+ transformer_path = None
68
+ vae_path = None
69
+ lora_path = None
70
+
71
+ # Other params
72
+ sample_size = [384, 672]
73
+ # V1.0 and V1.1 support up to 49 frames of video generation,
74
+ # while V1.5 supports up to 85 frames.
75
+ video_length = 49
76
+ fps = 8
77
+
78
+ # If you want to generate ultra long videos, please set partial_video_length as the length of each sub video segment
79
+ partial_video_length = None
80
+ overlap_video_length = 4
81
+
82
+ # Use torch.float16 if GPU does not support torch.bfloat16
83
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
84
+ weight_dtype = torch.bfloat16
85
+ # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
86
+ validation_image_start = "asset/1.png"
87
+ validation_image_end = None
88
+
89
+ # prompts
90
+ prompt = "The dog is shaking head. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
91
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
92
+ guidance_scale = 6.0
93
+ seed = 43
94
+ num_inference_steps = 50
95
+ lora_weight = 0.55
96
+ save_path = "samples/cogvideox-fun-videos_i2v"
97
+
98
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
99
+
100
+ transformer = CogVideoXTransformer3DModel.from_pretrained(
101
+ model_name,
102
+ subfolder="transformer",
103
+ low_cpu_mem_usage=True,
104
+ torch_dtype=weight_dtype,
105
+ ).to(weight_dtype)
106
+
107
+ if transformer_path is not None:
108
+ print(f"From checkpoint: {transformer_path}")
109
+ if transformer_path.endswith("safetensors"):
110
+ from safetensors.torch import load_file, safe_open
111
+ state_dict = load_file(transformer_path)
112
+ else:
113
+ state_dict = torch.load(transformer_path, map_location="cpu")
114
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
115
+
116
+ m, u = transformer.load_state_dict(state_dict, strict=False)
117
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
118
+
119
+ # Get Vae
120
+ vae = AutoencoderKLCogVideoX.from_pretrained(
121
+ model_name,
122
+ subfolder="vae"
123
+ ).to(weight_dtype)
124
+
125
+ if vae_path is not None:
126
+ print(f"From checkpoint: {vae_path}")
127
+ if vae_path.endswith("safetensors"):
128
+ from safetensors.torch import load_file, safe_open
129
+ state_dict = load_file(vae_path)
130
+ else:
131
+ state_dict = torch.load(vae_path, map_location="cpu")
132
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
133
+
134
+ m, u = vae.load_state_dict(state_dict, strict=False)
135
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
136
+
137
+ # Get tokenizer and text_encoder
138
+ tokenizer = T5Tokenizer.from_pretrained(
139
+ model_name, subfolder="tokenizer"
140
+ )
141
+ text_encoder = T5EncoderModel.from_pretrained(
142
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
143
+ )
144
+
145
+ # Get Scheduler
146
+ Chosen_Scheduler = scheduler_dict = {
147
+ "Euler": EulerDiscreteScheduler,
148
+ "Euler A": EulerAncestralDiscreteScheduler,
149
+ "DPM++": DPMSolverMultistepScheduler,
150
+ "PNDM": PNDMScheduler,
151
+ "DDIM_Cog": CogVideoXDDIMScheduler,
152
+ "DDIM_Origin": DDIMScheduler,
153
+ }[sampler_name]
154
+ scheduler = Chosen_Scheduler.from_pretrained(
155
+ model_name,
156
+ subfolder="scheduler"
157
+ )
158
+
159
+ if transformer.config.in_channels != vae.config.latent_channels:
160
+ pipeline = CogVideoXFunInpaintPipeline(
161
+ vae=vae,
162
+ tokenizer=tokenizer,
163
+ text_encoder=text_encoder,
164
+ transformer=transformer,
165
+ scheduler=scheduler,
166
+ )
167
+ else:
168
+ pipeline = CogVideoXFunPipeline(
169
+ vae=vae,
170
+ tokenizer=tokenizer,
171
+ text_encoder=text_encoder,
172
+ transformer=transformer,
173
+ scheduler=scheduler,
174
+ )
175
+ if ulysses_degree > 1 or ring_degree > 1:
176
+ from functools import partial
177
+ transformer.enable_multi_gpus_inference()
178
+ if fsdp_dit:
179
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
180
+ pipeline.transformer = shard_fn(pipeline.transformer)
181
+ print("Add FSDP DIT")
182
+ if fsdp_text_encoder:
183
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
184
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
185
+ print("Add FSDP TEXT ENCODER")
186
+
187
+ if compile_dit:
188
+ for i in range(len(pipeline.transformer.transformer_blocks)):
189
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
190
+ print("Add Compile")
191
+
192
+ if GPU_memory_mode == "sequential_cpu_offload":
193
+ pipeline.enable_sequential_cpu_offload(device=device)
194
+ elif GPU_memory_mode == "model_group_offload":
195
+ register_auto_device_hook(pipeline.transformer)
196
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
197
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
198
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
199
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
200
+ pipeline.enable_model_cpu_offload(device=device)
201
+ elif GPU_memory_mode == "model_cpu_offload":
202
+ pipeline.enable_model_cpu_offload(device=device)
203
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
204
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
205
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
206
+ pipeline.to(device=device)
207
+ else:
208
+ pipeline.to(device=device)
209
+
210
+ generator = torch.Generator(device=device).manual_seed(seed)
211
+
212
+ if lora_path is not None:
213
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
214
+
215
+ if partial_video_length is not None:
216
+ partial_video_length = int((partial_video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
217
+ latent_frames = (partial_video_length - 1) // vae.config.temporal_compression_ratio + 1
218
+ if partial_video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
219
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
220
+ partial_video_length += additional_frames * vae.config.temporal_compression_ratio
221
+
222
+ init_frames = 0
223
+ last_frames = init_frames + partial_video_length
224
+ while init_frames < video_length:
225
+ if last_frames >= video_length:
226
+ _partial_video_length = video_length - init_frames
227
+ _partial_video_length = int((_partial_video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1
228
+ latent_frames = (_partial_video_length - 1) // vae.config.temporal_compression_ratio + 1
229
+ if _partial_video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
230
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
231
+ _partial_video_length += additional_frames * vae.config.temporal_compression_ratio
232
+
233
+ if _partial_video_length <= 0:
234
+ break
235
+ else:
236
+ _partial_video_length = partial_video_length
237
+
238
+ input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image, None, video_length=_partial_video_length, sample_size=sample_size)
239
+
240
+ with torch.no_grad():
241
+ sample = pipeline(
242
+ prompt,
243
+ num_frames = _partial_video_length,
244
+ negative_prompt = negative_prompt,
245
+ height = sample_size[0],
246
+ width = sample_size[1],
247
+ generator = generator,
248
+ guidance_scale = guidance_scale,
249
+ num_inference_steps = num_inference_steps,
250
+
251
+ video = input_video,
252
+ mask_video = input_video_mask
253
+ ).videos
254
+
255
+ if init_frames != 0:
256
+ mix_ratio = torch.from_numpy(
257
+ np.array([float(_index) / float(overlap_video_length) for _index in range(overlap_video_length)], np.float32)
258
+ ).unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
259
+
260
+ new_sample[:, :, -overlap_video_length:] = new_sample[:, :, -overlap_video_length:] * (1 - mix_ratio) + \
261
+ sample[:, :, :overlap_video_length] * mix_ratio
262
+ new_sample = torch.cat([new_sample, sample[:, :, overlap_video_length:]], dim = 2)
263
+
264
+ sample = new_sample
265
+ else:
266
+ new_sample = sample
267
+
268
+ if last_frames >= video_length:
269
+ break
270
+
271
+ validation_image = [
272
+ Image.fromarray(
273
+ (sample[0, :, _index].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
274
+ ) for _index in range(-overlap_video_length, 0)
275
+ ]
276
+
277
+ init_frames = init_frames + _partial_video_length - overlap_video_length
278
+ last_frames = init_frames + _partial_video_length
279
+ else:
280
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
281
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
282
+ if video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
283
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
284
+ video_length += additional_frames * vae.config.temporal_compression_ratio
285
+ input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, validation_image_end, video_length=video_length, sample_size=sample_size)
286
+
287
+ with torch.no_grad():
288
+ sample = pipeline(
289
+ prompt,
290
+ num_frames = video_length,
291
+ negative_prompt = negative_prompt,
292
+ height = sample_size[0],
293
+ width = sample_size[1],
294
+ generator = generator,
295
+ guidance_scale = guidance_scale,
296
+ num_inference_steps = num_inference_steps,
297
+
298
+ video = input_video,
299
+ mask_video = input_video_mask
300
+ ).videos
301
+
302
+ if lora_path is not None:
303
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
304
+
305
+ def save_results():
306
+ if not os.path.exists(save_path):
307
+ os.makedirs(save_path, exist_ok=True)
308
+
309
+ index = len([path for path in os.listdir(save_path)]) + 1
310
+ prefix = str(index).zfill(8)
311
+ if video_length == 1:
312
+ video_path = os.path.join(save_path, prefix + ".png")
313
+
314
+ image = sample[0, :, 0]
315
+ image = image.transpose(0, 1).transpose(1, 2)
316
+ image = (image * 255).numpy().astype(np.uint8)
317
+ image = Image.fromarray(image)
318
+ image.save(video_path)
319
+ else:
320
+ video_path = os.path.join(save_path, prefix + ".mp4")
321
+ save_videos_grid(sample, video_path, fps=fps)
322
+
323
+ if ulysses_degree * ring_degree > 1:
324
+ import torch.distributed as dist
325
+ if dist.get_rank() == 0:
326
+ save_results()
327
+ else:
328
+ save_results()
vendor/VideoX-Fun/examples/cogvideox_fun/predict_t2v.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import (CogVideoXDDIMScheduler, DDIMScheduler,
7
+ DPMSolverMultistepScheduler,
8
+ EulerAncestralDiscreteScheduler, EulerDiscreteScheduler,
9
+ PNDMScheduler)
10
+ from PIL import Image
11
+ from transformers import T5EncoderModel
12
+
13
+ current_file_path = os.path.abspath(__file__)
14
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
15
+ for project_root in project_roots:
16
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
17
+
18
+ from videox_fun.models import (AutoencoderKLCogVideoX,
19
+ CogVideoXTransformer3DModel, T5EncoderModel,
20
+ T5Tokenizer)
21
+ from videox_fun.pipeline import (CogVideoXFunPipeline,
22
+ CogVideoXFunInpaintPipeline)
23
+ from videox_fun.utils import (register_auto_device_hook,
24
+ safe_enable_group_offload)
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
29
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
30
+
31
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
32
+ # model_full_load means that the entire model will be moved to the GPU.
33
+ #
34
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
35
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
36
+ #
37
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
38
+ #
39
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
40
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
41
+ #
42
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
43
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
44
+ #
45
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
46
+ # resulting in slower speeds but saving a large amount of GPU memory.
47
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
48
+ # Multi GPUs config
49
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
50
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
51
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
52
+ ulysses_degree = 1
53
+ ring_degree = 1
54
+ # Use FSDP to save more GPU memory in multi gpus.
55
+ fsdp_dit = False
56
+ fsdp_text_encoder = True
57
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
58
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
59
+ compile_dit = False
60
+
61
+ # model path
62
+ model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
63
+
64
+ # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" "DDIM_Cog" and "DDIM_Origin"
65
+ sampler_name = "DDIM_Origin"
66
+
67
+ # Load pretrained model if need
68
+ transformer_path = None
69
+ vae_path = None
70
+ lora_path = None
71
+
72
+ # Other params
73
+ sample_size = [384, 672]
74
+ # V1.0 and V1.1 support up to 49 frames of video generation,
75
+ # while V1.5 supports up to 85 frames.
76
+ video_length = 49
77
+ fps = 8
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ prompt = "A young woman with beautiful and clear eyes and blonde hair standing and white dress in a forest wearing a crown. She seems to be lost in thought, and the camera focuses on her face. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
83
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
84
+ guidance_scale = 6.0
85
+ seed = 43
86
+ num_inference_steps = 50
87
+ lora_weight = 0.55
88
+ save_path = "samples/cogvideox-fun-videos-t2v"
89
+
90
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
91
+
92
+ transformer = CogVideoXTransformer3DModel.from_pretrained(
93
+ model_name,
94
+ subfolder="transformer",
95
+ low_cpu_mem_usage=True,
96
+ torch_dtype=weight_dtype,
97
+ ).to(weight_dtype)
98
+
99
+ if transformer_path is not None:
100
+ print(f"From checkpoint: {transformer_path}")
101
+ if transformer_path.endswith("safetensors"):
102
+ from safetensors.torch import load_file, safe_open
103
+ state_dict = load_file(transformer_path)
104
+ else:
105
+ state_dict = torch.load(transformer_path, map_location="cpu")
106
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
107
+
108
+ m, u = transformer.load_state_dict(state_dict, strict=False)
109
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
110
+
111
+ # Get Vae
112
+ vae = AutoencoderKLCogVideoX.from_pretrained(
113
+ model_name,
114
+ subfolder="vae"
115
+ ).to(weight_dtype)
116
+
117
+ if vae_path is not None:
118
+ print(f"From checkpoint: {vae_path}")
119
+ if vae_path.endswith("safetensors"):
120
+ from safetensors.torch import load_file, safe_open
121
+ state_dict = load_file(vae_path)
122
+ else:
123
+ state_dict = torch.load(vae_path, map_location="cpu")
124
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
125
+
126
+ m, u = vae.load_state_dict(state_dict, strict=False)
127
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
128
+
129
+ # Get tokenizer and text_encoder
130
+ tokenizer = T5Tokenizer.from_pretrained(
131
+ model_name, subfolder="tokenizer"
132
+ )
133
+ text_encoder = T5EncoderModel.from_pretrained(
134
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
135
+ )
136
+
137
+ # Get Scheduler
138
+ Chosen_Scheduler = scheduler_dict = {
139
+ "Euler": EulerDiscreteScheduler,
140
+ "Euler A": EulerAncestralDiscreteScheduler,
141
+ "DPM++": DPMSolverMultistepScheduler,
142
+ "PNDM": PNDMScheduler,
143
+ "DDIM_Cog": CogVideoXDDIMScheduler,
144
+ "DDIM_Origin": DDIMScheduler,
145
+ }[sampler_name]
146
+ scheduler = Chosen_Scheduler.from_pretrained(
147
+ model_name,
148
+ subfolder="scheduler"
149
+ )
150
+
151
+ if transformer.config.in_channels != vae.config.latent_channels:
152
+ pipeline = CogVideoXFunInpaintPipeline(
153
+ vae=vae,
154
+ tokenizer=tokenizer,
155
+ text_encoder=text_encoder,
156
+ transformer=transformer,
157
+ scheduler=scheduler,
158
+ )
159
+ else:
160
+ pipeline = CogVideoXFunPipeline(
161
+ vae=vae,
162
+ tokenizer=tokenizer,
163
+ text_encoder=text_encoder,
164
+ transformer=transformer,
165
+ scheduler=scheduler,
166
+ )
167
+ if ulysses_degree > 1 or ring_degree > 1:
168
+ from functools import partial
169
+ transformer.enable_multi_gpus_inference()
170
+ if fsdp_dit:
171
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
172
+ pipeline.transformer = shard_fn(pipeline.transformer)
173
+ print("Add FSDP DIT")
174
+ if fsdp_text_encoder:
175
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
176
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
177
+ print("Add FSDP TEXT ENCODER")
178
+
179
+ if compile_dit:
180
+ for i in range(len(pipeline.transformer.transformer_blocks)):
181
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
182
+ print("Add Compile")
183
+
184
+ if GPU_memory_mode == "sequential_cpu_offload":
185
+ pipeline.enable_sequential_cpu_offload(device=device)
186
+ elif GPU_memory_mode == "model_group_offload":
187
+ register_auto_device_hook(pipeline.transformer)
188
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
189
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
190
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
191
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
192
+ pipeline.enable_model_cpu_offload(device=device)
193
+ elif GPU_memory_mode == "model_cpu_offload":
194
+ pipeline.enable_model_cpu_offload(device=device)
195
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
196
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
197
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
198
+ pipeline.to(device=device)
199
+ else:
200
+ pipeline.to(device=device)
201
+
202
+ generator = torch.Generator(device=device).manual_seed(seed)
203
+
204
+ if lora_path is not None:
205
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
206
+
207
+ with torch.no_grad():
208
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
209
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
210
+ if video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
211
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
212
+ video_length += additional_frames * vae.config.temporal_compression_ratio
213
+
214
+ if transformer.config.in_channels != vae.config.latent_channels:
215
+ input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size)
216
+
217
+ sample = pipeline(
218
+ prompt,
219
+ num_frames = video_length,
220
+ negative_prompt = negative_prompt,
221
+ height = sample_size[0],
222
+ width = sample_size[1],
223
+ generator = generator,
224
+ guidance_scale = guidance_scale,
225
+ num_inference_steps = num_inference_steps,
226
+
227
+ video = input_video,
228
+ mask_video = input_video_mask,
229
+ ).videos
230
+ else:
231
+ sample = pipeline(
232
+ prompt,
233
+ num_frames = video_length,
234
+ negative_prompt = negative_prompt,
235
+ height = sample_size[0],
236
+ width = sample_size[1],
237
+ generator = generator,
238
+ guidance_scale = guidance_scale,
239
+ num_inference_steps = num_inference_steps,
240
+ ).videos
241
+
242
+ if lora_path is not None:
243
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
244
+
245
+ def save_results():
246
+ if not os.path.exists(save_path):
247
+ os.makedirs(save_path, exist_ok=True)
248
+
249
+ index = len([path for path in os.listdir(save_path)]) + 1
250
+ prefix = str(index).zfill(8)
251
+ if video_length == 1:
252
+ video_path = os.path.join(save_path, prefix + ".png")
253
+
254
+ image = sample[0, :, 0]
255
+ image = image.transpose(0, 1).transpose(1, 2)
256
+ image = (image * 255).numpy().astype(np.uint8)
257
+ image = Image.fromarray(image)
258
+ image.save(video_path)
259
+ else:
260
+ video_path = os.path.join(save_path, prefix + ".mp4")
261
+ save_videos_grid(sample, video_path, fps=fps)
262
+
263
+ if ulysses_degree * ring_degree > 1:
264
+ import torch.distributed as dist
265
+ if dist.get_rank() == 0:
266
+ save_results()
267
+ else:
268
+ save_results()
vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import (CogVideoXDDIMScheduler, DDIMScheduler,
7
+ DPMSolverMultistepScheduler,
8
+ EulerAncestralDiscreteScheduler, EulerDiscreteScheduler,
9
+ PNDMScheduler)
10
+ from PIL import Image
11
+
12
+ current_file_path = os.path.abspath(__file__)
13
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
14
+ for project_root in project_roots:
15
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
16
+
17
+ from videox_fun.models import (AutoencoderKLCogVideoX,
18
+ CogVideoXTransformer3DModel, T5EncoderModel,
19
+ T5Tokenizer)
20
+ from videox_fun.pipeline import (CogVideoXFunPipeline,
21
+ CogVideoXFunInpaintPipeline)
22
+ from videox_fun.utils import (register_auto_device_hook,
23
+ safe_enable_group_offload)
24
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
28
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
29
+
30
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
31
+ # model_full_load means that the entire model will be moved to the GPU.
32
+ #
33
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
37
+ #
38
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
42
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
43
+ #
44
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
45
+ # resulting in slower speeds but saving a large amount of GPU memory.
46
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
47
+ # Multi GPUs config
48
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
49
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
50
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
51
+ ulysses_degree = 1
52
+ ring_degree = 1
53
+ # Use FSDP to save more GPU memory in multi gpus.
54
+ fsdp_dit = False
55
+ fsdp_text_encoder = True
56
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
57
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
58
+ compile_dit = False
59
+
60
+ # model path
61
+ model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-InP"
62
+
63
+ # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" "DDIM_Cog" and "DDIM_Origin"
64
+ sampler_name = "DDIM_Origin"
65
+
66
+ # Load pretrained model if need
67
+ transformer_path = None
68
+ vae_path = None
69
+ lora_path = None
70
+ # Other params
71
+ sample_size = [384, 672]
72
+ # V1.0 and V1.1 support up to 49 frames of video generation,
73
+ # while V1.5 supports up to 85 frames.
74
+ video_length = 49
75
+ fps = 8
76
+
77
+ # Use torch.float16 if GPU does not support torch.bfloat16
78
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
79
+ weight_dtype = torch.bfloat16
80
+ # If you are preparing to redraw the reference video, set validation_video and validation_video_mask.
81
+ # If you do not use validation_video_mask, the entire video will be redrawn;
82
+ # if you use validation_video_mask, only a portion of the video will be redrawn.
83
+ # Please set a larger denoise_strength when using validation_video_mask, such as 1.00 instead of 0.70
84
+ validation_video = "asset/1.mp4"
85
+ validation_video_mask = None
86
+ denoise_strength = 0.70
87
+
88
+ # prompts
89
+ prompt = "A cute cat is playing the guitar. "
90
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
91
+ guidance_scale = 6.0
92
+ seed = 43
93
+ num_inference_steps = 50
94
+ lora_weight = 0.55
95
+ save_path = "samples/cogvideox-fun-videos_v2v"
96
+
97
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
98
+
99
+ transformer = CogVideoXTransformer3DModel.from_pretrained(
100
+ model_name,
101
+ subfolder="transformer",
102
+ low_cpu_mem_usage=True,
103
+ torch_dtype=weight_dtype,
104
+ ).to(weight_dtype)
105
+
106
+ if transformer_path is not None:
107
+ print(f"From checkpoint: {transformer_path}")
108
+ if transformer_path.endswith("safetensors"):
109
+ from safetensors.torch import load_file, safe_open
110
+ state_dict = load_file(transformer_path)
111
+ else:
112
+ state_dict = torch.load(transformer_path, map_location="cpu")
113
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
114
+
115
+ m, u = transformer.load_state_dict(state_dict, strict=False)
116
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
117
+
118
+ # Get Vae
119
+ vae = AutoencoderKLCogVideoX.from_pretrained(
120
+ model_name,
121
+ subfolder="vae"
122
+ ).to(weight_dtype)
123
+
124
+ if vae_path is not None:
125
+ print(f"From checkpoint: {vae_path}")
126
+ if vae_path.endswith("safetensors"):
127
+ from safetensors.torch import load_file, safe_open
128
+ state_dict = load_file(vae_path)
129
+ else:
130
+ state_dict = torch.load(vae_path, map_location="cpu")
131
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
132
+
133
+ m, u = vae.load_state_dict(state_dict, strict=False)
134
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
135
+
136
+ # Get tokenizer and text_encoder
137
+ tokenizer = T5Tokenizer.from_pretrained(
138
+ model_name, subfolder="tokenizer"
139
+ )
140
+ text_encoder = T5EncoderModel.from_pretrained(
141
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
142
+ )
143
+
144
+ # Get Scheduler
145
+ Chosen_Scheduler = scheduler_dict = {
146
+ "Euler": EulerDiscreteScheduler,
147
+ "Euler A": EulerAncestralDiscreteScheduler,
148
+ "DPM++": DPMSolverMultistepScheduler,
149
+ "PNDM": PNDMScheduler,
150
+ "DDIM_Cog": CogVideoXDDIMScheduler,
151
+ "DDIM_Origin": DDIMScheduler,
152
+ }[sampler_name]
153
+ scheduler = Chosen_Scheduler.from_pretrained(
154
+ model_name,
155
+ subfolder="scheduler"
156
+ )
157
+
158
+ if transformer.config.in_channels != vae.config.latent_channels:
159
+ pipeline = CogVideoXFunInpaintPipeline(
160
+ vae=vae,
161
+ tokenizer=tokenizer,
162
+ text_encoder=text_encoder,
163
+ transformer=transformer,
164
+ scheduler=scheduler,
165
+ )
166
+ else:
167
+ pipeline = CogVideoXFunPipeline(
168
+ vae=vae,
169
+ tokenizer=tokenizer,
170
+ text_encoder=text_encoder,
171
+ transformer=transformer,
172
+ scheduler=scheduler,
173
+ )
174
+ if ulysses_degree > 1 or ring_degree > 1:
175
+ from functools import partial
176
+ transformer.enable_multi_gpus_inference()
177
+ if fsdp_dit:
178
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
179
+ pipeline.transformer = shard_fn(pipeline.transformer)
180
+ print("Add FSDP DIT")
181
+ if fsdp_text_encoder:
182
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
183
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
184
+ print("Add FSDP TEXT ENCODER")
185
+
186
+ if compile_dit:
187
+ for i in range(len(pipeline.transformer.transformer_blocks)):
188
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
189
+ print("Add Compile")
190
+
191
+ if GPU_memory_mode == "sequential_cpu_offload":
192
+ pipeline.enable_sequential_cpu_offload(device=device)
193
+ elif GPU_memory_mode == "model_group_offload":
194
+ register_auto_device_hook(pipeline.transformer)
195
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
196
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
197
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
198
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
199
+ pipeline.enable_model_cpu_offload(device=device)
200
+ elif GPU_memory_mode == "model_cpu_offload":
201
+ pipeline.enable_model_cpu_offload(device=device)
202
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
203
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
204
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
205
+ pipeline.to(device=device)
206
+ else:
207
+ pipeline.to(device=device)
208
+
209
+ generator = torch.Generator(device=device).manual_seed(seed)
210
+
211
+ if lora_path is not None:
212
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
213
+
214
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
215
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
216
+ if video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
217
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
218
+ video_length += additional_frames * vae.config.temporal_compression_ratio
219
+ input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(validation_video, video_length=video_length, sample_size=sample_size, validation_video_mask=validation_video_mask, fps=fps)
220
+
221
+ with torch.no_grad():
222
+ sample = pipeline(
223
+ prompt,
224
+ num_frames = video_length,
225
+ negative_prompt = negative_prompt,
226
+ height = sample_size[0],
227
+ width = sample_size[1],
228
+ generator = generator,
229
+ guidance_scale = guidance_scale,
230
+ num_inference_steps = num_inference_steps,
231
+
232
+ video = input_video,
233
+ mask_video = input_video_mask,
234
+ strength = denoise_strength,
235
+ ).videos
236
+
237
+ if lora_path is not None:
238
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
239
+
240
+ def save_results():
241
+ if not os.path.exists(save_path):
242
+ os.makedirs(save_path, exist_ok=True)
243
+
244
+ index = len([path for path in os.listdir(save_path)]) + 1
245
+ prefix = str(index).zfill(8)
246
+ if video_length == 1:
247
+ video_path = os.path.join(save_path, prefix + ".png")
248
+
249
+ image = sample[0, :, 0]
250
+ image = image.transpose(0, 1).transpose(1, 2)
251
+ image = (image * 255).numpy().astype(np.uint8)
252
+ image = Image.fromarray(image)
253
+ image.save(video_path)
254
+ else:
255
+ video_path = os.path.join(save_path, prefix + ".mp4")
256
+ save_videos_grid(sample, video_path, fps=fps)
257
+
258
+ if ulysses_degree * ring_degree > 1:
259
+ import torch.distributed as dist
260
+ if dist.get_rank() == 0:
261
+ save_results()
262
+ else:
263
+ save_results()
vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v_control.py ADDED
@@ -0,0 +1,248 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import cv2
5
+ import numpy as np
6
+ import torch
7
+ from diffusers import (CogVideoXDDIMScheduler, DDIMScheduler,
8
+ DPMSolverMultistepScheduler,
9
+ EulerAncestralDiscreteScheduler, EulerDiscreteScheduler,
10
+ PNDMScheduler)
11
+ from PIL import Image
12
+ from transformers import T5EncoderModel
13
+
14
+ current_file_path = os.path.abspath(__file__)
15
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
16
+ for project_root in project_roots:
17
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
18
+
19
+ from videox_fun.models import (AutoencoderKLCogVideoX,
20
+ CogVideoXTransformer3DModel, T5EncoderModel,
21
+ T5Tokenizer)
22
+ from videox_fun.pipeline import (CogVideoXFunControlPipeline,
23
+ CogVideoXFunInpaintPipeline)
24
+ from videox_fun.utils import (register_auto_device_hook,
25
+ safe_enable_group_offload)
26
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
27
+ convert_weight_dtype_wrapper)
28
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
29
+ from videox_fun.utils.utils import get_video_to_video_latent, save_videos_grid
30
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
31
+
32
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
33
+ # model_full_load means that the entire model will be moved to the GPU.
34
+ #
35
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
39
+ #
40
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
44
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
45
+ #
46
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
47
+ # resulting in slower speeds but saving a large amount of GPU memory.
48
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
49
+ # Multi GPUs config
50
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
51
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
52
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
53
+ ulysses_degree = 1
54
+ ring_degree = 1
55
+ # Use FSDP to save more GPU memory in multi gpus.
56
+ fsdp_dit = False
57
+ fsdp_text_encoder = True
58
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
59
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
60
+ compile_dit = False
61
+
62
+ # model path
63
+ model_name = "models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose"
64
+
65
+ # Choose the sampler in "Euler" "Euler A" "DPM++" "PNDM" "DDIM_Cog" and "DDIM_Origin"
66
+ sampler_name = "DDIM_Origin"
67
+
68
+ # Load pretrained model if need
69
+ transformer_path = None
70
+ vae_path = None
71
+ lora_path = None
72
+ # Other params
73
+ sample_size = [672, 384]
74
+ # V1.0 and V1.1 support up to 49 frames of video generation,
75
+ # while V1.5 supports up to 85 frames.
76
+ video_length = 49
77
+ fps = 8
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ control_video = "asset/pose.mp4"
83
+
84
+ # prompts
85
+ prompt = "A young woman with beautiful face, dressed in white, is moving her body. "
86
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
87
+ guidance_scale = 6.0
88
+ seed = 43
89
+ num_inference_steps = 50
90
+ lora_weight = 0.55
91
+ save_path = "samples/cogvideox-fun-videos_control"
92
+
93
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
94
+
95
+ transformer = CogVideoXTransformer3DModel.from_pretrained(
96
+ model_name,
97
+ subfolder="transformer",
98
+ low_cpu_mem_usage=True,
99
+ torch_dtype=weight_dtype,
100
+ ).to(weight_dtype)
101
+
102
+ if transformer_path is not None:
103
+ print(f"From checkpoint: {transformer_path}")
104
+ if transformer_path.endswith("safetensors"):
105
+ from safetensors.torch import load_file, safe_open
106
+ state_dict = load_file(transformer_path)
107
+ else:
108
+ state_dict = torch.load(transformer_path, map_location="cpu")
109
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
110
+
111
+ m, u = transformer.load_state_dict(state_dict, strict=False)
112
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
113
+
114
+ # Get Vae
115
+ vae = AutoencoderKLCogVideoX.from_pretrained(
116
+ model_name,
117
+ subfolder="vae"
118
+ ).to(weight_dtype)
119
+
120
+ if vae_path is not None:
121
+ print(f"From checkpoint: {vae_path}")
122
+ if vae_path.endswith("safetensors"):
123
+ from safetensors.torch import load_file, safe_open
124
+ state_dict = load_file(vae_path)
125
+ else:
126
+ state_dict = torch.load(vae_path, map_location="cpu")
127
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
128
+
129
+ m, u = vae.load_state_dict(state_dict, strict=False)
130
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
131
+
132
+ # Get tokenizer and text_encoder
133
+ tokenizer = T5Tokenizer.from_pretrained(
134
+ model_name, subfolder="tokenizer"
135
+ )
136
+ text_encoder = T5EncoderModel.from_pretrained(
137
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
138
+ )
139
+
140
+ # Get Scheduler
141
+ Chosen_Scheduler = scheduler_dict = {
142
+ "Euler": EulerDiscreteScheduler,
143
+ "Euler A": EulerAncestralDiscreteScheduler,
144
+ "DPM++": DPMSolverMultistepScheduler,
145
+ "PNDM": PNDMScheduler,
146
+ "DDIM_Cog": CogVideoXDDIMScheduler,
147
+ "DDIM_Origin": DDIMScheduler,
148
+ }[sampler_name]
149
+ scheduler = Chosen_Scheduler.from_pretrained(
150
+ model_name,
151
+ subfolder="scheduler"
152
+ )
153
+
154
+ pipeline = CogVideoXFunControlPipeline(
155
+ vae=vae,
156
+ tokenizer=tokenizer,
157
+ text_encoder=text_encoder,
158
+ transformer=transformer,
159
+ scheduler=scheduler,
160
+ )
161
+ if ulysses_degree > 1 or ring_degree > 1:
162
+ from functools import partial
163
+ transformer.enable_multi_gpus_inference()
164
+ if fsdp_dit:
165
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
166
+ pipeline.transformer = shard_fn(pipeline.transformer)
167
+ print("Add FSDP DIT")
168
+ if fsdp_text_encoder:
169
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
170
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
171
+ print("Add FSDP TEXT ENCODER")
172
+
173
+ if compile_dit:
174
+ for i in range(len(pipeline.transformer.transformer_blocks)):
175
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
176
+ print("Add Compile")
177
+
178
+ if GPU_memory_mode == "sequential_cpu_offload":
179
+ pipeline.enable_sequential_cpu_offload(device=device)
180
+ elif GPU_memory_mode == "model_group_offload":
181
+ register_auto_device_hook(pipeline.transformer)
182
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
183
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
184
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
185
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
186
+ pipeline.enable_model_cpu_offload(device=device)
187
+ elif GPU_memory_mode == "model_cpu_offload":
188
+ pipeline.enable_model_cpu_offload(device=device)
189
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
190
+ convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
191
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
192
+ pipeline.to(device=device)
193
+ else:
194
+ pipeline.to(device=device)
195
+
196
+ generator = torch.Generator(device=device).manual_seed(seed)
197
+
198
+ if lora_path is not None:
199
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
200
+
201
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
202
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
203
+ if video_length != 1 and transformer.config.patch_size_t is not None and latent_frames % transformer.config.patch_size_t != 0:
204
+ additional_frames = transformer.config.patch_size_t - latent_frames % transformer.config.patch_size_t
205
+ video_length += additional_frames * vae.config.temporal_compression_ratio
206
+ input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps)
207
+
208
+ with torch.no_grad():
209
+ sample = pipeline(
210
+ prompt,
211
+ num_frames = video_length,
212
+ negative_prompt = negative_prompt,
213
+ height = sample_size[0],
214
+ width = sample_size[1],
215
+ generator = generator,
216
+ guidance_scale = guidance_scale,
217
+ num_inference_steps = num_inference_steps,
218
+
219
+ control_video = input_video,
220
+ ).videos
221
+
222
+ if lora_path is not None:
223
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
224
+
225
+ def save_results():
226
+ if not os.path.exists(save_path):
227
+ os.makedirs(save_path, exist_ok=True)
228
+
229
+ index = len([path for path in os.listdir(save_path)]) + 1
230
+ prefix = str(index).zfill(8)
231
+ if video_length == 1:
232
+ video_path = os.path.join(save_path, prefix + ".png")
233
+
234
+ image = sample[0, :, 0]
235
+ image = image.transpose(0, 1).transpose(1, 2)
236
+ image = (image * 255).numpy().astype(np.uint8)
237
+ image = Image.fromarray(image)
238
+ image.save(video_path)
239
+ else:
240
+ video_path = os.path.join(save_path, prefix + ".mp4")
241
+ save_videos_grid(sample, video_path, fps=fps)
242
+
243
+ if ulysses_degree * ring_degree > 1:
244
+ import torch.distributed as dist
245
+ if dist.get_rank() == 0:
246
+ save_results()
247
+ else:
248
+ save_results()
vendor/VideoX-Fun/examples/ernie_image/predict_t2i.py ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import torch
5
+ from diffusers import FlowMatchEulerDiscreteScheduler
6
+
7
+ current_file_path = os.path.abspath(__file__)
8
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
9
+ for project_root in project_roots:
10
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
11
+
12
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
13
+ from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer,
14
+ ErnieImageTransformer2DModel, Mistral3Model)
15
+ from videox_fun.pipeline import ErnieImagePipeline
16
+ from videox_fun.utils import (register_auto_device_hook,
17
+ safe_enable_group_offload)
18
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
19
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
20
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
21
+ convert_weight_dtype_wrapper)
22
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
23
+
24
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
25
+ # model_full_load means that the entire model will be moved to the GPU.
26
+ #
27
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
28
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
29
+ #
30
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
31
+ #
32
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
33
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
34
+ #
35
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
36
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
37
+ #
38
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
39
+ # resulting in slower speeds but saving a large amount of GPU memory.
40
+ GPU_memory_mode = "model_cpu_offload"
41
+ # Multi GPUs config
42
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
43
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
44
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
45
+ ulysses_degree = 1
46
+ ring_degree = 1
47
+ # Use FSDP to save more GPU memory in multi gpus.
48
+ fsdp_dit = False
49
+ fsdp_text_encoder = False
50
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
51
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
52
+ compile_dit = False
53
+
54
+ # model path
55
+ model_name = "models/Diffusion_Transformer/ERNIE-Image"
56
+
57
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
58
+ sampler_name = "Flow"
59
+
60
+ # Load pretrained model if need
61
+ transformer_path = None
62
+ vae_path = None
63
+ lora_path = None
64
+
65
+ # Other params
66
+ sample_size = [1728, 992]
67
+
68
+ # Use torch.float16 if GPU does not support torch.bfloat16
69
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
70
+ weight_dtype = torch.bfloat16
71
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
72
+ negative_prompt = "低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"
73
+ guidance_scale = 4.5
74
+ seed = 43
75
+ num_inference_steps = 40
76
+ lora_weight = 0.55
77
+ save_path = "samples/ernie-image-t2i"
78
+
79
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
80
+
81
+ transformer = ErnieImageTransformer2DModel.from_pretrained(
82
+ model_name,
83
+ subfolder="transformer",
84
+ low_cpu_mem_usage=True,
85
+ torch_dtype=weight_dtype,
86
+ ).to(weight_dtype)
87
+
88
+ if transformer_path is not None:
89
+ print(f"From checkpoint: {transformer_path}")
90
+ if transformer_path.endswith("safetensors"):
91
+ from safetensors.torch import load_file, safe_open
92
+ state_dict = load_file(transformer_path)
93
+ else:
94
+ state_dict = torch.load(transformer_path, map_location="cpu")
95
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
96
+
97
+ m, u = transformer.load_state_dict(state_dict, strict=False)
98
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
99
+
100
+ # Get Vae
101
+ vae = AutoencoderKLFlux2.from_pretrained(
102
+ model_name,
103
+ subfolder="vae"
104
+ ).to(weight_dtype)
105
+
106
+ if vae_path is not None:
107
+ print(f"From checkpoint: {vae_path}")
108
+ if vae_path.endswith("safetensors"):
109
+ from safetensors.torch import load_file, safe_open
110
+ state_dict = load_file(vae_path)
111
+ else:
112
+ state_dict = torch.load(vae_path, map_location="cpu")
113
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
114
+
115
+ m, u = vae.load_state_dict(state_dict, strict=False)
116
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
117
+
118
+ # Get tokenizer and text_encoder
119
+ tokenizer = AutoTokenizer.from_pretrained(
120
+ model_name, subfolder="tokenizer"
121
+ )
122
+ text_encoder = Mistral3Model.from_pretrained(
123
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
124
+ )
125
+
126
+ # Get Scheduler
127
+ Chosen_Scheduler = scheduler_dict = {
128
+ "Flow": FlowMatchEulerDiscreteScheduler,
129
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
130
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
131
+ }[sampler_name]
132
+ scheduler = Chosen_Scheduler.from_pretrained(
133
+ model_name,
134
+ subfolder="scheduler"
135
+ )
136
+
137
+ pipeline = ErnieImagePipeline(
138
+ vae=vae,
139
+ tokenizer=tokenizer,
140
+ text_encoder=text_encoder,
141
+ transformer=transformer,
142
+ scheduler=scheduler,
143
+ )
144
+
145
+ if ulysses_degree > 1 or ring_degree > 1:
146
+ from functools import partial
147
+ transformer.enable_multi_gpus_inference()
148
+ if fsdp_dit:
149
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.layers))
150
+ pipeline.transformer = shard_fn(pipeline.transformer)
151
+ print("Add FSDP DIT")
152
+
153
+ if compile_dit:
154
+ for i in range(len(pipeline.transformer.layers)):
155
+ pipeline.transformer.layers[i] = torch.compile(pipeline.transformer.layers[i])
156
+ print("Add Compile")
157
+
158
+ if GPU_memory_mode == "sequential_cpu_offload":
159
+ pipeline.enable_sequential_cpu_offload(device=device)
160
+ elif GPU_memory_mode == "model_group_offload":
161
+ register_auto_device_hook(pipeline.transformer)
162
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
163
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
164
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
165
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
166
+ pipeline.enable_model_cpu_offload(device=device)
167
+ elif GPU_memory_mode == "model_cpu_offload":
168
+ pipeline.enable_model_cpu_offload(device=device)
169
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
170
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
171
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
172
+ pipeline.to(device=device)
173
+ else:
174
+ pipeline.to(device=device)
175
+
176
+ generator = torch.Generator(device=device).manual_seed(seed)
177
+
178
+ if lora_path is not None:
179
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
180
+
181
+ with torch.no_grad():
182
+ sample = pipeline(
183
+ prompt,
184
+ negative_prompt = negative_prompt,
185
+ height = sample_size[0],
186
+ width = sample_size[1],
187
+ generator = generator,
188
+ guidance_scale = guidance_scale,
189
+ num_inference_steps = num_inference_steps,
190
+ ).images
191
+
192
+ if lora_path is not None:
193
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
194
+
195
+ def save_results():
196
+ if not os.path.exists(save_path):
197
+ os.makedirs(save_path, exist_ok=True)
198
+
199
+ index = len([path for path in os.listdir(save_path)]) + 1
200
+ prefix = str(index).zfill(8)
201
+ video_path = os.path.join(save_path, prefix + ".png")
202
+ image = sample[0]
203
+ image.save(video_path)
204
+
205
+ if ulysses_degree * ring_degree > 1:
206
+ import torch.distributed as dist
207
+ if dist.get_rank() == 0:
208
+ save_results()
209
+ else:
210
+ save_results()
vendor/VideoX-Fun/examples/fantasytalking/predict_s2v.py ADDED
@@ -0,0 +1,335 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
17
+ AutoTokenizer, CLIPModel,
18
+ FantasyTalkingTransformer3DModel, FantasyTalkingAudioEncoder,
19
+ WanT5EncoderModel)
20
+ from videox_fun.models.cache_utils import get_teacache_coefficients
21
+ from videox_fun.pipeline import FantasyTalkingPipeline
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils import (register_auto_device_hook,
25
+ safe_enable_group_offload)
26
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
27
+ convert_weight_dtype_wrapper,
28
+ replace_parameters_by_name)
29
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
30
+ from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
31
+ get_image_to_video_latent,
32
+ get_video_to_video_latent,
33
+ merge_video_audio, save_videos_grid)
34
+
35
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
36
+ # model_full_load means that the entire model will be moved to the GPU.
37
+ #
38
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
42
+ #
43
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
44
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
45
+ #
46
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
47
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
48
+ #
49
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
50
+ # resulting in slower speeds but saving a large amount of GPU memory.
51
+ GPU_memory_mode = "sequential_cpu_offload"
52
+ # Multi GPUs config
53
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
54
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
55
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
56
+ ulysses_degree = 1
57
+ ring_degree = 1
58
+ # Use FSDP to save more GPU memory in multi gpus.
59
+ fsdp_dit = False
60
+ fsdp_text_encoder = True
61
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
62
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
63
+ compile_dit = False
64
+
65
+ # TeaCache config
66
+ enable_teacache = True
67
+ # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
68
+ # but it may cause slight differences between the generated content and the original content.
69
+ # # --------------------------------------------------------------------------------------------------- #
70
+ # | Model Name | threshold | Model Name | threshold |
71
+ # | Wan2.2-T2V-A14B | 0.10~0.15 | Wan2.2-I2V-A14B | 0.15~0.20 |
72
+ # # --------------------------------------------------------------------------------------------------- #
73
+ teacache_threshold = 0.10
74
+ # The number of steps to skip TeaCache at the beginning of the inference process, which can
75
+ # reduce the impact of TeaCache on generated video quality.
76
+ num_skip_start_steps = 5
77
+ # Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
78
+ teacache_offload = False
79
+
80
+ # Riflex config
81
+ enable_riflex = False
82
+ # Index of intrinsic frequency
83
+ riflex_k = 6
84
+
85
+ # Config and model path
86
+ config_path = "config/wan2.1/wan_civitai.yaml"
87
+ # model path
88
+ # Please Download https://modelscope.cn/models/AI-ModelScope/wav2vec2-base-960h/summary
89
+ # to models/Diffusion_Transformer/wav2vec2-base-960h for encoding audio.
90
+ model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
91
+ # audio encoder model path. If None, will use os.path.join(model_name, "audio_encoder")
92
+ model_name_audio = "models/Diffusion_Transformer/wav2vec2-base-960h"
93
+
94
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
95
+ sampler_name = "Flow"
96
+ # [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
97
+ # Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
98
+ shift = 5
99
+
100
+ # Load pretrained model if need
101
+ # The transformer_path is used for low noise model, the transformer_high_path is used for high noise model.
102
+ # The fantasytalking_model.ckpt can be downloaded in https://www.modelscope.cn/models/amap_cvlab/FantasyTalking/
103
+ transformer_path = "models/Personalized_Model/FantasyTalking/fantasytalking_model.ckpt"
104
+ vae_path = None
105
+ # Load lora model if need
106
+ lora_path = None
107
+
108
+ # Other params
109
+ sample_size = [832, 480]
110
+ video_length = 81
111
+ fps = 23
112
+
113
+ # Use torch.float16 if GPU does not support torch.bfloat16
114
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
115
+ weight_dtype = torch.bfloat16
116
+ # If you want to generate from text, please set the validation_image_start = None
117
+ validation_image_start = "asset/8.png"
118
+ audio_path = "asset/talk.wav"
119
+
120
+ # prompts
121
+ prompt = "一个女孩在海边说话。"
122
+ negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
123
+ guidance_scale = 4.5
124
+ audio_guide_scale = 4.0
125
+ seed = 43
126
+ num_inference_steps = 40
127
+ lora_weight = 0.55
128
+ save_path = "samples/fantasy-talking-videos-speech2v"
129
+
130
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
131
+ config = OmegaConf.load(config_path)
132
+
133
+ transformer = FantasyTalkingTransformer3DModel.from_pretrained(
134
+ os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
135
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
136
+ low_cpu_mem_usage=True,
137
+ torch_dtype=weight_dtype,
138
+ )
139
+
140
+ if transformer_path is not None:
141
+ print(f"From checkpoint: {transformer_path}")
142
+ if transformer_path.endswith("safetensors"):
143
+ from safetensors.torch import load_file, safe_open
144
+ state_dict = load_file(transformer_path)
145
+ else:
146
+ state_dict = torch.load(transformer_path, map_location="cpu")
147
+
148
+ if "audio_processor" in state_dict:
149
+ audio_processor_dict = state_dict["audio_processor"] if "audio_processor" in state_dict else state_dict
150
+ m, u = transformer.load_state_dict(audio_processor_dict, strict=False)
151
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
152
+
153
+ proj_model_dict = state_dict["proj_model"] if "proj_model" in state_dict else state_dict
154
+ proj_model_dict = {"proj_model." + k : v for k, v in proj_model_dict.items()}
155
+ m, u = transformer.load_state_dict(proj_model_dict, strict=False)
156
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
157
+ else:
158
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
159
+
160
+ m, u = transformer.load_state_dict(state_dict, strict=False)
161
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
162
+
163
+ # Get Vae
164
+ Chosen_AutoencoderKL = {
165
+ "AutoencoderKLWan": AutoencoderKLWan,
166
+ "AutoencoderKLWan3_8": AutoencoderKLWan3_8
167
+ }[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
168
+ vae = Chosen_AutoencoderKL.from_pretrained(
169
+ os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
170
+ additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
171
+ ).to(weight_dtype)
172
+
173
+ if vae_path is not None:
174
+ print(f"From checkpoint: {vae_path}")
175
+ if vae_path.endswith("safetensors"):
176
+ from safetensors.torch import load_file, safe_open
177
+ state_dict = load_file(vae_path)
178
+ else:
179
+ state_dict = torch.load(vae_path, map_location="cpu")
180
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
181
+
182
+ m, u = vae.load_state_dict(state_dict, strict=False)
183
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
184
+
185
+ # Get Tokenizer
186
+ tokenizer = AutoTokenizer.from_pretrained(
187
+ os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
188
+ )
189
+
190
+ # Get Text encoder
191
+ text_encoder = WanT5EncoderModel.from_pretrained(
192
+ os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
193
+ additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
194
+ low_cpu_mem_usage=True,
195
+ torch_dtype=weight_dtype,
196
+ )
197
+ text_encoder = text_encoder.eval()
198
+
199
+ # Get Clip Image Encoder
200
+ clip_image_encoder = CLIPModel.from_pretrained(
201
+ os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
202
+ ).to(weight_dtype)
203
+ clip_image_encoder = clip_image_encoder.eval()
204
+
205
+ audio_encoder_path = model_name_audio if model_name_audio is not None else os.path.join(model_name, "audio_encoder")
206
+ audio_encoder = FantasyTalkingAudioEncoder(audio_encoder_path)
207
+
208
+ # Get Scheduler
209
+ Chosen_Scheduler = scheduler_dict = {
210
+ "Flow": FlowMatchEulerDiscreteScheduler,
211
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
212
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
213
+ }[sampler_name]
214
+ if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
215
+ config['scheduler_kwargs']['shift'] = 1
216
+ scheduler = Chosen_Scheduler(
217
+ **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
218
+ )
219
+
220
+ # Get Pipeline
221
+ pipeline = FantasyTalkingPipeline(
222
+ transformer=transformer,
223
+ vae=vae,
224
+ tokenizer=tokenizer,
225
+ text_encoder=text_encoder,
226
+ scheduler=scheduler,
227
+ audio_encoder=audio_encoder,
228
+ clip_image_encoder=clip_image_encoder,
229
+ )
230
+ if ulysses_degree > 1 or ring_degree > 1:
231
+ from functools import partial
232
+ transformer.enable_multi_gpus_inference()
233
+ if fsdp_dit:
234
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
235
+ pipeline.transformer = shard_fn(pipeline.transformer)
236
+ print("Add FSDP DIT")
237
+ if fsdp_text_encoder:
238
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
239
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
240
+ print("Add FSDP TEXT ENCODER")
241
+
242
+ if compile_dit:
243
+ for i in range(len(pipeline.transformer.blocks)):
244
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
245
+ print("Add Compile")
246
+
247
+ if GPU_memory_mode == "sequential_cpu_offload":
248
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
249
+ transformer.freqs = transformer.freqs.to(device=device)
250
+ pipeline.enable_sequential_cpu_offload(device=device)
251
+ elif GPU_memory_mode == "model_group_offload":
252
+ register_auto_device_hook(pipeline.transformer)
253
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
254
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
255
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
256
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
257
+ pipeline.enable_model_cpu_offload(device=device)
258
+ elif GPU_memory_mode == "model_cpu_offload":
259
+ pipeline.enable_model_cpu_offload(device=device)
260
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
261
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
262
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
263
+ pipeline.to(device=device)
264
+ else:
265
+ pipeline.to(device=device)
266
+
267
+ coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
268
+ if coefficients is not None:
269
+ print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
270
+ pipeline.transformer.enable_teacache(
271
+ coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
272
+ )
273
+
274
+ generator = torch.Generator(device=device).manual_seed(seed)
275
+
276
+ if lora_path is not None:
277
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
278
+
279
+ with torch.no_grad():
280
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
281
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
282
+
283
+ if enable_riflex:
284
+ pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
285
+
286
+ input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
287
+
288
+ sample = pipeline(
289
+ prompt,
290
+ num_frames = video_length,
291
+ negative_prompt = negative_prompt,
292
+ height = sample_size[0],
293
+ width = sample_size[1],
294
+ generator = generator,
295
+ guidance_scale = guidance_scale,
296
+ audio_guide_scale = audio_guide_scale,
297
+ num_inference_steps = num_inference_steps,
298
+
299
+ video = input_video,
300
+ mask_video = input_video_mask,
301
+ clip_image = clip_image,
302
+ audio_path = audio_path,
303
+ shift = shift,
304
+ fps = fps
305
+ ).videos
306
+
307
+ if lora_path is not None:
308
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
309
+
310
+ def save_results():
311
+ if not os.path.exists(save_path):
312
+ os.makedirs(save_path, exist_ok=True)
313
+
314
+ index = len([path for path in os.listdir(save_path)]) + 1
315
+ prefix = str(index).zfill(8)
316
+ if video_length == 1:
317
+ video_path = os.path.join(save_path, prefix + ".png")
318
+
319
+ image = sample[0, :, 0]
320
+ image = image.transpose(0, 1).transpose(1, 2)
321
+ image = (image * 255).numpy().astype(np.uint8)
322
+ image = Image.fromarray(image)
323
+ image.save(video_path)
324
+ else:
325
+ video_path = os.path.join(save_path, prefix + ".mp4")
326
+ save_videos_grid(sample, video_path, fps=fps)
327
+
328
+ merge_video_audio(video_path=video_path, audio_path=audio_path)
329
+
330
+ if ulysses_degree * ring_degree > 1:
331
+ import torch.distributed as dist
332
+ if dist.get_rank() == 0:
333
+ save_results()
334
+ else:
335
+ save_results()
vendor/VideoX-Fun/examples/flashhead/predict_s2v.py ADDED
@@ -0,0 +1,262 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
17
+ FlashHeadTransformer3DModel, FlashHeadAudioEncoder)
18
+ from videox_fun.pipeline import FlashHeadPipeline
19
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
20
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
21
+ from videox_fun.utils import (register_auto_device_hook,
22
+ safe_enable_group_offload)
23
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
24
+ convert_weight_dtype_wrapper,
25
+ replace_parameters_by_name)
26
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
27
+ from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image,
28
+ get_video_to_video_latent,
29
+ merge_video_audio, save_videos_grid)
30
+
31
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
32
+ # model_full_load means that the entire model will be moved to the GPU.
33
+ #
34
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
35
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
36
+ #
37
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
38
+ #
39
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
40
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
41
+ #
42
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
43
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
44
+ #
45
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
46
+ # resulting in slower speeds but saving a large amount of GPU memory.
47
+ GPU_memory_mode = "model_full_load"
48
+ # Multi GPUs config
49
+ ulysses_degree = 1
50
+ ring_degree = 1
51
+ # Use FSDP to save more GPU memory in multi gpus.
52
+ fsdp_dit = False
53
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
54
+ # The compile_dit is not compatible with sequential_cpu_offload.
55
+ compile_dit = False
56
+
57
+ # Config and model path
58
+ config_path = "config/wan2.1/wan_civitai.yaml"
59
+ # model path
60
+ # Please Download https://modelscope.cn/models/AI-ModelScope/wav2vec2-base-960h/summary
61
+ model_name = "models/Diffusion_Transformer/SoulX-FlashHead-1_3B"
62
+ model_name_audio = "models/Diffusion_Transformer/wav2vec2-base-960h"
63
+
64
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
65
+ sampler_name = "Flow"
66
+ shift = 5.0
67
+ stochastic_sampling = True
68
+
69
+ # Load pretrained model if need
70
+ transformer_path = None
71
+ vae_path = None
72
+ lora_path = None
73
+
74
+ # Other params
75
+ sample_size = [512, 512]
76
+ segment_frame_length = 33
77
+ fps = 25
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ # The path of the reference image
83
+ ref_image = "asset/9.png"
84
+ # The path of the audio
85
+ audio_path = "asset/talk.wav"
86
+
87
+ # Audio guidance scale (FlashHead does not use text encoder, only audio conditioning)
88
+ audio_guide_scale = 1.0
89
+ seed = 42
90
+ num_inference_steps = 4
91
+ lora_weight = 0.55
92
+ save_path = "samples/flashhead-videos"
93
+
94
+ # FlashHead specific parameters
95
+ max_frames_num = 500
96
+ color_correction_strength = 1.0
97
+ use_apg = False
98
+ apg_momentum = 0.5
99
+ apg_norm_threshold = 1.0
100
+ audio_encode_mode = "stream"
101
+
102
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
103
+ config = OmegaConf.load(config_path)
104
+
105
+ transformer = FlashHeadTransformer3DModel.from_pretrained(
106
+ os.path.join(model_name, "Model_Pro", config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
107
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
108
+ low_cpu_mem_usage=True,
109
+ torch_dtype=weight_dtype,
110
+ )
111
+
112
+ if transformer_path is not None:
113
+ print(f"From checkpoint: {transformer_path}")
114
+ if transformer_path.endswith("safetensors"):
115
+ from safetensors.torch import load_file
116
+ state_dict = load_file(transformer_path)
117
+ else:
118
+ state_dict = torch.load(transformer_path, map_location="cpu")
119
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
120
+
121
+ m, u = transformer.load_state_dict(state_dict, strict=False)
122
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
123
+
124
+ # Get Vae
125
+ vae = AutoencoderKLWan.from_pretrained(
126
+ os.path.join(model_name, "VAE_Wan/Wan2.1_VAE.pth"),
127
+ additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
128
+ ).to(weight_dtype)
129
+
130
+ if vae_path is not None:
131
+ print(f"From checkpoint: {vae_path}")
132
+ if vae_path.endswith("safetensors"):
133
+ from safetensors.torch import load_file, safe_open
134
+ state_dict = load_file(vae_path)
135
+ else:
136
+ state_dict = torch.load(vae_path, map_location="cpu")
137
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
138
+
139
+ m, u = vae.load_state_dict(state_dict, strict=False)
140
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
141
+
142
+ # Initialize FlashHead audio encoder for real-time audio encoding
143
+ # Uses Wav2Vec2Model (not Wav2Vec2ForCTC) matching original FlashHead implementation
144
+ audio_encoder = FlashHeadAudioEncoder(
145
+ model_name_audio, "cpu"
146
+ )
147
+
148
+ # Get Scheduler
149
+ Chosen_Scheduler = scheduler_dict = {
150
+ "Flow": FlowMatchEulerDiscreteScheduler,
151
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
152
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
153
+ }[sampler_name]
154
+ if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
155
+ config['scheduler_kwargs']['shift'] = 1
156
+ scheduler = Chosen_Scheduler(
157
+ **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
158
+ )
159
+
160
+ # Get Pipeline (FlashHead does not use text encoder or clip image encoder)
161
+ pipeline = FlashHeadPipeline(
162
+ transformer=transformer,
163
+ vae=vae,
164
+ scheduler=scheduler,
165
+ audio_encoder=audio_encoder,
166
+ )
167
+ if ulysses_degree > 1 or ring_degree > 1:
168
+ from functools import partial
169
+ transformer.enable_multi_gpus_inference()
170
+ if fsdp_dit:
171
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
172
+ pipeline.transformer = shard_fn(pipeline.transformer)
173
+ print("Add FSDP DIT")
174
+
175
+ if compile_dit:
176
+ for i in range(len(pipeline.transformer.blocks)):
177
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
178
+ print("Add Compile")
179
+
180
+ if GPU_memory_mode == "sequential_cpu_offload":
181
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
182
+ transformer.freqs = transformer.freqs.to(device=device)
183
+ pipeline.enable_sequential_cpu_offload(device=device)
184
+ elif GPU_memory_mode == "model_group_offload":
185
+ register_auto_device_hook(pipeline.transformer)
186
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
187
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
188
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
189
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
190
+ pipeline.enable_model_cpu_offload(device=device)
191
+ elif GPU_memory_mode == "model_cpu_offload":
192
+ pipeline.enable_model_cpu_offload(device=device)
193
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
194
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
195
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
196
+ pipeline.to(device=device)
197
+ else:
198
+ pipeline.to(device=device)
199
+
200
+ generator = torch.Generator(device=device).manual_seed(seed)
201
+
202
+ if lora_path is not None:
203
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
204
+
205
+ with torch.no_grad():
206
+ # For FlashHead, (segment_frame_length - 1) must be divisible by 4
207
+ segment_frame_length = (segment_frame_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio + 1 if segment_frame_length != 1 else 1
208
+ latent_frames = (segment_frame_length - 1) // vae.config.temporal_compression_ratio + 1
209
+
210
+ # Prepare ref_image latent for FlashHead (no clip_image needed)
211
+ ref_image = get_image_latent(ref_image, sample_size=sample_size)
212
+
213
+ sample = pipeline(
214
+ segment_frame_length = segment_frame_length,
215
+ height = sample_size[0],
216
+ width = sample_size[1],
217
+ generator = generator,
218
+ audio_guide_scale = audio_guide_scale,
219
+ num_inference_steps = num_inference_steps,
220
+
221
+ ref_image = ref_image,
222
+ audio_path = audio_path,
223
+ audio_encode_mode = audio_encode_mode,
224
+ shift = shift,
225
+ fps = fps,
226
+ max_frames_num = max_frames_num,
227
+ color_correction_strength = color_correction_strength,
228
+ use_apg = use_apg,
229
+ apg_momentum = apg_momentum,
230
+ apg_norm_threshold = apg_norm_threshold,
231
+ stochastic_sampling = stochastic_sampling,
232
+ ).videos
233
+
234
+ if lora_path is not None:
235
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
236
+
237
+ def save_results():
238
+ if not os.path.exists(save_path):
239
+ os.makedirs(save_path, exist_ok=True)
240
+
241
+ index = len([path for path in os.listdir(save_path)]) + 1
242
+ prefix = str(index).zfill(8)
243
+ if sample.size()[2] == 1:
244
+ video_path = os.path.join(save_path, prefix + ".png")
245
+
246
+ image = sample[0, :, 0]
247
+ image = image.transpose(0, 1).transpose(1, 2)
248
+ image = (image * 255).numpy().astype(np.uint8)
249
+ image = Image.fromarray(image)
250
+ image.save(video_path)
251
+ else:
252
+ video_path = os.path.join(save_path, prefix + ".mp4")
253
+ save_videos_grid(sample, video_path, fps=fps)
254
+
255
+ merge_video_audio(video_path=video_path, audio_path=audio_path)
256
+
257
+ if ulysses_degree * ring_degree > 1:
258
+ import torch.distributed as dist
259
+ if dist.get_rank() == 0:
260
+ save_results()
261
+ else:
262
+ save_results()
vendor/VideoX-Fun/examples/flux/predict_t2i.py ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import torch
5
+ from diffusers import FlowMatchEulerDiscreteScheduler
6
+
7
+ current_file_path = os.path.abspath(__file__)
8
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
9
+ for project_root in project_roots:
10
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
11
+
12
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
13
+ from videox_fun.models import (AutoencoderKL, CLIPTextModel, CLIPTokenizer,
14
+ FluxTransformer2DModel, T5EncoderModel,
15
+ T5TokenizerFast)
16
+ from videox_fun.pipeline import FluxPipeline
17
+ from videox_fun.utils import (register_auto_device_hook,
18
+ safe_enable_group_offload)
19
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
20
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
21
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
22
+ convert_weight_dtype_wrapper)
23
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
24
+
25
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
26
+ # model_full_load means that the entire model will be moved to the GPU.
27
+ #
28
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
29
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
30
+ #
31
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
32
+ #
33
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
37
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
38
+ #
39
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
40
+ # resulting in slower speeds but saving a large amount of GPU memory.
41
+ GPU_memory_mode = "model_cpu_offload_and_qfloat8"
42
+ # Multi GPUs config
43
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
44
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
45
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
46
+ ulysses_degree = 1
47
+ ring_degree = 1
48
+ # Use FSDP to save more GPU memory in multi gpus.
49
+ fsdp_dit = False
50
+ fsdp_text_encoder = False
51
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
52
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
53
+ compile_dit = False
54
+
55
+ # model path
56
+ model_name = "models/Diffusion_Transformer/FLUX.1-dev"
57
+
58
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
59
+ sampler_name = "Flow"
60
+
61
+ # Load pretrained model if need
62
+ transformer_path = None
63
+ vae_path = None
64
+ lora_path = None
65
+
66
+ # Other params
67
+ sample_size = [1344, 768]
68
+
69
+ # Use torch.float16 if GPU does not support torch.bfloat16
70
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
71
+ weight_dtype = torch.bfloat16
72
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
73
+ negative_prompt = " "
74
+ guidance_scale = 1.0
75
+ seed = 43
76
+ num_inference_steps = 50
77
+ lora_weight = 0.55
78
+ save_path = "samples/flux-t2i"
79
+
80
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
81
+
82
+ transformer = FluxTransformer2DModel.from_pretrained(
83
+ model_name,
84
+ subfolder="transformer",
85
+ low_cpu_mem_usage=True,
86
+ torch_dtype=weight_dtype,
87
+ ).to(weight_dtype)
88
+
89
+ if transformer_path is not None:
90
+ print(f"From checkpoint: {transformer_path}")
91
+ if transformer_path.endswith("safetensors"):
92
+ from safetensors.torch import load_file, safe_open
93
+ state_dict = load_file(transformer_path)
94
+ else:
95
+ state_dict = torch.load(transformer_path, map_location="cpu")
96
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
97
+
98
+ m, u = transformer.load_state_dict(state_dict, strict=False)
99
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
100
+
101
+ # Get Vae
102
+ vae = AutoencoderKL.from_pretrained(
103
+ model_name,
104
+ subfolder="vae"
105
+ ).to(weight_dtype)
106
+
107
+ if vae_path is not None:
108
+ print(f"From checkpoint: {vae_path}")
109
+ if vae_path.endswith("safetensors"):
110
+ from safetensors.torch import load_file, safe_open
111
+ state_dict = load_file(vae_path)
112
+ else:
113
+ state_dict = torch.load(vae_path, map_location="cpu")
114
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
115
+
116
+ m, u = vae.load_state_dict(state_dict, strict=False)
117
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
118
+
119
+ # Get tokenizer and text_encoder
120
+ tokenizer = CLIPTokenizer.from_pretrained(
121
+ model_name, subfolder="tokenizer"
122
+ )
123
+ text_encoder = CLIPTextModel.from_pretrained(
124
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype
125
+ )
126
+
127
+ tokenizer_2 = T5TokenizerFast.from_pretrained(
128
+ model_name, subfolder="tokenizer_2"
129
+ )
130
+ text_encoder_2 = T5EncoderModel.from_pretrained(
131
+ model_name, subfolder="text_encoder_2", torch_dtype=weight_dtype
132
+ )
133
+
134
+ # Get Scheduler
135
+ Chosen_Scheduler = scheduler_dict = {
136
+ "Flow": FlowMatchEulerDiscreteScheduler,
137
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
138
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
139
+ }[sampler_name]
140
+ scheduler = Chosen_Scheduler.from_pretrained(
141
+ model_name,
142
+ subfolder="scheduler"
143
+ )
144
+
145
+ pipeline = FluxPipeline(
146
+ vae=vae,
147
+ tokenizer=tokenizer,
148
+ text_encoder=text_encoder,
149
+ tokenizer_2=tokenizer_2,
150
+ text_encoder_2=text_encoder_2,
151
+ transformer=transformer,
152
+ scheduler=scheduler,
153
+ )
154
+
155
+ if ulysses_degree > 1 or ring_degree > 1:
156
+ from functools import partial
157
+ transformer.enable_multi_gpus_inference()
158
+ if fsdp_dit:
159
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
160
+ pipeline.transformer = shard_fn(pipeline.transformer)
161
+ print("Add FSDP DIT")
162
+ if fsdp_text_encoder:
163
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.text_model.encoder.layers)
164
+ text_encoder = shard_fn(text_encoder)
165
+ print("Add FSDP TEXT ENCODER")
166
+
167
+ if compile_dit:
168
+ for i in range(len(pipeline.transformer.transformer_blocks)):
169
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
170
+ print("Add Compile")
171
+
172
+ if GPU_memory_mode == "sequential_cpu_offload":
173
+ pipeline.enable_sequential_cpu_offload(device=device)
174
+ elif GPU_memory_mode == "model_group_offload":
175
+ register_auto_device_hook(pipeline.transformer)
176
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
177
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
178
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
179
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
180
+ pipeline.enable_model_cpu_offload(device=device)
181
+ elif GPU_memory_mode == "model_cpu_offload":
182
+ pipeline.enable_model_cpu_offload(device=device)
183
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
184
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
185
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
186
+ pipeline.to(device=device)
187
+ else:
188
+ pipeline.to(device=device)
189
+
190
+ generator = torch.Generator(device=device).manual_seed(seed)
191
+
192
+ if lora_path is not None:
193
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
194
+
195
+ with torch.no_grad():
196
+ sample = pipeline(
197
+ prompt,
198
+ negative_prompt = negative_prompt,
199
+ height = sample_size[0],
200
+ width = sample_size[1],
201
+ generator = generator,
202
+ true_cfg_scale = guidance_scale,
203
+ num_inference_steps = num_inference_steps,
204
+ ).images
205
+
206
+ if lora_path is not None:
207
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
208
+
209
+ def save_results():
210
+ if not os.path.exists(save_path):
211
+ os.makedirs(save_path, exist_ok=True)
212
+
213
+ index = len([path for path in os.listdir(save_path)]) + 1
214
+ prefix = str(index).zfill(8)
215
+ video_path = os.path.join(save_path, prefix + ".png")
216
+ image = sample[0]
217
+ image.save(video_path)
218
+
219
+ if ulysses_degree * ring_degree > 1:
220
+ import torch.distributed as dist
221
+ if dist.get_rank() == 0:
222
+ save_results()
223
+ else:
224
+ save_results()
vendor/VideoX-Fun/examples/flux2/predict_t2i.py ADDED
@@ -0,0 +1,218 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import torch
5
+
6
+ from diffusers import (FlowMatchEulerDiscreteScheduler)
7
+
8
+ current_file_path = os.path.abspath(__file__)
9
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
10
+ for project_root in project_roots:
11
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
12
+
13
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
14
+ from videox_fun.models import (AutoencoderKLFlux2,
15
+ Mistral3ForConditionalGeneration,
16
+ PixtralProcessor, Flux2Transformer2DModel)
17
+ from videox_fun.models.cache_utils import get_teacache_coefficients
18
+ from videox_fun.pipeline import Flux2Pipeline
19
+ from videox_fun.utils import (register_auto_device_hook,
20
+ safe_enable_group_offload)
21
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
22
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
23
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
24
+ convert_weight_dtype_wrapper)
25
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
26
+
27
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
28
+ # model_full_load means that the entire model will be moved to the GPU.
29
+ #
30
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
31
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
32
+ #
33
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
34
+ #
35
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
39
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
40
+ #
41
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
42
+ # resulting in slower speeds but saving a large amount of GPU memory.
43
+ GPU_memory_mode = "sequential_cpu_offload"
44
+ # Multi GPUs config
45
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
46
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
47
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
48
+ ulysses_degree = 1
49
+ ring_degree = 1
50
+ # Use FSDP to save more GPU memory in multi gpus.
51
+ fsdp_dit = False
52
+ fsdp_text_encoder = False
53
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
54
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
55
+ compile_dit = False
56
+
57
+ # model path
58
+ model_name = "models/Diffusion_Transformer/FLUX.2-dev"
59
+
60
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
61
+ sampler_name = "Flow"
62
+
63
+ # Load pretrained model if need
64
+ transformer_path = None
65
+ vae_path = None
66
+ lora_path = None
67
+
68
+ # Other params
69
+ sample_size = [1344, 768]
70
+
71
+ # Use torch.float16 if GPU does not support torch.bfloat16
72
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
73
+ weight_dtype = torch.bfloat16
74
+ # Please use as detailed a prompt as possible to describe the object that needs to be generated.
75
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
76
+ negative_prompt = " "
77
+ guidance_scale = 4.00
78
+ seed = 43
79
+ num_inference_steps = 50
80
+ lora_weight = 0.55
81
+ save_path = "samples/flux2-t2i"
82
+
83
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
84
+
85
+ transformer = Flux2Transformer2DModel.from_pretrained(
86
+ model_name,
87
+ subfolder="transformer",
88
+ low_cpu_mem_usage=True,
89
+ torch_dtype=weight_dtype,
90
+ ).to(weight_dtype)
91
+
92
+ if transformer_path is not None:
93
+ print(f"From checkpoint: {transformer_path}")
94
+ if transformer_path.endswith("safetensors"):
95
+ from safetensors.torch import load_file, safe_open
96
+ state_dict = load_file(transformer_path)
97
+ else:
98
+ state_dict = torch.load(transformer_path, map_location="cpu")
99
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
100
+
101
+ m, u = transformer.load_state_dict(state_dict, strict=False)
102
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
103
+
104
+ # Get Vae
105
+ vae = AutoencoderKLFlux2.from_pretrained(
106
+ model_name,
107
+ subfolder="vae"
108
+ ).to(weight_dtype)
109
+
110
+ if vae_path is not None:
111
+ print(f"From checkpoint: {vae_path}")
112
+ if vae_path.endswith("safetensors"):
113
+ from safetensors.torch import load_file, safe_open
114
+ state_dict = load_file(vae_path)
115
+ else:
116
+ state_dict = torch.load(vae_path, map_location="cpu")
117
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
118
+
119
+ m, u = vae.load_state_dict(state_dict, strict=False)
120
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
121
+
122
+ # Get tokenizer and text_encoder
123
+ tokenizer = PixtralProcessor.from_pretrained(
124
+ model_name, subfolder="tokenizer"
125
+ )
126
+ text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
127
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
128
+ low_cpu_mem_usage=True,
129
+ )
130
+
131
+ # Get Scheduler
132
+ Chosen_Scheduler = scheduler_dict = {
133
+ "Flow": FlowMatchEulerDiscreteScheduler,
134
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
135
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
136
+ }[sampler_name]
137
+ scheduler = Chosen_Scheduler.from_pretrained(
138
+ model_name,
139
+ subfolder="scheduler"
140
+ )
141
+
142
+ pipeline = Flux2Pipeline(
143
+ vae=vae,
144
+ tokenizer=tokenizer,
145
+ text_encoder=text_encoder,
146
+ transformer=transformer,
147
+ scheduler=scheduler,
148
+ )
149
+
150
+ if ulysses_degree > 1 or ring_degree > 1:
151
+ from functools import partial
152
+ transformer.enable_multi_gpus_inference()
153
+ if fsdp_dit:
154
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
155
+ pipeline.transformer = shard_fn(pipeline.transformer)
156
+ print("Add FSDP DIT")
157
+ if fsdp_text_encoder:
158
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers)
159
+ text_encoder = shard_fn(text_encoder)
160
+ print("Add FSDP TEXT ENCODER")
161
+
162
+ if compile_dit:
163
+ for i in range(len(pipeline.transformer.transformer_blocks)):
164
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
165
+ print("Add Compile")
166
+
167
+ if GPU_memory_mode == "sequential_cpu_offload":
168
+ pipeline.enable_sequential_cpu_offload(device=device)
169
+ elif GPU_memory_mode == "model_group_offload":
170
+ register_auto_device_hook(pipeline.transformer)
171
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
172
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
173
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
174
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
175
+ pipeline.enable_model_cpu_offload(device=device)
176
+ elif GPU_memory_mode == "model_cpu_offload":
177
+ pipeline.enable_model_cpu_offload(device=device)
178
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
179
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
180
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
181
+ pipeline.to(device=device)
182
+ else:
183
+ pipeline.to(device=device)
184
+
185
+ generator = torch.Generator(device=device).manual_seed(seed)
186
+
187
+ if lora_path is not None:
188
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
189
+
190
+ with torch.no_grad():
191
+ sample = pipeline(
192
+ prompt = prompt,
193
+ height = sample_size[0],
194
+ width = sample_size[1],
195
+ generator = generator,
196
+ guidance_scale = guidance_scale,
197
+ num_inference_steps = num_inference_steps,
198
+ ).images
199
+
200
+ if lora_path is not None:
201
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
202
+
203
+ def save_results():
204
+ if not os.path.exists(save_path):
205
+ os.makedirs(save_path, exist_ok=True)
206
+
207
+ index = len([path for path in os.listdir(save_path)]) + 1
208
+ prefix = str(index).zfill(8)
209
+ video_path = os.path.join(save_path, prefix + ".png")
210
+ image = sample[0]
211
+ image.save(video_path)
212
+
213
+ if ulysses_degree * ring_degree > 1:
214
+ import torch.distributed as dist
215
+ if dist.get_rank() == 0:
216
+ save_results()
217
+ else:
218
+ save_results()
vendor/VideoX-Fun/examples/flux2_fun/predict_i2i_inpaint.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLFlux2,
17
+ Mistral3ForConditionalGeneration,
18
+ PixtralProcessor, Flux2ControlTransformer2DModel)
19
+ from videox_fun.models.cache_utils import get_teacache_coefficients
20
+ from videox_fun.pipeline import Flux2ControlPipeline
21
+ from videox_fun.utils import (register_auto_device_hook,
22
+ safe_enable_group_offload)
23
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
24
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent,
29
+ get_image_to_video_latent,
30
+ get_video_to_video_latent,
31
+ save_videos_grid)
32
+
33
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
34
+ # model_full_load means that the entire model will be moved to the GPU.
35
+ #
36
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
37
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
38
+ #
39
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
40
+ #
41
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
42
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
43
+ #
44
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
45
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
46
+ #
47
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
48
+ # resulting in slower speeds but saving a large amount of GPU memory.
49
+ GPU_memory_mode = "model_cpu_offload"
50
+ # Multi GPUs config
51
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
52
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
53
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
54
+ ulysses_degree = 1
55
+ ring_degree = 1
56
+ # Use FSDP to save more GPU memory in multi gpus.
57
+ fsdp_dit = False
58
+ fsdp_text_encoder = False
59
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
60
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
61
+ compile_dit = False
62
+
63
+ # Config and model path
64
+ config_path = "config/flux2/flux2_control.yaml"
65
+ # model path
66
+ model_name = "models/Diffusion_Transformer/FLUX.2-dev"
67
+
68
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
69
+ sampler_name = "Flow"
70
+
71
+ # Load pretrained model if need
72
+ transformer_path = "models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors"
73
+ vae_path = None
74
+ lora_path = None
75
+
76
+ # Other params
77
+ sample_size = [1728, 992]
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ image = None
83
+ control_image = None
84
+ inpaint_image = "asset/8.png"
85
+ mask_image = "asset/mask.png"
86
+ control_context_scale = 0.75
87
+
88
+ # Please use as detailed a prompt as possible to describe the object that needs to be generated.
89
+ prompt = "This is a panoramic portrait photo of a young man. He has medium-length flowing hair with a soft lavender-like color. He is wearing a white sleeveless shirt with a blue ribbon bow tied around the collar. He has a confident posture, with his left hand naturally hanging down and his right hand in his pocket, and his legs slightly apart. He looks straight at the camera. The sea breeze gently brushes his hair as he stands on the sunny seaside path, surrounded by blooming purple seaside flowers and smooth pebbles, with the sparkling sea and blue sky behind him. The scene presents a bright summer atmosphere, with soft and natural lighting, realistic details, and 8K ultra high definition image quality, clearly presenting fine textures such as clothing and hair."
90
+ negative_prompt = " "
91
+ guidance_scale = 4.00
92
+ seed = 43
93
+ num_inference_steps = 50
94
+ lora_weight = 0.55
95
+ save_path = "samples/flux2-t2i-control"
96
+
97
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
98
+ config = OmegaConf.load(config_path)
99
+
100
+ transformer = Flux2ControlTransformer2DModel.from_pretrained(
101
+ model_name,
102
+ subfolder="transformer",
103
+ low_cpu_mem_usage=True,
104
+ torch_dtype=weight_dtype,
105
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
106
+ ).to(weight_dtype)
107
+
108
+ if transformer_path is not None:
109
+ print(f"From checkpoint: {transformer_path}")
110
+ if transformer_path.endswith("safetensors"):
111
+ from safetensors.torch import load_file, safe_open
112
+ state_dict = load_file(transformer_path)
113
+ else:
114
+ state_dict = torch.load(transformer_path, map_location="cpu")
115
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
116
+
117
+ m, u = transformer.load_state_dict(state_dict, strict=False)
118
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
119
+
120
+ # Get Vae
121
+ vae = AutoencoderKLFlux2.from_pretrained(
122
+ model_name,
123
+ subfolder="vae"
124
+ ).to(weight_dtype)
125
+
126
+ if vae_path is not None:
127
+ print(f"From checkpoint: {vae_path}")
128
+ if vae_path.endswith("safetensors"):
129
+ from safetensors.torch import load_file, safe_open
130
+ state_dict = load_file(vae_path)
131
+ else:
132
+ state_dict = torch.load(vae_path, map_location="cpu")
133
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
134
+
135
+ m, u = vae.load_state_dict(state_dict, strict=False)
136
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
137
+
138
+ # Get tokenizer and text_encoder
139
+ tokenizer = PixtralProcessor.from_pretrained(
140
+ model_name, subfolder="tokenizer"
141
+ )
142
+ text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
143
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
144
+ low_cpu_mem_usage=True,
145
+ )
146
+
147
+ # Get Scheduler
148
+ Chosen_Scheduler = scheduler_dict = {
149
+ "Flow": FlowMatchEulerDiscreteScheduler,
150
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
151
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
152
+ }[sampler_name]
153
+ scheduler = Chosen_Scheduler.from_pretrained(
154
+ model_name,
155
+ subfolder="scheduler"
156
+ )
157
+
158
+ pipeline = Flux2ControlPipeline(
159
+ vae=vae,
160
+ tokenizer=tokenizer,
161
+ text_encoder=text_encoder,
162
+ transformer=transformer,
163
+ scheduler=scheduler,
164
+ )
165
+
166
+ if ulysses_degree > 1 or ring_degree > 1:
167
+ from functools import partial
168
+ transformer.enable_multi_gpus_inference()
169
+ if fsdp_dit:
170
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
171
+ pipeline.transformer = shard_fn(pipeline.transformer)
172
+ print("Add FSDP DIT")
173
+ if fsdp_text_encoder:
174
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers, ignored_modules=[text_encoder.language_model.embed_tokens], transformer_layer_cls_to_wrap=["MistralDecoderLayer", "PixtralTransformer"])
175
+ text_encoder = shard_fn(text_encoder)
176
+ print("Add FSDP TEXT ENCODER")
177
+
178
+ if compile_dit:
179
+ for i in range(len(pipeline.transformer.transformer_blocks)):
180
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
181
+ print("Add Compile")
182
+
183
+ if GPU_memory_mode == "sequential_cpu_offload":
184
+ pipeline.enable_sequential_cpu_offload(device=device)
185
+ elif GPU_memory_mode == "model_group_offload":
186
+ register_auto_device_hook(pipeline.transformer)
187
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
188
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
189
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
190
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
191
+ pipeline.enable_model_cpu_offload(device=device)
192
+ elif GPU_memory_mode == "model_cpu_offload":
193
+ pipeline.enable_model_cpu_offload(device=device)
194
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
195
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
196
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
197
+ pipeline.to(device=device)
198
+ else:
199
+ pipeline.to(device=device)
200
+
201
+ generator = torch.Generator(device=device).manual_seed(seed)
202
+
203
+ if lora_path is not None:
204
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
205
+
206
+ with torch.no_grad():
207
+ if image is not None:
208
+ if not isinstance(image, list):
209
+ image = get_image(image)
210
+ else:
211
+ image = [get_image(_image) for _image in image]
212
+
213
+ if inpaint_image is not None:
214
+ inpaint_image = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0]
215
+ else:
216
+ inpaint_image = torch.zeros([1, 3, sample_size[0], sample_size[1]])
217
+
218
+ if mask_image is not None:
219
+ mask_image = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0]
220
+ else:
221
+ mask_image = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255
222
+
223
+ if control_image is not None:
224
+ control_image = get_image_latent(control_image, sample_size=sample_size)[:, :, 0]
225
+
226
+ sample = pipeline(
227
+ prompt = prompt,
228
+ height = sample_size[0],
229
+ width = sample_size[1],
230
+ generator = generator,
231
+ guidance_scale = guidance_scale,
232
+ image = image,
233
+ inpaint_image = inpaint_image,
234
+ mask_image = mask_image,
235
+ control_image = control_image,
236
+ num_inference_steps = num_inference_steps,
237
+ control_context_scale = control_context_scale,
238
+ ).images
239
+
240
+ if lora_path is not None:
241
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
242
+
243
+ def save_results():
244
+ if not os.path.exists(save_path):
245
+ os.makedirs(save_path, exist_ok=True)
246
+
247
+ index = len([path for path in os.listdir(save_path)]) + 1
248
+ prefix = str(index).zfill(8)
249
+ video_path = os.path.join(save_path, prefix + ".png")
250
+ image = sample[0]
251
+ image.save(video_path)
252
+
253
+ if ulysses_degree * ring_degree > 1:
254
+ import torch.distributed as dist
255
+ if dist.get_rank() == 0:
256
+ save_results()
257
+ else:
258
+ save_results()
vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLFlux2,
17
+ Mistral3ForConditionalGeneration,
18
+ PixtralProcessor, Flux2ControlTransformer2DModel)
19
+ from videox_fun.models.cache_utils import get_teacache_coefficients
20
+ from videox_fun.pipeline import Flux2ControlPipeline
21
+ from videox_fun.utils import (register_auto_device_hook,
22
+ safe_enable_group_offload)
23
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
24
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent,
29
+ get_image_to_video_latent,
30
+ get_video_to_video_latent,
31
+ save_videos_grid)
32
+
33
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
34
+ # model_full_load means that the entire model will be moved to the GPU.
35
+ #
36
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
37
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
38
+ #
39
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
40
+ #
41
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
42
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
43
+ #
44
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
45
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
46
+ #
47
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
48
+ # resulting in slower speeds but saving a large amount of GPU memory.
49
+ GPU_memory_mode = "model_cpu_offload"
50
+ # Multi GPUs config
51
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
52
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
53
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
54
+ ulysses_degree = 1
55
+ ring_degree = 1
56
+ # Use FSDP to save more GPU memory in multi gpus.
57
+ fsdp_dit = False
58
+ fsdp_text_encoder = False
59
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
60
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
61
+ compile_dit = False
62
+
63
+ # Config and model path
64
+ config_path = "config/flux2/flux2_control.yaml"
65
+ # model path
66
+ model_name = "models/Diffusion_Transformer/FLUX.2-dev"
67
+
68
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
69
+ sampler_name = "Flow"
70
+
71
+ # Load pretrained model if need
72
+ transformer_path = "models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors"
73
+ vae_path = None
74
+ lora_path = None
75
+
76
+ # Other params
77
+ sample_size = [1728, 992]
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ image = None
83
+ control_image = "asset/pose.jpg"
84
+ inpaint_image = None
85
+ mask_image = None
86
+ control_context_scale = 0.75
87
+
88
+ # Please use as detailed a prompt as possible to describe the object that needs to be generated.
89
+ prompt = "This is a panoramic portrait photo of a young woman. She has flowing long hair and a soft lavender like color. She is wearing a white sleeveless dress with a blue ribbon bow tied around the collar. She has a confident posture, with her left hand naturally hanging down and her right hand in her pocket, and her legs slightly apart. Look straight at the camera. The sea breeze gently brushed her long hair, and they stood on the sunny seaside path, surrounded by blooming purple seaside flowers and smooth pebbles, with the sparkling sea and blue sky behind them. The screen presents a bright summer atmosphere, with soft and natural lighting, realistic details, and 8K ultra high definition image quality, clearly presenting fine textures such as clothing and hair. "
90
+ negative_prompt = " "
91
+ guidance_scale = 4.00
92
+ seed = 43
93
+ num_inference_steps = 50
94
+ lora_weight = 0.55
95
+ save_path = "samples/flux2-t2i-control"
96
+
97
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
98
+ config = OmegaConf.load(config_path)
99
+
100
+ transformer = Flux2ControlTransformer2DModel.from_pretrained(
101
+ model_name,
102
+ subfolder="transformer",
103
+ low_cpu_mem_usage=True,
104
+ torch_dtype=weight_dtype,
105
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
106
+ ).to(weight_dtype)
107
+
108
+ if transformer_path is not None:
109
+ print(f"From checkpoint: {transformer_path}")
110
+ if transformer_path.endswith("safetensors"):
111
+ from safetensors.torch import load_file, safe_open
112
+ state_dict = load_file(transformer_path)
113
+ else:
114
+ state_dict = torch.load(transformer_path, map_location="cpu")
115
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
116
+
117
+ m, u = transformer.load_state_dict(state_dict, strict=False)
118
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
119
+
120
+ # Get Vae
121
+ vae = AutoencoderKLFlux2.from_pretrained(
122
+ model_name,
123
+ subfolder="vae"
124
+ ).to(weight_dtype)
125
+
126
+ if vae_path is not None:
127
+ print(f"From checkpoint: {vae_path}")
128
+ if vae_path.endswith("safetensors"):
129
+ from safetensors.torch import load_file, safe_open
130
+ state_dict = load_file(vae_path)
131
+ else:
132
+ state_dict = torch.load(vae_path, map_location="cpu")
133
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
134
+
135
+ m, u = vae.load_state_dict(state_dict, strict=False)
136
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
137
+
138
+ # Get tokenizer and text_encoder
139
+ tokenizer = PixtralProcessor.from_pretrained(
140
+ model_name, subfolder="tokenizer"
141
+ )
142
+ text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
143
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
144
+ low_cpu_mem_usage=True,
145
+ )
146
+
147
+ # Get Scheduler
148
+ Chosen_Scheduler = scheduler_dict = {
149
+ "Flow": FlowMatchEulerDiscreteScheduler,
150
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
151
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
152
+ }[sampler_name]
153
+ scheduler = Chosen_Scheduler.from_pretrained(
154
+ model_name,
155
+ subfolder="scheduler"
156
+ )
157
+
158
+ pipeline = Flux2ControlPipeline(
159
+ vae=vae,
160
+ tokenizer=tokenizer,
161
+ text_encoder=text_encoder,
162
+ transformer=transformer,
163
+ scheduler=scheduler,
164
+ )
165
+
166
+ if ulysses_degree > 1 or ring_degree > 1:
167
+ from functools import partial
168
+ transformer.enable_multi_gpus_inference()
169
+ if fsdp_dit:
170
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
171
+ pipeline.transformer = shard_fn(pipeline.transformer)
172
+ print("Add FSDP DIT")
173
+ if fsdp_text_encoder:
174
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers, ignored_modules=[text_encoder.language_model.embed_tokens], transformer_layer_cls_to_wrap=["MistralDecoderLayer", "PixtralTransformer"])
175
+ text_encoder = shard_fn(text_encoder)
176
+ print("Add FSDP TEXT ENCODER")
177
+
178
+ if compile_dit:
179
+ for i in range(len(pipeline.transformer.transformer_blocks)):
180
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
181
+ print("Add Compile")
182
+
183
+ if GPU_memory_mode == "sequential_cpu_offload":
184
+ pipeline.enable_sequential_cpu_offload(device=device)
185
+ elif GPU_memory_mode == "model_group_offload":
186
+ register_auto_device_hook(pipeline.transformer)
187
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
188
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
189
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
190
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
191
+ pipeline.enable_model_cpu_offload(device=device)
192
+ elif GPU_memory_mode == "model_cpu_offload":
193
+ pipeline.enable_model_cpu_offload(device=device)
194
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
195
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
196
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
197
+ pipeline.to(device=device)
198
+ else:
199
+ pipeline.to(device=device)
200
+
201
+ generator = torch.Generator(device=device).manual_seed(seed)
202
+
203
+ if lora_path is not None:
204
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
205
+
206
+ with torch.no_grad():
207
+ if image is not None:
208
+ if not isinstance(image, list):
209
+ image = get_image(image)
210
+ else:
211
+ image = [get_image(_image) for _image in image]
212
+
213
+ if inpaint_image is not None:
214
+ inpaint_image = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0]
215
+ else:
216
+ inpaint_image = torch.zeros([1, 3, sample_size[0], sample_size[1]])
217
+
218
+ if mask_image is not None:
219
+ mask_image = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0]
220
+ else:
221
+ mask_image = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255
222
+
223
+ if control_image is not None:
224
+ control_image = get_image_latent(control_image, sample_size=sample_size)[:, :, 0]
225
+
226
+ sample = pipeline(
227
+ prompt = prompt,
228
+ height = sample_size[0],
229
+ width = sample_size[1],
230
+ generator = generator,
231
+ guidance_scale = guidance_scale,
232
+ image = image,
233
+ inpaint_image = inpaint_image,
234
+ mask_image = mask_image,
235
+ control_image = control_image,
236
+ num_inference_steps = num_inference_steps,
237
+ control_context_scale = control_context_scale,
238
+ ).images
239
+
240
+ if lora_path is not None:
241
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
242
+
243
+ def save_results():
244
+ if not os.path.exists(save_path):
245
+ os.makedirs(save_path, exist_ok=True)
246
+
247
+ index = len([path for path in os.listdir(save_path)]) + 1
248
+ prefix = str(index).zfill(8)
249
+ video_path = os.path.join(save_path, prefix + ".png")
250
+ image = sample[0]
251
+ image.save(video_path)
252
+
253
+ if ulysses_degree * ring_degree > 1:
254
+ import torch.distributed as dist
255
+ if dist.get_rank() == 0:
256
+ save_results()
257
+ else:
258
+ save_results()
vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control_ref.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLFlux2,
17
+ Mistral3ForConditionalGeneration,
18
+ PixtralProcessor, Flux2ControlTransformer2DModel)
19
+ from videox_fun.models.cache_utils import get_teacache_coefficients
20
+ from videox_fun.pipeline import Flux2ControlPipeline
21
+ from videox_fun.utils import (register_auto_device_hook,
22
+ safe_enable_group_offload)
23
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
24
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
26
+ convert_weight_dtype_wrapper)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent,
29
+ get_image_to_video_latent,
30
+ get_video_to_video_latent,
31
+ save_videos_grid)
32
+
33
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
34
+ # model_full_load means that the entire model will be moved to the GPU.
35
+ #
36
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
37
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
38
+ #
39
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
40
+ #
41
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
42
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
43
+ #
44
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
45
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
46
+ #
47
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
48
+ # resulting in slower speeds but saving a large amount of GPU memory.
49
+ GPU_memory_mode = "model_cpu_offload"
50
+ # Multi GPUs config
51
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
52
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
53
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
54
+ ulysses_degree = 1
55
+ ring_degree = 1
56
+ # Use FSDP to save more GPU memory in multi gpus.
57
+ fsdp_dit = False
58
+ fsdp_text_encoder = False
59
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
60
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
61
+ compile_dit = False
62
+
63
+ # Config and model path
64
+ config_path = "config/flux2/flux2_control.yaml"
65
+ # model path
66
+ model_name = "models/Diffusion_Transformer/FLUX.2-dev"
67
+
68
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
69
+ sampler_name = "Flow"
70
+
71
+ # Load pretrained model if need
72
+ transformer_path = "models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors"
73
+ vae_path = None
74
+ lora_path = None
75
+
76
+ # Other params
77
+ sample_size = [1728, 992]
78
+
79
+ # Use torch.float16 if GPU does not support torch.bfloat16
80
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
81
+ weight_dtype = torch.bfloat16
82
+ image = "asset/8.png"
83
+ control_image = "asset/pose.jpg"
84
+ inpaint_image = None
85
+ mask_image = None
86
+ control_context_scale = 0.75
87
+
88
+ # Please use as detailed a prompt as possible to describe the object that needs to be generated.
89
+ prompt = "This is a panoramic portrait photo of a young woman. She has flowing long hair and a soft lavender like color. She is wearing a white sleeveless dress with a blue ribbon bow tied around the collar. She has a confident posture, with her left hand naturally hanging down and her right hand in her pocket, and her legs slightly apart. Look straight at the camera. The sea breeze gently brushed her long hair, and they stood on the sunny seaside path, surrounded by blooming purple seaside flowers and smooth pebbles, with the sparkling sea and blue sky behind them. The screen presents a bright summer atmosphere, with soft and natural lighting, realistic details, and 8K ultra high definition image quality, clearly presenting fine textures such as clothing and hair. "
90
+ negative_prompt = " "
91
+ guidance_scale = 4.00
92
+ seed = 43
93
+ num_inference_steps = 50
94
+ lora_weight = 0.55
95
+ save_path = "samples/flux2-t2i-control"
96
+
97
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
98
+ config = OmegaConf.load(config_path)
99
+
100
+ transformer = Flux2ControlTransformer2DModel.from_pretrained(
101
+ model_name,
102
+ subfolder="transformer",
103
+ low_cpu_mem_usage=True,
104
+ torch_dtype=weight_dtype,
105
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
106
+ ).to(weight_dtype)
107
+
108
+ if transformer_path is not None:
109
+ print(f"From checkpoint: {transformer_path}")
110
+ if transformer_path.endswith("safetensors"):
111
+ from safetensors.torch import load_file, safe_open
112
+ state_dict = load_file(transformer_path)
113
+ else:
114
+ state_dict = torch.load(transformer_path, map_location="cpu")
115
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
116
+
117
+ m, u = transformer.load_state_dict(state_dict, strict=False)
118
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
119
+
120
+ # Get Vae
121
+ vae = AutoencoderKLFlux2.from_pretrained(
122
+ model_name,
123
+ subfolder="vae"
124
+ ).to(weight_dtype)
125
+
126
+ if vae_path is not None:
127
+ print(f"From checkpoint: {vae_path}")
128
+ if vae_path.endswith("safetensors"):
129
+ from safetensors.torch import load_file, safe_open
130
+ state_dict = load_file(vae_path)
131
+ else:
132
+ state_dict = torch.load(vae_path, map_location="cpu")
133
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
134
+
135
+ m, u = vae.load_state_dict(state_dict, strict=False)
136
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
137
+
138
+ # Get tokenizer and text_encoder
139
+ tokenizer = PixtralProcessor.from_pretrained(
140
+ model_name, subfolder="tokenizer"
141
+ )
142
+ text_encoder = Mistral3ForConditionalGeneration.from_pretrained(
143
+ model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
144
+ low_cpu_mem_usage=True,
145
+ )
146
+
147
+ # Get Scheduler
148
+ Chosen_Scheduler = scheduler_dict = {
149
+ "Flow": FlowMatchEulerDiscreteScheduler,
150
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
151
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
152
+ }[sampler_name]
153
+ scheduler = Chosen_Scheduler.from_pretrained(
154
+ model_name,
155
+ subfolder="scheduler"
156
+ )
157
+
158
+ pipeline = Flux2ControlPipeline(
159
+ vae=vae,
160
+ tokenizer=tokenizer,
161
+ text_encoder=text_encoder,
162
+ transformer=transformer,
163
+ scheduler=scheduler,
164
+ )
165
+
166
+ if ulysses_degree > 1 or ring_degree > 1:
167
+ from functools import partial
168
+ transformer.enable_multi_gpus_inference()
169
+ if fsdp_dit:
170
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
171
+ pipeline.transformer = shard_fn(pipeline.transformer)
172
+ print("Add FSDP DIT")
173
+ if fsdp_text_encoder:
174
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers, ignored_modules=[text_encoder.language_model.embed_tokens], transformer_layer_cls_to_wrap=["MistralDecoderLayer", "PixtralTransformer"])
175
+ text_encoder = shard_fn(text_encoder)
176
+ print("Add FSDP TEXT ENCODER")
177
+
178
+ if compile_dit:
179
+ for i in range(len(pipeline.transformer.transformer_blocks)):
180
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
181
+ print("Add Compile")
182
+
183
+ if GPU_memory_mode == "sequential_cpu_offload":
184
+ pipeline.enable_sequential_cpu_offload(device=device)
185
+ elif GPU_memory_mode == "model_group_offload":
186
+ register_auto_device_hook(pipeline.transformer)
187
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
188
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
189
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
190
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
191
+ pipeline.enable_model_cpu_offload(device=device)
192
+ elif GPU_memory_mode == "model_cpu_offload":
193
+ pipeline.enable_model_cpu_offload(device=device)
194
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
195
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
196
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
197
+ pipeline.to(device=device)
198
+ else:
199
+ pipeline.to(device=device)
200
+
201
+ generator = torch.Generator(device=device).manual_seed(seed)
202
+
203
+ if lora_path is not None:
204
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
205
+
206
+ with torch.no_grad():
207
+ if image is not None:
208
+ if not isinstance(image, list):
209
+ image = get_image(image)
210
+ else:
211
+ image = [get_image(_image) for _image in image]
212
+
213
+ if inpaint_image is not None:
214
+ inpaint_image = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0]
215
+ else:
216
+ inpaint_image = torch.zeros([1, 3, sample_size[0], sample_size[1]])
217
+
218
+ if mask_image is not None:
219
+ mask_image = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0]
220
+ else:
221
+ mask_image = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255
222
+
223
+ if control_image is not None:
224
+ control_image = get_image_latent(control_image, sample_size=sample_size)[:, :, 0]
225
+
226
+ sample = pipeline(
227
+ prompt = prompt,
228
+ height = sample_size[0],
229
+ width = sample_size[1],
230
+ generator = generator,
231
+ guidance_scale = guidance_scale,
232
+ image = image,
233
+ inpaint_image = inpaint_image,
234
+ mask_image = mask_image,
235
+ control_image = control_image,
236
+ num_inference_steps = num_inference_steps,
237
+ control_context_scale = control_context_scale,
238
+ ).images
239
+
240
+ if lora_path is not None:
241
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
242
+
243
+ def save_results():
244
+ if not os.path.exists(save_path):
245
+ os.makedirs(save_path, exist_ok=True)
246
+
247
+ index = len([path for path in os.listdir(save_path)]) + 1
248
+ prefix = str(index).zfill(8)
249
+ video_path = os.path.join(save_path, prefix + ".png")
250
+ image = sample[0]
251
+ image.save(video_path)
252
+
253
+ if ulysses_degree * ring_degree > 1:
254
+ import torch.distributed as dist
255
+ if dist.get_rank() == 0:
256
+ save_results()
257
+ else:
258
+ save_results()
vendor/VideoX-Fun/examples/hunyuanvideo/predict_i2v.py ADDED
@@ -0,0 +1,270 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from diffusers.utils import export_to_video
8
+ from omegaconf import OmegaConf
9
+ from PIL import Image
10
+
11
+ current_file_path = os.path.abspath(__file__)
12
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
13
+ for project_root in project_roots:
14
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
15
+
16
+ from diffusers.schedulers.scheduling_unipc_multistep import \
17
+ UniPCMultistepScheduler
18
+
19
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
20
+ from videox_fun.models import (AutoencoderKLHunyuanVideo, CLIPTextModel, CLIPImageProcessor,
21
+ CLIPTokenizer, HunyuanVideoTransformer3DModel,
22
+ LlavaForConditionalGeneration, LlamaTokenizerFast)
23
+ from videox_fun.models.cache_utils import get_teacache_coefficients
24
+ from videox_fun.pipeline import HunyuanVideoPipeline, HunyuanVideoI2VPipeline
25
+ from videox_fun.utils import (register_auto_device_hook,
26
+ safe_enable_group_offload)
27
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
28
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
29
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
30
+ convert_weight_dtype_wrapper,
31
+ replace_parameters_by_name)
32
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
33
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
34
+ save_videos_grid)
35
+ from videox_fun.utils.utils import get_image
36
+
37
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
38
+ # model_full_load means that the entire model will be moved to the GPU.
39
+ #
40
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
44
+ #
45
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
46
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
47
+ #
48
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
49
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
50
+ #
51
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
52
+ # resulting in slower speeds but saving a large amount of GPU memory.
53
+ GPU_memory_mode = "sequential_cpu_offload"
54
+ # Multi GPUs config
55
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
56
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
57
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
58
+ ulysses_degree = 1
59
+ ring_degree = 1
60
+ # Use FSDP to save more GPU memory in multi gpus.
61
+ fsdp_dit = False
62
+ fsdp_text_encoder = True
63
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
64
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
65
+ compile_dit = False
66
+
67
+ # model path
68
+ model_name = "models/Diffusion_Transformer/HunyuanVideo-I2V"
69
+
70
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
71
+ sampler_name = "Flow"
72
+
73
+ # Load pretrained model if need
74
+ transformer_path = None
75
+ vae_path = None
76
+ lora_path = None
77
+
78
+ # Other params
79
+ sample_size = [480, 832]
80
+ video_length = 81
81
+ fps = 16
82
+
83
+ # Use torch.float16 if GPU does not support torch.bfloat16
84
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
85
+ weight_dtype = torch.bfloat16
86
+ # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
87
+ validation_image_start = "asset/1.png"
88
+
89
+ # prompts
90
+ prompt = "The dog is shaking head. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
91
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
92
+ guidance_scale = 1.0
93
+ seed = 43
94
+ num_inference_steps = 40
95
+ lora_weight = 0.55
96
+ save_path = "samples/hunyuanvideo-videos-i2v"
97
+
98
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
99
+
100
+ transformer = HunyuanVideoTransformer3DModel.from_pretrained(
101
+ os.path.join(model_name, 'transformer'),
102
+ low_cpu_mem_usage=True,
103
+ torch_dtype=weight_dtype,
104
+ )
105
+
106
+ if transformer_path is not None:
107
+ print(f"From checkpoint: {transformer_path}")
108
+ if transformer_path.endswith("safetensors"):
109
+ from safetensors.torch import load_file, safe_open
110
+ state_dict = load_file(transformer_path)
111
+ else:
112
+ state_dict = torch.load(transformer_path, map_location="cpu")
113
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
114
+
115
+ m, u = transformer.load_state_dict(state_dict, strict=False)
116
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
117
+
118
+ # Get Vae
119
+ vae = AutoencoderKLHunyuanVideo.from_pretrained(
120
+ os.path.join(model_name, 'vae')
121
+ ).to(weight_dtype)
122
+
123
+ if vae_path is not None:
124
+ print(f"From checkpoint: {vae_path}")
125
+ if vae_path.endswith("safetensors"):
126
+ from safetensors.torch import load_file, safe_open
127
+ state_dict = load_file(vae_path)
128
+ else:
129
+ state_dict = torch.load(vae_path, map_location="cpu")
130
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
131
+
132
+ m, u = vae.load_state_dict(state_dict, strict=False)
133
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
134
+
135
+ # Get Tokenizer
136
+ tokenizer = LlamaTokenizerFast.from_pretrained(
137
+ os.path.join(model_name, 'tokenizer'),
138
+ )
139
+
140
+ # Get Text encoder
141
+ text_encoder = LlavaForConditionalGeneration.from_pretrained(
142
+ os.path.join(model_name, 'text_encoder'),
143
+ low_cpu_mem_usage=True,
144
+ torch_dtype=weight_dtype,
145
+ )
146
+
147
+ # Get Tokenizer 2
148
+ tokenizer_2 = CLIPTokenizer.from_pretrained(
149
+ os.path.join(model_name, 'tokenizer_2'),
150
+ )
151
+
152
+ # Get Text encoder 2
153
+ text_encoder_2 = CLIPTextModel.from_pretrained(
154
+ os.path.join(model_name, 'text_encoder_2'),
155
+ low_cpu_mem_usage=True,
156
+ torch_dtype=weight_dtype,
157
+ )
158
+
159
+ # Get Image Processor
160
+ image_processor = CLIPImageProcessor.from_pretrained(
161
+ os.path.join(model_name, 'image_processor'),
162
+ )
163
+
164
+ # Get Scheduler
165
+ Chosen_Scheduler = scheduler_dict = {
166
+ "Flow": FlowMatchEulerDiscreteScheduler,
167
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
168
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
169
+ }[sampler_name]
170
+ scheduler = Chosen_Scheduler.from_pretrained(
171
+ os.path.join(model_name, 'scheduler'),
172
+ )
173
+
174
+ # Get Pipeline
175
+ pipeline = HunyuanVideoI2VPipeline(
176
+ transformer=transformer,
177
+ vae=vae,
178
+ tokenizer=tokenizer,
179
+ text_encoder=text_encoder,
180
+ tokenizer_2=tokenizer_2,
181
+ text_encoder_2=text_encoder_2,
182
+ scheduler=scheduler,
183
+ image_processor=image_processor,
184
+ )
185
+ if ulysses_degree > 1 or ring_degree > 1:
186
+ from functools import partial
187
+ transformer.enable_multi_gpus_inference()
188
+ if fsdp_dit:
189
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
190
+ pipeline.transformer = shard_fn(pipeline.transformer)
191
+ print("Add FSDP DIT")
192
+ if fsdp_text_encoder:
193
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers)
194
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
195
+ print("Add FSDP TEXT ENCODER")
196
+
197
+ if compile_dit:
198
+ for i in range(len(pipeline.transformer.blocks)):
199
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
200
+ print("Add Compile")
201
+
202
+ if GPU_memory_mode == "sequential_cpu_offload":
203
+ pipeline.enable_sequential_cpu_offload(device=device)
204
+ elif GPU_memory_mode == "model_group_offload":
205
+ register_auto_device_hook(pipeline.transformer)
206
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
207
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
208
+ convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
209
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
210
+ pipeline.enable_model_cpu_offload(device=device)
211
+ elif GPU_memory_mode == "model_cpu_offload":
212
+ pipeline.enable_model_cpu_offload(device=device)
213
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
214
+ convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
215
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
216
+ pipeline.to(device=device)
217
+ else:
218
+ pipeline.to(device=device)
219
+
220
+ generator = torch.Generator(device=device).manual_seed(seed)
221
+
222
+ if lora_path is not None:
223
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
224
+
225
+ with torch.no_grad():
226
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
227
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
228
+
229
+ # open
230
+ image = get_image(validation_image_start)
231
+
232
+ sample = pipeline(
233
+ prompt,
234
+ image = image,
235
+ num_frames = video_length,
236
+ negative_prompt = negative_prompt,
237
+ height = sample_size[0],
238
+ width = sample_size[1],
239
+ generator = generator,
240
+ true_cfg_scale = guidance_scale,
241
+ num_inference_steps = num_inference_steps,
242
+ ).videos
243
+
244
+ if lora_path is not None:
245
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
246
+
247
+ def save_results():
248
+ if not os.path.exists(save_path):
249
+ os.makedirs(save_path, exist_ok=True)
250
+
251
+ index = len([path for path in os.listdir(save_path)]) + 1
252
+ prefix = str(index).zfill(8)
253
+ if video_length == 1:
254
+ video_path = os.path.join(save_path, prefix + ".png")
255
+
256
+ image = sample[0, :, 0]
257
+ image = image.transpose(0, 1).transpose(1, 2)
258
+ image = (image * 255).numpy().astype(np.uint8)
259
+ image = Image.fromarray(image)
260
+ image.save(video_path)
261
+ else:
262
+ video_path = os.path.join(save_path, prefix + ".mp4")
263
+ save_videos_grid(sample, video_path, fps=fps)
264
+
265
+ if ulysses_degree * ring_degree > 1:
266
+ import torch.distributed as dist
267
+ if dist.get_rank() == 0:
268
+ save_results()
269
+ else:
270
+ save_results()
vendor/VideoX-Fun/examples/hunyuanvideo/predict_t2v.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from diffusers.utils import export_to_video
8
+ from omegaconf import OmegaConf
9
+ from PIL import Image
10
+
11
+ current_file_path = os.path.abspath(__file__)
12
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
13
+ for project_root in project_roots:
14
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
15
+
16
+ from diffusers.schedulers.scheduling_unipc_multistep import \
17
+ UniPCMultistepScheduler
18
+
19
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
20
+ from videox_fun.models import (AutoencoderKLHunyuanVideo, CLIPTextModel,
21
+ CLIPTokenizer, HunyuanVideoTransformer3DModel,
22
+ LlamaModel, LlamaTokenizerFast)
23
+ from videox_fun.models.cache_utils import get_teacache_coefficients
24
+ from videox_fun.pipeline import HunyuanVideoPipeline
25
+ from videox_fun.utils import (register_auto_device_hook,
26
+ safe_enable_group_offload)
27
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
28
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
29
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
30
+ convert_weight_dtype_wrapper,
31
+ replace_parameters_by_name)
32
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
33
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
34
+ save_videos_grid)
35
+
36
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
37
+ # model_full_load means that the entire model will be moved to the GPU.
38
+ #
39
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
40
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
41
+ #
42
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
43
+ #
44
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
45
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
46
+ #
47
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
48
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
49
+ #
50
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
51
+ # resulting in slower speeds but saving a large amount of GPU memory.
52
+ GPU_memory_mode = "sequential_cpu_offload"
53
+ # Multi GPUs config
54
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
55
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
56
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
57
+ ulysses_degree = 1
58
+ ring_degree = 1
59
+ # Use FSDP to save more GPU memory in multi gpus.
60
+ fsdp_dit = False
61
+ fsdp_text_encoder = True
62
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
63
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
64
+ compile_dit = False
65
+
66
+ # model path
67
+ model_name = "models/Diffusion_Transformer/HunyuanVideo"
68
+
69
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
70
+ sampler_name = "Flow"
71
+
72
+ # Load pretrained model if need
73
+ transformer_path = None
74
+ vae_path = None
75
+ lora_path = None
76
+
77
+ # Other params
78
+ sample_size = [832, 480]
79
+ video_length = 81
80
+ fps = 16
81
+
82
+ # Use torch.float16 if GPU does not support torch.bfloat16
83
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
84
+ weight_dtype = torch.bfloat16
85
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
86
+ negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. "
87
+ guidance_scale = 1.0
88
+ seed = 43
89
+ num_inference_steps = 40
90
+ lora_weight = 0.55
91
+ save_path = "samples/hunyuanvideo-videos-t2v"
92
+
93
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
94
+
95
+ transformer = HunyuanVideoTransformer3DModel.from_pretrained(
96
+ os.path.join(model_name, 'transformer'),
97
+ low_cpu_mem_usage=True,
98
+ torch_dtype=weight_dtype,
99
+ )
100
+
101
+ if transformer_path is not None:
102
+ print(f"From checkpoint: {transformer_path}")
103
+ if transformer_path.endswith("safetensors"):
104
+ from safetensors.torch import load_file, safe_open
105
+ state_dict = load_file(transformer_path)
106
+ else:
107
+ state_dict = torch.load(transformer_path, map_location="cpu")
108
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
109
+
110
+ m, u = transformer.load_state_dict(state_dict, strict=False)
111
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
112
+
113
+ # Get Vae
114
+ vae = AutoencoderKLHunyuanVideo.from_pretrained(
115
+ os.path.join(model_name, 'vae')
116
+ ).to(weight_dtype)
117
+
118
+ if vae_path is not None:
119
+ print(f"From checkpoint: {vae_path}")
120
+ if vae_path.endswith("safetensors"):
121
+ from safetensors.torch import load_file, safe_open
122
+ state_dict = load_file(vae_path)
123
+ else:
124
+ state_dict = torch.load(vae_path, map_location="cpu")
125
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
126
+
127
+ m, u = vae.load_state_dict(state_dict, strict=False)
128
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
129
+
130
+ # Get Tokenizer
131
+ tokenizer = LlamaTokenizerFast.from_pretrained(
132
+ os.path.join(model_name, 'tokenizer'),
133
+ )
134
+
135
+ # Get Text encoder
136
+ text_encoder = LlamaModel.from_pretrained(
137
+ os.path.join(model_name, 'text_encoder'),
138
+ low_cpu_mem_usage=True,
139
+ torch_dtype=weight_dtype,
140
+ )
141
+
142
+ # Get Tokenizer 2
143
+ tokenizer_2 = CLIPTokenizer.from_pretrained(
144
+ os.path.join(model_name, 'tokenizer_2'),
145
+ )
146
+
147
+ # Get Text encoder 2
148
+ text_encoder_2 = CLIPTextModel.from_pretrained(
149
+ os.path.join(model_name, 'text_encoder_2'),
150
+ low_cpu_mem_usage=True,
151
+ torch_dtype=weight_dtype,
152
+ )
153
+
154
+ # Get Scheduler
155
+ Chosen_Scheduler = scheduler_dict = {
156
+ "Flow": FlowMatchEulerDiscreteScheduler,
157
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
158
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
159
+ }[sampler_name]
160
+ scheduler = Chosen_Scheduler.from_pretrained(
161
+ os.path.join(model_name, 'scheduler'),
162
+ )
163
+
164
+ # Get Pipeline
165
+ pipeline = HunyuanVideoPipeline(
166
+ transformer=transformer,
167
+ vae=vae,
168
+ tokenizer=tokenizer,
169
+ text_encoder=text_encoder,
170
+ tokenizer_2=tokenizer_2,
171
+ text_encoder_2=text_encoder_2,
172
+ scheduler=scheduler,
173
+ )
174
+ if ulysses_degree > 1 or ring_degree > 1:
175
+ from functools import partial
176
+ transformer.enable_multi_gpus_inference()
177
+ if fsdp_dit:
178
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.single_transformer_blocks))
179
+ pipeline.transformer = shard_fn(pipeline.transformer)
180
+ print("Add FSDP DIT")
181
+ if fsdp_text_encoder:
182
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.layers)
183
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
184
+ print("Add FSDP TEXT ENCODER")
185
+
186
+ if compile_dit:
187
+ for i in range(len(pipeline.transformer.blocks)):
188
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
189
+ print("Add Compile")
190
+
191
+ if GPU_memory_mode == "sequential_cpu_offload":
192
+ pipeline.enable_sequential_cpu_offload(device=device)
193
+ elif GPU_memory_mode == "model_group_offload":
194
+ register_auto_device_hook(pipeline.transformer)
195
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
196
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
197
+ convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
198
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
199
+ pipeline.enable_model_cpu_offload(device=device)
200
+ elif GPU_memory_mode == "model_cpu_offload":
201
+ pipeline.enable_model_cpu_offload(device=device)
202
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
203
+ convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
204
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
205
+ pipeline.to(device=device)
206
+ else:
207
+ pipeline.to(device=device)
208
+
209
+ generator = torch.Generator(device=device).manual_seed(seed)
210
+
211
+ if lora_path is not None:
212
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
213
+
214
+ with torch.no_grad():
215
+ video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
216
+ latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
217
+
218
+ sample = pipeline(
219
+ prompt,
220
+ num_frames = video_length,
221
+ negative_prompt = negative_prompt,
222
+ height = sample_size[0],
223
+ width = sample_size[1],
224
+ generator = generator,
225
+ true_cfg_scale = guidance_scale,
226
+ num_inference_steps = num_inference_steps,
227
+ ).videos
228
+
229
+ if lora_path is not None:
230
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
231
+
232
+ def save_results():
233
+ if not os.path.exists(save_path):
234
+ os.makedirs(save_path, exist_ok=True)
235
+
236
+ index = len([path for path in os.listdir(save_path)]) + 1
237
+ prefix = str(index).zfill(8)
238
+ if video_length == 1:
239
+ video_path = os.path.join(save_path, prefix + ".png")
240
+
241
+ image = sample[0, :, 0]
242
+ image = image.transpose(0, 1).transpose(1, 2)
243
+ image = (image * 255).numpy().astype(np.uint8)
244
+ image = Image.fromarray(image)
245
+ image.save(video_path)
246
+ else:
247
+ video_path = os.path.join(save_path, prefix + ".mp4")
248
+ save_videos_grid(sample, video_path, fps=fps)
249
+
250
+ if ulysses_degree * ring_degree > 1:
251
+ import torch.distributed as dist
252
+ if dist.get_rank() == 0:
253
+ save_results()
254
+ else:
255
+ save_results()
vendor/VideoX-Fun/examples/infinitetalk/predict_s2v.py ADDED
@@ -0,0 +1,319 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models.cache_utils import get_teacache_coefficients
17
+ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
18
+ AutoTokenizer, CLIPModel,
19
+ InfiniteTalkTransformer3DModel, InfiniteTalkAudioEncoder,
20
+ WanT5EncoderModel)
21
+ from videox_fun.pipeline import InfiniteTalkPipeline
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils import (register_auto_device_hook,
25
+ safe_enable_group_offload)
26
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
27
+ convert_weight_dtype_wrapper,
28
+ replace_parameters_by_name)
29
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
30
+ from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image,
31
+ get_video_to_video_latent,
32
+ merge_video_audio, save_videos_grid)
33
+
34
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
35
+ # model_full_load means that the entire model will be moved to the GPU.
36
+ #
37
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
38
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
39
+ #
40
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
41
+ #
42
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
43
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
44
+ #
45
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
46
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
47
+ #
48
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
49
+ # resulting in slower speeds but saving a large amount of GPU memory.
50
+ GPU_memory_mode = "sequential_cpu_offload"
51
+ # Multi GPUs config
52
+ ulysses_degree = 1
53
+ ring_degree = 1
54
+ # Use FSDP to save more GPU memory in multi gpus.
55
+ fsdp_dit = False
56
+ fsdp_text_encoder = True
57
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
58
+ # The compile_dit is not compatible with sequential_cpu_offload.
59
+ compile_dit = False
60
+
61
+ # Support TeaCache.
62
+ enable_teacache = True
63
+ # Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
64
+ # but it may cause slight differences between the generated content and the original content.
65
+ # # --------------------------------------------------------------------------------------------------- #
66
+ # | Model Name | threshold | Model Name | threshold | Model Name | threshold |
67
+ # | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
68
+ # | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
69
+ # # --------------------------------------------------------------------------------------------------- #
70
+ teacache_threshold = 0.20
71
+ # The number of steps to skip TeaCache at the beginning of the inference process, which can
72
+ # reduce the impact of TeaCache on generated video quality.
73
+ num_skip_start_steps = 5
74
+ # Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
75
+ teacache_offload = False
76
+
77
+ # Config and model path
78
+ config_path = "config/wan2.1/wan_civitai.yaml"
79
+ # model path
80
+ model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-480P"
81
+ model_name_audio = "models/Diffusion_Transformer/chinese-wav2vec2-base/"
82
+
83
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
84
+ sampler_name = "Flow"
85
+ shift = 5.0
86
+
87
+ # Load pretrained model if need
88
+ transformer_path = "models/Personalized_Model/infinitetalk.safetensors"
89
+ vae_path = None
90
+ lora_path = None
91
+
92
+ # Other params
93
+ sample_size = [832, 480]
94
+ segment_frame_length = 81
95
+ fps = 25
96
+
97
+ # Use torch.float16 if GPU does not support torch.bfloat16
98
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
99
+ weight_dtype = torch.bfloat16
100
+ # The path of the reference image
101
+ ref_image = "asset/8.png"
102
+ # The path of the audio
103
+ audio_path = "asset/talk.wav"
104
+
105
+ # prompts
106
+ prompt = "一个人在说话。"
107
+ negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
108
+ guidance_scale = 5.0
109
+ audio_guide_scale = 4.0
110
+ seed = 43
111
+ num_inference_steps = 40
112
+ lora_weight = 0.55
113
+ save_path = "samples/infitetalk-videos"
114
+
115
+ # InfiniteTalk specific parameters
116
+ max_frames_num = 500 # Total frames to generate
117
+ color_correction_strength = 1 # Color correction strength (0.0-1.0)
118
+ use_apg = False # Use Adaptive Projected Guidance
119
+ apg_momentum = 0.5
120
+ apg_norm_threshold = 1.0
121
+
122
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
123
+ config = OmegaConf.load(config_path)
124
+
125
+ transformer = InfiniteTalkTransformer3DModel.from_pretrained(
126
+ os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
127
+ transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
128
+ low_cpu_mem_usage=True,
129
+ torch_dtype=weight_dtype,
130
+ )
131
+
132
+ if transformer_path is not None:
133
+ print(f"From checkpoint: {transformer_path}")
134
+ if transformer_path.endswith("safetensors"):
135
+ from safetensors.torch import load_file
136
+ state_dict = load_file(transformer_path)
137
+ else:
138
+ state_dict = torch.load(transformer_path, map_location="cpu")
139
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
140
+
141
+ m, u = transformer.load_state_dict(state_dict, strict=False)
142
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
143
+
144
+ # Get Vae
145
+ vae = AutoencoderKLWan.from_pretrained(
146
+ os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
147
+ additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
148
+ ).to(weight_dtype)
149
+
150
+ if vae_path is not None:
151
+ print(f"From checkpoint: {vae_path}")
152
+ if vae_path.endswith("safetensors"):
153
+ from safetensors.torch import load_file, safe_open
154
+ state_dict = load_file(vae_path)
155
+ else:
156
+ state_dict = torch.load(vae_path, map_location="cpu")
157
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
158
+
159
+ m, u = vae.load_state_dict(state_dict, strict=False)
160
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
161
+
162
+ # Get Tokenizer
163
+ tokenizer = AutoTokenizer.from_pretrained(
164
+ os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
165
+ )
166
+
167
+ # Get Text encoder
168
+ text_encoder = WanT5EncoderModel.from_pretrained(
169
+ os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
170
+ additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
171
+ low_cpu_mem_usage=True,
172
+ torch_dtype=weight_dtype,
173
+ )
174
+ text_encoder = text_encoder.eval()
175
+
176
+ # Initialize InfiniteTalk audio encoder for real-time audio encoding
177
+ # Uses Wav2Vec2Model (not Wav2Vec2ForCTC) matching original InfiniteTalk implementation
178
+ audio_encoder = InfiniteTalkAudioEncoder(
179
+ model_name_audio, "cpu"
180
+ )
181
+
182
+ # Get Clip Image Encoder
183
+ clip_image_encoder = CLIPModel.from_pretrained(
184
+ os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
185
+ ).to(weight_dtype)
186
+ clip_image_encoder = clip_image_encoder.eval()
187
+
188
+ # Get Scheduler
189
+ Chosen_Scheduler = scheduler_dict = {
190
+ "Flow": FlowMatchEulerDiscreteScheduler,
191
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
192
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
193
+ }[sampler_name]
194
+ if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
195
+ config['scheduler_kwargs']['shift'] = 1
196
+ scheduler = Chosen_Scheduler(
197
+ **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
198
+ )
199
+
200
+ # Get Pipeline
201
+ pipeline = InfiniteTalkPipeline(
202
+ transformer=transformer,
203
+ vae=vae,
204
+ tokenizer=tokenizer,
205
+ text_encoder=text_encoder,
206
+ scheduler=scheduler,
207
+ audio_encoder=audio_encoder,
208
+ clip_image_encoder=clip_image_encoder,
209
+ )
210
+ if ulysses_degree > 1 or ring_degree > 1:
211
+ from functools import partial
212
+ transformer.enable_multi_gpus_inference()
213
+ if fsdp_dit:
214
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
215
+ pipeline.transformer = shard_fn(pipeline.transformer)
216
+ print("Add FSDP DIT")
217
+ if fsdp_text_encoder:
218
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
219
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
220
+ print("Add FSDP TEXT ENCODER")
221
+
222
+ if compile_dit:
223
+ for i in range(len(pipeline.transformer.blocks)):
224
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
225
+ print("Add Compile")
226
+
227
+ if GPU_memory_mode == "sequential_cpu_offload":
228
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
229
+ transformer.freqs = transformer.freqs.to(device=device)
230
+ pipeline.enable_sequential_cpu_offload(device=device)
231
+ elif GPU_memory_mode == "model_group_offload":
232
+ register_auto_device_hook(pipeline.transformer)
233
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
234
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
235
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
236
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
237
+ pipeline.enable_model_cpu_offload(device=device)
238
+ elif GPU_memory_mode == "model_cpu_offload":
239
+ pipeline.enable_model_cpu_offload(device=device)
240
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
241
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
242
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
243
+ pipeline.to(device=device)
244
+ else:
245
+ pipeline.to(device=device)
246
+
247
+ coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
248
+ if coefficients is not None:
249
+ print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
250
+ pipeline.transformer.enable_teacache(
251
+ coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
252
+ )
253
+
254
+ generator = torch.Generator(device=device).manual_seed(seed)
255
+
256
+ if lora_path is not None:
257
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
258
+
259
+ with torch.no_grad():
260
+ # For InfiniteTalk, (segment_frame_length - 1) must be divisible by 4
261
+ segment_frame_length = (segment_frame_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio + 1 if segment_frame_length != 1 else 1
262
+ latent_frames = (segment_frame_length - 1) // vae.config.temporal_compression_ratio + 1
263
+
264
+ # Prepare clip_image from original ref_image path
265
+ clip_image = get_image(ref_image)
266
+ ref_image = get_image_latent(ref_image, sample_size=sample_size)
267
+
268
+ sample = pipeline(
269
+ prompt,
270
+ segment_frame_length = segment_frame_length,
271
+ negative_prompt = negative_prompt,
272
+ height = sample_size[0],
273
+ width = sample_size[1],
274
+ generator = generator,
275
+ guidance_scale = guidance_scale,
276
+ audio_guide_scale = audio_guide_scale,
277
+ num_inference_steps = num_inference_steps,
278
+
279
+ ref_image = ref_image,
280
+ clip_image = clip_image, # Pass clip_image
281
+ audio_path = audio_path,
282
+ shift = shift,
283
+ fps = fps,
284
+ max_frames_num = max_frames_num,
285
+ color_correction_strength = color_correction_strength,
286
+ use_apg = use_apg,
287
+ apg_momentum = apg_momentum,
288
+ apg_norm_threshold = apg_norm_threshold,
289
+ ).videos
290
+
291
+ if lora_path is not None:
292
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
293
+
294
+ def save_results():
295
+ if not os.path.exists(save_path):
296
+ os.makedirs(save_path, exist_ok=True)
297
+
298
+ index = len([path for path in os.listdir(save_path)]) + 1
299
+ prefix = str(index).zfill(8)
300
+ if sample.size()[2] == 1:
301
+ video_path = os.path.join(save_path, prefix + ".png")
302
+
303
+ image = sample[0, :, 0]
304
+ image = image.transpose(0, 1).transpose(1, 2)
305
+ image = (image * 255).numpy().astype(np.uint8)
306
+ image = Image.fromarray(image)
307
+ image.save(video_path)
308
+ else:
309
+ video_path = os.path.join(save_path, prefix + ".mp4")
310
+ save_videos_grid(sample, video_path, fps=fps)
311
+
312
+ merge_video_audio(video_path=video_path, audio_path=audio_path)
313
+
314
+ if ulysses_degree * ring_degree > 1:
315
+ import torch.distributed as dist
316
+ if dist.get_rank() == 0:
317
+ save_results()
318
+ else:
319
+ save_results()
vendor/VideoX-Fun/examples/lens/predict_t2i.py ADDED
@@ -0,0 +1,226 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import torch
5
+ from diffusers import FlowMatchEulerDiscreteScheduler
6
+
7
+ current_file_path = os.path.abspath(__file__)
8
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
9
+ for project_root in project_roots:
10
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
11
+
12
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
13
+ from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer,
14
+ LensGptOssEncoder, LensTransformer2DModel)
15
+ from videox_fun.pipeline import LensPipeline
16
+ from videox_fun.utils import (register_auto_device_hook,
17
+ safe_enable_group_offload)
18
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
19
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
20
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
21
+ convert_weight_dtype_wrapper)
22
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
23
+
24
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
25
+ # model_full_load means that the entire model will be moved to the GPU.
26
+ #
27
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
28
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
29
+ #
30
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
31
+ #
32
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
33
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
34
+ #
35
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
36
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
37
+ #
38
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
39
+ # resulting in slower speeds but saving a large amount of GPU memory.
40
+ GPU_memory_mode = "model_cpu_offload"
41
+ # Multi GPUs config
42
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
43
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
44
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
45
+ ulysses_degree = 1
46
+ ring_degree = 1
47
+ # Use FSDP to save more GPU memory in multi gpus.
48
+ fsdp_dit = False
49
+ fsdp_text_encoder = False
50
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
51
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
52
+ compile_dit = False
53
+
54
+ # model path
55
+ model_name = "models/Diffusion_Transformer/Lens"
56
+
57
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
58
+ sampler_name = "Flow"
59
+
60
+ # Load pretrained model if need
61
+ transformer_path = None
62
+ vae_path = None
63
+ lora_path = None
64
+
65
+ # Other params
66
+ sample_size = [1728, 992]
67
+
68
+ # Use torch.float16 if GPU does not support torch.bfloat16
69
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
70
+ weight_dtype = torch.bfloat16
71
+ # Set to True on A100/V100 to dequantize MXFP4 GPT-OSS weights.
72
+ dequantize_mxfp4 = False
73
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
74
+ negative_prompt = " "
75
+ guidance_scale = 4.5
76
+ seed = 43
77
+ num_inference_steps = 40
78
+ lora_weight = 0.55
79
+ save_path = "samples/lens-t2i"
80
+
81
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
82
+
83
+ # Get transformer
84
+ transformer = LensTransformer2DModel.from_pretrained(
85
+ model_name,
86
+ subfolder="transformer",
87
+ low_cpu_mem_usage=True,
88
+ torch_dtype=weight_dtype,
89
+ ).to(weight_dtype)
90
+
91
+ if transformer_path is not None:
92
+ print(f"From checkpoint: {transformer_path}")
93
+ if transformer_path.endswith("safetensors"):
94
+ from safetensors.torch import load_file, safe_open
95
+ state_dict = load_file(transformer_path)
96
+ else:
97
+ state_dict = torch.load(transformer_path, map_location="cpu")
98
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
99
+
100
+ m, u = transformer.load_state_dict(state_dict, strict=False)
101
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
102
+
103
+ # Get Vae
104
+ vae = AutoencoderKLFlux2.from_pretrained(
105
+ model_name,
106
+ subfolder="vae",
107
+ ).to(weight_dtype)
108
+
109
+ if vae_path is not None:
110
+ print(f"From checkpoint: {vae_path}")
111
+ if vae_path.endswith("safetensors"):
112
+ from safetensors.torch import load_file, safe_open
113
+ state_dict = load_file(vae_path)
114
+ else:
115
+ state_dict = torch.load(vae_path, map_location="cpu")
116
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
117
+
118
+ m, u = vae.load_state_dict(state_dict, strict=False)
119
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
120
+
121
+ # Get tokenizer and text_encoder
122
+ tokenizer = AutoTokenizer.from_pretrained(
123
+ model_name, subfolder="tokenizer"
124
+ )
125
+ text_encoder_kwargs = {"subfolder": "text_encoder", "torch_dtype": weight_dtype}
126
+ try:
127
+ from transformers import Mxfp4Config
128
+ text_encoder_kwargs["quantization_config"] = Mxfp4Config(
129
+ dequantize=dequantize_mxfp4
130
+ )
131
+ except ImportError:
132
+ pass # Older transformers without Mxfp4Config
133
+
134
+ text_encoder = LensGptOssEncoder.from_pretrained(
135
+ model_name, **text_encoder_kwargs
136
+ )
137
+
138
+ # Get Scheduler
139
+ Chosen_Scheduler = scheduler_dict = {
140
+ "Flow": FlowMatchEulerDiscreteScheduler,
141
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
142
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
143
+ }[sampler_name]
144
+ scheduler = Chosen_Scheduler.from_pretrained(
145
+ model_name,
146
+ subfolder="scheduler"
147
+ )
148
+
149
+ pipeline = LensPipeline(
150
+ vae=vae,
151
+ tokenizer=tokenizer,
152
+ text_encoder=text_encoder,
153
+ transformer=transformer,
154
+ scheduler=scheduler,
155
+ )
156
+
157
+ if ulysses_degree > 1 or ring_degree > 1:
158
+ from functools import partial
159
+ transformer.enable_multi_gpus_inference()
160
+ if fsdp_dit:
161
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks))
162
+ pipeline.transformer = shard_fn(pipeline.transformer)
163
+ print("Add FSDP DIT")
164
+ if fsdp_text_encoder:
165
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(text_encoder.model.layers))
166
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
167
+ print("Add FSDP TEXT ENCODER")
168
+
169
+ if compile_dit:
170
+ for i in range(len(pipeline.transformer.transformer_blocks)):
171
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
172
+ print("Add Compile")
173
+
174
+ if GPU_memory_mode == "sequential_cpu_offload":
175
+ pipeline.enable_sequential_cpu_offload(device=device)
176
+ elif GPU_memory_mode == "model_group_offload":
177
+ register_auto_device_hook(pipeline.transformer)
178
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
179
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
180
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
181
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
182
+ pipeline.enable_model_cpu_offload(device=device)
183
+ elif GPU_memory_mode == "model_cpu_offload":
184
+ pipeline.enable_model_cpu_offload(device=device)
185
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
186
+ convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
187
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
188
+ pipeline.to(device=device)
189
+ else:
190
+ pipeline.to(device=device)
191
+
192
+ generator = torch.Generator(device=device).manual_seed(seed)
193
+
194
+ if lora_path is not None:
195
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
196
+
197
+ with torch.no_grad():
198
+ sample = pipeline(
199
+ prompt,
200
+ negative_prompt = negative_prompt,
201
+ height = sample_size[0],
202
+ width = sample_size[1],
203
+ generator = generator,
204
+ guidance_scale = guidance_scale,
205
+ num_inference_steps = num_inference_steps,
206
+ ).images
207
+
208
+ if lora_path is not None:
209
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
210
+
211
+ def save_results():
212
+ if not os.path.exists(save_path):
213
+ os.makedirs(save_path, exist_ok=True)
214
+
215
+ index = len([path for path in os.listdir(save_path)]) + 1
216
+ prefix = str(index).zfill(8)
217
+ video_path = os.path.join(save_path, prefix + ".png")
218
+ image = sample[0]
219
+ image.save(video_path)
220
+
221
+ if ulysses_degree * ring_degree > 1:
222
+ import torch.distributed as dist
223
+ if dist.get_rank() == 0:
224
+ save_results()
225
+ else:
226
+ save_results()
vendor/VideoX-Fun/examples/longcatvideo/predict_i2v.py ADDED
@@ -0,0 +1,247 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLLongCatVideo, UMT5EncoderModel, AutoTokenizer,
17
+ LongCatVideoTransformer3DModel)
18
+ from videox_fun.models.cache_utils import get_teacache_coefficients
19
+ from videox_fun.pipeline import LongCatVideoPipeline
20
+ from videox_fun.utils import (register_auto_device_hook,
21
+ safe_enable_group_offload)
22
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
23
+ convert_weight_dtype_wrapper)
24
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
25
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
26
+ save_videos_grid)
27
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
28
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
29
+
30
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
31
+ # model_full_load means that the entire model will be moved to the GPU.
32
+ #
33
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
37
+ #
38
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
42
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
43
+ #
44
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
45
+ # resulting in slower speeds but saving a large amount of GPU memory.
46
+ GPU_memory_mode = "model_group_offload"
47
+ # Multi GPUs config
48
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
49
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
50
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
51
+ ulysses_degree = 1
52
+ ring_degree = 1
53
+ # Use FSDP to save more GPU memory in multi gpus.
54
+ fsdp_dit = False
55
+ fsdp_text_encoder = True
56
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
57
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
58
+ compile_dit = False
59
+
60
+ # model path
61
+ model_name = "models/Diffusion_Transformer/LongCat-Video"
62
+
63
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
64
+ sampler_name = "Flow"
65
+
66
+ # Load pretrained model if need
67
+ transformer_path = None
68
+ vae_path = None
69
+ lora_path = None
70
+
71
+ # Other params
72
+ sample_size = [480, 832]
73
+ video_length = 81
74
+ fps = 16
75
+
76
+ # Use torch.float16 if GPU does not support torch.bfloat16
77
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
78
+ weight_dtype = torch.bfloat16
79
+ validation_image_start = "asset/1.png"
80
+
81
+ # Prompt
82
+ prompt = "The dog is shaking head. The video is of high quality, and the view is very clear. High quality, masterpiece, best quality, highres, ultra-detailed, fantastic."
83
+ negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
84
+ guidance_scale = 4.0
85
+ seed = 43
86
+ num_inference_steps = 25
87
+ lora_weight = 0.55
88
+ save_path = "samples/longcat-videos-i2v"
89
+
90
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
91
+
92
+ transformer = LongCatVideoTransformer3DModel.from_pretrained(
93
+ os.path.join(model_name, "dit"),
94
+ low_cpu_mem_usage=True,
95
+ torch_dtype=weight_dtype, cp_split_hw=[1, 1]
96
+ )
97
+
98
+ if transformer_path is not None:
99
+ print(f"From checkpoint: {transformer_path}")
100
+ if transformer_path.endswith("safetensors"):
101
+ from safetensors.torch import load_file, safe_open
102
+ state_dict = load_file(transformer_path)
103
+ else:
104
+ state_dict = torch.load(transformer_path, map_location="cpu")
105
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
106
+
107
+ m, u = transformer.load_state_dict(state_dict, strict=False)
108
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
109
+
110
+ # Get Vae
111
+ vae = AutoencoderKLLongCatVideo.from_pretrained(
112
+ os.path.join(model_name, "vae"),
113
+ ).to(weight_dtype)
114
+
115
+ if vae_path is not None:
116
+ print(f"From checkpoint: {vae_path}")
117
+ if vae_path.endswith("safetensors"):
118
+ from safetensors.torch import load_file, safe_open
119
+ state_dict = load_file(vae_path)
120
+ else:
121
+ state_dict = torch.load(vae_path, map_location="cpu")
122
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
123
+
124
+ m, u = vae.load_state_dict(state_dict, strict=False)
125
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
126
+
127
+ # Get Tokenizer
128
+ tokenizer = AutoTokenizer.from_pretrained(
129
+ os.path.join(model_name, "tokenizer"),
130
+ )
131
+
132
+ # Get Text encoder
133
+ text_encoder = UMT5EncoderModel.from_pretrained(
134
+ os.path.join(model_name, "text_encoder"),
135
+ low_cpu_mem_usage=True,
136
+ torch_dtype=weight_dtype,
137
+ )
138
+
139
+ # Get Scheduler
140
+ Chosen_Scheduler = scheduler_dict = {
141
+ "Flow": FlowMatchEulerDiscreteScheduler,
142
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
143
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
144
+ }[sampler_name]
145
+ scheduler = Chosen_Scheduler.from_pretrained(
146
+ model_name,
147
+ subfolder="scheduler"
148
+ )
149
+
150
+ # Get Pipeline
151
+ pipeline = LongCatVideoPipeline(
152
+ transformer=transformer,
153
+ vae=vae,
154
+ tokenizer=tokenizer,
155
+ text_encoder=text_encoder,
156
+ scheduler=scheduler,
157
+ )
158
+ if ulysses_degree > 1 or ring_degree > 1:
159
+ from functools import partial
160
+ transformer.enable_multi_gpus_inference()
161
+ if fsdp_dit:
162
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
163
+ pipeline.transformer = shard_fn(pipeline.transformer)
164
+ print("Add FSDP DIT")
165
+ if fsdp_text_encoder:
166
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
167
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
168
+ print("Add FSDP TEXT ENCODER")
169
+
170
+ if compile_dit:
171
+ for i in range(len(pipeline.transformer.blocks)):
172
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
173
+ print("Add Compile")
174
+
175
+ if GPU_memory_mode == "sequential_cpu_offload":
176
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
177
+ pipeline.enable_sequential_cpu_offload(device=device)
178
+ elif GPU_memory_mode == "model_group_offload":
179
+ register_auto_device_hook(pipeline.transformer)
180
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
181
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
182
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
183
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
184
+ pipeline.enable_model_cpu_offload(device=device)
185
+ elif GPU_memory_mode == "model_cpu_offload":
186
+ pipeline.enable_model_cpu_offload(device=device)
187
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
188
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
189
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
190
+ pipeline.to(device=device)
191
+ else:
192
+ pipeline.to(device=device)
193
+
194
+ generator = torch.Generator(device=device).manual_seed(seed)
195
+
196
+ if lora_path is not None:
197
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
198
+
199
+ with torch.no_grad():
200
+ video_length = int((video_length - 1) // vae.scale_factor_temporal * vae.scale_factor_temporal) + 1 if video_length != 1 else 1
201
+ latent_frames = (video_length - 1) // vae.scale_factor_temporal + 1
202
+
203
+ if validation_image_start is not None:
204
+ input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
205
+ else:
206
+ input_video, input_video_mask, clip_image = None, None, None
207
+
208
+ sample = pipeline(
209
+ prompt,
210
+ num_frames = video_length,
211
+ negative_prompt = negative_prompt,
212
+ height = sample_size[0],
213
+ width = sample_size[1],
214
+ generator = generator,
215
+ guidance_scale = guidance_scale,
216
+ num_inference_steps = num_inference_steps,
217
+ video = input_video,
218
+ mask_video = input_video_mask,
219
+ ).videos
220
+
221
+ if lora_path is not None:
222
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
223
+
224
+ def save_results():
225
+ if not os.path.exists(save_path):
226
+ os.makedirs(save_path, exist_ok=True)
227
+
228
+ index = len([path for path in os.listdir(save_path)]) + 1
229
+ prefix = str(index).zfill(8)
230
+ if video_length == 1:
231
+ video_path = os.path.join(save_path, prefix + ".png")
232
+
233
+ image = sample[0, :, 0]
234
+ image = image.transpose(0, 1).transpose(1, 2)
235
+ image = (image * 255).numpy().astype(np.uint8)
236
+ image = Image.fromarray(image)
237
+ image.save(video_path)
238
+ else:
239
+ video_path = os.path.join(save_path, prefix + ".mp4")
240
+ save_videos_grid(sample, video_path, fps=fps)
241
+
242
+ if ulysses_degree * ring_degree > 1:
243
+ import torch.distributed as dist
244
+ if dist.get_rank() == 0:
245
+ save_results()
246
+ else:
247
+ save_results()
vendor/VideoX-Fun/examples/longcatvideo/predict_s2v_avatar.py ADDED
@@ -0,0 +1,293 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ from pathlib import Path
4
+
5
+ import numpy as np
6
+ import torch
7
+ from audio_separator.separator import Separator
8
+ from diffusers import FlowMatchEulerDiscreteScheduler
9
+ from omegaconf import OmegaConf
10
+ from PIL import Image
11
+
12
+ current_file_path = os.path.abspath(__file__)
13
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
14
+ for project_root in project_roots:
15
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
16
+
17
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
18
+ from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
19
+ LongCatVideoAudioEncoder,
20
+ LongCatVideoAvatarTransformer3DModel,
21
+ UMT5EncoderModel)
22
+ from videox_fun.models.cache_utils import get_teacache_coefficients
23
+ from videox_fun.pipeline import LongCatVideoAvatarPipeline
24
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
25
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
26
+ from videox_fun.utils import (register_auto_device_hook,
27
+ safe_enable_group_offload)
28
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
29
+ convert_weight_dtype_wrapper,
30
+ replace_parameters_by_name)
31
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
32
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
33
+ merge_video_audio, save_videos_grid)
34
+
35
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
36
+ # model_full_load means that the entire model will be moved to the GPU.
37
+ #
38
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
42
+ #
43
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
44
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
45
+ #
46
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
47
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
48
+ #
49
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
50
+ # resulting in slower speeds but saving a large amount of GPU memory.
51
+ GPU_memory_mode = "model_group_offload"
52
+ # Multi GPUs config
53
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
54
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
55
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
56
+ ulysses_degree = 1
57
+ ring_degree = 1
58
+ # Use FSDP to save more GPU memory in multi gpus.
59
+ fsdp_dit = False
60
+ fsdp_text_encoder = True
61
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
62
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
63
+ compile_dit = False
64
+
65
+ # model path
66
+ model_name = "models/Diffusion_Transformer/LongCat-Video"
67
+ model_name_avatar = "models/Diffusion_Transformer/LongCat-Video-Avatar"
68
+
69
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
70
+ sampler_name = "Flow"
71
+
72
+ # Load pretrained model if need
73
+ transformer_path = None
74
+ vae_path = None
75
+ lora_path = None
76
+
77
+ # Other params
78
+ sample_size = [832, 480]
79
+ video_length = 81
80
+ fps = 16
81
+
82
+ # Start Image
83
+ validation_image_start = "asset/8.png"
84
+
85
+ # Audio params
86
+ audio_path = "asset/talk.wav"
87
+ use_audio_vocal_separator = False
88
+
89
+ # Use torch.float16 if GPU does not support torch.bfloat16
90
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
91
+ weight_dtype = torch.bfloat16
92
+ # Prompt
93
+ prompt = "A young woman with long flowing purple hair stands by the seaside on a sunny day, singing. Wearing a white sleeveless dress with a navy blue bow at the collar, her hair gently sways in the ocean breeze. The sparkling sea, blue sky with white clouds, and pink wildflowers along the shore create a beautiful and vibrant scene."
94
+ negative_prompt = "Close-up, Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
95
+ guidance_scale = 4.5
96
+ seed = 43
97
+ num_inference_steps = 25
98
+ lora_weight = 0.55
99
+ save_path = "samples/longcat-avatar-videos-t2v"
100
+
101
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
102
+
103
+ transformer = LongCatVideoAvatarTransformer3DModel.from_pretrained(
104
+ os.path.join(model_name_avatar, "avatar_single"),
105
+ low_cpu_mem_usage=True,
106
+ torch_dtype=weight_dtype, cp_split_hw=[1, 1]
107
+ )
108
+
109
+ if transformer_path is not None:
110
+ print(f"From checkpoint: {transformer_path}")
111
+ if transformer_path.endswith("safetensors"):
112
+ from safetensors.torch import load_file, safe_open
113
+ state_dict = load_file(transformer_path)
114
+ else:
115
+ state_dict = torch.load(transformer_path, map_location="cpu")
116
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
117
+
118
+ m, u = transformer.load_state_dict(state_dict, strict=False)
119
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
120
+
121
+ # Get Vae
122
+ vae = AutoencoderKLLongCatVideo.from_pretrained(
123
+ os.path.join(model_name, "vae"),
124
+ ).to(weight_dtype)
125
+
126
+ if vae_path is not None:
127
+ print(f"From checkpoint: {vae_path}")
128
+ if vae_path.endswith("safetensors"):
129
+ from safetensors.torch import load_file, safe_open
130
+ state_dict = load_file(vae_path)
131
+ else:
132
+ state_dict = torch.load(vae_path, map_location="cpu")
133
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
134
+
135
+ m, u = vae.load_state_dict(state_dict, strict=False)
136
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
137
+
138
+ # Get Tokenizer
139
+ tokenizer = AutoTokenizer.from_pretrained(
140
+ os.path.join(model_name, "tokenizer"),
141
+ )
142
+
143
+ # Get Text encoder
144
+ text_encoder = UMT5EncoderModel.from_pretrained(
145
+ os.path.join(model_name, "text_encoder"),
146
+ low_cpu_mem_usage=True,
147
+ torch_dtype=weight_dtype,
148
+ )
149
+
150
+ # Get Audio encoder (for avatar mode)
151
+ audio_encoder = LongCatVideoAudioEncoder(
152
+ os.path.join(model_name_avatar, 'chinese-wav2vec2-base')
153
+ )
154
+ audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
155
+
156
+ # Get Scheduler
157
+ Chosen_Scheduler = scheduler_dict = {
158
+ "Flow": FlowMatchEulerDiscreteScheduler,
159
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
160
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
161
+ }[sampler_name]
162
+ scheduler = Chosen_Scheduler.from_pretrained(
163
+ model_name,
164
+ subfolder="scheduler"
165
+ )
166
+
167
+ # Get Pipeline
168
+ pipeline = LongCatVideoAvatarPipeline(
169
+ transformer=transformer,
170
+ vae=vae,
171
+ tokenizer=tokenizer,
172
+ text_encoder=text_encoder,
173
+ scheduler=scheduler,
174
+ audio_encoder=audio_encoder,
175
+ )
176
+ if ulysses_degree > 1 or ring_degree > 1:
177
+ from functools import partial
178
+ transformer.enable_multi_gpus_inference()
179
+ if fsdp_dit:
180
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
181
+ pipeline.transformer = shard_fn(pipeline.transformer)
182
+ print("Add FSDP DIT")
183
+ if fsdp_text_encoder:
184
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
185
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
186
+ print("Add FSDP TEXT ENCODER")
187
+
188
+ if compile_dit:
189
+ for i in range(len(pipeline.transformer.blocks)):
190
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
191
+ print("Add Compile")
192
+
193
+ if GPU_memory_mode == "sequential_cpu_offload":
194
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
195
+ pipeline.enable_sequential_cpu_offload(device=device)
196
+ elif GPU_memory_mode == "model_group_offload":
197
+ register_auto_device_hook(pipeline.transformer)
198
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
199
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
200
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
201
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
202
+ pipeline.enable_model_cpu_offload(device=device)
203
+ elif GPU_memory_mode == "model_cpu_offload":
204
+ pipeline.enable_model_cpu_offload(device=device)
205
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
206
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
207
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
208
+ pipeline.to(device=device)
209
+ else:
210
+ pipeline.to(device=device)
211
+
212
+ generator = torch.Generator(device=device).manual_seed(seed)
213
+
214
+ if lora_path is not None:
215
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
216
+
217
+ # Get Vocal separator
218
+ if use_audio_vocal_separator:
219
+ vocal_separator_path = os.path.join(model_name_avatar, 'vocal_separator/Kim_Vocal_2.onnx')
220
+ audio_output_dir_temp = Path("./audio_temp_file")
221
+ audio_output_dir_temp.mkdir(parents=True, exist_ok=True)
222
+
223
+ vocal_separator = Separator(
224
+ output_dir=audio_output_dir_temp / "vocals",
225
+ output_single_stem="vocals",
226
+ model_file_dir=os.path.dirname(vocal_separator_path),
227
+ )
228
+ vocal_separator.load_model(os.path.basename(vocal_separator_path))
229
+
230
+ # Process audio if provided
231
+ audio_emb = None
232
+ if audio_path is not None:
233
+ # Extract vocal from audio
234
+ outputs = vocal_separator.separate(audio_path)
235
+ if len(outputs) > 0:
236
+ temp_vocal_path = audio_output_dir_temp / "vocals" / outputs[0]
237
+ temp_vocal_path = temp_vocal_path.resolve().as_posix()
238
+ audio_path = temp_vocal_path
239
+
240
+ with torch.no_grad():
241
+ video_length = int((video_length - 1) // vae.scale_factor_temporal * vae.scale_factor_temporal) + 1 if video_length != 1 else 1
242
+ latent_frames = (video_length - 1) // vae.scale_factor_temporal + 1
243
+
244
+ if validation_image_start is not None:
245
+ input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
246
+ else:
247
+ input_video, input_video_mask, clip_image = None, None, None
248
+
249
+ sample = pipeline(
250
+ prompt = prompt,
251
+ num_frames = video_length,
252
+ negative_prompt = negative_prompt,
253
+ height = sample_size[0],
254
+ width = sample_size[1],
255
+ generator = generator,
256
+ guidance_scale = guidance_scale,
257
+ num_inference_steps = num_inference_steps,
258
+
259
+ audio_path = audio_path,
260
+ video = input_video,
261
+ mask_video = input_video_mask,
262
+ fps = fps,
263
+ ).videos
264
+
265
+ if lora_path is not None:
266
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
267
+
268
+ def save_results():
269
+ if not os.path.exists(save_path):
270
+ os.makedirs(save_path, exist_ok=True)
271
+
272
+ index = len([path for path in os.listdir(save_path)]) + 1
273
+ prefix = str(index).zfill(8)
274
+ if video_length == 1:
275
+ video_path = os.path.join(save_path, prefix + ".png")
276
+
277
+ image = sample[0, :, 0]
278
+ image = image.transpose(0, 1).transpose(1, 2)
279
+ image = (image * 255).numpy().astype(np.uint8)
280
+ image = Image.fromarray(image)
281
+ image.save(video_path)
282
+ else:
283
+ video_path = os.path.join(save_path, prefix + ".mp4")
284
+ save_videos_grid(sample, video_path, fps=fps)
285
+
286
+ merge_video_audio(video_path=video_path, audio_path=audio_path)
287
+
288
+ if ulysses_degree * ring_degree > 1:
289
+ import torch.distributed as dist
290
+ if dist.get_rank() == 0:
291
+ save_results()
292
+ else:
293
+ save_results()
vendor/VideoX-Fun/examples/longcatvideo/predict_t2v.py ADDED
@@ -0,0 +1,239 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from omegaconf import OmegaConf
8
+ from PIL import Image
9
+
10
+ current_file_path = os.path.abspath(__file__)
11
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
12
+ for project_root in project_roots:
13
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
14
+
15
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
16
+ from videox_fun.models import (AutoencoderKLLongCatVideo, UMT5EncoderModel, AutoTokenizer,
17
+ LongCatVideoTransformer3DModel)
18
+ from videox_fun.models.cache_utils import get_teacache_coefficients
19
+ from videox_fun.pipeline import LongCatVideoPipeline
20
+ from videox_fun.utils import (register_auto_device_hook,
21
+ safe_enable_group_offload)
22
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
23
+ convert_weight_dtype_wrapper)
24
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
25
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
26
+ save_videos_grid)
27
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
28
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
29
+
30
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
31
+ # model_full_load means that the entire model will be moved to the GPU.
32
+ #
33
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
37
+ #
38
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
42
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
43
+ #
44
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
45
+ # resulting in slower speeds but saving a large amount of GPU memory.
46
+ GPU_memory_mode = "model_group_offload"
47
+ # Multi GPUs config
48
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
49
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
50
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
51
+ ulysses_degree = 1
52
+ ring_degree = 1
53
+ # Use FSDP to save more GPU memory in multi gpus.
54
+ fsdp_dit = False
55
+ fsdp_text_encoder = True
56
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
57
+ # The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
58
+ compile_dit = False
59
+
60
+ # model path
61
+ model_name = "models/Diffusion_Transformer/LongCat-Video"
62
+
63
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
64
+ sampler_name = "Flow"
65
+
66
+ # Load pretrained model if need
67
+ transformer_path = None
68
+ vae_path = None
69
+ lora_path = None
70
+
71
+ # Other params
72
+ sample_size = [832, 480]
73
+ video_length = 81
74
+ fps = 16
75
+
76
+ # Use torch.float16 if GPU does not support torch.bfloat16
77
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
78
+ weight_dtype = torch.bfloat16
79
+ # Prompt
80
+ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_background, jewelry, lips, long_hair, looking_at_viewer, nose, piercing, realistic, red_lips, solo, upper_body"
81
+ negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
82
+ guidance_scale = 4.0
83
+ seed = 43
84
+ num_inference_steps = 25
85
+ lora_weight = 0.55
86
+ save_path = "samples/longcat-videos-t2v"
87
+
88
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
89
+
90
+ transformer = LongCatVideoTransformer3DModel.from_pretrained(
91
+ os.path.join(model_name, "dit"),
92
+ low_cpu_mem_usage=True,
93
+ torch_dtype=weight_dtype, cp_split_hw=[1, 1]
94
+ )
95
+
96
+ if transformer_path is not None:
97
+ print(f"From checkpoint: {transformer_path}")
98
+ if transformer_path.endswith("safetensors"):
99
+ from safetensors.torch import load_file, safe_open
100
+ state_dict = load_file(transformer_path)
101
+ else:
102
+ state_dict = torch.load(transformer_path, map_location="cpu")
103
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
104
+
105
+ m, u = transformer.load_state_dict(state_dict, strict=False)
106
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
107
+
108
+ # Get Vae
109
+ vae = AutoencoderKLLongCatVideo.from_pretrained(
110
+ os.path.join(model_name, "vae"),
111
+ ).to(weight_dtype)
112
+
113
+ if vae_path is not None:
114
+ print(f"From checkpoint: {vae_path}")
115
+ if vae_path.endswith("safetensors"):
116
+ from safetensors.torch import load_file, safe_open
117
+ state_dict = load_file(vae_path)
118
+ else:
119
+ state_dict = torch.load(vae_path, map_location="cpu")
120
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
121
+
122
+ m, u = vae.load_state_dict(state_dict, strict=False)
123
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
124
+
125
+ # Get Tokenizer
126
+ tokenizer = AutoTokenizer.from_pretrained(
127
+ os.path.join(model_name, "tokenizer"),
128
+ )
129
+
130
+ # Get Text encoder
131
+ text_encoder = UMT5EncoderModel.from_pretrained(
132
+ os.path.join(model_name, "text_encoder"),
133
+ low_cpu_mem_usage=True,
134
+ torch_dtype=weight_dtype,
135
+ )
136
+
137
+ # Get Scheduler
138
+ Chosen_Scheduler = scheduler_dict = {
139
+ "Flow": FlowMatchEulerDiscreteScheduler,
140
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
141
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
142
+ }[sampler_name]
143
+ scheduler = Chosen_Scheduler.from_pretrained(
144
+ model_name,
145
+ subfolder="scheduler"
146
+ )
147
+
148
+ # Get Pipeline
149
+ pipeline = LongCatVideoPipeline(
150
+ transformer=transformer,
151
+ vae=vae,
152
+ tokenizer=tokenizer,
153
+ text_encoder=text_encoder,
154
+ scheduler=scheduler,
155
+ )
156
+
157
+ if ulysses_degree > 1 or ring_degree > 1:
158
+ from functools import partial
159
+ transformer.enable_multi_gpus_inference()
160
+ if fsdp_dit:
161
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
162
+ pipeline.transformer = shard_fn(pipeline.transformer)
163
+ print("Add FSDP DIT")
164
+ if fsdp_text_encoder:
165
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
166
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
167
+ print("Add FSDP TEXT ENCODER")
168
+
169
+ if compile_dit:
170
+ for i in range(len(pipeline.transformer.blocks)):
171
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
172
+ print("Add Compile")
173
+
174
+ if GPU_memory_mode == "sequential_cpu_offload":
175
+ replace_parameters_by_name(transformer, ["modulation",], device=device)
176
+ pipeline.enable_sequential_cpu_offload(device=device)
177
+ elif GPU_memory_mode == "model_group_offload":
178
+ register_auto_device_hook(pipeline.transformer)
179
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
180
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
181
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
182
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
183
+ pipeline.enable_model_cpu_offload(device=device)
184
+ elif GPU_memory_mode == "model_cpu_offload":
185
+ pipeline.enable_model_cpu_offload(device=device)
186
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
187
+ convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
188
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
189
+ pipeline.to(device=device)
190
+ else:
191
+ pipeline.to(device=device)
192
+
193
+ generator = torch.Generator(device=device).manual_seed(seed)
194
+
195
+ if lora_path is not None:
196
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
197
+
198
+ with torch.no_grad():
199
+ video_length = int((video_length - 1) // vae.scale_factor_temporal * vae.scale_factor_temporal) + 1 if video_length != 1 else 1
200
+ latent_frames = (video_length - 1) // vae.scale_factor_temporal + 1
201
+
202
+ sample = pipeline(
203
+ prompt,
204
+ num_frames = video_length,
205
+ negative_prompt = negative_prompt,
206
+ height = sample_size[0],
207
+ width = sample_size[1],
208
+ generator = generator,
209
+ guidance_scale = guidance_scale,
210
+ num_inference_steps = num_inference_steps,
211
+ ).videos
212
+
213
+ if lora_path is not None:
214
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
215
+
216
+ def save_results():
217
+ if not os.path.exists(save_path):
218
+ os.makedirs(save_path, exist_ok=True)
219
+
220
+ index = len([path for path in os.listdir(save_path)]) + 1
221
+ prefix = str(index).zfill(8)
222
+ if video_length == 1:
223
+ video_path = os.path.join(save_path, prefix + ".png")
224
+
225
+ image = sample[0, :, 0]
226
+ image = image.transpose(0, 1).transpose(1, 2)
227
+ image = (image * 255).numpy().astype(np.uint8)
228
+ image = Image.fromarray(image)
229
+ image.save(video_path)
230
+ else:
231
+ video_path = os.path.join(save_path, prefix + ".mp4")
232
+ save_videos_grid(sample, video_path, fps=fps)
233
+
234
+ if ulysses_degree * ring_degree > 1:
235
+ import torch.distributed as dist
236
+ if dist.get_rank() == 0:
237
+ save_results()
238
+ else:
239
+ save_results()
vendor/VideoX-Fun/examples/ltx2.3/predict_i2v.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
15
+ Gemma3ForConditionalGeneration,
16
+ GemmaTokenizerFast, LTX2TextConnectors, Gemma3Processor,
17
+ LTX2VideoTransformer3DModel, LTX2VocoderWithBWE)
18
+ from videox_fun.pipeline import LTX2I2VPipeline
19
+ from videox_fun.utils import (register_auto_device_hook,
20
+ safe_enable_group_offload)
21
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
25
+ convert_weight_dtype_wrapper,
26
+ replace_parameters_by_name)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
29
+ save_videos_grid,
30
+ save_videos_with_audio_grid)
31
+
32
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
33
+ # model_full_load means that the entire model will be moved to the GPU.
34
+ #
35
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
39
+ #
40
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
44
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
45
+ #
46
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
47
+ # resulting in slower speeds but saving a large amount of GPU memory.
48
+ GPU_memory_mode = "model_group_offload"
49
+ # Multi GPUs config
50
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
51
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
52
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
53
+ ulysses_degree = 1
54
+ ring_degree = 1
55
+ # Use FSDP to save more GPU memory in multi gpus.
56
+ fsdp_dit = False
57
+ fsdp_text_encoder = False
58
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
59
+ # The compile_dit is not compatible with sequential_cpu_offload.
60
+ compile_dit = False
61
+
62
+ # model path
63
+ model_name = "models/Diffusion_Transformer/LTX-2.3-Diffusers"
64
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
65
+ sampler_name = "Flow"
66
+
67
+ # Load pretrained model if need
68
+ transformer_path = None
69
+ vae_path = None
70
+ lora_path = None
71
+
72
+ # Other params
73
+ sample_size = [512, 768]
74
+ video_length = 121
75
+ fps = 24
76
+
77
+ # Use torch.float16 if GPU does not support torch.bfloat16
78
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
79
+ weight_dtype = torch.bfloat16
80
+ # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
81
+ validation_image_start = "asset/1.png"
82
+
83
+ # prompts
84
+ prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
85
+ negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
86
+ # CFG guidance scale for video and audio modality
87
+ guidance_scale = 3.0
88
+ audio_guidance_scale = 7.0
89
+ # Spatio-Temporal Guidance (STG) scale for video and audio
90
+ stg_scale = 1.0
91
+ audio_stg_scale = 1.0
92
+ # Modality isolation guidance scale for video and audio
93
+ modality_scale = 3.0
94
+ audio_modality_scale = 3.0
95
+ # Guidance rescale factor for video and audio to prevent overexposure
96
+ guidance_rescale = 0.7
97
+ audio_guidance_rescale = 0.7
98
+ spatio_temporal_guidance_blocks = [28]
99
+ seed = 43
100
+ num_inference_steps = 50
101
+ lora_weight = 0.55
102
+ save_path = "samples/ltx2-videos-i2v"
103
+
104
+ # Audio sample rate will be read from vocoder config
105
+ audio_sample_rate = 24000
106
+
107
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
108
+
109
+ # Transformer
110
+ transformer = LTX2VideoTransformer3DModel.from_pretrained(
111
+ model_name,
112
+ subfolder="transformer",
113
+ low_cpu_mem_usage=True,
114
+ torch_dtype=weight_dtype,
115
+ )
116
+
117
+ if transformer_path is not None:
118
+ print(f"From checkpoint: {transformer_path}")
119
+ if transformer_path.endswith("safetensors"):
120
+ from safetensors.torch import load_file, safe_open
121
+ state_dict = load_file(transformer_path)
122
+ else:
123
+ state_dict = torch.load(transformer_path, map_location="cpu")
124
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
125
+
126
+ m, u = transformer.load_state_dict(state_dict, strict=False)
127
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
128
+
129
+ # Video VAE
130
+ vae = AutoencoderKLLTX2Video.from_pretrained(
131
+ model_name,
132
+ subfolder="vae",
133
+ torch_dtype=weight_dtype,
134
+ )
135
+
136
+ if vae_path is not None:
137
+ print(f"From checkpoint: {vae_path}")
138
+ if vae_path.endswith("safetensors"):
139
+ from safetensors.torch import load_file, safe_open
140
+ state_dict = load_file(vae_path)
141
+ else:
142
+ state_dict = torch.load(vae_path, map_location="cpu")
143
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
144
+
145
+ m, u = vae.load_state_dict(state_dict, strict=False)
146
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
147
+
148
+ # Audio VAE
149
+ audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
150
+ model_name,
151
+ subfolder="audio_vae",
152
+ torch_dtype=weight_dtype,
153
+ )
154
+
155
+ # Get Processor
156
+ processor = Gemma3Processor.from_pretrained(
157
+ model_name,
158
+ subfolder="processor",
159
+ )
160
+
161
+ # Get Tokenizer
162
+ tokenizer = processor.tokenizer
163
+
164
+ # Get Text encoder
165
+ text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
166
+ model_name,
167
+ subfolder="text_encoder",
168
+ low_cpu_mem_usage=True,
169
+ torch_dtype=weight_dtype,
170
+ )
171
+ text_encoder = text_encoder.eval()
172
+
173
+ # Connectors
174
+ connectors = LTX2TextConnectors.from_pretrained(
175
+ model_name,
176
+ subfolder="connectors",
177
+ torch_dtype=weight_dtype,
178
+ )
179
+
180
+ # Vocoder
181
+ vocoder = LTX2VocoderWithBWE.from_pretrained(
182
+ model_name,
183
+ subfolder="vocoder",
184
+ torch_dtype=weight_dtype,
185
+ )
186
+
187
+ # Get Scheduler
188
+ Chosen_Scheduler = {
189
+ "Flow": FlowMatchEulerDiscreteScheduler,
190
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
191
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
192
+ }[sampler_name]
193
+ scheduler = Chosen_Scheduler.from_pretrained(
194
+ model_name,
195
+ subfolder="scheduler"
196
+ )
197
+
198
+ pipeline = LTX2I2VPipeline(
199
+ scheduler=scheduler,
200
+ vae=vae,
201
+ audio_vae=audio_vae,
202
+ text_encoder=text_encoder,
203
+ tokenizer=tokenizer,
204
+ processor=processor,
205
+ connectors=connectors,
206
+ transformer=transformer,
207
+ vocoder=vocoder,
208
+ )
209
+
210
+ if ulysses_degree > 1 or ring_degree > 1:
211
+ from functools import partial
212
+ transformer.enable_multi_gpus_inference()
213
+ if fsdp_dit:
214
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
215
+ module_to_wrapper=list(transformer.transformer_blocks))
216
+ pipeline.transformer = shard_fn(pipeline.transformer)
217
+ print("Add FSDP DIT")
218
+ if fsdp_text_encoder:
219
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
220
+ module_to_wrapper=text_encoder.language_model.layers)
221
+ text_encoder = shard_fn(text_encoder)
222
+ print("Add FSDP TEXT ENCODER")
223
+
224
+ if compile_dit:
225
+ for i in range(len(pipeline.transformer.transformer_blocks)):
226
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
227
+ print("Add Compile")
228
+
229
+ if GPU_memory_mode == "sequential_cpu_offload":
230
+ pipeline.enable_sequential_cpu_offload(device=device)
231
+ elif GPU_memory_mode == "model_group_offload":
232
+ register_auto_device_hook(pipeline.transformer)
233
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
234
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
235
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
236
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
237
+ pipeline.enable_model_cpu_offload(device=device)
238
+ elif GPU_memory_mode == "model_cpu_offload":
239
+ pipeline.enable_model_cpu_offload(device=device)
240
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
241
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
242
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
243
+ pipeline.to(device=device)
244
+ else:
245
+ pipeline.to(device=device)
246
+
247
+ generator = torch.Generator(device=device).manual_seed(seed)
248
+
249
+ if lora_path is not None:
250
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
251
+
252
+ with torch.no_grad():
253
+ output = pipeline(
254
+ image=Image.open(validation_image_start),
255
+ prompt=prompt,
256
+ negative_prompt=negative_prompt,
257
+ height=sample_size[0],
258
+ width=sample_size[1],
259
+ num_frames=video_length,
260
+ frame_rate=fps,
261
+ num_inference_steps=num_inference_steps,
262
+ guidance_scale=guidance_scale,
263
+ stg_scale=stg_scale,
264
+ modality_scale=modality_scale,
265
+ guidance_rescale=guidance_rescale,
266
+ audio_guidance_scale=audio_guidance_scale,
267
+ audio_stg_scale=audio_stg_scale,
268
+ audio_modality_scale=audio_modality_scale,
269
+ audio_guidance_rescale=audio_guidance_rescale,
270
+ spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks,
271
+ generator=generator,
272
+ output_type="pt",
273
+ )
274
+
275
+ if lora_path is not None:
276
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
277
+
278
+ sample = output.videos
279
+ audio = output.audio
280
+
281
+ def save_results():
282
+ if not os.path.exists(save_path):
283
+ os.makedirs(save_path, exist_ok=True)
284
+
285
+ index = len([path for path in os.listdir(save_path)]) + 1
286
+ prefix = str(index).zfill(8)
287
+ if video_length == 1:
288
+ video_path = os.path.join(save_path, prefix + ".png")
289
+
290
+ image = sample[0, :, 0]
291
+ image = image.transpose(0, 1).transpose(1, 2)
292
+ image = (image * 255).numpy().astype(np.uint8)
293
+ image = Image.fromarray(image)
294
+ image.save(video_path)
295
+ else:
296
+ video_path = os.path.join(save_path, prefix + ".mp4")
297
+ sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
298
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
299
+
300
+ if ulysses_degree * ring_degree > 1:
301
+ import torch.distributed as dist
302
+ if dist.get_rank() == 0:
303
+ save_results()
304
+ else:
305
+ save_results()
vendor/VideoX-Fun/examples/ltx2.3/predict_t2v.py ADDED
@@ -0,0 +1,300 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
15
+ Gemma3ForConditionalGeneration,
16
+ GemmaTokenizerFast, LTX2TextConnectors, Gemma3Processor,
17
+ LTX2VideoTransformer3DModel, LTX2VocoderWithBWE)
18
+ from videox_fun.pipeline import LTX2Pipeline
19
+ from videox_fun.utils import (register_auto_device_hook,
20
+ safe_enable_group_offload)
21
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
25
+ convert_weight_dtype_wrapper,
26
+ replace_parameters_by_name)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
29
+ save_videos_grid,
30
+ save_videos_with_audio_grid)
31
+
32
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
33
+ # model_full_load means that the entire model will be moved to the GPU.
34
+ #
35
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
39
+ #
40
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
44
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
45
+ #
46
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
47
+ # resulting in slower speeds but saving a large amount of GPU memory.
48
+ GPU_memory_mode = "model_group_offload"
49
+ # Multi GPUs config
50
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
51
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
52
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
53
+ ulysses_degree = 1
54
+ ring_degree = 1
55
+ # Use FSDP to save more GPU memory in multi gpus.
56
+ fsdp_dit = False
57
+ fsdp_text_encoder = False
58
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
59
+ # The compile_dit is not compatible with sequential_cpu_offload.
60
+ compile_dit = False
61
+
62
+ # model path
63
+ model_name = "models/Diffusion_Transformer/LTX-2.3-Diffusers"
64
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
65
+ sampler_name = "Flow"
66
+
67
+ # Load pretrained model if need
68
+ transformer_path = None
69
+ vae_path = None
70
+ lora_path = None
71
+
72
+ # Other params
73
+ sample_size = [512, 768]
74
+ video_length = 121
75
+ fps = 24
76
+
77
+ # Use torch.float16 if GPU does not support torch.bfloat16
78
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
79
+ weight_dtype = torch.bfloat16
80
+ prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
81
+ negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
82
+ # CFG guidance scale for video and audio modality
83
+ guidance_scale = 3.0
84
+ audio_guidance_scale = 7.0
85
+ # Spatio-Temporal Guidance (STG) scale for video and audio
86
+ stg_scale = 1.0
87
+ audio_stg_scale = 1.0
88
+ # Modality isolation guidance scale for video and audio
89
+ modality_scale = 3.0
90
+ audio_modality_scale = 3.0
91
+ # Guidance rescale factor for video and audio to prevent overexposure
92
+ guidance_rescale = 0.7
93
+ audio_guidance_rescale = 0.7
94
+ spatio_temporal_guidance_blocks = [28]
95
+ seed = 43
96
+ num_inference_steps = 50
97
+ lora_weight = 0.55
98
+ save_path = "samples/ltx2-videos-t2v"
99
+
100
+ # Audio sample rate will be read from vocoder config
101
+ audio_sample_rate = 24000
102
+
103
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
104
+
105
+ # Transformer
106
+ transformer = LTX2VideoTransformer3DModel.from_pretrained(
107
+ model_name,
108
+ subfolder="transformer",
109
+ low_cpu_mem_usage=True,
110
+ torch_dtype=weight_dtype,
111
+ )
112
+
113
+ if transformer_path is not None:
114
+ print(f"From checkpoint: {transformer_path}")
115
+ if transformer_path.endswith("safetensors"):
116
+ from safetensors.torch import load_file, safe_open
117
+ state_dict = load_file(transformer_path)
118
+ else:
119
+ state_dict = torch.load(transformer_path, map_location="cpu")
120
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
121
+
122
+ m, u = transformer.load_state_dict(state_dict, strict=False)
123
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
124
+
125
+ # Video VAE
126
+ vae = AutoencoderKLLTX2Video.from_pretrained(
127
+ model_name,
128
+ subfolder="vae",
129
+ torch_dtype=weight_dtype,
130
+ )
131
+
132
+ if vae_path is not None:
133
+ print(f"From checkpoint: {vae_path}")
134
+ if vae_path.endswith("safetensors"):
135
+ from safetensors.torch import load_file, safe_open
136
+ state_dict = load_file(vae_path)
137
+ else:
138
+ state_dict = torch.load(vae_path, map_location="cpu")
139
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
140
+
141
+ m, u = vae.load_state_dict(state_dict, strict=False)
142
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
143
+
144
+ # Audio VAE
145
+ audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
146
+ model_name,
147
+ subfolder="audio_vae",
148
+ torch_dtype=weight_dtype,
149
+ )
150
+
151
+ # Get Processor
152
+ processor = Gemma3Processor.from_pretrained(
153
+ model_name,
154
+ subfolder="processor",
155
+ )
156
+
157
+ # Get Tokenizer
158
+ tokenizer = processor.tokenizer
159
+
160
+ # Get Text encoder
161
+ text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
162
+ model_name,
163
+ subfolder="text_encoder",
164
+ low_cpu_mem_usage=True,
165
+ torch_dtype=weight_dtype,
166
+ )
167
+ text_encoder = text_encoder.eval()
168
+
169
+ # Connectors
170
+ connectors = LTX2TextConnectors.from_pretrained(
171
+ model_name,
172
+ subfolder="connectors",
173
+ torch_dtype=weight_dtype,
174
+ )
175
+
176
+ # Vocoder
177
+ vocoder = LTX2VocoderWithBWE.from_pretrained(
178
+ model_name,
179
+ subfolder="vocoder",
180
+ torch_dtype=weight_dtype,
181
+ )
182
+
183
+ # Get Scheduler
184
+ Chosen_Scheduler = {
185
+ "Flow": FlowMatchEulerDiscreteScheduler,
186
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
187
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
188
+ }[sampler_name]
189
+ scheduler = Chosen_Scheduler.from_pretrained(
190
+ model_name,
191
+ subfolder="scheduler"
192
+ )
193
+
194
+ pipeline = LTX2Pipeline(
195
+ scheduler=scheduler,
196
+ vae=vae,
197
+ audio_vae=audio_vae,
198
+ text_encoder=text_encoder,
199
+ tokenizer=tokenizer,
200
+ processor=processor,
201
+ connectors=connectors,
202
+ transformer=transformer,
203
+ vocoder=vocoder,
204
+ )
205
+
206
+ if ulysses_degree > 1 or ring_degree > 1:
207
+ from functools import partial
208
+ transformer.enable_multi_gpus_inference()
209
+ if fsdp_dit:
210
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
211
+ module_to_wrapper=list(transformer.transformer_blocks))
212
+ pipeline.transformer = shard_fn(pipeline.transformer)
213
+ print("Add FSDP DIT")
214
+ if fsdp_text_encoder:
215
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
216
+ module_to_wrapper=text_encoder.language_model.layers)
217
+ text_encoder = shard_fn(text_encoder)
218
+ print("Add FSDP TEXT ENCODER")
219
+
220
+ if compile_dit:
221
+ for i in range(len(pipeline.transformer.transformer_blocks)):
222
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
223
+ print("Add Compile")
224
+
225
+ if GPU_memory_mode == "sequential_cpu_offload":
226
+ pipeline.enable_sequential_cpu_offload(device=device)
227
+ elif GPU_memory_mode == "model_group_offload":
228
+ register_auto_device_hook(pipeline.transformer)
229
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
230
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
231
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
232
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
233
+ pipeline.enable_model_cpu_offload(device=device)
234
+ elif GPU_memory_mode == "model_cpu_offload":
235
+ pipeline.enable_model_cpu_offload(device=device)
236
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
237
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
238
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
239
+ pipeline.to(device=device)
240
+ else:
241
+ pipeline.to(device=device)
242
+
243
+ generator = torch.Generator(device=device).manual_seed(seed)
244
+
245
+ if lora_path is not None:
246
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
247
+
248
+ with torch.no_grad():
249
+ output = pipeline(
250
+ prompt=prompt,
251
+ negative_prompt=negative_prompt,
252
+ height=sample_size[0],
253
+ width=sample_size[1],
254
+ num_frames=video_length,
255
+ frame_rate=fps,
256
+ num_inference_steps=num_inference_steps,
257
+ guidance_scale=guidance_scale,
258
+ stg_scale=stg_scale,
259
+ modality_scale=modality_scale,
260
+ guidance_rescale=guidance_rescale,
261
+ audio_guidance_scale=audio_guidance_scale,
262
+ audio_stg_scale=audio_stg_scale,
263
+ audio_modality_scale=audio_modality_scale,
264
+ audio_guidance_rescale=audio_guidance_rescale,
265
+ spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks,
266
+ generator=generator,
267
+ output_type="pt",
268
+ )
269
+
270
+ if lora_path is not None:
271
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
272
+
273
+ sample = output.videos
274
+ audio = output.audio
275
+
276
+ def save_results():
277
+ if not os.path.exists(save_path):
278
+ os.makedirs(save_path, exist_ok=True)
279
+
280
+ index = len([path for path in os.listdir(save_path)]) + 1
281
+ prefix = str(index).zfill(8)
282
+ if video_length == 1:
283
+ video_path = os.path.join(save_path, prefix + ".png")
284
+
285
+ image = sample[0, :, 0]
286
+ image = image.transpose(0, 1).transpose(1, 2)
287
+ image = (image * 255).numpy().astype(np.uint8)
288
+ image = Image.fromarray(image)
289
+ image.save(video_path)
290
+ else:
291
+ video_path = os.path.join(save_path, prefix + ".mp4")
292
+ sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
293
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
294
+
295
+ if ulysses_degree * ring_degree > 1:
296
+ import torch.distributed as dist
297
+ if dist.get_rank() == 0:
298
+ save_results()
299
+ else:
300
+ save_results()
vendor/VideoX-Fun/examples/ltx2/predict_i2v.py ADDED
@@ -0,0 +1,281 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
15
+ Gemma3ForConditionalGeneration,
16
+ GemmaTokenizerFast, LTX2TextConnectors,
17
+ LTX2VideoTransformer3DModel, LTX2Vocoder)
18
+ from videox_fun.pipeline import LTX2I2VPipeline
19
+ from videox_fun.utils import (register_auto_device_hook,
20
+ safe_enable_group_offload)
21
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
25
+ convert_weight_dtype_wrapper,
26
+ replace_parameters_by_name)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
29
+ save_videos_grid,
30
+ save_videos_with_audio_grid)
31
+
32
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
33
+ # model_full_load means that the entire model will be moved to the GPU.
34
+ #
35
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
39
+ #
40
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
44
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
45
+ #
46
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
47
+ # resulting in slower speeds but saving a large amount of GPU memory.
48
+ GPU_memory_mode = "sequential_cpu_offload"
49
+ # Multi GPUs config
50
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
51
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
52
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
53
+ ulysses_degree = 1
54
+ ring_degree = 1
55
+ # Use FSDP to save more GPU memory in multi gpus.
56
+ fsdp_dit = False
57
+ fsdp_text_encoder = False
58
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
59
+ # The compile_dit is not compatible with sequential_cpu_offload.
60
+ compile_dit = False
61
+
62
+ # model path
63
+ model_name = "models/Diffusion_Transformer/LTX-2"
64
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
65
+ sampler_name = "Flow"
66
+
67
+ # Load pretrained model if need
68
+ transformer_path = None
69
+ vae_path = None
70
+ lora_path = None
71
+
72
+ # Other params
73
+ sample_size = [480, 832]
74
+ video_length = 121
75
+ fps = 24
76
+
77
+ # Use torch.float16 if GPU does not support torch.bfloat16
78
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
79
+ weight_dtype = torch.bfloat16
80
+ # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
81
+ validation_image_start = "asset/1.png"
82
+
83
+ # prompts
84
+ prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
85
+ negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
86
+ guidance_scale = 6.0
87
+ seed = 43
88
+ num_inference_steps = 50
89
+ lora_weight = 0.55
90
+ save_path = "samples/ltx2-videos-i2v"
91
+
92
+ # Audio sample rate will be read from vocoder config
93
+ audio_sample_rate = 24000
94
+
95
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
96
+
97
+ # Transformer
98
+ transformer = LTX2VideoTransformer3DModel.from_pretrained(
99
+ model_name,
100
+ subfolder="transformer",
101
+ low_cpu_mem_usage=True,
102
+ torch_dtype=weight_dtype,
103
+ )
104
+
105
+ if transformer_path is not None:
106
+ print(f"From checkpoint: {transformer_path}")
107
+ if transformer_path.endswith("safetensors"):
108
+ from safetensors.torch import load_file, safe_open
109
+ state_dict = load_file(transformer_path)
110
+ else:
111
+ state_dict = torch.load(transformer_path, map_location="cpu")
112
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
113
+
114
+ m, u = transformer.load_state_dict(state_dict, strict=False)
115
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
116
+
117
+ # Video VAE
118
+ vae = AutoencoderKLLTX2Video.from_pretrained(
119
+ model_name,
120
+ subfolder="vae",
121
+ torch_dtype=weight_dtype,
122
+ )
123
+
124
+ if vae_path is not None:
125
+ print(f"From checkpoint: {vae_path}")
126
+ if vae_path.endswith("safetensors"):
127
+ from safetensors.torch import load_file, safe_open
128
+ state_dict = load_file(vae_path)
129
+ else:
130
+ state_dict = torch.load(vae_path, map_location="cpu")
131
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
132
+
133
+ m, u = vae.load_state_dict(state_dict, strict=False)
134
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
135
+
136
+ # Audio VAE
137
+ audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
138
+ model_name,
139
+ subfolder="audio_vae",
140
+ torch_dtype=weight_dtype,
141
+ )
142
+
143
+ # Get Tokenizer
144
+ tokenizer = GemmaTokenizerFast.from_pretrained(
145
+ model_name,
146
+ subfolder="tokenizer",
147
+ )
148
+
149
+ # Get Text encoder
150
+ text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
151
+ model_name,
152
+ subfolder="text_encoder",
153
+ low_cpu_mem_usage=True,
154
+ torch_dtype=weight_dtype,
155
+ )
156
+ text_encoder = text_encoder.eval()
157
+
158
+ # Connectors
159
+ connectors = LTX2TextConnectors.from_pretrained(
160
+ model_name,
161
+ subfolder="connectors",
162
+ torch_dtype=weight_dtype,
163
+ )
164
+
165
+ # Vocoder
166
+ vocoder = LTX2Vocoder.from_pretrained(
167
+ model_name,
168
+ subfolder="vocoder",
169
+ torch_dtype=weight_dtype,
170
+ )
171
+
172
+ # Get Scheduler
173
+ Chosen_Scheduler = {
174
+ "Flow": FlowMatchEulerDiscreteScheduler,
175
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
176
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
177
+ }[sampler_name]
178
+ scheduler = Chosen_Scheduler.from_pretrained(
179
+ model_name,
180
+ subfolder="scheduler"
181
+ )
182
+
183
+ pipeline = LTX2I2VPipeline(
184
+ scheduler=scheduler,
185
+ vae=vae,
186
+ audio_vae=audio_vae,
187
+ text_encoder=text_encoder,
188
+ tokenizer=tokenizer,
189
+ connectors=connectors,
190
+ transformer=transformer,
191
+ vocoder=vocoder,
192
+ )
193
+
194
+ if ulysses_degree > 1 or ring_degree > 1:
195
+ from functools import partial
196
+ transformer.enable_multi_gpus_inference()
197
+ if fsdp_dit:
198
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
199
+ module_to_wrapper=list(transformer.transformer_blocks))
200
+ pipeline.transformer = shard_fn(pipeline.transformer)
201
+ print("Add FSDP DIT")
202
+ if fsdp_text_encoder:
203
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
204
+ module_to_wrapper=text_encoder.language_model.layers)
205
+ text_encoder = shard_fn(text_encoder)
206
+ print("Add FSDP TEXT ENCODER")
207
+
208
+ if compile_dit:
209
+ for i in range(len(pipeline.transformer.transformer_blocks)):
210
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
211
+ print("Add Compile")
212
+
213
+ if GPU_memory_mode == "sequential_cpu_offload":
214
+ pipeline.enable_sequential_cpu_offload(device=device)
215
+ elif GPU_memory_mode == "model_group_offload":
216
+ register_auto_device_hook(pipeline.transformer)
217
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
218
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
219
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
220
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
221
+ pipeline.enable_model_cpu_offload(device=device)
222
+ elif GPU_memory_mode == "model_cpu_offload":
223
+ pipeline.enable_model_cpu_offload(device=device)
224
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
225
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
226
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
227
+ pipeline.to(device=device)
228
+ else:
229
+ pipeline.to(device=device)
230
+
231
+ generator = torch.Generator(device=device).manual_seed(seed)
232
+
233
+ if lora_path is not None:
234
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
235
+
236
+ with torch.no_grad():
237
+ output = pipeline(
238
+ image=Image.open(validation_image_start),
239
+ prompt=prompt,
240
+ negative_prompt=negative_prompt,
241
+ height=sample_size[0],
242
+ width=sample_size[1],
243
+ num_frames=video_length,
244
+ frame_rate=fps,
245
+ num_inference_steps=num_inference_steps,
246
+ guidance_scale=guidance_scale,
247
+ generator=generator,
248
+ output_type="pt",
249
+ )
250
+
251
+ if lora_path is not None:
252
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
253
+
254
+ sample = output.videos
255
+ audio = output.audio
256
+
257
+ def save_results():
258
+ if not os.path.exists(save_path):
259
+ os.makedirs(save_path, exist_ok=True)
260
+
261
+ index = len([path for path in os.listdir(save_path)]) + 1
262
+ prefix = str(index).zfill(8)
263
+ if video_length == 1:
264
+ video_path = os.path.join(save_path, prefix + ".png")
265
+
266
+ image = sample[0, :, 0]
267
+ image = image.transpose(0, 1).transpose(1, 2)
268
+ image = (image * 255).numpy().astype(np.uint8)
269
+ image = Image.fromarray(image)
270
+ image.save(video_path)
271
+ else:
272
+ video_path = os.path.join(save_path, prefix + ".mp4")
273
+ sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
274
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
275
+
276
+ if ulysses_degree * ring_degree > 1:
277
+ import torch.distributed as dist
278
+ if dist.get_rank() == 0:
279
+ save_results()
280
+ else:
281
+ save_results()
vendor/VideoX-Fun/examples/ltx2/predict_i2v_upsample.py ADDED
@@ -0,0 +1,326 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
15
+ Gemma3ForConditionalGeneration,
16
+ GemmaTokenizerFast, LTX2LatentUpsamplerModel,
17
+ LTX2TextConnectors,
18
+ LTX2VideoTransformer3DModel, LTX2Vocoder)
19
+ from videox_fun.pipeline import LTX2I2VPipeline, LTX2LatentUpsamplePipeline
20
+ from videox_fun.utils import (register_auto_device_hook,
21
+ safe_enable_group_offload)
22
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
23
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
24
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
25
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
26
+ convert_weight_dtype_wrapper,
27
+ replace_parameters_by_name)
28
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
29
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
30
+ save_videos_grid,
31
+ save_videos_with_audio_grid)
32
+
33
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
34
+ # model_full_load means that the entire model will be moved to the GPU.
35
+ #
36
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
37
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
38
+ #
39
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
40
+ #
41
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
42
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
43
+ #
44
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
45
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
46
+ #
47
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
48
+ # resulting in slower speeds but saving a large amount of GPU memory.
49
+ GPU_memory_mode = "sequential_cpu_offload"
50
+ # Multi GPUs config
51
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
52
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
53
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
54
+ ulysses_degree = 1
55
+ ring_degree = 1
56
+ # Use FSDP to save more GPU memory in multi gpus.
57
+ fsdp_dit = False
58
+ fsdp_text_encoder = False
59
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
60
+ # The compile_dit is not compatible with sequential_cpu_offload.
61
+ compile_dit = False
62
+
63
+ # model path
64
+ model_name = "models/Diffusion_Transformer/LTX-2"
65
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
66
+ sampler_name = "Flow"
67
+
68
+ # Load pretrained model if need
69
+ transformer_path = None
70
+ vae_path = None
71
+ lora_path = None
72
+ latent_upsampler_path = None
73
+
74
+ # Other params
75
+ sample_size = [480, 832]
76
+ video_length = 121
77
+ fps = 24
78
+ # Latent upsampler config
79
+ enable_latent_upsample = True
80
+
81
+ # Use torch.float16 if GPU does not support torch.bfloat16
82
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
83
+ weight_dtype = torch.bfloat16
84
+ # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
85
+ validation_image_start = "asset/1.png"
86
+
87
+ # prompts
88
+ prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
89
+ negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
90
+ guidance_scale = 6.0
91
+ seed = 43
92
+ num_inference_steps = 50
93
+ lora_weight = 0.55
94
+ save_path = "samples/ltx2-videos-i2v"
95
+
96
+ # Audio sample rate will be read from vocoder config
97
+ audio_sample_rate = 24000
98
+
99
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
100
+
101
+ # Transformer
102
+ transformer = LTX2VideoTransformer3DModel.from_pretrained(
103
+ model_name,
104
+ subfolder="transformer",
105
+ low_cpu_mem_usage=True,
106
+ torch_dtype=weight_dtype,
107
+ )
108
+
109
+ if transformer_path is not None:
110
+ print(f"From checkpoint: {transformer_path}")
111
+ if transformer_path.endswith("safetensors"):
112
+ from safetensors.torch import load_file, safe_open
113
+ state_dict = load_file(transformer_path)
114
+ else:
115
+ state_dict = torch.load(transformer_path, map_location="cpu")
116
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
117
+
118
+ m, u = transformer.load_state_dict(state_dict, strict=False)
119
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
120
+
121
+ # Video VAE
122
+ vae = AutoencoderKLLTX2Video.from_pretrained(
123
+ model_name,
124
+ subfolder="vae",
125
+ torch_dtype=weight_dtype,
126
+ )
127
+
128
+ if vae_path is not None:
129
+ print(f"From checkpoint: {vae_path}")
130
+ if vae_path.endswith("safetensors"):
131
+ from safetensors.torch import load_file, safe_open
132
+ state_dict = load_file(vae_path)
133
+ else:
134
+ state_dict = torch.load(vae_path, map_location="cpu")
135
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
136
+
137
+ m, u = vae.load_state_dict(state_dict, strict=False)
138
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
139
+
140
+ # Audio VAE
141
+ audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
142
+ model_name,
143
+ subfolder="audio_vae",
144
+ torch_dtype=weight_dtype,
145
+ )
146
+
147
+ # Get Tokenizer
148
+ tokenizer = GemmaTokenizerFast.from_pretrained(
149
+ model_name,
150
+ subfolder="tokenizer",
151
+ )
152
+
153
+ # Get Text encoder
154
+ text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
155
+ model_name,
156
+ subfolder="text_encoder",
157
+ low_cpu_mem_usage=True,
158
+ torch_dtype=weight_dtype,
159
+ )
160
+ text_encoder = text_encoder.eval()
161
+
162
+ # Connectors
163
+ connectors = LTX2TextConnectors.from_pretrained(
164
+ model_name,
165
+ subfolder="connectors",
166
+ torch_dtype=weight_dtype,
167
+ )
168
+
169
+ # Vocoder
170
+ vocoder = LTX2Vocoder.from_pretrained(
171
+ model_name,
172
+ subfolder="vocoder",
173
+ torch_dtype=weight_dtype,
174
+ )
175
+
176
+ # Get Scheduler
177
+ Chosen_Scheduler = {
178
+ "Flow": FlowMatchEulerDiscreteScheduler,
179
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
180
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
181
+ }[sampler_name]
182
+ scheduler = Chosen_Scheduler.from_pretrained(
183
+ model_name,
184
+ subfolder="scheduler"
185
+ )
186
+
187
+ pipeline = LTX2I2VPipeline(
188
+ scheduler=scheduler,
189
+ vae=vae,
190
+ audio_vae=audio_vae,
191
+ text_encoder=text_encoder,
192
+ tokenizer=tokenizer,
193
+ connectors=connectors,
194
+ transformer=transformer,
195
+ vocoder=vocoder,
196
+ )
197
+
198
+ if ulysses_degree > 1 or ring_degree > 1:
199
+ from functools import partial
200
+ transformer.enable_multi_gpus_inference()
201
+ if fsdp_dit:
202
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
203
+ module_to_wrapper=list(transformer.transformer_blocks))
204
+ pipeline.transformer = shard_fn(pipeline.transformer)
205
+ print("Add FSDP DIT")
206
+ if fsdp_text_encoder:
207
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
208
+ module_to_wrapper=text_encoder.language_model.layers)
209
+ text_encoder = shard_fn(text_encoder)
210
+ print("Add FSDP TEXT ENCODER")
211
+
212
+ if compile_dit:
213
+ for i in range(len(pipeline.transformer.transformer_blocks)):
214
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
215
+ print("Add Compile")
216
+
217
+ if GPU_memory_mode == "sequential_cpu_offload":
218
+ pipeline.enable_sequential_cpu_offload(device=device)
219
+ elif GPU_memory_mode == "model_group_offload":
220
+ register_auto_device_hook(pipeline.transformer)
221
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
222
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
223
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
224
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
225
+ pipeline.enable_model_cpu_offload(device=device)
226
+ elif GPU_memory_mode == "model_cpu_offload":
227
+ pipeline.enable_model_cpu_offload(device=device)
228
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
229
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
230
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
231
+ pipeline.to(device=device)
232
+ else:
233
+ pipeline.to(device=device)
234
+
235
+ generator = torch.Generator(device=device).manual_seed(seed)
236
+
237
+ if lora_path is not None:
238
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
239
+
240
+ with torch.no_grad():
241
+ output = pipeline(
242
+ image=Image.open(validation_image_start),
243
+ prompt=prompt,
244
+ negative_prompt=negative_prompt,
245
+ height=sample_size[0],
246
+ width=sample_size[1],
247
+ num_frames=video_length,
248
+ frame_rate=fps,
249
+ num_inference_steps=num_inference_steps,
250
+ guidance_scale=guidance_scale,
251
+ generator=generator,
252
+ output_type="latent" if enable_latent_upsample else "pt",
253
+ )
254
+
255
+ if lora_path is not None:
256
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
257
+
258
+ if enable_latent_upsample:
259
+ # Load latent upsampler model
260
+ latent_upsampler = LTX2LatentUpsamplerModel.from_pretrained(
261
+ model_name, subfolder="latent_upsampler", torch_dtype=weight_dtype,
262
+ )
263
+ if latent_upsampler_path is not None:
264
+ print(f"From latent_upsampler checkpoint: {latent_upsampler_path}")
265
+ if latent_upsampler_path.endswith("safetensors"):
266
+ from safetensors.torch import load_file
267
+ state_dict = load_file(latent_upsampler_path)
268
+ else:
269
+ state_dict = torch.load(latent_upsampler_path, map_location="cpu")
270
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
271
+ m, u = latent_upsampler.load_state_dict(state_dict, strict=False)
272
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
273
+
274
+ upsample_pipeline = LTX2LatentUpsamplePipeline(
275
+ vae=pipeline.vae,
276
+ latent_upsampler=latent_upsampler,
277
+ )
278
+ upsample_pipeline.vae.enable_tiling()
279
+ upsample_pipeline.to(device=device, dtype=weight_dtype)
280
+
281
+ # output_type="latent" returns denormalized (raw) video latents [B, C, F, H, W]
282
+ # and raw audio latents [B, C, L, M]; decode audio manually
283
+ audio_latents = output.audio.to(device=device, dtype=pipeline.audio_vae.dtype)
284
+ mel = pipeline.audio_vae.decode(audio_latents, return_dict=False)[0]
285
+ audio = pipeline.vocoder(mel).cpu().float()
286
+
287
+ # Pass video latents directly to upsample pipeline (skip decode→re-encode roundtrip)
288
+ with torch.no_grad():
289
+ upsampled = upsample_pipeline(
290
+ latents=output.videos,
291
+ height=sample_size[0],
292
+ width=sample_size[1],
293
+ num_frames=video_length,
294
+ output_type="pt",
295
+ return_dict=False,
296
+ )
297
+ sample = upsampled[0]
298
+ else:
299
+ sample = output.videos
300
+ audio = output.audio
301
+
302
+ def save_results():
303
+ if not os.path.exists(save_path):
304
+ os.makedirs(save_path, exist_ok=True)
305
+
306
+ index = len([path for path in os.listdir(save_path)]) + 1
307
+ prefix = str(index).zfill(8)
308
+ if video_length == 1:
309
+ video_path = os.path.join(save_path, prefix + ".png")
310
+
311
+ image = sample[0, :, 0]
312
+ image = image.transpose(0, 1).transpose(1, 2)
313
+ image = (image * 255).numpy().astype(np.uint8)
314
+ image = Image.fromarray(image)
315
+ image.save(video_path)
316
+ else:
317
+ video_path = os.path.join(save_path, prefix + ".mp4")
318
+ sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
319
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
320
+
321
+ if ulysses_degree * ring_degree > 1:
322
+ import torch.distributed as dist
323
+ if dist.get_rank() == 0:
324
+ save_results()
325
+ else:
326
+ save_results()
vendor/VideoX-Fun/examples/ltx2/predict_t2v.py ADDED
@@ -0,0 +1,276 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
15
+ Gemma3ForConditionalGeneration,
16
+ GemmaTokenizerFast, LTX2TextConnectors,
17
+ LTX2VideoTransformer3DModel, LTX2Vocoder)
18
+ from videox_fun.pipeline import LTX2Pipeline
19
+ from videox_fun.utils import (register_auto_device_hook,
20
+ safe_enable_group_offload)
21
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
25
+ convert_weight_dtype_wrapper,
26
+ replace_parameters_by_name)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
29
+ save_videos_grid,
30
+ save_videos_with_audio_grid)
31
+
32
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
33
+ # model_full_load means that the entire model will be moved to the GPU.
34
+ #
35
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
36
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
37
+ #
38
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
39
+ #
40
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
41
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
42
+ #
43
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
44
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
45
+ #
46
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
47
+ # resulting in slower speeds but saving a large amount of GPU memory.
48
+ GPU_memory_mode = "sequential_cpu_offload"
49
+ # Multi GPUs config
50
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
51
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
52
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
53
+ ulysses_degree = 1
54
+ ring_degree = 1
55
+ # Use FSDP to save more GPU memory in multi gpus.
56
+ fsdp_dit = False
57
+ fsdp_text_encoder = False
58
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
59
+ # The compile_dit is not compatible with sequential_cpu_offload.
60
+ compile_dit = False
61
+
62
+ # model path
63
+ model_name = "models/Diffusion_Transformer/LTX-2"
64
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
65
+ sampler_name = "Flow"
66
+
67
+ # Load pretrained model if need
68
+ transformer_path = None
69
+ vae_path = None
70
+ lora_path = None
71
+
72
+ # Other params
73
+ sample_size = [512, 768]
74
+ video_length = 121
75
+ fps = 24
76
+
77
+ # Use torch.float16 if GPU does not support torch.bfloat16
78
+ # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
79
+ weight_dtype = torch.bfloat16
80
+ prompt = "A brown dog barks on a sofa, sitting on a light-colored couch in a cozy room. Behind the dog, there is a framed painting on a shelf, surrounded by pink flowers. "
81
+ negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
82
+ guidance_scale = 6.0
83
+ seed = 43
84
+ num_inference_steps = 50
85
+ lora_weight = 0.55
86
+ save_path = "samples/ltx2-videos-t2v"
87
+
88
+ # Audio sample rate will be read from vocoder config
89
+ audio_sample_rate = 24000
90
+
91
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
92
+
93
+ # Transformer
94
+ transformer = LTX2VideoTransformer3DModel.from_pretrained(
95
+ model_name,
96
+ subfolder="transformer",
97
+ low_cpu_mem_usage=True,
98
+ torch_dtype=weight_dtype,
99
+ )
100
+
101
+ if transformer_path is not None:
102
+ print(f"From checkpoint: {transformer_path}")
103
+ if transformer_path.endswith("safetensors"):
104
+ from safetensors.torch import load_file, safe_open
105
+ state_dict = load_file(transformer_path)
106
+ else:
107
+ state_dict = torch.load(transformer_path, map_location="cpu")
108
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
109
+
110
+ m, u = transformer.load_state_dict(state_dict, strict=False)
111
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
112
+
113
+ # Video VAE
114
+ vae = AutoencoderKLLTX2Video.from_pretrained(
115
+ model_name,
116
+ subfolder="vae",
117
+ torch_dtype=weight_dtype,
118
+ )
119
+
120
+ if vae_path is not None:
121
+ print(f"From checkpoint: {vae_path}")
122
+ if vae_path.endswith("safetensors"):
123
+ from safetensors.torch import load_file, safe_open
124
+ state_dict = load_file(vae_path)
125
+ else:
126
+ state_dict = torch.load(vae_path, map_location="cpu")
127
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
128
+
129
+ m, u = vae.load_state_dict(state_dict, strict=False)
130
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
131
+
132
+ # Audio VAE
133
+ audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
134
+ model_name,
135
+ subfolder="audio_vae",
136
+ torch_dtype=weight_dtype,
137
+ )
138
+
139
+ # Get Tokenizer
140
+ tokenizer = GemmaTokenizerFast.from_pretrained(
141
+ model_name,
142
+ subfolder="tokenizer",
143
+ )
144
+
145
+ # Get Text encoder
146
+ text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
147
+ model_name,
148
+ subfolder="text_encoder",
149
+ low_cpu_mem_usage=True,
150
+ torch_dtype=weight_dtype,
151
+ )
152
+ text_encoder = text_encoder.eval()
153
+
154
+ # Connectors
155
+ connectors = LTX2TextConnectors.from_pretrained(
156
+ model_name,
157
+ subfolder="connectors",
158
+ torch_dtype=weight_dtype,
159
+ )
160
+
161
+ # Vocoder
162
+ vocoder = LTX2Vocoder.from_pretrained(
163
+ model_name,
164
+ subfolder="vocoder",
165
+ torch_dtype=weight_dtype,
166
+ )
167
+
168
+ # Get Scheduler
169
+ Chosen_Scheduler = {
170
+ "Flow": FlowMatchEulerDiscreteScheduler,
171
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
172
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
173
+ }[sampler_name]
174
+ scheduler = Chosen_Scheduler.from_pretrained(
175
+ model_name,
176
+ subfolder="scheduler"
177
+ )
178
+
179
+ pipeline = LTX2Pipeline(
180
+ scheduler=scheduler,
181
+ vae=vae,
182
+ audio_vae=audio_vae,
183
+ text_encoder=text_encoder,
184
+ tokenizer=tokenizer,
185
+ connectors=connectors,
186
+ transformer=transformer,
187
+ vocoder=vocoder,
188
+ )
189
+
190
+ if ulysses_degree > 1 or ring_degree > 1:
191
+ from functools import partial
192
+ transformer.enable_multi_gpus_inference()
193
+ if fsdp_dit:
194
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
195
+ module_to_wrapper=list(transformer.transformer_blocks))
196
+ pipeline.transformer = shard_fn(pipeline.transformer)
197
+ print("Add FSDP DIT")
198
+ if fsdp_text_encoder:
199
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
200
+ module_to_wrapper=text_encoder.language_model.layers)
201
+ text_encoder = shard_fn(text_encoder)
202
+ print("Add FSDP TEXT ENCODER")
203
+
204
+ if compile_dit:
205
+ for i in range(len(pipeline.transformer.transformer_blocks)):
206
+ pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
207
+ print("Add Compile")
208
+
209
+ if GPU_memory_mode == "sequential_cpu_offload":
210
+ pipeline.enable_sequential_cpu_offload(device=device)
211
+ elif GPU_memory_mode == "model_group_offload":
212
+ register_auto_device_hook(pipeline.transformer)
213
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
214
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
215
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
216
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
217
+ pipeline.enable_model_cpu_offload(device=device)
218
+ elif GPU_memory_mode == "model_cpu_offload":
219
+ pipeline.enable_model_cpu_offload(device=device)
220
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
221
+ convert_model_weight_to_float8(transformer, exclude_module_name=["scale_shift_table", "audio_scale_shift_table", "video_a2v_cross_attn_scale_shift_table", "audio_a2v_cross_attn_scale_shift_table", ""], device=device)
222
+ convert_weight_dtype_wrapper(transformer, weight_dtype)
223
+ pipeline.to(device=device)
224
+ else:
225
+ pipeline.to(device=device)
226
+
227
+ generator = torch.Generator(device=device).manual_seed(seed)
228
+
229
+ if lora_path is not None:
230
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
231
+
232
+ with torch.no_grad():
233
+ output = pipeline(
234
+ prompt=prompt,
235
+ negative_prompt=negative_prompt,
236
+ height=sample_size[0],
237
+ width=sample_size[1],
238
+ num_frames=video_length,
239
+ frame_rate=fps,
240
+ num_inference_steps=num_inference_steps,
241
+ guidance_scale=guidance_scale,
242
+ generator=generator,
243
+ output_type="pt",
244
+ )
245
+
246
+ if lora_path is not None:
247
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
248
+
249
+ sample = output.videos
250
+ audio = output.audio
251
+
252
+ def save_results():
253
+ if not os.path.exists(save_path):
254
+ os.makedirs(save_path, exist_ok=True)
255
+
256
+ index = len([path for path in os.listdir(save_path)]) + 1
257
+ prefix = str(index).zfill(8)
258
+ if video_length == 1:
259
+ video_path = os.path.join(save_path, prefix + ".png")
260
+
261
+ image = sample[0, :, 0]
262
+ image = image.transpose(0, 1).transpose(1, 2)
263
+ image = (image * 255).numpy().astype(np.uint8)
264
+ image = Image.fromarray(image)
265
+ image.save(video_path)
266
+ else:
267
+ video_path = os.path.join(save_path, prefix + ".mp4")
268
+ sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
269
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
270
+
271
+ if ulysses_degree * ring_degree > 1:
272
+ import torch.distributed as dist
273
+ if dist.get_rank() == 0:
274
+ save_results()
275
+ else:
276
+ save_results()
vendor/VideoX-Fun/examples/mova/predict_i2v.py ADDED
@@ -0,0 +1,380 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from diffusers import FlowMatchEulerDiscreteScheduler
7
+ from PIL import Image
8
+
9
+ current_file_path = os.path.abspath(__file__)
10
+ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
11
+ for project_root in project_roots:
12
+ sys.path.insert(0, project_root) if project_root not in sys.path else None
13
+
14
+ from videox_fun.dist import set_multi_gpus_devices, shard_model
15
+ from videox_fun.models import (AutoencoderKLMOVAAudio, AutoencoderKLWan,
16
+ AutoTokenizer, MOVADualTowerConditionalBridge,
17
+ UMT5EncoderModel, WanAudioTransformer3DModel,
18
+ WanTransformer3DModel)
19
+ from videox_fun.pipeline import MOVAPipeline
20
+ from videox_fun.utils import (register_auto_device_hook,
21
+ safe_enable_group_offload)
22
+ from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
23
+ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
24
+ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
25
+ convert_weight_dtype_wrapper,
26
+ replace_parameters_by_name)
27
+ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
28
+ from videox_fun.utils.utils import save_videos_with_audio_grid
29
+
30
+ # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload].
31
+ # model_full_load means that the entire model will be moved to the GPU.
32
+ #
33
+ # model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
34
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
35
+ #
36
+ # model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
37
+ #
38
+ # model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
39
+ # and the transformer model has been quantized to float8, which can save more GPU memory.
40
+ #
41
+ # model_group_offload transfers internal layer groups between CPU/CUDA,
42
+ # balancing memory efficiency and speed between full-module and leaf-level offloading methods.
43
+ #
44
+ # sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
45
+ # resulting in slower speeds but saving a large amount of GPU memory.
46
+ GPU_memory_mode = "sequential_cpu_offload"
47
+ # Multi GPUs config
48
+ # Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
49
+ # For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
50
+ # If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
51
+ ulysses_degree = 1
52
+ ring_degree = 1
53
+ # Use FSDP to save more GPU memory in multi gpus.
54
+ fsdp_dit = False
55
+ fsdp_text_encoder = True
56
+ # Compile will give a speedup in fixed resolution and need a little GPU memory.
57
+ # The compile_dit is not compatible with sequential_cpu_offload.
58
+ compile_dit = False
59
+
60
+ # model path
61
+ model_name = "models/Diffusion_Transformer/MOVA-360p"
62
+
63
+ # Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
64
+ sampler_name = "Flow"
65
+ boundary_ratio = 0.9
66
+
67
+ # Load pretrained model if need
68
+ # The transformer_path is used for low noise model, the transformer_high_path is used for high noise model.
69
+ transformer_path = None
70
+ transformer_high_path = None
71
+ transformer_audio_path = None
72
+ bridge_path = None
73
+ vae_path = None
74
+ audio_vae_path = None
75
+ # Load lora model if need
76
+ # The lora_path is used for low noise model, the lora_high_path is used for high noise model.
77
+ lora_path = None
78
+ lora_high_path = None
79
+
80
+ # Other params
81
+ sample_size = [640, 352]
82
+ video_length = 81
83
+ fps = 24
84
+
85
+ # Use torch.float16 if GPU does not support torch.bfloat16
86
+ # Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
87
+ weight_dtype = torch.bfloat16
88
+
89
+ # Input image for I2V
90
+ validation_image = "asset/8.png"
91
+
92
+ # prompts
93
+ prompt = "Medium shot of a girl by the ocean. She starts with a bright smile, then gently nods her head while speaking. Her mouth moves naturally to say: \"Hi, nice to meet you.\" She maintains eye contact throughout. The background shows calm waves. Smooth motion, cinematic quality, realistic facial expressions."
94
+ negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指"
95
+ guidance_scale = 5.0
96
+ seed = 43
97
+ num_inference_steps = 50
98
+ # The lora_weight is used for low noise model, the lora_high_weight is used for high noise model.
99
+ lora_weight = 0.55
100
+ lora_high_weight = 0.55
101
+ save_path = "samples/mova-videos-i2v"
102
+
103
+ device = set_multi_gpus_devices(ulysses_degree, ring_degree)
104
+
105
+ # The from_pretrained method automatically converts WanModel config to WanTransformer3DModel config
106
+ print("Loading Video DiT (High Noise) with WanTransformer3DModel...")
107
+ transformer = WanTransformer3DModel.from_pretrained(
108
+ model_name,
109
+ subfolder="video_dit_2",
110
+ low_cpu_mem_usage=True,
111
+ torch_dtype=weight_dtype,
112
+ )
113
+
114
+ if transformer_path is not None:
115
+ print(f"From checkpoint: {transformer_path}")
116
+ if transformer_path.endswith("safetensors"):
117
+ from safetensors.torch import load_file
118
+ state_dict = load_file(transformer_path)
119
+ else:
120
+ state_dict = torch.load(transformer_path, map_location="cpu")
121
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
122
+ m, u = transformer.load_state_dict(state_dict, strict=False)
123
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
124
+
125
+ # Video DiT 2 (Low Noise) - Using WanTransformer3DModel
126
+ print("Loading Video DiT 2 (Low Noise) with WanTransformer3DModel...")
127
+ transformer_2 = WanTransformer3DModel.from_pretrained(
128
+ model_name,
129
+ subfolder="video_dit",
130
+ low_cpu_mem_usage=True,
131
+ torch_dtype=weight_dtype,
132
+ )
133
+
134
+ if transformer_high_path is not None:
135
+ print(f"From checkpoint: {transformer_high_path}")
136
+ if transformer_high_path.endswith("safetensors"):
137
+ from safetensors.torch import load_file
138
+ state_dict = load_file(transformer_high_path)
139
+ else:
140
+ state_dict = torch.load(transformer_high_path, map_location="cpu")
141
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
142
+ m, u = transformer_2.load_state_dict(state_dict, strict=False)
143
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
144
+
145
+ # Audio DiT - Using WanAudioTransformer3DModel
146
+ print("Loading Audio DiT with WanAudioTransformer3DModel...")
147
+ transformer_audio = WanAudioTransformer3DModel.from_pretrained(
148
+ model_name,
149
+ subfolder="audio_dit",
150
+ low_cpu_mem_usage=True,
151
+ torch_dtype=weight_dtype,
152
+ )
153
+
154
+ if transformer_audio_path is not None:
155
+ print(f"From checkpoint: {transformer_audio_path}")
156
+ if transformer_audio_path.endswith("safetensors"):
157
+ from safetensors.torch import load_file
158
+ state_dict = load_file(transformer_audio_path)
159
+ else:
160
+ state_dict = torch.load(transformer_audio_path, map_location="cpu")
161
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
162
+ m, u = transformer_audio.load_state_dict(state_dict, strict=False)
163
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
164
+
165
+ # Dual Tower Bridge
166
+ print("Loading Dual Tower Bridge...")
167
+ dual_tower_bridge = MOVADualTowerConditionalBridge.from_pretrained(
168
+ model_name,
169
+ subfolder="dual_tower_bridge",
170
+ low_cpu_mem_usage=True,
171
+ torch_dtype=weight_dtype,
172
+ )
173
+
174
+ if bridge_path is not None:
175
+ print(f"From checkpoint: {bridge_path}")
176
+ if bridge_path.endswith("safetensors"):
177
+ from safetensors.torch import load_file
178
+ state_dict = load_file(bridge_path)
179
+ else:
180
+ state_dict = torch.load(bridge_path, map_location="cpu")
181
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
182
+ m, u = dual_tower_bridge.load_state_dict(state_dict, strict=False)
183
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
184
+
185
+ # Video VAE
186
+ print("Loading Video VAE...")
187
+ vae = AutoencoderKLWan.from_pretrained(
188
+ os.path.join(model_name, "video_vae/diffusion_pytorch_model.safetensors")
189
+ ).to(weight_dtype)
190
+
191
+ if vae_path is not None:
192
+ print(f"From checkpoint: {vae_path}")
193
+ if vae_path.endswith("safetensors"):
194
+ from safetensors.torch import load_file
195
+ state_dict = load_file(vae_path)
196
+ else:
197
+ state_dict = torch.load(vae_path, map_location="cpu")
198
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
199
+ m, u = vae.load_state_dict(state_dict, strict=False)
200
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
201
+
202
+ audio_vae = AutoencoderKLMOVAAudio.from_pretrained(
203
+ model_name,
204
+ subfolder="audio_vae",
205
+ torch_dtype=torch.float32,
206
+ )
207
+
208
+ if audio_vae_path is not None:
209
+ print(f"From checkpoint: {audio_vae_path}")
210
+ if audio_vae_path.endswith("safetensors"):
211
+ from safetensors.torch import load_file
212
+ state_dict = load_file(audio_vae_path)
213
+ else:
214
+ state_dict = torch.load(audio_vae_path, map_location="cpu")
215
+ state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
216
+ m, u = audio_vae.load_state_dict(state_dict, strict=False)
217
+ print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
218
+
219
+ # Get Tokenizer
220
+ print("Loading Tokenizer...")
221
+ tokenizer = AutoTokenizer.from_pretrained(
222
+ model_name,
223
+ subfolder="tokenizer",
224
+ )
225
+
226
+ # Get Text Encoder
227
+ print("Loading Text Encoder...")
228
+ text_encoder = UMT5EncoderModel.from_pretrained(
229
+ model_name,
230
+ subfolder="text_encoder",
231
+ low_cpu_mem_usage=True,
232
+ torch_dtype=weight_dtype,
233
+ )
234
+ text_encoder = text_encoder.eval()
235
+
236
+ # Get Scheduler
237
+ print("Loading Scheduler...")
238
+ Chosen_Scheduler = {
239
+ "Flow": FlowMatchEulerDiscreteScheduler,
240
+ "Flow_Unipc": FlowUniPCMultistepScheduler,
241
+ "Flow_DPM++": FlowDPMSolverMultistepScheduler,
242
+ }[sampler_name]
243
+ scheduler = Chosen_Scheduler.from_pretrained(
244
+ model_name,
245
+ subfolder="scheduler"
246
+ )
247
+
248
+ # Build Pipeline
249
+ print("Building MOVAPipeline Pipeline...")
250
+ pipeline = MOVAPipeline(
251
+ vae=vae,
252
+ audio_vae=audio_vae,
253
+ text_encoder=text_encoder,
254
+ tokenizer=tokenizer,
255
+ scheduler=scheduler,
256
+ transformer=transformer,
257
+ transformer_2=transformer_2,
258
+ transformer_audio=transformer_audio,
259
+ dual_tower_bridge=dual_tower_bridge,
260
+ audio_vae_type="dac",
261
+ )
262
+
263
+ if ulysses_degree > 1 or ring_degree > 1:
264
+ from functools import partial
265
+
266
+ # Enable multi-GPU inference for visual transformers
267
+ transformer.enable_multi_gpus_inference()
268
+ transformer_2.enable_multi_gpus_inference()
269
+
270
+ if fsdp_dit:
271
+ # Apply FSDP to visual transformer blocks
272
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
273
+ pipeline.transformer = shard_fn(pipeline.transformer)
274
+ pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
275
+ print("Add FSDP DIT")
276
+
277
+ if fsdp_text_encoder:
278
+ shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
279
+ pipeline.text_encoder = shard_fn(pipeline.text_encoder)
280
+ print("Add FSDP TEXT ENCODER")
281
+
282
+ if compile_dit:
283
+ # Compile MOVAModel blocks
284
+ # NOTE: compile_dit is not compatible with fsdp_dit
285
+ if fsdp_dit:
286
+ print("WARNING: compile_dit is not compatible with fsdp_dit. Disabling compile.")
287
+ else:
288
+ for i in range(len(pipeline.transformer.blocks)):
289
+ pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
290
+ for i in range(len(pipeline.transformer_2.blocks)):
291
+ pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
292
+ for i in range(len(pipeline.transformer_audio.blocks)):
293
+ pipeline.transformer_audio.blocks[i] = torch.compile(pipeline.transformer_audio.blocks[i])
294
+ print("Add Compile")
295
+
296
+ if GPU_memory_mode == "sequential_cpu_offload":
297
+ replace_parameters_by_name(pipeline.transformer, ["modulation",], device=device)
298
+ replace_parameters_by_name(pipeline.transformer_2, ["modulation",], device=device)
299
+ pipeline.transformer.freqs = pipeline.transformer.freqs.to(device=device)
300
+ pipeline.transformer_2.freqs = pipeline.transformer_2.freqs.to(device=device)
301
+ pipeline.enable_sequential_cpu_offload(device=device)
302
+ elif GPU_memory_mode == "model_group_offload":
303
+ register_auto_device_hook(pipeline.transformer)
304
+ safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
305
+ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
306
+ convert_model_weight_to_float8(pipeline.transformer, exclude_module_name=["modulation",], device=device)
307
+ convert_model_weight_to_float8(pipeline.transformer_2, exclude_module_name=["modulation",], device=device)
308
+ convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
309
+ convert_weight_dtype_wrapper(pipeline.transformer_2, weight_dtype)
310
+ pipeline.enable_model_cpu_offload(device=device)
311
+ elif GPU_memory_mode == "model_cpu_offload":
312
+ pipeline.enable_model_cpu_offload(device=device)
313
+ elif GPU_memory_mode == "model_full_load_and_qfloat8":
314
+ convert_model_weight_to_float8(pipeline.transformer, exclude_module_name=["modulation",], device=device)
315
+ convert_model_weight_to_float8(pipeline.transformer_2, exclude_module_name=["modulation",], device=device)
316
+ convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
317
+ convert_weight_dtype_wrapper(pipeline.transformer_2, weight_dtype)
318
+ pipeline.to(device=device)
319
+ else:
320
+ pipeline.to(device=device)
321
+
322
+ generator = torch.Generator(device=device).manual_seed(seed)
323
+
324
+ if lora_path is not None:
325
+ pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
326
+ pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
327
+
328
+ # Run inference
329
+ print("Running inference...")
330
+ with torch.no_grad():
331
+ image = Image.open(validation_image).convert("RGB")
332
+ output = pipeline(
333
+ prompt=prompt,
334
+ image=image,
335
+ negative_prompt=negative_prompt,
336
+ height=sample_size[0],
337
+ width=sample_size[1],
338
+ num_frames=video_length,
339
+ frame_rate=fps,
340
+ num_inference_steps=num_inference_steps,
341
+ guidance_scale=guidance_scale,
342
+ generator=generator,
343
+ boundary=boundary_ratio,
344
+ )
345
+
346
+ if lora_path is not None:
347
+ pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
348
+ pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
349
+
350
+ sample = output.videos
351
+ audio = output.audio
352
+
353
+ # Get audio sample rate from pipeline
354
+ audio_sample_rate = pipeline.audio_sample_rate
355
+
356
+ def save_results():
357
+ if not os.path.exists(save_path):
358
+ os.makedirs(save_path, exist_ok=True)
359
+
360
+ index = len([path for path in os.listdir(save_path)]) + 1
361
+ prefix = str(index).zfill(8)
362
+ if video_length == 1:
363
+ video_path = os.path.join(save_path, prefix + ".png")
364
+
365
+ image = sample[0, :, 0]
366
+ image = image.transpose(0, 1).transpose(1, 2)
367
+ image = (image * 255).numpy().astype(np.uint8)
368
+ image = Image.fromarray(image)
369
+ image.save(video_path)
370
+ else:
371
+ video_path = os.path.join(save_path, prefix + ".mp4")
372
+ sr = getattr(pipeline.audio_vae.config, "output_sampling_rate", audio_sample_rate)
373
+ save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
374
+
375
+ if ulysses_degree > 1 or ring_degree > 1:
376
+ import torch.distributed as dist
377
+ if dist.get_rank() == 0:
378
+ save_results()
379
+ else:
380
+ save_results()