Fun CN: fp8 stream DiT + local bnb4 TE + xlarge (no bf16 host dump / no remote TE)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +6 -0
- README.md +19 -9
- app.py +335 -67
- requirements.txt +1 -0
- vendor/VideoX-Fun/LICENSE +201 -0
- vendor/VideoX-Fun/README.md +733 -0
- vendor/VideoX-Fun/config/flux2/flux2_control.yaml +5 -0
- vendor/VideoX-Fun/config/qwenimage/qwenimage_control.yaml +5 -0
- vendor/VideoX-Fun/config/wan2.1/wan_civitai.yaml +39 -0
- vendor/VideoX-Fun/config/wan2.2/wan_civitai_5b.yaml +41 -0
- vendor/VideoX-Fun/config/wan2.2/wan_civitai_animate.yaml +41 -0
- vendor/VideoX-Fun/config/wan2.2/wan_civitai_i2v.yaml +43 -0
- vendor/VideoX-Fun/config/wan2.2/wan_civitai_s2v.yaml +44 -0
- vendor/VideoX-Fun/config/wan2.2/wan_civitai_t2v.yaml +43 -0
- vendor/VideoX-Fun/config/z_image/z_image_control.yaml +5 -0
- vendor/VideoX-Fun/config/z_image/z_image_control_2.0.yaml +8 -0
- vendor/VideoX-Fun/config/z_image/z_image_control_2.1.yaml +8 -0
- vendor/VideoX-Fun/config/z_image/z_image_control_2.1_lite.yaml +8 -0
- vendor/VideoX-Fun/config/zero_stage2_config.json +16 -0
- vendor/VideoX-Fun/config/zero_stage3_config.json +28 -0
- vendor/VideoX-Fun/config/zero_stage3_config_cpu_offload.json +28 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/app.py +73 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/launch_api.py +90 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/post_infer.py +150 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/post_infer_queue.py +145 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/predict_i2v.py +328 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/predict_t2v.py +268 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v.py +263 -0
- vendor/VideoX-Fun/examples/cogvideox_fun/predict_v2v_control.py +248 -0
- vendor/VideoX-Fun/examples/ernie_image/predict_t2i.py +210 -0
- vendor/VideoX-Fun/examples/fantasytalking/predict_s2v.py +335 -0
- vendor/VideoX-Fun/examples/flashhead/predict_s2v.py +262 -0
- vendor/VideoX-Fun/examples/flux/predict_t2i.py +224 -0
- vendor/VideoX-Fun/examples/flux2/predict_t2i.py +218 -0
- vendor/VideoX-Fun/examples/flux2_fun/predict_i2i_inpaint.py +258 -0
- vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control.py +258 -0
- vendor/VideoX-Fun/examples/flux2_fun/predict_t2i_control_ref.py +258 -0
- vendor/VideoX-Fun/examples/hunyuanvideo/predict_i2v.py +270 -0
- vendor/VideoX-Fun/examples/hunyuanvideo/predict_t2v.py +255 -0
- vendor/VideoX-Fun/examples/infinitetalk/predict_s2v.py +319 -0
- vendor/VideoX-Fun/examples/lens/predict_t2i.py +226 -0
- vendor/VideoX-Fun/examples/longcatvideo/predict_i2v.py +247 -0
- vendor/VideoX-Fun/examples/longcatvideo/predict_s2v_avatar.py +293 -0
- vendor/VideoX-Fun/examples/longcatvideo/predict_t2v.py +239 -0
- vendor/VideoX-Fun/examples/ltx2.3/predict_i2v.py +305 -0
- vendor/VideoX-Fun/examples/ltx2.3/predict_t2v.py +300 -0
- vendor/VideoX-Fun/examples/ltx2/predict_i2v.py +281 -0
- vendor/VideoX-Fun/examples/ltx2/predict_i2v_upsample.py +326 -0
- vendor/VideoX-Fun/examples/ltx2/predict_t2v.py +276 -0
- 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 (
|
| 20 |
|
| 21 |
| Piece | Value |
|
| 22 |
|---|---|
|
| 23 |
-
| Painter
|
| 24 |
-
|
|
| 25 |
-
|
|
| 26 |
-
|
|
| 27 |
-
|
|
| 28 |
-
|
|
|
|
|
| 29 |
|
| 30 |
Requires Space secret **`HF_TOKEN`** (gated FLUX.2-dev license accepted on the account).
|
| 31 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
## API
|
| 33 |
|
| 34 |
-
Same client contract
|
| 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 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 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 |
-
|
|
|
|
| 124 |
if GPU_SIZE not in ("large", "xlarge"):
|
| 125 |
-
GPU_SIZE = "
|
| 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
|
| 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
|
| 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 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 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 =
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
)
|
| 289 |
convert_weight_dtype_wrapper(transformer, WEIGHT_DTYPE)
|
| 290 |
-
|
| 291 |
-
elif MEM_MODE == "sequential_cpu_offload":
|
| 292 |
_PIPE.enable_sequential_cpu_offload(device=device)
|
| 293 |
-
elif MEM_MODE
|
| 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"-
|
|
|
|
| 403 |
f"- GPU: `size={GPU_SIZE}` duration={GPU_DURATION}s mem=`{MEM_MODE}`\n"
|
| 404 |
-
"-
|
| 405 |
-
"-
|
| 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 |
+
[](https://huggingface.co/spaces/alibaba-pai/CogVideoX-Fun-5b)
|
| 7 |
+
|
| 8 |
+
Wan-Fun:
|
| 9 |
+
[](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 |
+

|
| 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 |
+
[](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 |
+

|
| 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()
|