Author SHA1 Message Date
bubbliiiing d1c581bb17 Merge branch 'main' into speed 2026-09-04 21:47:41 +08:00
hkz 968f0e2192 Add PDD for MiniMax-H3 (#515) 2026-09-04 17:02:21 +08:00
谢翊凡and谢翊凡 43739895a1 fix: use dtype instead of deprecated torch_dtype (#514)
* fix: use dtype instead of deprecated torch_dtype for transformers >= 4.56

config.torch_dtype and the torch_dtype keyword argument were deprecated in transformers 4.56 (PR #39782). Use dtype when accepted and fall back to torch_dtype via try/except TypeError for older versions.

* fix: use dtype instead of deprecated torch_dtype for transformers >= 4.56

config.torch_dtype and the torch_dtype keyword argument were deprecated in transformers 4.56 (PR #39782). Pass dtype based on the installed transformers version (packaging.version), falling back to torch_dtype on older versions.

---------

Co-authored-by: 谢翊凡 <xyf5432@users.noreply.github.com>
2026-09-02 16:47:39 +08:00
bubbliiiing d31de92872 Merge branch 'main' into speed 2026-08-26 17:22:10 +08:00
Bubbliiiing 6f3fb60dad Update scale fp8 with fsdp (#508) 2026-08-25 12:35:36 +08:00
bubbliiiing 35bfe679dc Merge branch 'main' into speed 2026-08-24 22:26:30 +08:00
Bubbliiiing b0acf916c2 Update Lingbot Worl, Lingbot Video, Minimax-H3 and Minimax-H3 Control (#506) 2026-08-24 16:47:38 +08:00
hkz 6787dc8ed4 Add DFD for Wan2.1 and Wan2.2 (#505) 2026-08-13 16:56:53 +08:00
bubbliiiing d555ea056b Update speed test 2026-07-24 16:59:48 +08:00
bubbliiiing b6e5a32b3e Merge branch 'main' into speed 2026-07-24 16:30:07 +08:00
hkz 248ab0ac0e Add Causal-Focing for Wan2.1 (#500) 2026-07-22 17:37:05 +08:00
Bubbliiiing 403f1f7b78 Fix Bug in Self Forcing in Multi-Gpus Infernece && Update Training Code and Docs && Update Reward Models (#498) 2026-07-14 10:30:54 +08:00
Bubbliiiing 1fd9ed9208 Ode training && Update Lens model && Update LTX2 upsampler (#497) 2026-06-09 15:20:04 +08:00
Bubbliiiing 2b5596b8e6 Update Self-Forcing and Ernie Image (#490) 2026-05-25 17:39:17 +08:00
Bubbliiiing 804a4258e2 Update Flash Head && Update Readmes (#491) 2026-05-12 16:48:47 +08:00
Tyler Luan 8fb0bb165c Align with the official FlashHead implementation (#488)
- Update predict_s2v.py with improved inference configuration
- Align pipeline_flashhead.py with official streaming, scheduling, conditioning, and motion-cache behavior
2026-05-08 14:04:09 +08:00
Bubbliiiing 0266bab98b Update Model Group Offload and Readmes (#487) 2026-05-06 10:48:08 +08:00
Bubbliiiing 199a43b544 Update READMEs, add LongCatVideo multi-GPU inference, and fix import bugs (#485) 2026-04-24 15:54:22 +08:00
Bubbliiiing 34036517a8 Added LTX-2.3 support, optimized memory usage by casting e and e0 types in Wan-based models, and updated READMEs for image training and digital human models. (#483) 2026-04-16 20:45:16 +08:00
Bubbliiiing 54bf97ab66 Fix bug in FA3 and LTX2 && Update FA4 support && Update Infinitalk && Update MOVA && Update FlashHead (#479) 2026-04-07 14:08:39 +08:00
Bubbliiiing 4a86483cc2 Reformat S2V models && Update LTX-2 (#476) 2026-03-20 10:45:57 +08:00
Bubbliiiing ad72867c0f Fix Bug in Group Offload when diffusers version is low. (#474) 2026-03-10 15:02:16 +08:00
Bubbliiiing 5202421e7c Update Z Image Turbo Control 2602 (#460) 2026-03-04 16:17:04 +08:00
Sense_wangandhaosenwang1018 745acc1f47 Fix: Replace 74 bare excepts with except Exception (#458)
Co-authored-by: haosenwang1018 <haosenwang1018@users.noreply.github.com>
2026-02-26 16:24:51 +08:00
Bubbliiiing a4b60f40bf Update Z Image Control Tile (#457) 2026-02-25 11:12:57 +08:00
Bubbliiiing 7a085c9535 Update Longcat Avatar (#453) 2026-02-13 10:14:14 +08:00
Bubbliiiing 1ca4162447 Support validations in all models && Update Saving Code && Preparing Code (#452) 2026-02-10 17:23:23 +08:00
Bubbliiiing ee44fbc950 Lora Loading Log and Z Image Distill (#450) 2026-02-04 16:55:24 +08:00
bubbliiiing 67ab552885 Merge branch 'main' into speed 2025-08-01 14:45:57 +08:00
bubbliiiing 413f831ced Update Wan2.2 Speed 2025-07-31 19:28:17 +08:00
bubbliiiing f4ffd1c64b Update Wan2.2 Speed 2025-07-31 15:55:06 +08:00
bubbliiiing ed42563dd3 Update Wan2.2 Speed 2025-07-31 15:00:46 +08:00
bubbliiiing 8ee48e420c Update Wan2.2 Speed 2025-07-31 14:31:08 +08:00
bubbliiiing b757b8edae Merge branch 'main' into speed 2025-07-31 14:17:20 +08:00
bubbliiiing 924dd8528a Update auto_tile_batch_size args 2025-07-30 19:39:31 +08:00
bubbliiiing e258d4158b Update ui 2025-07-30 14:20:47 +08:00
bubbliiiing 48c8288323 Update api and ui 2025-07-30 12:53:45 +08:00
bubbliiiing 1aacfe6bca Update Predict Code 2025-07-30 11:04:49 +08:00
bubbliiiing 288e88eb46 Update Wan2.2 UI 2025-07-30 10:54:08 +08:00
bubbliiiing dfb8ca04a7 Update Wan2.2 UI 2025-07-30 10:52:02 +08:00
bubbliiiing 499eef12c3 Update Wan2.2 UI 2025-07-30 10:48:59 +08:00
bubbliiiing fa8623b22b Update speed sh 2025-06-03 02:42:30 +00:00
bubbliiiing 61fb59833f Update i2v speed 2025-06-02 16:24:49 +00:00
bubbliiiing 4b0b009fd3 Update i2v speed 2025-06-02 15:46:37 +00:00
bubbliiiing 518c2bdc7e Merge branch 'main' into speed 2025-05-29 03:35:06 +00:00
bubbliiiing 6188e66fc4 Merge branch 'main' into speed 2025-05-26 10:41:24 +00:00
bubbliiiing 7c9d655822 Merge branch 'main' into speed 2025-05-19 03:13:34 +00:00
bubbliiiing 873f622dae Speed Test 2025-05-16 02:42:35 +00:00
577 changed files with 211135 additions and 14923 deletions
+12 -10
View File
@@ -1,17 +1,19 @@
# Byte-compiled / optimized / DLL files
models*
output*
logs*
taming*
samples*
datasets*
asset*
# Used in VideoX-Fun
_*
logs*
/models*
/output*
/logs*
/taming*
/samples*
/datasets*
/asset*
/repo*
/scripts_demo*
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
scripts_demo*
# C extensions
*.so
+128
View File
@@ -0,0 +1,128 @@
---
name: integrating-models
description: Guides adding, porting, or onboarding a diffusion model (transformer/VAE/encoder, inference pipeline, training script, config) into the VideoX-Fun repository by mirroring the closest existing model family and maximizing reuse of the repository's existing code and shared infrastructure. Use when integrating a new model/architecture, or when creating predict_*.py inference scripts, scripts/*/train*.py training scripts, pipeline_*.py, config/*.yaml, or model definitions under videox_fun/models/.
---
# Integrating Models into VideoX-Fun
## Core rule: maximize reuse of existing repo code — mirror, extend, never reinvent
**Prime directive: reuse this repository's existing code to the maximum.** Nearly every building block you need already exists in `videox_fun/` or in a sibling model family. Your job is to **find it, import it, and extend it** — not to write a parallel implementation. A new file should be mostly reused structure plus the genuinely model-specific delta; the less new code you write, the better.
**Reuse-first protocol — before writing ANY new function / class / util:**
1. **Search the repo first.** Grep `videox_fun/` and the closest family for an existing equivalent (weight loader, scheduler, sampler, offload, attention, LoRA, fp8, dataset, dist helper, save/metric util). If one exists → **import and reuse it**. If it is 80% right → **extend / parameterize it**, do not fork it.
2. **Only if nothing exists** may you add new code — and then put it in the shared layer (`videox_fun/utils`, `videox_fun/data`, `videox_fun/dist`) so the next model reuses it too, instead of burying it in a family folder.
3. **Never copy-paste** a util into a new file (that creates drift); import the single source of truth.
**Mirror the closest family.** Every model follows the **same layered template**. Integrating a model means finding the closest existing family and mirroring its structure, changing only what genuinely differs:
1. Pick the closest existing family by task type (t2v / i2v / v2v-control / s2v / t2i / edit / distill): `wan2.1`, `wan2.1_fun`, `wan2.2`, `qwenimage`, `flux2`, `minimax_h3`, `ltx2`, `longcatvideo`, `cogvideox_fun`, `z_image`, etc.
2. Read that family end-to-end across all layers:
- `examples/<family>/predict_*.py` (inference entry)
- `scripts/<family>/train*.py` + `*.sh` + `README_TRAIN*.md` (training)
- `videox_fun/pipeline/pipeline_<family>*.py` (pipeline)
- `videox_fun/models/<family>_*.py` (model definitions)
- `config/<family>/*.yaml` (config)
3. Copy that structure and adapt. Keep names, argument sets, control flow, and reuse points identical in shape.
Writing a bespoke pipeline, weight loader, trainer, sampler, dataset, or offload scheme from scratch is a **failure mode**. If you are tempted to, **stop** and check the Reuse inventory below first.
## Repository layout (where each layer lives)
| Layer | Location | What it is |
|-------|----------|------------|
| Model definitions | `videox_fun/models/<family>_*.py` | Transformer / VAE / text-audio-image encoders. Diffusers `ModelMixin`+`ConfigMixin`, `@register_to_config`, custom `from_pretrained`. |
| Model registry | `videox_fun/models/__init__.py` | Imports every model class. **Must be updated** for a new model. |
| Inference pipelines | `videox_fun/pipeline/pipeline_<family>*.py` | `<Family>Pipeline(DiffusionPipeline)` with `__call__`. |
| Pipeline registry | `videox_fun/pipeline/__init__.py` | Imports every pipeline + aliases. **Must be updated.** |
| Configs (optional) | `config/<family>/*.yaml` | OmegaConf YAML for civitai/custom layouts; a standard diffusers-layout checkpoint can load without one. |
| Inference entry scripts | `examples/<family>/predict_*.py` | User-facing, config-block-at-top runnable scripts. |
| Inference services | `examples/<family>/{app.py,launch_api.py,post_infer*.py}` | Gradio UI / API server / batch inference. |
| Training scripts | `scripts/<family>/train*.py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py`, ... |
| Training launchers | `scripts/<family>/train*.sh` | `accelerate launch` / DeepSpeed command with full arg list. |
| Training docs | `scripts/<family>/README_TRAIN*.md` | Bilingual pairs: `README_TRAIN.md` + `README_TRAIN_zh-CN.md`. |
| Shared: schedulers/utils | `videox_fun/utils/` | `fm_solvers`, `fm_solvers_unipc`, `lora_utils`, `fp8_optimization`, `group_offload`, `utils.py`. |
| Shared: distributed | `videox_fun/dist/` | `fsdp.shard_model`, `fuser.set_multi_gpus_devices`, `<family>_xfuser` sequence-parallel attention. |
| Shared: data | `videox_fun/data/` | Datasets (`ImageVideoDataset`, `VideoDataset`, ...) + bucket/aspect-ratio samplers. |
| Demo / test datasets | `datasets/X-Fun-*-Demo/` | Ready-made smoke-test data, downloaded via `modelscope download --dataset PAI/<name>`; each ships several `metadata*.json` variants. **The only test data to use** (see reference.md §8). |
| Preprocessing (data gen) | `scripts/<family>/generate_*.py` / `train_preprocess.py` (+ `.sh`) | Offline multi-GPU generation of cached training data (latents / ODE pairs / embeddings) → per-sample `.safetensors` + `outputs.json`, loaded by `ImageVideoSafetensorsDataset`. |
| ComfyUI nodes | `comfyui/<family>/nodes.py` | Optional node integration mirroring the pipeline. |
## Integration workflow
Copy this checklist and track progress:
```
Integration Progress:
- [ ] Step 0: Choose the closest family to mirror; read it across all layers
- [ ] Step 1: Model definitions in videox_fun/models/ + register in models/__init__.py
- [ ] Step 2: Pipeline in videox_fun/pipeline/ + register in pipeline/__init__.py
- [ ] Step 3: Config YAML in config/<family>/
- [ ] Step 4: Inference script(s) in examples/<family>/predict_*.py
- [ ] Step 5: Training script(s) in scripts/<family>/train*.py + .sh
- [ ] Step 6: Training docs README_TRAIN.md + README_TRAIN_zh-CN.md
- [ ] Step 7: Reuse audit + verification (incl. smoke test on the matching demo dataset)
```
**Step 0 — Choose the mirror.** Match by task and architecture. A new control model mirrors an existing `*_fun`/`*_control` family; a new audio/talking model mirrors `minimax_h3`/`longcatvideo`/`infinitetalk`; a new image model mirrors `qwenimage`/`flux2`/`z_image`.
**Step 1 — Model.** Create `videox_fun/models/<family>_transformer3d.py` (or `2d`), `<family>_vae.py`, encoders as needed. Mirror the class shape: `class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin)`, `_supports_gradient_checkpointing = True`, `@register_to_config __init__`, and a `from_pretrained` that supports `transformer_additional_kwargs`, `dict_mapping`, `low_cpu_mem_usage`, and missing-key init. Add imports to `videox_fun/models/__init__.py`.
**Step 2 — Pipeline.** Create `videox_fun/pipeline/pipeline_<family>.py`. Mirror `pipeline_wan.py`: module-level `retrieve_timesteps`, a `<Family>PipelineOutput(BaseOutput)` dataclass, `<Family>Pipeline(DiffusionPipeline)` with `model_cpu_offload_seq`, `_callback_tensor_inputs`, `__init__(vae, tokenizer, text_encoder, transformer, scheduler, ...)`, `encode_prompt`, and `__call__`. Add imports/aliases to `videox_fun/pipeline/__init__.py`.
**Step 3 — Config (optional).** A YAML under `config/<family>/` is **not always required**. It is needed mainly for **civitai-format / custom single-file layouts** — to supply `transformer_additional_kwargs`, `dict_mapping` (civitai key → `__init__` kwarg), component subpaths, and `vae_kwargs`/`text_encoder_kwargs`/`scheduler_kwargs`/`image_encoder_kwargs`. For a **standard diffusers-layout** checkpoint (`model_index.json` + per-subfolder `config.json`), load directly via `from_pretrained(model_name, subfolder=...)` with no YAML — mirror `examples/minimax_h3_fun/predict_v2v_control.py`, which guards `if config_path is not None:`. When you do add a YAML, load it via `OmegaConf.load(config_path)` and spread into `from_pretrained` instead of hardcoding those values.
**Step 4 — Inference script.** Create `examples/<family>/predict_<task>.py` following the exact template (config block at top → component loading → scheduler dict → pipeline construction → multi-GPU/FSDP/compile → `GPU_memory_mode` branching → TeaCache → LoRA merge → inference → `save_results`). See [examples.md](examples.md).
**Step 5 — Training script.** Create `scripts/<family>/train.py` (+ `train_lora.py` etc.). Mirror the shared structure: license header, `sys.path` bootstrap, imports from `videox_fun`, `log_validation()` that **reuses the inference Pipeline**, `parse_args()` (reuse the existing shared argument set), `main()`. Add a `train.sh` launcher. Reuse `videox_fun.data` datasets/samplers — do not write a new dataset.
**Step 6 — Docs.** Write `README_TRAIN.md` and `README_TRAIN_zh-CN.md` as an aligned bilingual pair (same structure, same commands/params, matching section order).
**Step 7 — Reuse audit + verification.** Confirm you reused shared infra (below), smoke-test the new train/predict path on the **matching official demo dataset** under `datasets/X-Fun-*-Demo/` (pick by task and metadata variant — see reference.md §8), then run the verification checklist. Never invent an ad-hoc test set and never leave `datasets/internal_datasets/` placeholders in shipped scripts/docs.
## Reuse inventory (use these, do not reimplement)
**Reuse-first catalog: import from here instead of reimplementing. If a helper you need is not listed, grep `videox_fun/` and the closest family before writing your own.**
- **Schedulers**: `FlowMatchEulerDiscreteScheduler`, `videox_fun.utils.fm_solvers.FlowDPMSolverMultistepScheduler`, `fm_solvers_unipc.FlowUniPCMultistepScheduler`. Selected via a `sampler_name` dict.
- **LoRA**: `videox_fun.utils.lora_utils` — `merge_lora`, `unmerge_lora`, `create_network`, `convert_peft_lora_to_kohya_lora`.
- **FP8 / quantization**: `videox_fun.utils.fp8_optimization` — `convert_model_weight_to_float8`, `convert_weight_dtype_wrapper`, `replace_parameters_by_name`.
- **Offloading**: `videox_fun.utils.group_offload` — `register_auto_device_hook`, `safe_enable_group_offload`; plus pipeline `enable_sequential_cpu_offload` / `enable_model_cpu_offload` / `.to(device)`.
- **Distributed**: `videox_fun.dist` — `set_multi_gpus_devices`, `shard_model` (FSDP), `<family>_xfuser` sequence-parallel attention processors, `enable_multi_gpus_inference()`.
- **IO / helpers**: `videox_fun.utils.utils` — `save_videos_grid`, `save_videos_with_audio_grid`, `get_image_to_video_latent`, `get_video_to_video_latent`, `get_image_latent`, `filter_kwargs`, `calculate_dimensions`.
- **Data**: `videox_fun.data` — `ImageVideoDataset`, `VideoDataset`, `ImageVideoControlDataset`, `VideoSpeechDataset`, bucket/aspect-ratio samplers, `get_closest_ratio`, `get_random_mask`.
- **Caching / speedups**: TeaCache (`models/cache_utils`, `get_teacache_coefficients`, `transformer.enable_teacache`), `enable_cfg_skip`, Riflex (`enable_riflex`), `torch.compile` on `transformer.blocks`.
- **Preprocessing (data gen, multi-GPU)**: mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py` — `accelerate launch` + `Accelerator` (interleaved rank sharding), config-driven `from_pretrained` for the teacher/VAE/text-encoder, `safetensors.torch.save_file` per sample + `outputs.json` index, consumed by `videox_fun.data.ImageVideoSafetensorsDataset`. Store as **safetensors only — never LMDB or `.pt`** (see reference.md §10).
## Non-negotiable conventions
- **Maximize reuse of existing repo code**: import existing `videox_fun/` helpers and mirror the closest family; never fork or copy-paste a util, and never write a parallel pipeline / loader / scheduler / sampler / offload. Genuinely-new shared code goes in `videox_fun/{utils,data,dist}` (so the next model reuses it), not buried in a family folder.
- **`sys.path` bootstrap**: every runnable script starts with the 3-level `project_roots` loop inserting into `sys.path` before importing `videox_fun`.
- **Config-driven loading (YAML optional)**: a `config/<family>/*.yaml` is required for civitai-format/custom layouts (it supplies `transformer_additional_kwargs`/`dict_mapping`/subpaths); it is **optional for standard diffusers-layout checkpoints**, which load directly via `from_pretrained(model_name, subfolder=...)`. When a YAML is used, don't hardcode the values it provides.
- **`GPU_memory_mode`**: support the standard six modes — `model_full_load`, `model_full_load_and_qfloat8`, `model_cpu_offload`, `model_cpu_offload_and_qfloat8`, `model_group_offload`, `sequential_cpu_offload` — with the exact branching order used in existing `predict_*.py`.
- **Naming**: files `<family>_transformer3d.py` / `<family>_vae.py` / `pipeline_<family>.py`; classes `<Family>Transformer3DModel` / `AutoencoderKL<Family>` / `<Family>Pipeline`.
- **Resolution args**: drive canvas size with a single square `--video_sample_size` (`type=int`, height = width); never `--video_sample_height` / `--video_sample_width`. For a fixed non-square shape add `--fix_sample_size` (`nargs=2, type=int`, `[height, width]`) that overrides the square size, and derive the effective height/width once in `parse_args()` (see reference.md §5).
- **Registries**: a model is not integrated until it is imported in BOTH `videox_fun/models/__init__.py` and `videox_fun/pipeline/__init__.py`.
- **Two weight formats**: support `civitai` and `diffusers` via config `format` + `dict_mapping` (maps civitai keys such as `in_dim`→`in_channels`, `dim`→`hidden_size`).
- **Bilingual docs**: training READMEs ship as EN + `_zh-CN` pairs with aligned structure and identical commands/params.
- **Test data = official demo datasets**: smoke tests, `log_validation` checks, launcher `.sh` defaults, and doc examples all point at `datasets/X-Fun-*-Demo/` (ModelScope `PAI/<name>`), with the metadata variant matching the task — `metadata_add_width_height.json` by default, `_add_objects.json` for VACE/subject-reference, `_add_wav.json` for audio-visual joint models, `metadata_lingbot_video_add_width_height.json` for `lingbot_video`. Selection matrix: reference.md §8.
- **Preprocessing = offline data generation, multi-GPU + safetensors**: cached training data (latents / ODE pairs / embeddings) is produced by `accelerate launch` scripts like `generate_ode_pairs.py` (interleaved rank sharding, resume by skipping existing files, `wait_for_everyone`, rank-0 JSON index) and saved with `safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`. **Never single-GPU / `cuda:0`; never LMDB or `.pt`/`torch.save` pickles for preprocessed data.** See reference.md §10.
## Verification checklist
- [ ] New model classes imported in `videox_fun/models/__init__.py`
- [ ] New pipeline(s) imported in `videox_fun/pipeline/__init__.py`
- [ ] Config YAML present **only if** the checkpoint is civitai-format/custom-layout; a diffusers-layout model may load directly via `from_pretrained(model_name, subfolder=...)` with no YAML. When a YAML is used, it drives component loading (no hardcoded kwargs)
- [ ] `predict_*.py` mirrors an existing script: `sys.path` bootstrap, config block, scheduler dict, `GPU_memory_mode` branching, LoRA merge, `save_results`
- [ ] `train*.py` reuses `videox_fun.data` + shared args, and `log_validation()` reuses the inference Pipeline
- [ ] `train*.sh` launcher provided (`accelerate launch` / DeepSpeed)
- [ ] Shared infra reused (schedulers / lora_utils / fp8 / group_offload / dist / utils / data) — nothing reimplemented
- [ ] Any offline data-generation/preprocessing script runs multi-GPU (`accelerate launch` + `Accelerator`) and saves cached tensors as **safetensors + `outputs.json`** for `ImageVideoSafetensorsDataset` — never LMDB or `.pt`
- [ ] `README_TRAIN.md` + `README_TRAIN_zh-CN.md` aligned pair present
- [ ] Smoke test / doc examples use the matching `datasets/X-Fun-*-Demo` dataset and the correct `metadata*.json` variant — no `internal_datasets` placeholders (reference.md §8)
- [ ] Optional: ComfyUI node in `comfyui/<family>/nodes.py` mirrors the pipeline
## Additional resources
- Detailed file-by-file conventions, class/method shapes, and the model-loading internals: [reference.md](reference.md)
- **Dataset & sampler selection matrix** (which `videox_fun.data` dataset/loader each training task uses), **demo-dataset / metadata-variant selection matrix** (which `datasets/X-Fun-*-Demo` to smoke-test with), **inference task matrix** (which pipeline each `predict_<task>.py` uses), and **multi-GPU preprocessing patterns**: [reference.md](reference.md) §8–§10
- Concrete skeletons (config YAML, `predict_*.py`, pipeline class, training script + DataLoader): [examples.md](examples.md)
+510
View File
@@ -0,0 +1,510 @@
# VideoX-Fun Integration Skeletons
Starting templates. **Always open the mirrored family's real file and adapt it** — these skeletons show shape and required reuse points, not full implementations. Replace `<family>` / `<Family>` / `<task>`.
## Config — `config/<family>/<variant>.yaml` (optional)
> **Not always required.** Author a YAML only for civitai-format / custom single-file layouts. A standard diffusers-layout checkpoint (`model_index.json` + per-subfolder `config.json`) loads directly via `from_pretrained(model_name, subfolder=...)` with no YAML — set `config_path = None` and guard `if config_path is not None:` (see `examples/minimax_h3_fun/predict_v2v_control.py`).
```yaml
format: civitai
pipeline: <Family>
transformer_additional_kwargs:
transformer_subpath: ./
dict_mapping:
in_dim: in_channels
dim: hidden_size
vae_kwargs:
vae_subpath: <Family>_VAE.pth
temporal_compression_ratio: 4
spatial_compression_ratio: 8
text_encoder_kwargs:
text_encoder_subpath: <text_encoder>.pth
tokenizer_subpath: <tokenizer_id>
text_length: 512
scheduler_kwargs:
scheduler_subpath: null
num_train_timesteps: 1000
shift: 5.0
# Only for i2v / models with a CLIP image encoder:
image_encoder_kwargs:
image_encoder_subpath: <image_encoder>.pth
```
## Inference — `examples/<family>/predict_<task>.py`
```python
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
from transformers import AutoTokenizer
# --- sys.path bootstrap (required, before importing videox_fun) ---
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKL<Family>, <Family>TextEncoder,
<Family>Transformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import <Family>Pipeline
from videox_fun.utils import register_auto_device_hook, safe_enable_group_offload
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# --- user config block (keep the conventional order + comments) ---
GPU_memory_mode = "sequential_cpu_offload"
ulysses_degree = 1
ring_degree = 1
fsdp_dit = False
fsdp_text_encoder = True
compile_dit = False
enable_teacache = True
teacache_threshold = 0.10
num_skip_start_steps = 5
teacache_offload = False
cfg_skip_ratio = 0
enable_riflex = False
riflex_k = 6
config_path = "config/<family>/<variant>.yaml"
model_name = "models/Diffusion_Transformer/<Family>-Model"
sampler_name = "Flow"
shift = 3
transformer_path = None
vae_path = None
lora_path = None
sample_size = [480, 832]
video_length = 81
fps = 16
weight_dtype = torch.bfloat16
prompt = "..."
negative_prompt = "..."
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/<family>-<task>"
# --- device + config (config_path may be None for a diffusers-layout checkpoint) ---
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path) # or guard: if config_path is not None: ... (then load components via subfolder=...)
# --- components (when a YAML is used, paths/kwargs come from config; otherwise pass subfolder=... directly) ---
transformer = <Family>Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True, torch_dtype=weight_dtype,
)
# optional transformer_path / vae_path override -> load_state_dict(strict=False) + print missing/unexpected
vae = AutoencoderKL<Family>.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')))
text_encoder = <Family>TextEncoder.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True, torch_dtype=weight_dtype).eval()
# --- scheduler selection dict ---
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler(**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])))
# --- pipeline ---
pipeline = <Family>Pipeline(vae=vae, tokenizer=tokenizer, text_encoder=text_encoder,
transformer=transformer, scheduler=scheduler)
# --- multi-gpu / fsdp / compile ---
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
pipeline.transformer = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.transformer)
if fsdp_text_encoder:
pipeline.text_encoder = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.text_encoder)
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
# --- GPU_memory_mode branching (keep this exact order) ---
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
# --- teacache / cfg_skip / riflex / lora ---
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
pipeline.transformer.enable_teacache(coefficients, num_inference_steps, teacache_threshold,
num_skip_start_steps=num_skip_start_steps, offload=teacache_offload)
if cfg_skip_ratio is not None:
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
# --- inference ---
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
if enable_riflex:
pipeline.transformer.enable_riflex(k=riflex_k, L_test=(video_length - 1) // vae.config.temporal_compression_ratio + 1)
sample = pipeline(prompt, num_frames=video_length, negative_prompt=negative_prompt,
height=sample_size[0], width=sample_size[1], generator=generator,
guidance_scale=guidance_scale, num_inference_steps=num_inference_steps,
shift=shift).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
# --- save (rank 0 only when multi-gpu) ---
def save_results():
os.makedirs(save_path, exist_ok=True)
prefix = str(len(os.listdir(save_path)) + 1).zfill(8)
if video_length == 1:
image = (sample[0, :, 0].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
Image.fromarray(image).save(os.path.join(save_path, prefix + ".png"))
else:
save_videos_grid(sample, os.path.join(save_path, prefix + ".mp4"), fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
```
For i2v, gate the CLIP image encoder and pass `video`/`mask_video`:
```python
if transformer.config.in_channels != vae.config.latent_channels:
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder'))).to(weight_dtype).eval()
input_video, input_video_mask, _ = get_image_to_video_latent(start_image, None, video_length=video_length, sample_size=sample_size)
# pipeline = <Family>InpaintPipeline(..., clip_image_encoder=clip_image_encoder)
# sample = pipeline(..., video=input_video, mask_video=input_video_mask).videos
```
## Pipeline class — `videox_fun/pipeline/pipeline_<family>.py`
```python
from dataclasses import dataclass
from typing import List, Optional, Union
import torch
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from ..models import AutoencoderKL<Family>, <Family>Transformer3DModel
from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler, get_sampling_sigmas
from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
logger = logging.get_logger(__name__)
EXAMPLE_DOC_STRING = """Examples:\n```python\npass\n```"""
# reuse retrieve_timesteps verbatim from pipeline_wan.py
@dataclass
class <Family>PipelineOutput(BaseOutput):
videos: torch.Tensor
class <Family>Pipeline(DiffusionPipeline):
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler):
super().__init__()
self.register_modules(tokenizer=tokenizer, text_encoder=text_encoder, vae=vae,
transformer=transformer, scheduler=scheduler)
# video_processor / vae_scale_factor / etc. as in pipeline_wan.py
def encode_prompt(self, prompt, negative_prompt, device, num_videos_per_prompt=1, ...):
... # mirror pipeline_wan.py
def prepare_latents(self, batch_size, num_channels_latents, height, width, num_frames, dtype, device, generator, latents=None):
...
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(self, prompt, negative_prompt=None, height=480, width=832, num_frames=81,
num_inference_steps=50, guidance_scale=6.0, generator=None, shift=1.0,
callback_on_step_end=None, return_dict=True, **kwargs) -> Union[<Family>PipelineOutput, tuple]:
# 1. encode_prompt 2. prepare_latents 3. retrieve_timesteps
# 4. denoising loop with guidance 5. vae.decode 6. return <Family>PipelineOutput(videos=...)
...
```
Then register in `videox_fun/pipeline/__init__.py`:
```python
from .pipeline_<family> import <Family>Pipeline
```
## Model class — `videox_fun/models/<family>_transformer3d.py`
```python
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
from .attention_utils import attention # unified FA/SDPA backend — do not hand-roll SDPA
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
_supports_gradient_checkpointing = True
@register_to_config
def __init__(self, model_type='t2v', in_dim=16, dim=2048, ffn_dim=8192,
num_heads=16, num_layers=32, in_channels=16, hidden_size=2048, ...):
super().__init__()
...
def _set_gradient_checkpointing(self, *args, **kwargs):
self.gradient_checkpointing = True
def enable_multi_gpus_inference(self): ... # route attn through dist/<family>_xfuser.py
def enable_teacache(self, ...): ...
def enable_cfg_skip(self, ...): ...
def forward(self, x, timestep, context, ...): ...
@classmethod
def from_pretrained(cls, pretrained_model_path, subfolder=None,
transformer_additional_kwargs=None, low_cpu_mem_usage=False,
torch_dtype=torch.bfloat16):
... # mirror wan_transformer3d.py: config.json -> dict_mapping -> init_empty_weights
# -> load .bin/.safetensors -> shape-filter -> initialize missing keys -> load
```
Then register in `videox_fun/models/__init__.py`:
```python
from .<family>_transformer3d import <Family>Transformer3DModel
from .<family>_vae import AutoencoderKL<Family>
```
## Training — `scripts/<family>/train.py` (key reuse points)
```python
"""Modified from https://github.com/huggingface/diffusers/.../train_text_to_image.py"""
import argparse, gc, logging, math, os, sys
import accelerate, diffusers, torch, transformers
from accelerate import Accelerator
from diffusers.optimization import get_scheduler
from omegaconf import OmegaConf
# same sys.path bootstrap as predict scripts
from videox_fun.data import (ASPECT_RATIO_512, AspectRatioBatchImageVideoSampler,
ImageVideoDataset, ImageVideoSampler, RandomSampler,
get_closest_ratio, get_random_mask)
from videox_fun.models import AutoencoderKL<Family>, <Family>Transformer3DModel
from videox_fun.pipeline import <Family>Pipeline # REUSED for validation
from videox_fun.utils.lora_utils import create_network # for train_lora
from videox_fun.utils.utils import save_videos_grid, get_image_to_video_latent
def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config,
accelerator, weight_dtype, global_step):
# build <Family>Pipeline from accelerator.unwrap_model(transformer3d),
# run validation_prompts, save_videos_grid to output_dir/sample/. Reuse the pipeline.
...
def parse_args():
parser = argparse.ArgumentParser(...)
# reuse the shared arg surface: --config_path, --pretrained_model_name_or_path,
# --train_data_dir, --train_data_meta, --video_sample_n_frames, --train_batch_size,
# --gradient_accumulation_steps, --learning_rate, --lr_scheduler, --checkpointing_steps,
# --output_dir, --mixed_precision, --gradient_checkpointing, --enable_bucket,
# --train_mode, --trainable_modules, --validation_prompts ... (add only what's needed)
return parser.parse_args()
def main():
args = parse_args()
accelerator = Accelerator(mixed_precision=args.mixed_precision, ...)
config = OmegaConf.load(args.config_path)
# load transformer/vae/text_encoder via config
# --- Dataset: pick by task (see reference.md §8) ---
# T2V/I2V base + inpaint -> ImageVideoDataset(enable_inpaint = args.train_mode != "normal")
# Control -> ImageVideoControlDataset(enable_camera_info = ...)
# Image edit -> ImageEditDataset
# Speech/audio (S2V) -> VideoSpeechDataset / VideoSpeechControlDataset
# Animate -> VideoAnimateDataset
# Distill text / GRPO / DPO -> TextDataset
# Smoke-test on the matching official demo dataset (reference.md §8), e.g.
# datasets/X-Fun-Videos-Demo + metadata_add_width_height.json for T2V/I2V.
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride,
video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size, enable_bucket=args.enable_bucket,
enable_inpaint=True if args.train_mode != "normal" else False)
# --- Sampler + DataLoader: branch on enable_bucket (see reference.md §8) ---
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
if args.enable_bucket:
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
batch_sampler = AspectRatioBatchImageVideoSampler(
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
aspect_ratios=aspect_ratio_sample_size)
def collate_fn(examples):
new_examples = {"pixel_values": [], "text": []}
if args.train_mode != "normal":
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
# get_closest_ratio -> Resize/CenterCrop/Normalize -> stack; masks via get_random_mask
return new_examples
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
else:
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
# trainable-module filtering or create_network for LoRA
# optimizer + get_scheduler; accelerator.prepare; checkpoint hooks
# training loop: timestep sampling -> transformer forward -> loss -> backward
# periodic log_validation(...); final save weights / LoRA
...
if __name__ == "__main__":
main()
```
## Launcher — `scripts/<family>/train.sh`
```bash
export MODEL_NAME="models/Diffusion_Transformer/<Family>-Model"
# Test data = the official demo dataset matching the task (reference.md §8). Download once, e.g.:
# modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
# T2I -> X-Fun-Images-Demo | control -> X-Fun-{Videos,Images}-Controls-Demo
# S2V -> X-Fun-Videos-Audios-Demo | image edit -> X-Fun-Images-Edit-Demo
export DATASET_NAME="datasets/X-Fun-Videos-Demo/" # = train_data_dir (data_root); media live under train/
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json" # = train_data_meta: [{"file_path","text","type","width","height"}] — see reference.md §8
# Metadata variants: VACE/subject-ref -> metadata_add_width_height_add_objects.json (X-Fun-Videos-Controls-Demo);
# audio-visual joint -> metadata_add_width_height_add_wav.json; lingbot_video -> metadata_lingbot_video_add_width_height.json
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/<family>/train.py \
--config_path="config/<family>/<variant>.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--checkpointing_steps=50 \
--output_dir="output_dir_<family>" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--enable_bucket \
--low_vram \
--train_mode="normal" \
--trainable_modules "."
```
## Preprocessing (data gen) — `scripts/<family>/generate_<...>.py`
Offline generation of cached training data (latents / ODE-trajectory pairs / prompt embeddings). **Always multi-GPU** (`accelerate launch` + `Accelerator`) and **always safetensors** (`safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`) — never LMDB, never `.pt`. Mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py`:
```python
# ...license header + sys.path bootstrap...
import argparse, json, math, os, torch
from accelerate import Accelerator
from omegaconf import OmegaConf
from safetensors.torch import save_file
from tqdm import tqdm
from videox_fun.models import AutoencoderKLWan, WanT5EncoderModel, WanTransformer3DModel # reuse repo models
from videox_fun.utils.utils import save_videos_grid # reuse repo IO
def main():
args = parse_args() # --pretrained_model_name_or_path --config_path --caption_path --output_folder
# --num_inference_steps --guidance_scale --shift --mixed_precision ...
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
torch.set_grad_enabled(False) # inference-only
torch.backends.cuda.matmul.allow_tf32 = True
config = OmegaConf.load(args.config_path) # config-driven loading (Section 3)
weight_dtype = {"fp16": torch.float16, "bf16": torch.bfloat16}.get(accelerator.mixed_precision, torch.float32)
text_encoder = WanT5EncoderModel.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), torch_dtype=weight_dtype).to(device).eval()
vae = AutoencoderKLWan.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['vae_kwargs'])).to(device, dtype=weight_dtype).eval()
transformer = WanTransformer3DModel.from_pretrained(..., transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])).to(device, dtype=weight_dtype).eval()
prompts = [l.rstrip() for l in open(args.caption_path, encoding="utf-8") if l.strip()]
os.makedirs(args.output_folder, exist_ok=True)
total_per_rank = math.ceil(len(prompts) / world_size)
for index in tqdm(range(total_per_rank), disable=rank != 0, desc="Generating"):
prompt_index = index * world_size + rank # interleaved multi-GPU shard
if prompt_index >= len(prompts):
continue
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(out_path): # resume: skip already-done samples
continue
prompt = prompts[prompt_index]
# ... encode prompt, sample noise, run the teacher ODE (CFG), collect latents ...
save_file( # safetensors ONLY (no lmdb / no .pt)
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
out_path, metadata={"prompt": prompt},
)
accelerator.wait_for_everyone()
if accelerator.is_main_process: # rank-0 writes the JSON index
entries = [{"file_path": os.path.join(args.output_folder, f"{i:05d}.safetensors")}
for i in range(len(prompts))
if os.path.exists(os.path.join(args.output_folder, f"{i:05d}.safetensors"))]
json.dump(entries, open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
if __name__ == "__main__":
main()
```
Launcher (`generate_<...>.sh`) — `accelerate launch` uses every visible GPU:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/<family>/*.yaml" \
--caption_path="datasets/prompts.txt" \
--output_folder="datasets/<family>_ode_pairs" \
--num_inference_steps=48 --guidance_scale=6.0 --shift=8.0
```
Training then reads the cache with `ImageVideoSafetensorsDataset(ann_path=".../outputs.json")` (single-file mode `{"file_path": ...}`, or per-tensor mode via `--save_per_tensor`). See reference.md §10.
> Dataset *curation* (scoring/filtering/captioning under `videox_fun/video_caption/`) is a different activity: also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM tensor-parallel) but writes csv/jsonl metadata, not safetensors. See reference.md §10 “Related but different”.
+414
View File
@@ -0,0 +1,414 @@
# VideoX-Fun Integration Reference
Detailed conventions per layer. Read the mirrored family's real files alongside this — the existing code is always the source of truth.
## 1. Model definitions — `videox_fun/models/<family>_*.py`
### File naming
- Transformer / DiT: `<family>_transformer3d.py` (video) or `<family>_transformer2d.py` (image). Variants append a suffix: `_control`, `_s2v`, `_vace`, `_animate`, `_self_forcing`, `_avatar`.
- VAE: `<family>_vae.py` → class `AutoencoderKL<Family>`.
- Encoders: `<family>_text_encoder.py`, `<family>_audio_encoder.py`, `<family>_image_encoder.py`.
### Class shape (mirror `wan_transformer3d.py`)
```python
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
_supports_gradient_checkpointing = True
@register_to_config
def __init__(self, model_type='t2v', patch_size=(1,2,2), in_dim=16, dim=2048,
ffn_dim=8192, num_heads=16, num_layers=32, in_channels=16,
hidden_size=2048, ...):
super().__init__()
...
```
- Keep BOTH civitai names (`in_dim`, `dim`, `ffn_dim`) and diffusers aliases (`in_channels`, `hidden_size`) in `__init__` so either format maps cleanly.
- Implement `_set_gradient_checkpointing(self, *args, **kwargs)`.
- Attention must go through `videox_fun.models.attention_utils.attention` (backend-agnostic), not a hand-rolled `scaled_dot_product_attention`.
- Multi-GPU: expose `enable_multi_gpus_inference()` and route attention through the family's `dist/<family>_xfuser.py` processor.
- Speedups live on the model: `enable_teacache(...)`, `enable_cfg_skip(...)`, `enable_riflex(...)`.
### `from_pretrained` internals (do not simplify)
The custom classmethod must keep these behaviors (see `wan_transformer3d.py::from_pretrained`):
1. Accept `transformer_additional_kwargs`, `subfolder`, `low_cpu_mem_usage`, `torch_dtype`.
2. Read `config.json`; auto-convert foreign configs (e.g. diffsynth `has_image_input`) via a `_convert_from_*_config` helper.
3. Apply `dict_mapping`: pop it from kwargs, then for each `key: target` set `kwargs[target] = config[key]`.
4. Under `low_cpu_mem_usage`, build with `accelerate.init_empty_weights()`, load `.bin`/`.safetensors` (single file or glob all shards), and **filter by exact shape match** before loading.
5. Initialize missing keys deliberately: zero-init control/audio projections (`after_proj`, `before_proj`, `processor.k_proj/v_proj`, `audio_injector`, `cond_encoder`, ...), ones for norms, xavier for ≥2D weights, so new branches start as no-ops.
### Registry — `videox_fun/models/__init__.py`
Add an import line for every new public class, grouped with the family. Wrap optional-dependency imports in `try/except` with a helpful upgrade message (see the Qwen2.5-VL / Mistral3 blocks at the top).
## 2. Pipelines — `videox_fun/pipeline/pipeline_<family>*.py`
Mirror `pipeline_wan.py`. Required pieces:
- Module-level `retrieve_timesteps(scheduler, num_inference_steps, device, timesteps, sigmas, **kwargs)` (copied from diffusers) — reuse verbatim.
- `EXAMPLE_DOC_STRING` for the `@replace_example_docstring` decorator.
- Output dataclass:
```python
@dataclass
class <Family>PipelineOutput(BaseOutput):
videos: torch.Tensor
```
- Pipeline class:
```python
class <Family>Pipeline(DiffusionPipeline):
_optional_component = [...]
model_cpu_offload_seq = "text_encoder->transformer->vae" # order matters for offload
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler, ...): ...
def encode_prompt(...): ...
def prepare_latents(...): ...
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(self, prompt, negative_prompt=..., height=..., width=...,
num_frames=..., num_inference_steps=..., guidance_scale=...,
generator=None, ..., return_dict=True) -> Union[<Family>PipelineOutput, Tuple]: ...
```
- Import schedulers from `..utils.fm_solvers` / `..utils.fm_solvers_unipc`, models from `..models`.
- Separate pipelines per task: base (`pipeline_<family>.py`), inpaint/i2v (`_inpaint`), control (`_control`), s2v, etc. Register all in `videox_fun/pipeline/__init__.py`, adding convenience aliases (e.g. `WanI2VPipeline = WanFunInpaintPipeline`) where existing code expects them.
## 3. Config — `config/<family>/<name>.yaml` (optional)
**The YAML is not mandatory.** Decide by checkpoint layout:
- **Required** for civitai-format / custom single-file layouts, where weights and key names are not diffusers-native. The YAML supplies `transformer_additional_kwargs` (incl. `dict_mapping` mapping civitai config keys → model `__init__` kwargs), component `*_subpath`s, and `vae/text_encoder/scheduler/image_encoder` kwargs.
- **Optional** for a standard diffusers-layout checkpoint (`model_index.json` + each subfolder carrying its own `config.json`). Load components directly: `<Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", low_cpu_mem_usage=True, torch_dtype=...)`, `AutoencoderKL<Family>.from_pretrained(model_name, subfolder="vae")`, etc. Guard the config path exactly like `examples/minimax_h3_fun/predict_v2v_control.py`:
```python
transformer_load_kwargs = {}
if config_path is not None:
from omegaconf import OmegaConf
config = OmegaConf.load(config_path)
transformer_load_kwargs.update(OmegaConf.to_container(config["transformer_additional_kwargs"], resolve=True))
transformer = <Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", **transformer_load_kwargs, ...)
```
When you do use a YAML, the canonical schema is below (see `config/wan2.1/wan_civitai.yaml`):
```yaml
format: civitai # or diffusers — selects weight-key handling
pipeline: Wan # family label consumed by API/ComfyUI loaders
transformer_additional_kwargs:
transformer_subpath: ./ # subfolder under model_name holding the DiT
dict_mapping: # civitai config key -> model __init__ kwarg
in_dim: in_channels
dim: hidden_size
vae_kwargs:
vae_subpath: Wan2.1_VAE.pth
temporal_compression_ratio: 4
spatial_compression_ratio: 8
text_encoder_kwargs:
text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
tokenizer_subpath: google/umt5-xxl
text_length: 512
...
scheduler_kwargs:
scheduler_subpath: null
num_train_timesteps: 1000
shift: 5.0
...
image_encoder_kwargs: # only for i2v / models with a CLIP image encoder
image_encoder_subpath: models_clip_...pth
```
Every `*_subpath` is joined onto `model_name` in scripts. Load with `OmegaConf.load` and pass `OmegaConf.to_container(config['<section>'])` into `from_pretrained`. Use `filter_kwargs(Cls, OmegaConf.to_container(config['scheduler_kwargs']))` to build schedulers.
## 4. Inference scripts — `examples/<family>/predict_<task>.py`
Anatomy, top to bottom (see `examples/wan2.1_fun/predict_t2v.py`):
1. **`sys.path` bootstrap** (before importing `videox_fun`):
```python
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
```
2. **User config block** as top-level variables with explanatory comments, in the conventional order: `GPU_memory_mode`, `ulysses_degree`/`ring_degree`, `fsdp_dit`/`fsdp_text_encoder`, `compile_dit`, TeaCache (`enable_teacache`, `teacache_threshold`, `num_skip_start_steps`, `teacache_offload`), `cfg_skip_ratio`, Riflex (`enable_riflex`, `riflex_k`), `config_path`, `model_name`, `sampler_name`, `shift`, `transformer_path`/`vae_path`/`lora_path`, `sample_size`, `video_length`, `fps`, `weight_dtype`, `prompt`/`negative_prompt`, `guidance_scale`, `seed`, `num_inference_steps`, `lora_weight`, `save_path`.
3. **Device + config**: `device = set_multi_gpus_devices(ulysses_degree, ring_degree)`; then either `config = OmegaConf.load(config_path)` (civitai/custom layout) **or** guard `if config_path is not None:` and load components directly from a diffusers-layout checkpoint (see §3).
4. **Component loading**: transformer (`from_pretrained(..., transformer_additional_kwargs=...)`), optional `transformer_path`/`vae_path` override with `load_state_dict(strict=False)` + missing/unexpected key print, vae, tokenizer, text_encoder, and clip image encoder gated by `transformer.config.in_channels != vae.config.latent_channels`.
5. **Scheduler selection dict**: `{"Flow": FlowMatchEulerDiscreteScheduler, "Flow_Unipc": FlowUniPCMultistepScheduler, "Flow_DPM++": FlowDPMSolverMultistepScheduler}[sampler_name]`; build with `filter_kwargs`.
6. **Pipeline construction**: choose base vs inpaint/i2v/control pipeline by the model's channel condition.
7. **Multi-GPU / FSDP / compile**: if `ulysses_degree>1 or ring_degree>1` call `transformer.enable_multi_gpus_inference()` and optionally `shard_model`; if `compile_dit`, `torch.compile` each `transformer.blocks[i]`.
8. **`GPU_memory_mode` branching** — keep this exact order:
```python
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
```
9. **TeaCache / cfg_skip / Riflex** enablement, `generator = torch.Generator(device).manual_seed(seed)`, LoRA `merge_lora`.
10. **Inference** under `torch.no_grad()`; align `video_length` to `vae.config.temporal_compression_ratio`; pass `video`/`mask_video` for i2v via `get_image_to_video_latent`.
11. **`save_results()`**: `save_videos_grid(sample, path, fps=fps)` for video, PIL save for a single frame; only rank 0 saves when multi-GPU. LoRA `unmerge_lora` after.
Other entry points to mirror when needed: `app.py` (Gradio), `launch_api.py` (API server backed by `videox_fun/api`), `post_infer*.py` (batch/queue inference).
## 5. Training scripts — `scripts/<family>/train*.py`
Mirror `scripts/wan2.1_fun/train.py`. Structure:
1. Diffusers-derived license header + `"""Modified from ..."""` note.
2. Third-party imports, then the **same `sys.path` bootstrap**, then `from videox_fun.data/models/pipeline/utils import ...`.
3. Helper funcs: `filter_kwargs`, `resize_mask`, `linear_decay`, `generate_timestep_with_lognorm`.
4. **`log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step)`** — builds the **inference Pipeline** from the live (unwrapped) transformer and runs it to produce sample videos under `output_dir/sample/`. Wrapped in try/except; handles DeepSpeed (`transformer3d.config` swap) and restores VAE/text-encoder placement (`low_vram`). **Reuse the pipeline; never write a separate sampler.**
5. **`parse_args()`** — reuse the shared argument surface: `--config_path`, `--pretrained_model_name_or_path`, `--train_data_dir`, `--train_data_meta`, `--image_sample_size`/`--video_sample_size`/`--token_sample_size`, `--video_sample_n_frames`, `--video_sample_stride`, `--train_batch_size`, `--gradient_accumulation_steps`, `--learning_rate`, `--lr_scheduler`, `--lr_warmup_steps`, `--checkpointing_steps`, `--output_dir`, `--mixed_precision`, `--gradient_checkpointing`, `--enable_bucket`, `--random_hw_adapt`, `--training_with_video_token_length`, `--uniform_sampling`, `--low_vram`, `--train_mode`, `--trainable_modules`, LoRA args (`--use_lora`, `--rank`, ...), `--validation_prompts`/`--validation_paths`. Add new args only when the family genuinely needs them.
6. **`main()`** — Accelerator setup, DeepSpeed/FSDP zero-stage handling (auto-sets `save_state`), model loading via config, dataset + bucket sampler from `videox_fun.data`, trainable-module filtering / LoRA network via `create_network`, optimizer + `get_scheduler`, `accelerator.prepare`, checkpoint save/load hooks, training loop with timestep sampling, loss, `log_validation` at intervals, and final weight/LoRA save.
### Resolution args — `--video_sample_size` (+ `--fix_sample_size`)
Canvas resolution is always driven by a **single square** `--video_sample_size` (`type=int`, height = width) — never by separate `--video_sample_height` / `--video_sample_width`. When a **fixed non-square shape** is required, add `--fix_sample_size` (`nargs=2, type=int, default=None`, `[height, width]`) that overrides the square size; mirror `scripts/wan2.2_fun/train_lora.py`, `scripts/z_image/train_distill.py`. Derive the effective `height` / `width` once in `parse_args()` and reuse them everywhere downstream:
```python
parser.add_argument("--video_sample_size", type=int, default=1280)
parser.add_argument("--fix_sample_size", nargs=2, type=int, default=None,
help="Fix Sample size [height, width] to override `--video_sample_size` with a fixed non-square shape.")
...
if args.fix_sample_size is not None:
args.video_sample_height, args.video_sample_width = args.fix_sample_size
else:
args.video_sample_height = args.video_sample_width = args.video_sample_size
```
In bucket datasets `--fix_sample_size` also forces `random_hw_adapt=False` / `training_with_video_token_length=False` and bumps `video_sample_size = max(max(fix_sample_size), video_sample_size)`; in data-free scripts (e.g. `scripts/minimax_h3/train_pdd_lora.py`) it simply pins the generation canvas. Always validate the size against the patch/VAE constraint (minimax_h3: `% 32`). The `.sh` launcher passes it space-separated (`nargs=2`): `--fix_sample_size 768 1344`.
### Launcher — `scripts/<family>/train*.sh`
`export MODEL_NAME/DATASET_NAME/DATASET_META_NAME`, then `accelerate launch --mixed_precision="bf16" scripts/<family>/train.py --config_path=... <full arg list>`. Include commented I2V/control variants and DeepSpeed/NCCL notes as the existing scripts do.
### Docs — `README_TRAIN.md` + `README_TRAIN_zh-CN.md`
Aligned bilingual pair: identical section order, identical commands and parameter tables; only the prose language differs. Follow the top-level section order used across existing training READMEs.
## 6. Shared infrastructure map (reuse, never reimplement)
| Need | Import from |
|------|-------------|
| Flow/DPM/UniPC schedulers | `diffusers`, `videox_fun.utils.fm_solvers`, `videox_fun.utils.fm_solvers_unipc` |
| LoRA create/merge/unmerge/convert | `videox_fun.utils.lora_utils` |
| FP8 quantization | `videox_fun.utils.fp8_optimization` |
| Group / leaf offload hooks | `videox_fun.utils.group_offload` |
| Multi-GPU device + FSDP shard + seq-parallel attn | `videox_fun.dist` |
| Save video/audio, image→video latents, kwarg filter, dimension calc | `videox_fun.utils.utils` |
| Datasets + bucket/aspect-ratio samplers + masks | `videox_fun.data` |
| TeaCache coefficients | `videox_fun.models.cache_utils` |
## 7. Naming quick reference
| Concept | Convention | Example |
|---------|-----------|---------|
| Model file | `<family>_transformer3d.py` | `wan_transformer3d.py` |
| Model class | `<Family>Transformer3DModel` | `WanTransformer3DModel` |
| VAE class | `AutoencoderKL<Family>` | `AutoencoderKLWan` |
| Pipeline file | `pipeline_<family>.py` | `pipeline_wan.py` |
| Pipeline class | `<Family>Pipeline` | `WanPipeline` / `WanFunInpaintPipeline` |
| Config | `config/<family>/<variant>.yaml` | `config/wan2.1/wan_civitai.yaml` |
| Inference | `examples/<family>/predict_<task>.py` | `predict_t2v.py`, `predict_i2v.py`, `predict_v2v_control.py` |
| Training | `scripts/<family>/train[_<variant>].py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py` |
## 8. Training data pipeline — dataset & sampler selection
Pick the dataset by **task / `train_mode`**, then the sampler by **`enable_bucket`** and dataset type. All datasets/samplers come from `videox_fun.data` — never write a new one.
### Annotation format — the `train_data_meta` file (`metadata.json` / `.csv`)
Every dataset class reads an annotation file (`args.train_data_meta`) that indexes the media under `args.train_data_dir` (`data_root`). `ImageVideoDataset` accepts **`.json`** (a top-level array of records) or **`.csv`** (`csv.DictReader`; the header row is the field names). Each record for ordinary image/video training:
| Field | Required | Meaning |
|-------|----------|---------|
| `file_path` | yes | Media path, resolved **relative to `train_data_dir`** via `os.path.join(data_root, file_path)`. If `data_root is None`, `file_path` is used as-is. |
| `text` | yes | Caption / prompt. Dropped to `""` with probability `text_drop_ratio` (default `0.1`) for classifier-free guidance. |
| `type` | no | `"video"` or `"image"`; **defaults to `"image"`** when the key is absent (`data_info.get('type', 'image')`). |
```json
[
{"file_path": "train/00000000.mp4", "text": "A young woman gently turns her head to the right ...", "type": "video"},
{"file_path": "train/00000001.jpg", "text": "a dog running on the beach", "type": "image"}
]
```
The directory layout matches the index — media in a `train/` subdir, the annotation file beside it. Ready-made examples ship in `datasets/X-Fun-Videos-Demo/` (`train/*.mp4` + `metadata.json`) and `datasets/X-Fun-Images-Demo/`. The equivalent `.csv`:
```csv
file_path,text,type
train/00000000.mp4,"A young woman gently turns her head to the right ...",video
train/00000001.jpg,"a dog running on the beach",image
```
**Variant datasets append extra fields to this same record shape**, each consumed by its own class (see the table below) — e.g. camera-pose adds `action_path` (`LingbotImageVideoDataset`), object/VACE/S2V variants add object fields (`object_file_path` / `objects`). The demo folders also ship several augmented metadata variants (next subsection). Always read the target class's `get_batch` for the exact fields it consumes.
### Ready-made demo datasets — the standard test data (never invent a test set)
Smoke tests, `log_validation` checks, and doc examples all run on the official demo datasets under `datasets/`, downloaded from ModelScope as `PAI/<name>`:
```bash
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
Pick the demo by **task**, matching the dataset class in the table below:
| Demo dataset (`datasets/...`) | Contents | Extra metadata fields | Task it tests | Dataset class |
|-------------------------------|----------|----------------------|---------------|---------------|
| `X-Fun-Videos-Demo` | 16 videos (832×480) in `train/` | — | T2V / I2V base + inpaint, distill | `ImageVideoDataset` |
| `X-Fun-Videos-Controls-Demo` | 16 videos in `train/` + `canny/` + `object/<video_id>/` + `wav/` | `control_file_path`, `object_file_path` (list), `audio_path` | V2V control, VACE, S2V-with-control | `ImageVideoControlDataset`, `VideoSpeechControlDataset` |
| `X-Fun-Videos-Audios-Demo` | 17 video/audio pairs: `train/` (1280×720) + `wav/` (16 kHz mono) + `pose/` | `audio_path`, `control_file_path` | Speech-driven S2V / avatar / talking-head | `VideoSpeechDataset` |
| `X-Fun-Images-Demo` | 19 images in `train/` | — | T2I full fine-tune + LoRA (z_image / flux2 / qwenimage / lens / ernie) | `ImageVideoDataset` |
| `X-Fun-Images-Controls-Demo` | 19 images in `train/` + `canny/` | `control_file_path` | Image control / ControlNet / i2i inpaint | `ImageVideoControlDataset` |
| `X-Fun-Images-Edit-Demo` | 21 records: `source/souce-<id>/` (multi-source supported) → `train/` | `source_file_path` (**list**) | Image edit (Qwen-Image-Edit family) | `ImageEditDataset` |
| `X-Fun-Videos-Lingbot-Demo` | video + `intrinsics.npy` / `poses.npy` | camera pose / action | Camera-pose world model (`lingbot_world`) | `LingbotImageVideoDataset` |
**Which metadata file to point `--train_data_meta` at** (each demo ships several variants beside the media):
| Metadata file | Use when |
|---------------|----------|
| `metadata.json` | Base format only (`file_path` / `text` / `type`) — fine for a minimal check |
| `metadata_add_width_height.json` | **Default choice.** Adds `width` / `height` so bucketing doesn't decode media (matters on slow storage such as OSS). Used by non-VACE control / S2V training too |
| `metadata_add_width_height_add_objects.json` | VACE / subject-reference training (`object_file_path` list → `object/<video_id>/`; shuffled at train time) |
| `metadata_add_width_height_add_wav.json` | Audio-visual joint models (e.g. `minimax_h3_fun` control training): `audio_path` → `wav/`. Keep the `.sh` launcher and the README on the same file |
| `metadata_lingbot_video_add_width_height.json` | `lingbot_video` — `text` is already a structured JSON caption (lives in `X-Fun-Videos-Demo`) |
| `metadata_origin.json` | Pre-processing original kept for reference; not used for training |
Regenerate the width/height variant with the shipped helper when adding your own media:
`python scripts/process_json_add_width_and_height.py --input_file datasets/<Demo>/metadata.json --output_file datasets/<Demo>/metadata_add_width_height.json`.
`audio_path` optionality differs per class (`videox_fun/data/dataset_video.py`): `VideoSpeechDataset` reads `video_dict['audio_path']` directly, so it is **required**; `VideoSpeechControlDataset` uses `.get('audio_path')` and **falls back to the video file's own audio track** when the field is absent.
### Dataset by task (all take `train_data_meta, train_data_dir, ...`)
| Task / mode | Dataset class | Used by | Key kwargs |
|-------------|--------------|---------|-----------|
| T2V / I2V base (`normal` + inpaint) | `ImageVideoDataset` | `train.py`, `train_lora.py`, t2i `train.py` | `enable_inpaint = train_mode != "normal"`, `video_sample_size/stride/n_frames`, `image_sample_size`, `video_repeat` |
| Image T2I (qwenimage/flux/z_image) | `ImageVideoDataset` | `scripts/<img>/train.py` | `image_sample_size` |
| Control (canny/pose/depth/camera) | `ImageVideoControlDataset` | `train_control*.py`, `train_control_distill.py` | `enable_camera_info = train_mode == "control_camera_ref"` |
| Image Edit (source→target) | `ImageEditDataset` | `qwenimage/train_edit*.py` | `image_sample_size` |
| Speech/audio-driven (S2V, avatar, talking) | `VideoSpeechDataset` | `mova`, `ltx2`, `minimax_h3`, `fantasytalking`, `infinitetalk`, `flashhead`, `longcatvideo/train_avatar*` | audio + video fields |
| S2V **with control** | `VideoSpeechControlDataset` | `wan2.2/train_s2v*.py`, `minimax_h3_fun/train_control*` | audio + control |
| Motion/pose animate | `VideoAnimateDataset` | `wan2.2/train_animate*.py` | motion/pose driven |
| Distill text-only branch, GRPO, DPO | `TextDataset` | `train_distill*.py` (text branch), `z_image/train_grpo_lora.py`, `train_dpo_lora.py` | reads only the `text` field; `text_drop_ratio` |
| Precomputed latents (ODE pairs) | `ImageVideoSafetensorsDataset` | `wan2.1_self_forcing/train_ode.py` | `data_root` |
| Camera-pose conditioning | `LingbotImageVideoDataset` | `lingbot_world/train.py` | `intrinsics.npy` / `poses.npy` |
| Video-only (VAE/TAEHV distill) | `VideoDataset` | `taehv/train_taehv.py` | `sample_size/stride/n_frames`, `enable_inpaint=False` |
### Sampler by condition
| Condition | Sampler | Shape |
|-----------|---------|-------|
| `enable_bucket=True` (default; image+video) | `AspectRatioBatchImageVideoSampler` | `sampler=RandomSampler(ds, generator=g), dataset=train_dataset.dataset, batch_size, train_folder=args.train_data_dir, drop_last=True, aspect_ratios=aspect_ratio_sample_size` |
| `enable_bucket=False` | `ImageVideoSampler` | `ImageVideoSampler(RandomSampler(ds, generator=g), train_dataset, batch_size)` |
| `TextDataset` (distill text branch / GRPO / DPO) | `BatchSampler` (plain) | `BatchSampler(RandomSampler(ds, generator=g), batch_size, drop_last=True)`; GRPO adds `k_repeat=args.num_image_per_prompt` |
| video-only bucket (available, not used by current scripts) | `AspectRatioBatchSampler` | — |
| image-only bucket (available, not used by current scripts) | `AspectRatioBatchImageSampler` | — |
`aspect_ratio_sample_size` is built from `ASPECT_RATIO_512` scaled by `args.video_sample_size`; `get_closest_ratio` picks the bucket inside `collate_fn`.
### Universal DataLoader creation pattern
```python
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
if args.enable_bucket:
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
batch_sampler = AspectRatioBatchImageVideoSampler(
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
aspect_ratios=aspect_ratio_sample_size)
def collate_fn(examples):
new_examples = {"pixel_values": [], "text": []}
if args.train_mode != "normal": # inpaint/i2v adds mask fields
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
# bucket via get_closest_ratio -> transform (Resize/CenterCrop/Normalize) -> stack
# masked branch uses get_random_mask(...)
return new_examples
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
else:
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler,
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
```
`collate_fn` receives the `examples` **list** (not a `batch` dict); build every batch-level field (`text`, `pixel_values`, masks) explicitly from `examples` into `new_examples`. When `--enable_text_encoder_in_dataloader`, encode prompts inside `collate_fn` and emit `encoder_hidden_states` / `encoder_attention_mask`.
## 9. Inference task matrix — predict script → pipeline → inputs
Pick the pipeline by **task**; the `predict_<task>.py` name and its inputs follow the same convention across families.
| Task | `predict_<task>.py` | Pipeline (family example) | Extra `__call__` inputs | Input helper |
|------|--------------------|---------------------------|-------------------------|--------------|
| Text→Video | `predict_t2v.py` | `WanPipeline`, `Wan2_2Pipeline`, `CogVideoXFunPipeline`, `LongCatVideoPipeline`, `LTX2Pipeline` | `prompt` only | — |
| Image→Video | `predict_i2v.py` | `WanI2VPipeline`(=`WanFunInpaintPipeline`), `Wan2_2FunInpaintPipeline`, `Wan2_2I2VPipeline`, `HunyuanVideoI2VPipeline` | `video`, `mask_video` | `get_image_to_video_latent(start_image, end_image, video_length, sample_size)` |
| Text+Image→Video (5B) | `predict_ti2v.py` | `Wan2_2TI2VPipeline` | `prompt` (+ optional image) | `get_image_to_video_latent` |
| Video→Video Control | `predict_v2v_control.py` | `WanFunControlPipeline`, `Wan2_2FunControlPipeline` | `control_video` | `get_video_to_video_latent(control_video, ...)` |
| Control + reference | `predict_v2v_control_ref.py` | `WanFunControlPipeline` | `control_video` + `ref_image` | `get_video_to_video_latent` + `get_image_latent` |
| Control + camera | `predict_v2v_control_camera.py` | `WanFunControlPipeline` | `control_video` + camera pose | — |
| VACE (control/mask/i2v/s2v) | `predict_v2v_control.py`, `predict_v2v_mask.py`, `predict_s2v.py`, `predict_i2v.py` | `WanVacePipeline`, `Wan2_2VaceFunPipeline` | control/mask/ref | — |
| Speech→Video (audio) | `predict_s2v.py` | `Wan2_2S2VPipeline`, `MiniMaxH3Pipeline`, `InfiniteTalkPipeline`, `FantasyTalkingPipeline`, `FlashHeadPipeline`, `MOVAPipeline`, `LongCatVideoAvatarPipeline` | `audio` + reference image | — |
| Animate (motion/pose) | `predict_animate.py` | `Wan2_2AnimatePipeline` | motion/pose video + ref | — |
| Subject reference | `predict_s2v.py` (phantom) | `WanFunPhantomPipeline` | reference images | — |
| Text→Image | `predict_t2i.py` | `QwenImagePipeline`, `Flux2Pipeline`, `ZImagePipeline`, `LensPipeline`, `ErnieImagePipeline` | `prompt` | — |
| Image Control (t2i) | `predict_t2i_control.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline`, `QwenImageControlNetPipeline` | `control_image` | — |
| Inpaint (i2i) | `predict_i2i_inpaint.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline` | `image` + `mask` | — |
| Image Edit | `predict_t2i_edit.py`, `predict_t2i_edit_plus.py` | `QwenImageEditPipeline`, `QwenImageEditPlusPipeline` | source image + instruction | — |
| Layered edit | `predict_i2i_layered.py` | `QwenImageLayeredPipeline` | image | — |
| Camera-pose world | `predict_i2v.py` (lingbot_world) | `Wan2_2I2VPipeline`, `WanFunLingbotWorldFastPipeline` | image + camera pose | — |
| Latent upsample | `predict_i2v_upsample.py` | `LTX2LatentUpsamplePipeline`, `WanLatentUpsamplePipeline` | low-res latent/video | — |
| AR / streaming distill | `predict_t2v_stream.py` | `WanSelfForcingPipeline` | prompt (streamed) | — |
### Predict-script variant suffixes (same task, different backend/model)
| Suffix | Meaning |
|--------|---------|
| `_tae` | Fast decode via `AutoencoderTinyWan` (TAEHV) instead of the full VAE |
| `_2.2vae` | Uses the Wan2.2 VAE (`AutoencoderKLWan3_8`) |
| `_5b` | 5B-parameter model variant |
| `turbo` / distill | Distilled model, few-step inference (e.g. `predict_turbo_*.py`) |
| `_refine` | Two-stage refine pass |
| `_ref` / `_camera` | Adds reference-image / camera conditioning |
All variants keep the identical config block, `GPU_memory_mode` branching, and `save_results()` from Section 4 — only the loaded VAE/transformer and pipeline class change.
## 10. Preprocessing — offline training-data generation (multi-GPU + safetensors)
Here "preprocessing" means **generating/caching training data offline** with the teacher / VAE / text-encoder — latents, ODE-trajectory pairs, prompt/text embeddings — so training just reads cached tensors instead of re-encoding every step. Canonical example: `scripts/wan2.1_self_forcing/generate_ode_pairs.py` (+ `generate_ode_pairs.sh`); the loader-side contract is `ImageVideoSafetensorsDataset` in `videox_fun/data/dataset_image_video.py`. Two rules are non-negotiable.
### Rule 1 — multi-GPU is mandatory
Never a single-GPU / hardcoded `cuda:0` loop. Launch with `accelerate launch` and shard work across ranks by interleaving:
```python
from accelerate import Accelerator
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
torch.set_grad_enabled(False) # inference-only
total_per_rank = math.ceil(len(prompts) / world_size)
for index in tqdm(range(total_per_rank), disable=rank != 0):
prompt_index = index * world_size + rank # interleaved shard
if prompt_index >= len(prompts):
continue
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(out_path): # resume-friendly
continue
... # encode prompt / run teacher ODE / collect latents
accelerator.wait_for_everyone()
if accelerator.is_main_process: # write the JSON index once, on rank 0
json.dump([{"file_path": p} for p in all_safetensor_paths],
open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
```
Launcher (`.sh`): `accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py --pretrained_model_name_or_path=... --config_path=config/<family>/*.yaml --output_folder=datasets/<...> ...`. Reuse `videox_fun.models` + config-driven `from_pretrained` (Section 3) and `videox_fun.utils.utils.save_videos_grid` for sample previews — do not write a new loader.
### Rule 2 — store as safetensors; do NOT use LMDB or `.pt`
Save every cached tensor with `safetensors.torch.save_file`, one `.safetensors` per sample (or per tensor), plus a JSON index of `{"file_path": ...}` entries:
```python
from safetensors.torch import save_file
save_file(
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
out_path, # f"{prompt_index:05d}.safetensors"
metadata={"prompt": prompt},
)
```
`ImageVideoSafetensorsDataset(ann_path, data_root=None)` reads that JSON and supports two layouts:
- **Single-file (default)**: `{"file_path": "scene.safetensors"}` — whole state dict in one archive.
- **Per-tensor (`--save_per_tensor`)**: `{"file_path": "scene_dir", "latents": ".../latents.safetensors", "prompt_embeds": ".../prompt_embeds.safetensors"}` — each key loaded and merged.
**Do not** cache preprocessed data in **LMDB** or as **`.pt`/`.pth` `torch.save` pickles**. safetensors is the repo-wide standard (also used for LoRA/weight saving), is pickle-free/safe, memory-maps fast, and is exactly what `ImageVideoSafetensorsDataset` loads. (Scope: this governs cached *data tensors*; accelerate optimizer/scheduler/scaler `.pt` states written during training checkpoints are a separate mechanism and unaffected.)
### Related but different — dataset curation
Scoring / filtering / captioning under `videox_fun/video_caption/` (`compute_*.py`, `internvl2_video_recaptioning.py`) is dataset *curation*, not latent caching. It is also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM `tensor_parallel_size=device_count()`), but writes csv/jsonl **metadata** (not tensors), so Rule 2 does not apply there.
+28 -2
View File
@@ -11,13 +11,14 @@ Wan-Fun:
English | [简体中文](./README_zh-CN.md) | [日本語](./README_ja-JP.md)
# Table of Contents
- [Table of Contents](#table-of-contents)
- [Introduction](#introduction)
- [Quick Start](#quick-start)
- [Video Result](#video-result)
- [How to use](#how-to-use)
- [How to Use](#how-to-use)
- [Model zoo](#model-zoo)
- [Reference](#reference)
- [Citation](#citation)
- [Limitations and Risks](#limitations-and-risks)
- [License](#license)
# Introduction
@@ -699,6 +700,31 @@ V1.1:
- ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper
- CameraCtrl: https://github.com/hehao13/CameraCtrl
# Citation
If you use VideoX-Fun in your research or project, please cite it as follows:
```bibtex
@misc{aigc_apps_VideoX_Fun_2026,
author = {aigc-apps},
title = {VideoX-Fun: A Video Generation Pipeline for Diffusion Transformer},
year = {2026},
publisher = {GitHub},
url = {https://github.com/aigc-apps/VideoX-Fun}
}
```
# Limitations and Risks
- Generated videos may have artifacts or quality issues, especially in complex scenes.
- The model may struggle with fine details, text rendering, or specific artistic styles.
- Performance varies with input prompt quality, resolution, and other parameters.
- The technology could be misused to create misleading content (e.g., deepfakes). Users are responsible for ethical use.
- The model may reflect biases present in the training data.
- Users should respect privacy and copyright when using real people's images or videos.
We encourage responsible use and recommend implementing safeguards in production environments.
# License
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
+27 -1
View File
@@ -11,13 +11,14 @@ Wan-Fun:
[English](./README.md) | [简体中文](./README_zh-CN.md) | 日本語
# 目次
- [目次](#目次)
- [紹介](#紹介)
- [クイックスタート](#クイックスタート)
- [ビデオ結果](#ビデオ結果)
- [使用方法](#使用方法)
- [モデルの場所](#モデルの場所)
- [参考文献](#参考文献)
- [引用](#引用)
- [制限とリスク](#制限とリスク)
- [ライセンス](#ライセンス)
# 紹介
@@ -699,6 +700,31 @@ V1.1:
- ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper
- CameraCtrl: https://github.com/hehao13/CameraCtrl
# 引用
研究やプロジェクトでVideoX-Funを使用する場合は、以下の形式で引用してください:
```bibtex
@misc{aigc_apps_VideoX_Fun_2026,
author = {aigc-apps},
title = {VideoX-Fun: A Video Generation Pipeline for Diffusion Transformer},
year = {2026},
publisher = {GitHub},
url = {https://github.com/aigc-apps/VideoX-Fun}
}
```
# 制限とリスク
- 生成された動画には、特に複雑なシーンでアーティファクトや品質の問題がある場合があります。
- モデルは、細かい詳細、テキストのレンダリング、または特定の芸術スタイルで苦労する場合があります。
- パフォーマンスは、入力プロンプトの品質、解像度、その他のパラメータによって異なります。
- この技術は、誤解を招くコンテンツ(例:ディープフェイク)を作成するために悪用される可能性があります。ユーザーは倫理的な使用に責任を持ちます。
- モデルは、トレーニングデータに存在するバイアスを反映する可能性があります。
- ユーザーは、実在の人物の画像や動画を使用する際、プライバシーと著作権を尊重する必要があります。
責任ある使用を推奨し、本番環境でのセーフガードの実装をお勧めします。
# ライセンス
このプロジェクトは[Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE)の下でライセンスされています。
+27 -1
View File
@@ -11,13 +11,14 @@ Wan-Fun:
[English](./README.md) | 简体中文 | [日本語](./README_ja-JP.md)
# 目录
- [目录](#目录)
- [简介](#简介)
- [快速启动](#快速启动)
- [视频作品](#视频作品)
- [如何使用](#如何使用)
- [模型地址](#模型地址)
- [参考文献](#参考文献)
- [引用](#引用)
- [限制与风险](#限制与风险)
- [许可证](#许可证)
# 简介
@@ -689,6 +690,31 @@ V1.1:
- ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper
- CameraCtrl: https://github.com/hehao13/CameraCtrl
# 引用
如果您在研究或项目中使用了 VideoX-Fun,请按以下格式引用:
```bibtex
@misc{aigc_apps_VideoX_Fun_2026,
author = {aigc-apps},
title = {VideoX-Fun: A Video Generation Pipeline for Diffusion Transformer},
year = {2026},
publisher = {GitHub},
url = {https://github.com/aigc-apps/VideoX-Fun}
}
```
# 限制与风险
- 生成的视频可能存在伪影或质量问题,尤其在复杂场景中。
- 模型在处理精细细节、文字渲染或特定艺术风格时可能有困难。
- 性能因输入提示词质量、分辨率等参数而异。
- 该技术可能被滥用于创建误导性内容(如深度伪造)。用户需对道德使用负责。
- 模型可能反映训练数据中存在的偏见。
- 用户在使用真人图片或视频时应尊重隐私和版权。
我们鼓励负责任地使用该技术,并建议在生产环境中实施安全措施。
# 许可证
本项目采用 [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 349 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.

After

Width:  |  Height:  |  Size: 150 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.
+1 -2
View File
@@ -15,8 +15,7 @@ from diffusers import EulerDiscreteScheduler
from einops import rearrange
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKLCogVideoX,
CogVideoXTransformer3DModel, T5EncoderModel,
T5Tokenizer)
+3 -4
View File
@@ -26,8 +26,7 @@ else:
from diffusers.models.modeling_utils import \
load_model_dict_into_meta
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKLFlux2,
Flux2ControlTransformer2DModel,
Flux2Transformer2DModel,
@@ -598,10 +597,10 @@ class LoadFlux2TextEncoderModel:
[os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check
try:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="flux2_tokenizer")
except:
except Exception:
try:
tokenizer_path = os.path.join(search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="FLUX.2-dev"), "tokenizer")
except:
except Exception:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Mistral-Nemo-Instruct-2407")
tokenizer = PixtralProcessor.from_pretrained(tokenizer_path)
+5 -6
View File
@@ -25,8 +25,7 @@ else:
from diffusers.models.modeling_utils import \
load_model_dict_into_meta
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKLQwenImage, Qwen2_5_VLConfig,
Qwen2_5_VLForConditionalGeneration,
Qwen2Tokenizer, Qwen2VLProcessor,
@@ -476,10 +475,10 @@ class LoadQwenImageTextEncoderModel:
[os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check
try:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="qwen2_tokenizer")
except:
except Exception:
try:
tokenizer_path = os.path.join(search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Qwen-Image"), "tokenizer")
except:
except Exception:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Qwen2.5-VL-7B-Instruct")
tokenizer = Qwen2Tokenizer.from_pretrained(tokenizer_path)
@@ -504,10 +503,10 @@ class LoadQwenImageProcessor:
[os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check
try:
processor_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="qwen2_processor")
except:
except Exception:
try:
processor_path = os.path.join(search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Qwen-Image-Edit"), "processor")
except:
except Exception:
processor_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Qwen2.5-VL-7B-Instruct")
# Get processor
+1 -2
View File
@@ -17,8 +17,7 @@ from einops import rearrange
from omegaconf import OmegaConf
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel, WanT5EncoderModel,
WanTransformer3DModel)
+2 -3
View File
@@ -16,9 +16,8 @@ from einops import rearrange
from omegaconf import OmegaConf
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data.dataset_image_video import process_pose_params
from ...videox_fun.data import (ASPECT_RATIO_512, get_closest_ratio,
process_pose_params)
from ...videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
from ...videox_fun.models.cache_utils import get_teacache_coefficients
+1 -2
View File
@@ -17,8 +17,7 @@ from einops import rearrange
from omegaconf import OmegaConf
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
+2 -3
View File
@@ -16,9 +16,8 @@ from einops import rearrange
from omegaconf import OmegaConf
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data.dataset_image_video import process_pose_params
from ...videox_fun.data import (ASPECT_RATIO_512, get_closest_ratio,
process_pose_params)
from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
+2 -3
View File
@@ -17,9 +17,8 @@ from einops import rearrange
from omegaconf import OmegaConf
from PIL import Image
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data.dataset_image_video import process_pose_params
from ...videox_fun.data import (ASPECT_RATIO_512, get_closest_ratio,
process_pose_params)
from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
VaceWanTransformer3DModel, WanT5EncoderModel)
+5 -5
View File
@@ -18,8 +18,7 @@ from omegaconf import OmegaConf
from PIL import Image
from safetensors.torch import load_file
from ...videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
get_closest_ratio)
from ...videox_fun.data import ASPECT_RATIO_512, get_closest_ratio
from ...videox_fun.models import (AutoencoderKL, AutoTokenizer,
Qwen2VLProcessor, Qwen3Config,
Qwen3ForCausalLM,
@@ -435,10 +434,10 @@ class LoadZImageTextEncoderModel:
[os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "models/Diffusion_Transformer")] # Possible folder names to check
try:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="qwen3_tokenizer")
except:
except Exception:
try:
tokenizer_path = os.path.join(search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Z-Image-Turbo"), "tokenizer")
except:
except Exception:
tokenizer_path = search_sub_dir_in_possible_folders(possible_folders, sub_dir_name="Qwen3-4B")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
@@ -504,8 +503,8 @@ class CombineZImagePipeline:
)
pipeline.remove_all_hooks()
safe_remove_group_offloading(pipeline)
undo_convert_weight_dtype_wrapper(transformer)
pipeline.to(device=offload_device)
transformer = transformer.to(weight_dtype)
if GPU_memory_mode == "sequential_cpu_offload":
@@ -747,6 +746,7 @@ class LoadZImageControlNetInPipeline:
# Remove hooks
funmodels["pipeline"].remove_all_hooks()
safe_remove_group_offloading(funmodels["pipeline"])
# Load config
config_path = f"{script_directory}/config/{config}"
@@ -0,0 +1,614 @@
{
"id": "dcf2fcac-6293-4a86-b30b-f63e420177f2",
"revision": 0,
"last_node_id": 102,
"last_link_id": 96,
"nodes": [
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": [
210,
88
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 91,
"type": "LoadZImageTextEncoderModel",
"pos": [
283.53765869140625,
-280.6837463378906
],
"size": [
407.4130859375,
102
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder",
"type": "TextEncoderModel",
"links": [
80
]
},
{
"name": "tokenizer",
"type": "Tokenizer",
"links": [
81
]
}
],
"properties": {
"Node name for S&R": "LoadZImageTextEncoderModel"
},
"widgets_values": [
"qwen_3_4b.safetensors",
"bf16"
]
},
{
"id": 75,
"type": "FunTextBox",
"pos": [
250,
-50
],
"size": [
383.54010009765625,
156.71620178222656
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"slot_index": 0,
"links": [
88
]
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "FunTextBox"
},
"widgets_values": [
"A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\n\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\n\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock.\n"
]
},
{
"id": 73,
"type": "FunTextBox",
"pos": [
250,
160
],
"size": [
383.7149963378906,
183.83506774902344
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"slot_index": 0,
"links": [
89
]
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "FunTextBox"
},
"widgets_values": [
""
]
},
{
"id": 97,
"type": "Note",
"pos": [
-354.4680507215508,
-433.6570714778354
],
"size": [
598.1623727144193,
233.48501180428053
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"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].\nmodel_full_load means that the entire model will be moved to the GPU.\n\nmodel_full_load_and_qfloat8 means that the entire model will be moved to the GPU,\nand the transformer model has been quantized to float8, which can save more GPU memory. \n\nmodel_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.\n\nmodel_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, \nand the transformer model has been quantized to float8, which can save more GPU memory. \n\nsequential_cpu_offload means that each layer of the model will be moved to the CPU after use, \nresulting in slower speeds but saving a large amount of GPU memory."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 93,
"type": "LoadZImageVAEModel",
"pos": [
1168.1599508804559,
-314.8650476692169
],
"size": [
377.8583984375,
84.69844055175781
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "VAEModel",
"links": [
79
]
}
],
"properties": {
"Node name for S&R": "LoadZImageVAEModel"
},
"widgets_values": [
"ae.safetensors",
"bf16"
]
},
{
"id": 96,
"type": "CombineZImagePipeline",
"pos": [
790.3572998046875,
-328.7134094238281
],
"size": [
342.5804748535156,
162
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "transformer",
"type": "TransformerModel",
"link": 96
},
{
"name": "vae",
"type": "VAEModel",
"link": 79
},
{
"name": "text_encoder",
"type": "TextEncoderModel",
"link": 80
},
{
"name": "tokenizer",
"type": "Tokenizer",
"link": 81
},
{
"name": "processor",
"shape": 7,
"type": "Processor",
"link": null
},
{
"name": "model_name",
"type": "STRING",
"widget": {
"name": "model_name"
},
"link": 82
}
],
"outputs": [
{
"name": "funmodels",
"type": "FunModels",
"links": [
87
]
}
],
"properties": {
"Node name for S&R": "CombineZImagePipeline"
},
"widgets_values": [
"",
"sequential_cpu_offload"
]
},
{
"id": 88,
"type": "PreviewImage",
"pos": [
1049.1402001998824,
-52.699751832945104
],
"size": [
366.56134033203125,
415.4429626464844
],
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 90
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 100,
"type": "LoadImage",
"pos": [
396.48551767952716,
419.158511044569
],
"size": [
270,
314.00000000000006
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
86
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"z-image-turbo_00004_.png",
"image"
]
},
{
"id": 92,
"type": "LoadZImageTransformerModel",
"pos": [
275.9798278808594,
-465.2391052246094
],
"size": [
416.3677673339844,
106.13789367675781
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "transformer",
"type": "TransformerModel",
"links": [
95
]
},
{
"name": "model_name",
"type": "STRING",
"links": [
82
]
}
],
"properties": {
"Node name for S&R": "LoadZImageTransformerModel"
},
"widgets_values": [
"z_image_bf16.safetensors",
"bf16"
]
},
{
"id": 102,
"type": "LoadZImageControlNetInModel",
"pos": [
779.793189390101,
-457.3825558553134
],
"size": [
589.8698159570357,
82
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "transformer",
"type": "TransformerModel",
"link": 95
}
],
"outputs": [
{
"name": "transformer",
"type": "TransformerModel",
"links": [
96
]
}
],
"properties": {
"Node name for S&R": "LoadZImageControlNetInModel"
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Fun-Controlnet-Tile-2.1.safetensors"
]
},
{
"id": 99,
"type": "ZImageControlSampler",
"pos": [
727.2482831521481,
-52.41010988674983
],
"size": [
270,
350
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "funmodels",
"type": "FunModels",
"link": 87
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 88
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 89
},
{
"name": "control_image",
"shape": 7,
"type": "IMAGE",
"link": 86
},
{
"name": "inpaint_image",
"shape": 7,
"type": "IMAGE",
"link": null
},
{
"name": "mask_image",
"shape": 7,
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
90
]
}
],
"properties": {
"Node name for S&R": "ZImageControlSampler"
},
"widgets_values": [
1824,
2416,
43,
"fixed",
25,
4.5,
"Flow",
3,
0.85
]
}
],
"links": [
[
79,
93,
0,
96,
1,
"VAEModel"
],
[
80,
91,
0,
96,
2,
"TextEncoderModel"
],
[
81,
91,
1,
96,
3,
"Tokenizer"
],
[
82,
92,
1,
96,
5,
"STRING"
],
[
86,
100,
0,
99,
3,
"IMAGE"
],
[
87,
96,
0,
99,
0,
"FunModels"
],
[
88,
75,
0,
99,
1,
"STRING_PROMPT"
],
[
89,
73,
0,
99,
2,
"STRING_PROMPT"
],
[
90,
99,
0,
88,
0,
"IMAGE"
],
[
95,
92,
0,
102,
0,
"TransformerModel"
],
[
96,
102,
0,
96,
0,
"TransformerModel"
]
],
"groups": [
{
"id": 1,
"title": "Load Model",
"bounding": [
227.96267700195312,
-546.4359741210938,
1350.4793699732413,
404.87677206390265
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"id": 2,
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.6477940671634007,
"offset": [
411.8286437242131,
622.7162385081058
]
},
"frontendVersion": "1.36.11",
"workflowRendererVersion": "LG",
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
},
"node_versions": {
"CogVideoX-Fun": "93aa7b2530dccd1e91c625eee439a5e24f8ffa04",
"comfy-core": "0.6.0"
}
},
"version": 0.4
}
@@ -0,0 +1,614 @@
{
"id": "dcf2fcac-6293-4a86-b30b-f63e420177f2",
"revision": 0,
"last_node_id": 102,
"last_link_id": 96,
"nodes": [
{
"id": 78,
"type": "Note",
"pos": [
18,
-46
],
"size": [
210,
88
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"You can write prompt here\n(你可以在此填写提示词)"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 91,
"type": "LoadZImageTextEncoderModel",
"pos": [
283.53765869140625,
-280.6837463378906
],
"size": [
407.4130859375,
102
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder",
"type": "TextEncoderModel",
"links": [
80
]
},
{
"name": "tokenizer",
"type": "Tokenizer",
"links": [
81
]
}
],
"properties": {
"Node name for S&R": "LoadZImageTextEncoderModel"
},
"widgets_values": [
"qwen_3_4b.safetensors",
"bf16"
]
},
{
"id": 75,
"type": "FunTextBox",
"pos": [
250,
-50
],
"size": [
383.54010009765625,
156.71620178222656
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"slot_index": 0,
"links": [
88
]
}
],
"title": "Positive Prompt(正向提示词)",
"properties": {
"Node name for S&R": "FunTextBox"
},
"widgets_values": [
"A photo of Sakura, a 17-year-old high school student from Japan, captured in a candid, high-fidelity cinematic moment on a rainy evening. She is squatting low on the rain-slicked asphalt of an urban sidewalk, holding a transparent vinyl umbrella with a white handle resting over her shoulder in one hand, her other hand resting on her knee. The clear plastic canopy is streaked with rivulets of water and beaded with droplets that catch the ambient city light. A profound, silent interaction defines the scene: Sakura is looking directly downward, her expression gentle and focused, locking eyes with a small black cat sitting on the wet ground in front of her.\n\nSakura has long, lustrous black hair styled in a precise hime cut with blunt bangs across her forehead and sidelocks framing her cheeks, damp strands clinging subtly to her jacket, with a single red ribbon tied on the left side. Her visible pores on her nose, and a soft sheen of moisture on her cheeks. She wears a dark navy sailor-style school uniform (seifuku) featuring a white collar with red linear detailing and a bright red necktie loosely knotted at the chest; a simple black choker encircles her neck. The uniform jacket has oversized sleeves. Her lower body features a short, dark pleated miniskirt that fans slightly over clean white ankle socks that provide a stark contrast to the wet asphalt, ending in dark leather loafers that gleam with moisture.\n\nThe black cat sits upright in a shallow puddle, its short fur slicked by the rain, tilting its head back to stare intently up into Sakura's face, establishing a clear line of sight. The background is anchored by a large, illuminated red vending machine standing against the darkness, its cool bluish-white interior light spilling onto Sakura's profile and the umbrella. The ground reflects the red chassis and the neon streetlights in distorted patches on the wet pavement. Additional cool rain streaks fall through the frame, some caught in sharp focus and others blurred into vertical lines against the background lights. The scene is rendered with a wide-aperture lens creating a shallow depth of field, keeping the girl and cat in sharp focus while softening the background into gentle bokeh, with the texture of fine-grain 35mm film stock.\n"
]
},
{
"id": 73,
"type": "FunTextBox",
"pos": [
250,
160
],
"size": [
383.7149963378906,
183.83506774902344
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "prompt",
"type": "STRING_PROMPT",
"slot_index": 0,
"links": [
89
]
}
],
"title": "Negtive Prompt(反向提示词)",
"properties": {
"Node name for S&R": "FunTextBox"
},
"widgets_values": [
""
]
},
{
"id": 97,
"type": "Note",
"pos": [
-354.4680507215508,
-433.6570714778354
],
"size": [
598.1623727144193,
233.48501180428053
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"text": ""
},
"widgets_values": [
"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].\nmodel_full_load means that the entire model will be moved to the GPU.\n\nmodel_full_load_and_qfloat8 means that the entire model will be moved to the GPU,\nand the transformer model has been quantized to float8, which can save more GPU memory. \n\nmodel_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.\n\nmodel_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, \nand the transformer model has been quantized to float8, which can save more GPU memory. \n\nsequential_cpu_offload means that each layer of the model will be moved to the CPU after use, \nresulting in slower speeds but saving a large amount of GPU memory."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 93,
"type": "LoadZImageVAEModel",
"pos": [
1168.1599508804559,
-314.8650476692169
],
"size": [
377.8583984375,
84.69844055175781
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "VAEModel",
"links": [
79
]
}
],
"properties": {
"Node name for S&R": "LoadZImageVAEModel"
},
"widgets_values": [
"ae.safetensors",
"bf16"
]
},
{
"id": 88,
"type": "PreviewImage",
"pos": [
1049.1402001998824,
-52.699751832945104
],
"size": [
366.56134033203125,
415.4429626464844
],
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 90
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 100,
"type": "LoadImage",
"pos": [
396.48551767952716,
419.158511044569
],
"size": [
270,
314.00000000000006
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
86
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"z-image-turbo_00004_.png",
"image"
]
},
{
"id": 92,
"type": "LoadZImageTransformerModel",
"pos": [
275.9798278808594,
-465.2391052246094
],
"size": [
416.3677673339844,
106.13789367675781
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "transformer",
"type": "TransformerModel",
"links": [
95
]
},
{
"name": "model_name",
"type": "STRING",
"links": [
82
]
}
],
"properties": {
"Node name for S&R": "LoadZImageTransformerModel"
},
"widgets_values": [
"z_image_bf16.safetensors",
"bf16"
]
},
{
"id": 99,
"type": "ZImageControlSampler",
"pos": [
727.2482831521481,
-52.41010988674983
],
"size": [
270,
350
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "funmodels",
"type": "FunModels",
"link": 87
},
{
"name": "prompt",
"type": "STRING_PROMPT",
"link": 88
},
{
"name": "negative_prompt",
"type": "STRING_PROMPT",
"link": 89
},
{
"name": "control_image",
"shape": 7,
"type": "IMAGE",
"link": 86
},
{
"name": "inpaint_image",
"shape": 7,
"type": "IMAGE",
"link": null
},
{
"name": "mask_image",
"shape": 7,
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
90
]
}
],
"properties": {
"Node name for S&R": "ZImageControlSampler"
},
"widgets_values": [
1824,
2416,
43,
"fixed",
25,
4.5,
"Flow",
3,
0.85
]
},
{
"id": 102,
"type": "LoadZImageControlNetInModel",
"pos": [
779.793189390101,
-457.3825558553134
],
"size": [
589.8698159570357,
82
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "transformer",
"type": "TransformerModel",
"link": 95
}
],
"outputs": [
{
"name": "transformer",
"type": "TransformerModel",
"links": [
96
]
}
],
"properties": {
"Node name for S&R": "LoadZImageControlNetInModel"
},
"widgets_values": [
"z_image/z_image_control_2.1_lite.yaml",
"Z-Image-Fun-Controlnet-Tile-2.1-lite.safetensors"
]
},
{
"id": 96,
"type": "CombineZImagePipeline",
"pos": [
790.3572998046875,
-328.7134094238281
],
"size": [
342.5804748535156,
162
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "transformer",
"type": "TransformerModel",
"link": 96
},
{
"name": "vae",
"type": "VAEModel",
"link": 79
},
{
"name": "text_encoder",
"type": "TextEncoderModel",
"link": 80
},
{
"name": "tokenizer",
"type": "Tokenizer",
"link": 81
},
{
"name": "processor",
"shape": 7,
"type": "Processor",
"link": null
},
{
"name": "model_name",
"type": "STRING",
"widget": {
"name": "model_name"
},
"link": 82
}
],
"outputs": [
{
"name": "funmodels",
"type": "FunModels",
"links": [
87
]
}
],
"properties": {
"Node name for S&R": "CombineZImagePipeline"
},
"widgets_values": [
"",
"sequential_cpu_offload"
]
}
],
"links": [
[
79,
93,
0,
96,
1,
"VAEModel"
],
[
80,
91,
0,
96,
2,
"TextEncoderModel"
],
[
81,
91,
1,
96,
3,
"Tokenizer"
],
[
82,
92,
1,
96,
5,
"STRING"
],
[
86,
100,
0,
99,
3,
"IMAGE"
],
[
87,
96,
0,
99,
0,
"FunModels"
],
[
88,
75,
0,
99,
1,
"STRING_PROMPT"
],
[
89,
73,
0,
99,
2,
"STRING_PROMPT"
],
[
90,
99,
0,
88,
0,
"IMAGE"
],
[
95,
92,
0,
102,
0,
"TransformerModel"
],
[
96,
102,
0,
96,
0,
"TransformerModel"
]
],
"groups": [
{
"id": 1,
"title": "Load Model",
"bounding": [
227.96267700195312,
-546.4359741210938,
1350.4793699732413,
404.87677206390265
],
"color": "#b06634",
"font_size": 24,
"flags": {}
},
{
"id": 2,
"title": "Prompts",
"bounding": [
218,
-127,
450,
483
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.6477940671634007,
"offset": [
475.03594149909674,
811.7170336004714
]
},
"frontendVersion": "1.34.9",
"workflowRendererVersion": "LG",
"workspace_info": {
"id": "776b62b4-bd17-4ed3-9923-b7aad000b1ea"
},
"node_versions": {
"CogVideoX-Fun": "07dd34b942f866d5f95e8b812b6082d359079260",
"comfy-core": "0.6.0"
}
},
"version": 0.4
}
@@ -300,7 +300,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors"
]
},
{
@@ -535,7 +535,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1_lite.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-lite-2601-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-lite-2602-8steps.safetensors"
]
}
],
@@ -396,7 +396,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors"
]
},
{
@@ -300,7 +300,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors"
]
},
{
@@ -300,7 +300,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors"
]
},
{
@@ -469,7 +469,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1_lite.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-lite-2601-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-lite-2602-8steps.safetensors"
]
}
],
@@ -300,7 +300,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors"
]
},
{
@@ -255,7 +255,7 @@
},
"widgets_values": [
"",
"model_cpu_offload"
"sequential_cpu_offload"
]
},
{
@@ -396,7 +396,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Tile-2.1-8steps.safetensors"
"Z-Image-Turbo-Fun-Controlnet-Tile-2.1-2601-8steps.safetensors"
]
},
{
@@ -225,7 +225,7 @@
},
"widgets_values": [
"z_image/z_image_control_2.1.yaml",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-8steps.safetensors",
"Z-Image-Turbo-Fun-Controlnet-Union-2.1-2602-8steps.safetensors",
"transformer"
]
},
@@ -0,0 +1,6 @@
format: diffusers
pipeline: minimax-h3
transformer_additional_kwargs:
control_blocks_places: [0, 10, 20, 30, 40]
control_in_dim: 49
control_apply_audio: false
@@ -0,0 +1,6 @@
format: diffusers
pipeline: minimax-h3
transformer_additional_kwargs:
control_blocks_places: [0, 10, 20, 30, 40]
control_in_dim: 24
control_apply_audio: false
+8
View File
@@ -20,6 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunInpaintPipeline,
CogVideoXFunPipeline)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
@@ -36,6 +38,9 @@ from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
@@ -186,6 +191,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+8
View File
@@ -20,6 +20,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
@@ -37,6 +39,9 @@ from videox_fun.dist import set_multi_gpus_devices, shard_model
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
@@ -178,6 +183,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+8
View File
@@ -19,6 +19,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
@@ -36,6 +38,9 @@ from videox_fun.dist import set_multi_gpus_devices, shard_model
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
@@ -185,6 +190,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -21,6 +21,8 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
T5Tokenizer)
from videox_fun.pipeline import (CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
@@ -38,6 +40,9 @@ from videox_fun.dist import set_multi_gpus_devices, shard_model
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
@@ -172,6 +177,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=[], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+210
View File
@@ -0,0 +1,210 @@
import os
import sys
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer,
ErnieImageTransformer2DModel, Mistral3Model)
from videox_fun.pipeline import ErnieImagePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/ERNIE-Image"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [1728, 992]
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
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"
negative_prompt = "低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"
guidance_scale = 4.5
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/ernie-image-t2i"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = ErnieImageTransformer2DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLFlux2.from_pretrained(
model_name,
subfolder="vae"
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get tokenizer and text_encoder
tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
text_encoder = Mistral3Model.from_pretrained(
model_name, subfolder="text_encoder", torch_dtype=weight_dtype
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = ErnieImagePipeline(
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
transformer=transformer,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.layers))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if compile_dit:
for i in range(len(pipeline.transformer.layers)):
pipeline.transformer.layers[i] = torch.compile(pipeline.transformer.layers[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
sample = pipeline(
prompt,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
).images
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0]
image.save(video_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+16 -13
View File
@@ -21,6 +21,8 @@ from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import FantasyTalkingPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
@@ -41,6 +43,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -72,10 +77,6 @@ num_skip_start_steps = 5
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
teacache_offload = False
# Skip some cfg steps in inference
# Recommended to be set between 0.00 and 0.25
cfg_skip_ratio = 0
# Riflex config
enable_riflex = False
# Index of intrinsic frequency
@@ -85,8 +86,10 @@ riflex_k = 6
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
# Please Download https://modelscope.cn/models/AI-ModelScope/wav2vec2-base-960h/summary
# to models/Diffusion_Transformer/Wan2.1-I2V-14B-720P/audio_encoder for encoding audio.
# to models/Diffusion_Transformer/wav2vec2-base-960h for encoding audio.
model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
# audio encoder model path. If None, will use os.path.join(model_name, "audio_encoder")
model_name_audio = "models/Diffusion_Transformer/wav2vec2-base-960h"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
@@ -97,7 +100,7 @@ shift = 5
# Load pretrained model if need
# The transformer_path is used for low noise model, the transformer_high_path is used for high noise model.
# The fantasytalking_model.ckpt can be downloaded in https://www.modelscope.cn/models/amap_cvlab/FantasyTalking/
transformer_path = "models/Personalized_Model/fantasytalking_model.ckpt"
transformer_path = "models/Personalized_Model/FantasyTalking/fantasytalking_model.ckpt"
vae_path = None
# Load lora model if need
lora_path = None
@@ -118,6 +121,7 @@ audio_path = "asset/talk.wav"
prompt = "一个女孩在海边说话。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 4.5
audio_guide_scale = 4.0
seed = 43
num_inference_steps = 40
lora_weight = 0.55
@@ -198,9 +202,8 @@ clip_image_encoder = CLIPModel.from_pretrained(
).to(weight_dtype)
clip_image_encoder = clip_image_encoder.eval()
audio_encoder = FantasyTalkingAudioEncoder(
os.path.join(model_name, "audio_encoder")
)
audio_encoder_path = model_name_audio if model_name_audio is not None else os.path.join(model_name, "audio_encoder")
audio_encoder = FantasyTalkingAudioEncoder(audio_encoder_path)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
@@ -245,6 +248,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -265,10 +271,6 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
@@ -291,6 +293,7 @@ with torch.no_grad():
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
audio_guide_scale = audio_guide_scale,
num_inference_steps = num_inference_steps,
video = input_video,
+262
View File
@@ -0,0 +1,262 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
FlashHeadTransformer3DModel, FlashHeadAudioEncoder)
from videox_fun.pipeline import FlashHeadPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image,
get_video_to_video_latent,
merge_video_audio, save_videos_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_full_load"
# Multi GPUs config
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
# Please Download https://modelscope.cn/models/AI-ModelScope/wav2vec2-base-960h/summary
model_name = "models/Diffusion_Transformer/SoulX-FlashHead-1_3B"
model_name_audio = "models/Diffusion_Transformer/wav2vec2-base-960h"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
shift = 5.0
stochastic_sampling = True
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [512, 512]
segment_frame_length = 33
fps = 25
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# The path of the reference image
ref_image = "asset/9.png"
# The path of the audio
audio_path = "asset/talk.wav"
# Audio guidance scale (FlashHead does not use text encoder, only audio conditioning)
audio_guide_scale = 1.0
seed = 42
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/flashhead-videos"
# FlashHead specific parameters
max_frames_num = 500
color_correction_strength = 1.0
use_apg = False
apg_momentum = 0.5
apg_norm_threshold = 1.0
audio_encode_mode = "stream"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
transformer = FlashHeadTransformer3DModel.from_pretrained(
os.path.join(model_name, "Model_Pro", config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, "VAE_Wan/Wan2.1_VAE.pth"),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Initialize FlashHead audio encoder for real-time audio encoding
# Uses Wav2Vec2Model (not Wav2Vec2ForCTC) matching original FlashHead implementation
audio_encoder = FlashHeadAudioEncoder(
model_name_audio, "cpu"
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline (FlashHead does not use text encoder or clip image encoder)
pipeline = FlashHeadPipeline(
transformer=transformer,
vae=vae,
scheduler=scheduler,
audio_encoder=audio_encoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
# For FlashHead, (segment_frame_length - 1) must be divisible by 4
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
latent_frames = (segment_frame_length - 1) // vae.config.temporal_compression_ratio + 1
# Prepare ref_image latent for FlashHead (no clip_image needed)
ref_image = get_image_latent(ref_image, sample_size=sample_size)
sample = pipeline(
segment_frame_length = segment_frame_length,
height = sample_size[0],
width = sample_size[1],
generator = generator,
audio_guide_scale = audio_guide_scale,
num_inference_steps = num_inference_steps,
ref_image = ref_image,
audio_path = audio_path,
audio_encode_mode = audio_encode_mode,
shift = shift,
fps = fps,
max_frames_num = max_frames_num,
color_correction_strength = color_correction_strength,
use_apg = use_apg,
apg_momentum = apg_momentum,
apg_norm_threshold = apg_norm_threshold,
stochastic_sampling = stochastic_sampling,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if sample.size()[2] == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
merge_video_audio(video_path=video_path, audio_path=audio_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+9 -1
View File
@@ -14,6 +14,8 @@ from videox_fun.models import (AutoencoderKL, CLIPTextModel, CLIPTokenizer,
FluxTransformer2DModel, T5EncoderModel,
T5TokenizerFast)
from videox_fun.pipeline import FluxPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
@@ -31,6 +33,9 @@ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload_and_qfloat8"
@@ -65,7 +70,7 @@ sample_size = [1344, 768]
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
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"
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. "
negative_prompt = " "
guidance_scale = 1.0
seed = 43
num_inference_steps = 50
@@ -166,6 +171,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+8
View File
@@ -22,6 +22,8 @@ from videox_fun.models import (AutoencoderKLHunyuanVideo, CLIPTextModel, CLIPIma
LlavaForConditionalGeneration, LlamaTokenizerFast)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import HunyuanVideoPipeline, HunyuanVideoI2VPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
@@ -43,6 +45,9 @@ from videox_fun.utils.utils import get_image
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -196,6 +201,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+8
View File
@@ -22,6 +22,8 @@ from videox_fun.models import (AutoencoderKLHunyuanVideo, CLIPTextModel,
LlamaModel, LlamaTokenizerFast)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import HunyuanVideoPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
@@ -42,6 +44,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -185,6 +190,9 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["x_embedder", "context_embedder", "time_text_embed", "rope", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+319
View File
@@ -0,0 +1,319 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
InfiniteTalkTransformer3DModel, InfiniteTalkAudioEncoder,
WanT5EncoderModel)
from videox_fun.pipeline import InfiniteTalkPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image,
get_video_to_video_latent,
merge_video_audio, save_videos_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | 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 |
# | 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 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.20
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
num_skip_start_steps = 5
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
teacache_offload = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-I2V-14B-480P"
model_name_audio = "models/Diffusion_Transformer/chinese-wav2vec2-base/"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
shift = 5.0
# Load pretrained model if need
transformer_path = "models/Personalized_Model/infinitetalk.safetensors"
vae_path = None
lora_path = None
# Other params
sample_size = [832, 480]
segment_frame_length = 81
fps = 25
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# The path of the reference image
ref_image = "asset/8.png"
# The path of the audio
audio_path = "asset/talk.wav"
# prompts
prompt = "一个人在说话。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 5.0
audio_guide_scale = 4.0
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/infitetalk-videos"
# InfiniteTalk specific parameters
max_frames_num = 500 # Total frames to generate
color_correction_strength = 1 # Color correction strength (0.0-1.0)
use_apg = False # Use Adaptive Projected Guidance
apg_momentum = 0.5
apg_norm_threshold = 1.0
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
transformer = InfiniteTalkTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Initialize InfiniteTalk audio encoder for real-time audio encoding
# Uses Wav2Vec2Model (not Wav2Vec2ForCTC) matching original InfiniteTalk implementation
audio_encoder = InfiniteTalkAudioEncoder(
model_name_audio, "cpu"
)
# Get Clip Image Encoder
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
).to(weight_dtype)
clip_image_encoder = clip_image_encoder.eval()
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = InfiniteTalkPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
audio_encoder=audio_encoder,
clip_image_encoder=clip_image_encoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
# For InfiniteTalk, (segment_frame_length - 1) must be divisible by 4
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
latent_frames = (segment_frame_length - 1) // vae.config.temporal_compression_ratio + 1
# Prepare clip_image from original ref_image path
clip_image = get_image(ref_image)
ref_image = get_image_latent(ref_image, sample_size=sample_size)
sample = pipeline(
prompt,
segment_frame_length = segment_frame_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
audio_guide_scale = audio_guide_scale,
num_inference_steps = num_inference_steps,
ref_image = ref_image,
clip_image = clip_image, # Pass clip_image
audio_path = audio_path,
shift = shift,
fps = fps,
max_frames_num = max_frames_num,
color_correction_strength = color_correction_strength,
use_apg = use_apg,
apg_momentum = apg_momentum,
apg_norm_threshold = apg_norm_threshold,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if sample.size()[2] == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
merge_video_audio(video_path=video_path, audio_path=audio_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+226
View File
@@ -0,0 +1,226 @@
import os
import sys
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLFlux2, AutoTokenizer,
LensGptOssEncoder, LensTransformer2DModel)
from videox_fun.pipeline import LensPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/Lens"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [1728, 992]
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Set to True on A100/V100 to dequantize MXFP4 GPT-OSS weights.
dequantize_mxfp4 = False
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"
negative_prompt = " "
guidance_scale = 4.5
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/lens-t2i"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Get transformer
transformer = LensTransformer2DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
).to(weight_dtype)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLFlux2.from_pretrained(
model_name,
subfolder="vae",
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get tokenizer and text_encoder
tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
text_encoder_kwargs = {"subfolder": "text_encoder", "torch_dtype": weight_dtype}
try:
from transformers import Mxfp4Config
text_encoder_kwargs["quantization_config"] = Mxfp4Config(
dequantize=dequantize_mxfp4
)
except ImportError:
pass # Older transformers without Mxfp4Config
text_encoder = LensGptOssEncoder.from_pretrained(
model_name, **text_encoder_kwargs
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LensPipeline(
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
transformer=transformer,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=list(text_encoder.model.layers))
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["img_in", "txt_in", "timestep"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
sample = pipeline(
prompt,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
).images
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0]
image.save(video_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+253
View File
@@ -0,0 +1,253 @@
import os
import sys
import numpy as np
import torch
from PIL import Image
from transformers import AutoProcessor
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLQwenImage,
LingBotVideoTransformer3DModel,
Qwen3VLForConditionalGeneration)
from videox_fun.models.lingbot_video_rewriter import ensure_json_caption
from videox_fun.pipeline import LingBotVideoI2VPipeline
from videox_fun.pipeline.pipeline_lingbot_video import DEFAULT_NEGATIVE_PROMPT
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
# Sequence parallelism shards the video tokens across ranks and keeps the text tokens
# replicated, so the video token count (T/pF * H/16 * W/16) must be divisible by
# ulysses_degree * ring_degree.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
# Config and model path
# model path
model_name = "models/Diffusion_Transformer/lingbot-video-dense-1.3b"
# Rewriter weights: the base VLM and the rewriter LoRA used to rewrite the
# plain prompt into the structured JSON caption the DiT expects.
rewriter_base_model = "models/Diffusion_Transformer/Qwen3.6-27B"
rewriter_lora_path = "models/Diffusion_Transformer/lingbot-video-rewriter-lora"
# Only "Flow_Unipc" is supported: LingBot-Video ships and was trained with FlowUniPCMultistepScheduler.
sampler_name = "Flow_Unipc"
# Flow shift. 3.0 is the officially recommended value for both dense and MoE models.
shift = 3.0
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# The condition image is used twice: as Qwen3-VL visual input and as a clean
# first-frame latent injected into the diffusion latent (ti2v).
validation_image = "asset/1.png"
# prompts
# Write a plain natural-language prompt: it is ALWAYS rewritten into the
# structured JSON caption the DiT expects by the official prompt rewriter
# (EXPAND -> MAP, Qwen3.6-27B base + rewriter LoRA). For ti2v the same first
# frame is fed to the rewriter. Direct JSON/hand-written input is not a
# supported path; the rewrite result is cached under save_path.
prompt = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。"
negative_prompt = DEFAULT_NEGATIVE_PROMPT
guidance_scale = 3.0
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/lingbot-video-i2v"
# Rewrite the prompt before loading any generation model (the rewriter's 27B
# base VLM is freed right after, so it never coexists with the DiT on GPU).
# ti2v: the same first frame is fed to the rewriter.
prompt = ensure_json_caption(
prompt, mode="ti2v", duration=round(video_length / fps, 2),
first_frame=validation_image,
cache_file=os.path.join(save_path, "caption_cache.json"),
base=rewriter_base_model, adapter=rewriter_lora_path,
)
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = LingBotVideoTransformer3DModel.from_pretrained(
os.path.join(model_name, "transformer"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Re-apply the fp32-sensitive-module cast (norm / router / modulation stay fp32).
transformer = transformer.to(weight_dtype)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae (diffusers-format QwenImage VAE, Wan-style 16ch causal VAE)
vae = AutoencoderKLQwenImage.from_pretrained(
model_name,
subfolder="vae",
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Processor (Qwen3-VL tokenizer + image processor)
processor = AutoProcessor.from_pretrained(
os.path.join(model_name, "processor"),
)
# Get Text encoder (Qwen3-VL)
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow_Unipc": FlowUniPCMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
# Get Pipeline
pipeline = LingBotVideoI2VPipeline(
transformer=transformer,
vae=vae,
text_encoder=text_encoder,
processor=processor,
scheduler=scheduler,
)
if GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // pipeline.vae_scale_factor_temporal * pipeline.vae_scale_factor_temporal) + 1 if video_length != 1 else 1
image = Image.open(validation_image).convert("RGB")
sample = pipeline(
prompt,
image = image,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
shift = shift,
num_inference_steps = num_inference_steps,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
# count outputs only: caption_cache.json must not shift the index
index = len([path for path in os.listdir(save_path) if path.endswith((".mp4", ".png"))]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+254
View File
@@ -0,0 +1,254 @@
import os
import sys
import numpy as np
import torch
from PIL import Image
from transformers import AutoProcessor
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLQwenImage,
LingBotVideoTransformer3DModel,
Qwen3VLForConditionalGeneration)
from videox_fun.pipeline import LingBotVideoPipeline
from videox_fun.pipeline.pipeline_lingbot_video import DEFAULT_NEGATIVE_PROMPT
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_grid
from videox_fun.models.lingbot_video_rewriter import ensure_json_caption
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
# Sequence parallelism shards the video tokens across ranks and keeps the text tokens
# replicated, so the video token count (T/pF * H/16 * W/16) must be divisible by
# ulysses_degree * ring_degree.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
# Config and model path
# model path
model_name = "models/Diffusion_Transformer/lingbot-video-dense-1.3b/"
# Rewriter weights: the base VLM and the rewriter LoRA used to rewrite the
# plain prompt into the structured JSON caption the DiT expects.
rewriter_base_model = "models/Diffusion_Transformer/Qwen3.6-27B"
rewriter_lora_path = "models/Diffusion_Transformer/lingbot-video-rewriter-lora"
# Only "Flow_Unipc" is supported: LingBot-Video ships and was trained with FlowUniPCMultistepScheduler.
sampler_name = "Flow_Unipc"
# Flow shift. 3.0 is the officially recommended value for both dense and MoE models.
shift = 3.0
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
# video_length 1 generates a still image (t2i); videos must be 4n+1 frames.
video_length = 81
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# prompts
# Write a plain natural-language prompt: it is ALWAYS rewritten into the
# structured JSON caption the DiT expects by the official prompt rewriter
# (EXPAND -> MAP, Qwen3.6-27B base + rewriter LoRA). Direct JSON/hand-written
# input is not a supported path; the rewrite result is cached under save_path
# so re-runs with the same prompt skip the rewrite.
prompt = (
"A young musician sits on a weathered wooden stool in a sunlit rehearsal room, "
"steadily strumming an acoustic guitar. Warm golden-hour light streams through "
"tall windows, dust motes drifting slowly in the air; the wooden floor shows "
"visible wear and the walls carry acoustic panels. The camera slowly orbits "
"from a side profile to a frontal view at eye level, keeping the musician "
"centered in the frame."
)
negative_prompt = DEFAULT_NEGATIVE_PROMPT
guidance_scale = 3.0
seed = 43
num_inference_steps = 40
lora_weight = 0.55
save_path = "samples/lingbot-video-t2v"
# Rewrite the prompt before loading any generation model (the rewriter's 27B
# base VLM is freed right after, so it never coexists with the DiT on GPU).
prompt = ensure_json_caption(
prompt, mode="t2v", duration=round(video_length / fps, 2),
cache_file=os.path.join(save_path, "caption_cache.json"),
base=rewriter_base_model, adapter=rewriter_lora_path,
)
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = LingBotVideoTransformer3DModel.from_pretrained(
os.path.join(model_name, "transformer"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Re-apply the fp32-sensitive-module cast (norm / router / modulation stay fp32).
transformer = transformer.to(weight_dtype)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae (diffusers-format QwenImage VAE, Wan-style 16ch causal VAE)
vae = AutoencoderKLQwenImage.from_pretrained(
model_name,
subfolder="vae",
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Processor (Qwen3-VL tokenizer + image processor)
processor = AutoProcessor.from_pretrained(
os.path.join(model_name, "processor"),
)
# Get Text encoder (Qwen3-VL)
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow_Unipc": FlowUniPCMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
# Get Pipeline
pipeline = LingBotVideoPipeline(
transformer=transformer,
vae=vae,
text_encoder=text_encoder,
processor=processor,
scheduler=scheduler,
)
if GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // pipeline.vae_scale_factor_temporal * pipeline.vae_scale_factor_temporal) + 1 if video_length != 1 else 1
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
shift = shift,
num_inference_steps = num_inference_steps,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
# count outputs only: caption_cache.json must not shift the index
index = len([path for path in os.listdir(save_path) if path.endswith((".mp4", ".png"))]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
@@ -0,0 +1,363 @@
import gc
import os
import sys
import torch
import torch.nn.functional as F
from transformers import AutoProcessor
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLQwenImage,
LingBotVideoTransformer3DModel,
Qwen3VLForConditionalGeneration)
from videox_fun.pipeline import LingBotVideoPipeline
from videox_fun.pipeline.pipeline_lingbot_video import (DEFAULT_NEGATIVE_PROMPT,
prepare_refiner_latent)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.utils import save_videos_grid
from videox_fun.models.lingbot_video_rewriter import ensure_json_caption
# Two-stage LingBot-Video t2v: the base DiT samples at a low resolution, then the
# "refiner" DiT re-noises the upsampled latent to sigma = refiner_t_thresh and
# denoises it at the target resolution.
#
# The two DiTs are loaded and freed one at a time, so a single GPU only ever holds
# one 30B transformer (the MoE base and refiner are ~60GB each in bfloat16).
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
# Sequence parallelism shards the video tokens across ranks and keeps the text tokens
# replicated, so the video token count (T/pF * H/16 * W/16) must be divisible by
# ulysses_degree * ring_degree, at both the base and the refiner resolution.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
# Config and model path
# model path
# The refiner ships only with the MoE 30B-A3B model, as its "refiner" subfolder.
model_name = "models/Diffusion_Transformer/lingbot-video-moe-30b-a3b"
refiner_model_name = model_name
# Subfolders of the base and refiner DiT inside the model root.
transformer_subpath = "transformer"
refiner_subpath = "refiner"
# Rewriter weights: the base VLM and the rewriter LoRA used to rewrite the
# plain prompt into the structured JSON caption the DiT expects.
rewriter_base_model = "models/Diffusion_Transformer/Qwen3.6-27B"
rewriter_lora_path = "models/Diffusion_Transformer/lingbot-video-rewriter-lora"
# Only "Flow_Unipc" is supported: LingBot-Video ships and was trained with FlowUniPCMultistepScheduler.
sampler_name = "Flow_Unipc"
# Flow shift. 3.0 is the officially recommended value for both dense and MoE models.
shift = 3.0
# Load pretrained model if need
transformer_path = None
refiner_path = None
vae_path = None
# Base stage params. video_length must be 1 or 4n+1; 121 frames is 5s at 24 fps.
sample_size = [480, 832]
video_length = 81
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# prompts
# Write a plain natural-language prompt: it is ALWAYS rewritten into the
# structured JSON caption the DiT expects by the official prompt rewriter
# (EXPAND -> MAP, Qwen3.6-27B base + rewriter LoRA). Direct JSON/hand-written
# input is not a supported path; the rewrite result is cached under save_path.
prompt = (
"A young musician sits on a weathered wooden stool in a sunlit rehearsal room, "
"steadily strumming an acoustic guitar. Warm golden-hour light streams through "
"tall windows, dust motes drifting slowly in the air. The camera slowly orbits "
"from a side profile to a frontal view at eye level, keeping the musician "
"centered in the frame."
)
negative_prompt = DEFAULT_NEGATIVE_PROMPT
guidance_scale = 3.0
seed = 43
num_inference_steps = 40
save_path = "samples/lingbot-video-t2v-refine"
# Refiner stage params (official defaults). The refiner attends over the full
# high-resolution latent, so 1088x1920 is heavy on a single GPU: prefer fewer
# frames there, or shard the DiT across GPUs.
refiner_sample_size = [1088, 1920]
refiner_steps = 8
refiner_guidance_scale = 3.0
refiner_shift = 3.0
# Re-noise level: the refiner only walks the schedule from this sigma down to 0.
refiner_t_thresh = 0.85
# Extra low-noise steps appended after the truncated schedule.
refiner_sigma_tail_steps = 2
# Rewrite the prompt before loading any generation model (the rewriter's 27B
# base VLM is freed right after, so it never coexists with the DiT on GPU).
prompt = ensure_json_caption(
prompt, mode="t2v", duration=round(video_length / fps, 2),
cache_file=os.path.join(save_path, "caption_cache.json"),
base=rewriter_base_model, adapter=rewriter_lora_path,
)
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
def load_transformer(root, subpath, checkpoint_path):
transformer = LingBotVideoTransformer3DModel.from_pretrained(
os.path.join(root, subpath),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Re-apply the fp32-sensitive-module cast (norm / router / modulation stay fp32).
transformer = transformer.to(weight_dtype)
if checkpoint_path is not None:
print(f"From checkpoint: {checkpoint_path}")
if checkpoint_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(checkpoint_path)
else:
state_dict = torch.load(checkpoint_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
return transformer
def next_index():
if not os.path.exists(save_path):
return 1
return len([path for path in os.listdir(save_path) if path.endswith("_base.mp4")]) + 1
def save_results(sample, index, prefix, save_fps):
# Both stages of a run share an index: 00000001_base.mp4 / 00000001_refined.mp4.
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
video_path = os.path.join(save_path, f"{str(index).zfill(8)}_{prefix}.mp4")
save_videos_grid(sample, video_path, fps=save_fps)
return video_path
# Get Vae (diffusers-format QwenImage VAE, Wan-style 16ch causal VAE), shared by both stages
vae = AutoencoderKLQwenImage.from_pretrained(
model_name,
subfolder="vae",
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Processor (Qwen3-VL tokenizer + image processor)
processor = AutoProcessor.from_pretrained(
os.path.join(model_name, "processor"),
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow_Unipc": FlowUniPCMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
# Stage 0: encode the prompts once. Both stages condition on the same text, so the
# text encoder is released before either 30B DiT is loaded.
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
encode_pipeline = LingBotVideoPipeline(
transformer=None,
vae=vae,
text_encoder=text_encoder,
processor=processor,
scheduler=scheduler,
)
encode_pipeline.text_encoder.to(device)
with torch.no_grad():
prompt_embeds, prompt_mask = encode_pipeline.encode_prompt(prompt, device=device)
negative_prompt_embeds, negative_prompt_mask = encode_pipeline.encode_prompt(negative_prompt, device=device)
del encode_pipeline, text_encoder
gc.collect()
torch.cuda.empty_cache()
# Stage 1: base sampling at sample_size
transformer = load_transformer(model_name, transformer_subpath, transformer_path)
pipeline = LingBotVideoPipeline(
transformer=transformer,
vae=vae,
text_encoder=None,
processor=processor,
scheduler=scheduler,
)
if GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
generator = torch.Generator(device=device).manual_seed(seed)
with torch.no_grad():
video_length = int((video_length - 1) // pipeline.vae_scale_factor_temporal * pipeline.vae_scale_factor_temporal) + 1 if video_length != 1 else 1
base_sample = pipeline(
prompt,
num_frames = video_length,
prompt_embeds = prompt_embeds,
prompt_mask = prompt_mask,
negative_prompt_embeds = negative_prompt_embeds,
negative_prompt_mask = negative_prompt_mask,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
shift = shift,
num_inference_steps = num_inference_steps,
).videos
save_index = next_index()
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
base_path = save_results(base_sample, save_index, "base", fps)
else:
base_path = save_results(base_sample, save_index, "base", fps)
del pipeline, transformer
gc.collect()
torch.cuda.empty_cache()
# Stage 2: refinement at refiner_sample_size. Unlike the reference runner, the base
# frames are refined in memory instead of being re-read from the saved mp4, which
# skips a lossy encode/decode round trip.
refiner = load_transformer(refiner_model_name, refiner_subpath, refiner_path)
refiner_pipeline = LingBotVideoPipeline(
transformer=refiner,
vae=vae,
text_encoder=None,
processor=processor,
scheduler=scheduler,
)
if GPU_memory_mode == "model_group_offload":
register_auto_device_hook(refiner_pipeline.transformer)
safe_enable_group_offload(refiner_pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(refiner, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(refiner, weight_dtype)
refiner_pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
refiner_pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(refiner, exclude_module_name=["time_embedder", "time_modulation", "text_embedder", "norm", "router", "scale_shift_table", "proj_out"], device=device)
convert_weight_dtype_wrapper(refiner, weight_dtype)
refiner_pipeline.to(device=device)
else:
refiner_pipeline.to(device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
refiner.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
refiner_pipeline.transformer = shard_fn(refiner_pipeline.transformer)
print("Add FSDP DIT")
refiner_generator = torch.Generator(device=device).manual_seed(seed)
with torch.no_grad():
# video: [B, C, T, H, W] in [0, 1]
bsz, channels, frames, _height, _width = base_sample.shape
flat = base_sample.permute(0, 2, 1, 3, 4).reshape(bsz * frames, channels, _height, _width)
resized = F.interpolate(flat, size=(refiner_sample_size[0], refiner_sample_size[1]), mode="bicubic", align_corners=False).clamp(0.0, 1.0)
lowres_video = resized.reshape(bsz, frames, channels, refiner_sample_size[0], refiner_sample_size[1]).permute(0, 2, 1, 3, 4).contiguous()
x_up = refiner_pipeline.encode_video_latent(lowres_video, generator=refiner_generator)
noise = torch.randn(x_up.shape, device=x_up.device, dtype=x_up.dtype, generator=refiner_generator)
initial_latent = prepare_refiner_latent(x_up, noise, refiner_t_thresh)
del lowres_video, x_up, noise
# The refiner's unconditional branch uses zeroed conditions rather than the
# negative prompt (null_cond_clone_zero in the reference implementation).
refiner_sample = refiner_pipeline(
prompt,
num_frames = video_length,
prompt_embeds = prompt_embeds,
prompt_mask = prompt_mask,
negative_prompt_embeds = torch.zeros_like(prompt_embeds),
negative_prompt_mask = prompt_mask.clone(),
height = refiner_sample_size[0],
width = refiner_sample_size[1],
latents = initial_latent,
generator = refiner_generator,
guidance_scale = refiner_guidance_scale,
shift = refiner_shift,
num_inference_steps = refiner_steps,
t_thresh = refiner_t_thresh,
refiner_sigma_tail_steps = refiner_sigma_tail_steps,
).videos
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
refined_path = save_results(refiner_sample, save_index, "refined", fps)
print(f"base: {base_path}\nrefined: {refined_path}")
else:
refined_path = save_results(refiner_sample, save_index, "refined", fps)
print(f"base: {base_path}\nrefined: {refined_path}")
+372
View File
@@ -0,0 +1,372 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, WanT5EncoderModel,
WanTransformer3DModel_LingbotWorld)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2I2VPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.data.utils import prepare_lingbot_dit_cond_dict
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# TeaCache config
enable_teacache = True
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
num_skip_start_steps = 5
# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory.
teacache_offload = False
# Skip some cfg steps in inference
# Recommended to be set between 0.00 and 0.25
cfg_skip_ratio = 0
# Config and model path (the lingbot model reuses the Wan2.2 I2V layout).
config_path = "config/wan2.2/wan_civitai_i2v.yaml"
model_name = "models/Diffusion_Transformer/lingbot-world-base-cam"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow_Unipc"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# For 480p generation, a shift of 3.0 is recommended.
shift = 5
# Load pretrained model if need
transformer_path = None
transformer_high_path = None
vae_path = None
lora_path = None
lora_high_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Camera trajectory (poses.npy / intrinsics.npy) + reference image + prompt.
action_path = "asset/lingbot_demo"
validation_image_start = "asset/lingbot_demo/image.jpg"
# prompts
prompt = "The video presents a soaring journey through a fantasy jungle. The wind whips past the rider's blue hands gripping the reins, causing the leather straps to vibrate. The ancient gothic castle approaches steadily, its stone details becoming clearer against the backdrop of floating islands and distant waterfalls."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 5.0
seed = 43
num_inference_steps = 40
lora_weight = 0.55
lora_high_weight = 0.55
save_path = "samples/lingbot-world-i2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
transformer = WanTransformer3DModel_LingbotWorld.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe":
transformer_2 = WanTransformer3DModel_LingbotWorld.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
transformer_2 = None
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
if transformer_2 is not None:
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
Chosen_AutoencoderKL = {
"AutoencoderKLWan": AutoencoderKLWan,
"AutoencoderKLWan3_8": AutoencoderKLWan3_8
}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
vae = Chosen_AutoencoderKL.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline (reuse the standard Wan2.2 I2V pipeline unchanged).
pipeline = Wan2_2I2VPipeline(
transformer=transformer,
transformer_2=transformer_2,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if transformer_2 is not None:
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
if transformer_2 is not None:
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
if transformer_2 is not None:
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
if transformer_2 is not None:
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
if transformer_2 is not None:
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
# Prepare the camera condition. The trajectory length may shrink video_length.
dit_cond_dict, video_length = prepare_lingbot_dit_cond_dict(
action_path=action_path,
frame_num=video_length,
height=sample_size[0],
width=sample_size[1],
device=device,
dtype=weight_dtype,
control_type=getattr(transformer, "control_type", "cam"),
vae_stride=(vae.config.temporal_compression_ratio, vae.config.spatial_compression_ratio, vae.config.spatial_compression_ratio),
patch_size=transformer.config.patch_size,
)
# Feed the camera condition to the transformers so the standard pipeline can
# be reused without any modification.
pipeline.transformer.dit_cond_dict = dit_cond_dict
if transformer_2 is not None:
pipeline.transformer_2.dit_cond_dict = dit_cond_dict
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
boundary = boundary,
video = input_video,
mask_video = input_video_mask,
shift = shift,
).videos
# Clear the camera condition after generation.
pipeline.transformer.dit_cond_dict = None
if transformer_2 is not None:
pipeline.transformer_2.dit_cond_dict = None
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+318
View File
@@ -0,0 +1,318 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_LingbotWorldFast)
from videox_fun.pipeline import WanFunLingbotWorldFastPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.utils import filter_kwargs, save_videos_grid
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path.
# The lingbot-world fast checkpoint only ships the transformer (16 shards at the
# repo root); its VAE / T5 / tokenizer are sourced from the base-cam repo, which
# keeps the raw Wan2.1 layout used by ``config/wan2.1/wan_civitai.yaml``.
config_path = "config/wan2.1/wan_civitai.yaml"
transformer_name = "models/Diffusion_Transformer/lingbot-world-fast"
model_name = "models/Diffusion_Transformer/lingbot-world-base-cam"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++".
# The distilled fast model is trained on the FlowUniPC schedule (int64 timesteps,
# shift applied inside set_timesteps); the tuned ``timesteps_index`` only matches
# that grid, so "Flow_Unipc" is required to reproduce the reference quality.
sampler_name = "Flow_Unipc"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics of the
# distilled few-step flow-matching schedule. The official lingbot-world fast
# model is calibrated with sample_shift=10.0 (wan_i2v_A14B.py); generate_fast.py
# passes cfg.sample_shift (=10.0), NOT the generate() signature default of 5.0.
# Using 5.0 here builds the wrong sigma grid ([999,957,899,702] instead of
# [999,978,947,825]), so every renoise step is off-distribution and the decoded
# frames look grainy/painterly. 10.0 reproduces the reference quality.
shift = 10.0
# stochastic_sampling=True selects the native calibrated few-step schedule
# (fixed-index FlowUniPC grid the fast model was distilled on) — the correct
# path. False falls back to the generic scheduler dispatch (off-distribution
# for the fast weights).
stochastic_sampling = True
# Load pretrained transformer weights (optional override).
transformer_path = None
vae_path = None
# LoRA path (optional). The fast model is a single distilled transformer, so only
# one LoRA is used here (no high-noise counterpart like the MoE predict_i2v.py).
lora_path = None
# Camera control type - 'cam' (6-dim plücker).
control_type = "cam"
# Self-Forcing causal inference config.
# `num_frame_per_block`: 3 = chunk-wise:
# 1 = frame-wise:
num_frame_per_block = 3
# Local attention window size (-1 for global attention).
local_attn_size = -1
sink_size = 0
# Other params
# The reference derives the output resolution from a pixel-area budget and the
# input image aspect ratio (image2video_fast.py), rather than forcing a fixed
# size. ``sample_size`` here is only used as the area budget (max_area =
# sample_size[0] * sample_size[1]); the actual height/width are computed per
# image below so the aspect ratio is preserved (no stretch).
sample_size = [480, 832]
video_length = 81
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Camera trajectory (poses.npy / intrinsics.npy) + reference image + prompt.
action_path = "asset/lingbot_demo"
validation_image_start = "asset/lingbot_demo/image.jpg"
# prompts
prompt = "The video presents a soaring journey through a fantasy jungle. The wind whips past the rider's blue hands gripping the reins, causing the leather straps to vibrate. The ancient gothic castle approaches steadily, its stone details becoming clearer against the backdrop of floating islands and distant waterfalls."
# negative_prompt / guidance_scale mirror predict_i2v.py. The distilled few-step
# model is trained WITHOUT classifier-free guidance, so guidance_scale is left at
# 1.0 (CFG disabled, native behavior). Set it > 1.0 to enable CFG with the
# negative prompt (off-distribution, may degrade quality).
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 1.0
seed = 43
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/lingbot-world-i2v-fast"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference + camera control support.
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer_additional_kwargs['sink_size'] = sink_size
transformer_additional_kwargs['control_type'] = control_type
transformer_additional_kwargs['cross_attn_type'] = 'cross_attn'
transformer = WanTransformer3DModel_LingbotWorldFast.from_pretrained(
transformer_name,
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler_kwargs = OmegaConf.to_container(config['scheduler_kwargs'])
# The lingbot-world fast reference builds FlowUniPCMultistepScheduler with shift=1
# and applies the real shift only inside set_timesteps. The shared config carries
# shift=5.0 (used by other pipelines), which would double-shift the sigma grid
# here, so pin the constructor shift to 1 for the UniPC path; the runtime shift
# is forwarded to set_timesteps by the pipeline instead.
if Chosen_Scheduler is FlowUniPCMultistepScheduler:
scheduler_kwargs['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, scheduler_kwargs)
)
# Get Pipeline
pipeline = WanFunLingbotWorldFastPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
image = Image.open(validation_image_start).convert("RGB")
# Aspect-preserving resolution from the area budget (matches the reference):
# keep width*height close to max_area while snapping to the VAE stride *
# patch-size grid, so the input aspect ratio is preserved instead of forcing
# a fixed sample_size (which would stretch a 16:9 frame).
max_area = sample_size[0] * sample_size[1]
vae_stride = vae.config.spatial_compression_ratio
patch = pipeline.transformer.config.patch_size[1]
aspect_ratio = image.height / image.width
lat_h = round(np.sqrt(max_area * aspect_ratio) // vae_stride // patch * patch)
lat_w = round(np.sqrt(max_area / aspect_ratio) // vae_stride // patch * patch)
height = lat_h * vae_stride
width = lat_w * vae_stride
sample = pipeline(
prompt,
image = image,
negative_prompt = negative_prompt,
action_path = action_path,
control_type = control_type,
height = height,
width = width,
num_frames = video_length,
num_frame_per_block = num_frame_per_block,
num_inference_steps = num_inference_steps,
guidance_scale = guidance_scale,
stochastic_sampling = stochastic_sampling,
shift = shift,
generator = generator,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+37 -5
View File
@@ -17,6 +17,8 @@ from videox_fun.models import (AutoencoderKLLongCatVideo, UMT5EncoderModel, Auto
LongCatVideoTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import LongCatVideoPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
@@ -36,9 +38,21 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
@@ -69,11 +83,11 @@ prompt = "The dog is shaking head. The video is of high quality, an
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"
guidance_scale = 4.0
seed = 43
num_inference_steps = 50
num_inference_steps = 25
lora_weight = 0.55
save_path = "samples/longcat-videos-i2v"
device = set_multi_gpus_devices(1, 1)
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = LongCatVideoTransformer3DModel.from_pretrained(
os.path.join(model_name, "dit"),
@@ -141,6 +155,17 @@ pipeline = LongCatVideoPipeline(
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
@@ -149,8 +174,10 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -212,4 +239,9 @@ def save_results():
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
save_results()
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+293
View File
@@ -0,0 +1,293 @@
import os
import sys
from pathlib import Path
import numpy as np
import torch
from audio_separator.separator import Separator
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLLongCatVideo, AutoTokenizer,
LongCatVideoAudioEncoder,
LongCatVideoAvatarTransformer3DModel,
UMT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import LongCatVideoAvatarPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
merge_video_audio, save_videos_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LongCat-Video"
model_name_avatar = "models/Diffusion_Transformer/LongCat-Video-Avatar"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [832, 480]
video_length = 81
fps = 16
# Start Image
validation_image_start = "asset/8.png"
# Audio params
audio_path = "asset/talk.wav"
use_audio_vocal_separator = False
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Prompt
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."
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"
guidance_scale = 4.5
seed = 43
num_inference_steps = 25
lora_weight = 0.55
save_path = "samples/longcat-avatar-videos-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = LongCatVideoAvatarTransformer3DModel.from_pretrained(
os.path.join(model_name_avatar, "avatar_single"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype, cp_split_hw=[1, 1]
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLLongCatVideo.from_pretrained(
os.path.join(model_name, "vae"),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, "tokenizer"),
)
# Get Text encoder
text_encoder = UMT5EncoderModel.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Audio encoder (for avatar mode)
audio_encoder = LongCatVideoAudioEncoder(
os.path.join(model_name_avatar, 'chinese-wav2vec2-base')
)
audio_encoder.audio_encoder.feature_extractor._freeze_parameters()
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
# Get Pipeline
pipeline = LongCatVideoAvatarPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
audio_encoder=audio_encoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
# Get Vocal separator
if use_audio_vocal_separator:
vocal_separator_path = os.path.join(model_name_avatar, 'vocal_separator/Kim_Vocal_2.onnx')
audio_output_dir_temp = Path("./audio_temp_file")
audio_output_dir_temp.mkdir(parents=True, exist_ok=True)
vocal_separator = Separator(
output_dir=audio_output_dir_temp / "vocals",
output_single_stem="vocals",
model_file_dir=os.path.dirname(vocal_separator_path),
)
vocal_separator.load_model(os.path.basename(vocal_separator_path))
# Process audio if provided
audio_emb = None
if audio_path is not None:
# Extract vocal from audio
outputs = vocal_separator.separate(audio_path)
if len(outputs) > 0:
temp_vocal_path = audio_output_dir_temp / "vocals" / outputs[0]
temp_vocal_path = temp_vocal_path.resolve().as_posix()
audio_path = temp_vocal_path
with torch.no_grad():
video_length = int((video_length - 1) // vae.scale_factor_temporal * vae.scale_factor_temporal) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.scale_factor_temporal + 1
if validation_image_start is not None:
input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
else:
input_video, input_video_mask, clip_image = None, None, None
sample = pipeline(
prompt = prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
audio_path = audio_path,
video = input_video,
mask_video = input_video_mask,
fps = fps,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
merge_video_audio(video_path=video_path, audio_path=audio_path)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+38 -5
View File
@@ -17,6 +17,8 @@ from videox_fun.models import (AutoencoderKLLongCatVideo, UMT5EncoderModel, Auto
LongCatVideoTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import LongCatVideoPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
@@ -36,9 +38,21 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
@@ -67,11 +81,11 @@ prompt = "1girl, black_hair, brown_eyes, earrings, freckles, grey_b
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"
guidance_scale = 4.0
seed = 43
num_inference_steps = 50
num_inference_steps = 25
lora_weight = 0.55
save_path = "samples/longcat-videos-t2v"
device = set_multi_gpus_devices(1, 1)
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
transformer = LongCatVideoTransformer3DModel.from_pretrained(
os.path.join(model_name, "dit"),
@@ -140,6 +154,18 @@ pipeline = LongCatVideoPipeline(
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
@@ -147,8 +173,10 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -203,4 +231,9 @@ def save_results():
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
save_results()
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+305
View File
@@ -0,0 +1,305 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors, Gemma3Processor,
LTX2VideoTransformer3DModel, LTX2VocoderWithBWE)
from videox_fun.pipeline import LTX2I2VPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid,
save_videos_with_audio_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2.3-Diffusers"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [512, 768]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
validation_image_start = "asset/1.png"
# prompts
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. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
# CFG guidance scale for video and audio modality
guidance_scale = 3.0
audio_guidance_scale = 7.0
# Spatio-Temporal Guidance (STG) scale for video and audio
stg_scale = 1.0
audio_stg_scale = 1.0
# Modality isolation guidance scale for video and audio
modality_scale = 3.0
audio_modality_scale = 3.0
# Guidance rescale factor for video and audio to prevent overexposure
guidance_rescale = 0.7
audio_guidance_rescale = 0.7
spatio_temporal_guidance_blocks = [28]
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-i2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Processor
processor = Gemma3Processor.from_pretrained(
model_name,
subfolder="processor",
)
# Get Tokenizer
tokenizer = processor.tokenizer
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2VocoderWithBWE.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2I2VPipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=text_encoder.language_model.layers)
text_encoder = shard_fn(text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
image=Image.open(validation_image_start),
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
stg_scale=stg_scale,
modality_scale=modality_scale,
guidance_rescale=guidance_rescale,
audio_guidance_scale=audio_guidance_scale,
audio_stg_scale=audio_stg_scale,
audio_modality_scale=audio_modality_scale,
audio_guidance_rescale=audio_guidance_rescale,
spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+300
View File
@@ -0,0 +1,300 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors, Gemma3Processor,
LTX2VideoTransformer3DModel, LTX2VocoderWithBWE)
from videox_fun.pipeline import LTX2Pipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid,
save_videos_with_audio_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2.3-Diffusers"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [512, 768]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
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. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
# CFG guidance scale for video and audio modality
guidance_scale = 3.0
audio_guidance_scale = 7.0
# Spatio-Temporal Guidance (STG) scale for video and audio
stg_scale = 1.0
audio_stg_scale = 1.0
# Modality isolation guidance scale for video and audio
modality_scale = 3.0
audio_modality_scale = 3.0
# Guidance rescale factor for video and audio to prevent overexposure
guidance_rescale = 0.7
audio_guidance_rescale = 0.7
spatio_temporal_guidance_blocks = [28]
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-t2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Processor
processor = Gemma3Processor.from_pretrained(
model_name,
subfolder="processor",
)
# Get Tokenizer
tokenizer = processor.tokenizer
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2VocoderWithBWE.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2Pipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=text_encoder.language_model.layers)
text_encoder = shard_fn(text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
stg_scale=stg_scale,
modality_scale=modality_scale,
guidance_rescale=guidance_rescale,
audio_guidance_scale=audio_guidance_scale,
audio_stg_scale=audio_stg_scale,
audio_modality_scale=audio_modality_scale,
audio_guidance_rescale=audio_guidance_rescale,
spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+281
View File
@@ -0,0 +1,281 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors,
LTX2VideoTransformer3DModel, LTX2Vocoder)
from videox_fun.pipeline import LTX2I2VPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid,
save_videos_with_audio_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
validation_image_start = "asset/1.png"
# prompts
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. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-i2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Tokenizer
tokenizer = GemmaTokenizerFast.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2Vocoder.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2I2VPipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=text_encoder.language_model.layers)
text_encoder = shard_fn(text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
image=Image.open(validation_image_start),
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+326
View File
@@ -0,0 +1,326 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2LatentUpsamplerModel,
LTX2TextConnectors,
LTX2VideoTransformer3DModel, LTX2Vocoder)
from videox_fun.pipeline import LTX2I2VPipeline, LTX2LatentUpsamplePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid,
save_videos_with_audio_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
latent_upsampler_path = None
# Other params
sample_size = [480, 832]
video_length = 121
fps = 24
# Latent upsampler config
enable_latent_upsample = True
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# If you want to generate from text, please set the validation_image_start = None and validation_image_end = None
validation_image_start = "asset/1.png"
# prompts
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. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-i2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Tokenizer
tokenizer = GemmaTokenizerFast.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2Vocoder.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2I2VPipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=text_encoder.language_model.layers)
text_encoder = shard_fn(text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
image=Image.open(validation_image_start),
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
output_type="latent" if enable_latent_upsample else "pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if enable_latent_upsample:
# Load latent upsampler model
latent_upsampler = LTX2LatentUpsamplerModel.from_pretrained(
model_name, subfolder="latent_upsampler", torch_dtype=weight_dtype,
)
if latent_upsampler_path is not None:
print(f"From latent_upsampler checkpoint: {latent_upsampler_path}")
if latent_upsampler_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(latent_upsampler_path)
else:
state_dict = torch.load(latent_upsampler_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = latent_upsampler.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
upsample_pipeline = LTX2LatentUpsamplePipeline(
vae=pipeline.vae,
latent_upsampler=latent_upsampler,
)
upsample_pipeline.vae.enable_tiling()
upsample_pipeline.to(device=device, dtype=weight_dtype)
# output_type="latent" returns denormalized (raw) video latents [B, C, F, H, W]
# and raw audio latents [B, C, L, M]; decode audio manually
audio_latents = output.audio.to(device=device, dtype=pipeline.audio_vae.dtype)
mel = pipeline.audio_vae.decode(audio_latents, return_dict=False)[0]
audio = pipeline.vocoder(mel).cpu().float()
# Pass video latents directly to upsample pipeline (skip decode→re-encode roundtrip)
with torch.no_grad():
upsampled = upsample_pipeline(
latents=output.videos,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
output_type="pt",
return_dict=False,
)
sample = upsampled[0]
else:
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+276
View File
@@ -0,0 +1,276 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.models import (AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video,
Gemma3ForConditionalGeneration,
GemmaTokenizerFast, LTX2TextConnectors,
LTX2VideoTransformer3DModel, LTX2Vocoder)
from videox_fun.pipeline import LTX2Pipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid,
save_videos_with_audio_grid)
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = False
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/LTX-2"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# Load pretrained model if need
transformer_path = None
vae_path = None
lora_path = None
# Other params
sample_size = [512, 768]
video_length = 121
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
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. "
negative_prompt = "worst quality, inconsistent motion, blurry, jittery, distorted, static, low quality, artifacts"
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/ltx2-videos-t2v"
# Audio sample rate will be read from vocoder config
audio_sample_rate = 24000
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# Transformer
transformer = LTX2VideoTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
vae = AutoencoderKLLTX2Video.from_pretrained(
model_name,
subfolder="vae",
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE
audio_vae = AutoencoderKLLTX2Audio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=weight_dtype,
)
# Get Tokenizer
tokenizer = GemmaTokenizerFast.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text encoder
text_encoder = Gemma3ForConditionalGeneration.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Connectors
connectors = LTX2TextConnectors.from_pretrained(
model_name,
subfolder="connectors",
torch_dtype=weight_dtype,
)
# Vocoder
vocoder = LTX2Vocoder.from_pretrained(
model_name,
subfolder="vocoder",
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
pipeline = LTX2Pipeline(
scheduler=scheduler,
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
connectors=connectors,
transformer=transformer,
vocoder=vocoder,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(transformer.transformer_blocks))
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=text_encoder.language_model.layers)
text_encoder = shard_fn(text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
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)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
output = pipeline(
prompt=prompt,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.vocoder.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+306
View File
@@ -0,0 +1,306 @@
import os
import sys
import numpy as np
import torch
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
MiniMaxH3Transformer3DModel, Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
Qwen3VLProcessor)
from videox_fun.pipeline import MiniMaxH3Pipeline
from videox_fun.utils import (MiniMaxH3Scheduler, register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_with_audio_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus. The Qwen3-VL conditioner is ~62 GB, so with fsdp_dit alone every
# rank still replicates it; fsdp_text_encoder shards it too. Note it must wrap the inner `text_encoder.model`
# (Qwen3VLModel): encode_prompt calls that submodule directly, so a wrap on the top-level module would never fire.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# Load pretrained model if need
# A full finetune goes in `transformer_path`, either as the `transformer` folder a training checkpoint writes
# (`output_dir_minimax_h3/checkpoint-N/transformer`, config.json included) or as a single safetensors file. A LoRA
# goes in `lora_path`: handed to `transformer_path` it would match no key at all and load nothing. A PDD LoRA
# (parallel decoder) goes in `pdd_lora_path` and cannot be combined with `lora_path`.
transformer_path = None
vae_path = None
lora_path = None
pdd_lora_path = None
# Other params
# MiniMax-H3 generates at a fixed 24 fps, only accepts multiples of 32 as height / width, and snaps video_length up
# to the next 17 * n + 5 the video VAE can decode (the duration has to stay between 5 and 15 seconds).
# Leave sample_size as None to follow the aspect ratio of the first keyframe, which is what the model was released
# for; the first keyframe is stretched onto that canvas.
sample_size = [704, 1280]
video_length = 124
fps = 24
# The keyframe the video starts from, and the one it ends on. Either can be left as None: only an end frame generates
# *up to* that frame, and neither of them is a plain text-to-video request (see predict_t2v.py).
validation_image_start = "asset/1.png"
validation_image_end = None
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。"
seed = 43
# Number of denoising steps, i.e. of model evaluations: num_inference_steps = 40 runs 40 of them.
num_inference_steps = 40
# The released checkpoint is guidance-distilled: leave guidance_scale at 1 to run one forward pass per step
# with no CFG. A value above 1 enables classifier-free guidance with a negative_prompt, running two passes.
guidance_scale = 1
# The exponential sigma shifts of the two schedules. None keeps the ones of the checkpoint (12.0 video, 3.0 audio).
flow_shift = None
audio_flow_shift = None
lora_weight = 0.55
save_path = "samples/minimax-h3-videos-i2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# `model_name` may point either at a converted diffusers layout or at an *original* MiniMax-H3 partition (e.g.
# `MiniMax-H3/FL2VA`); the original shards are converted on the fly while loading, no intermediate copy on disk.
# Transformer
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if os.path.isdir(transformer_path):
# A training checkpoint's `transformer` folder carries its own config.json, so the loader restores the
# mixed-precision contract of the checkpoint (`_keep_in_fp32_modules`) by itself.
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
transformer_path,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# `strict=False` accepts a file whose keys belong to another model — a LoRA checkpoint, say — by loading
# nothing at all and silently generating with the base weights, so an unexpected key is a hard error.
assert len(u) == 0, (
f"{transformer_path} holds {len(u)} key(s) the transformer does not have, e.g. {u[:3]}. A LoRA "
"checkpoint belongs in `lora_path`, not `transformer_path`."
)
pdd_config = None
if pdd_lora_path is not None:
if lora_path is not None:
raise ValueError("`lora_path` and `pdd_lora_path` cannot be used together.")
from videox_fun.models.minimax_h3_pdd import (load_pdd_lora,
pdd_num_inference_steps,
pdd_step_callback)
pdd_config = load_pdd_lora(transformer, pdd_lora_path)
num_inference_steps = pdd_num_inference_steps(pdd_config, num_inference_steps, teacher_default=40)
# Video VAE. The released weights are float32 and the decode runs under float16 autocast, so the VAE is not
# downcast even when the rest of the pipeline is bfloat16 (this is also how the training scripts load it).
vae = AutoencoderKLMiniMaxH3.from_pretrained(
model_name,
subfolder="vae",
low_cpu_mem_usage=True,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE, waveform in / waveform out: MiniMax-H3 has no separate vocoder. Float32 as released, like the video VAE.
audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained(
model_name,
subfolder="audio_vae",
low_cpu_mem_usage=True,
)
# Get Tokenizer and Processor
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(model_name, "tokenizer"))
processor = Qwen3VLProcessor.from_pretrained(os.path.join(model_name, "processor"))
# Get Text encoder. MiniMax-H3 reads the unnormalized hidden state after the 50th decoder layer of Qwen3-VL.
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Schedulers. MiniMax-H3 steps the video and the audio latents down two schedules inside one transformer call.
scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="scheduler")
audio_scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="audio_scheduler")
pipeline = MiniMaxH3Pipeline(
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
transformer=transformer,
scheduler=scheduler,
audio_scheduler=audio_scheduler,
)
# The float32 modules of the mixed-precision checkpoint stay untouched by the float8 quantization.
fp8_exclude_module_name = [
"proj_in", "audio_proj_in", "context_embedder", "time_embedder", "time_proj",
"token_refiner", "norm_out", "proj_out", "audio_proj_out",
]
use_qfloat8 = "qfloat8" in GPU_memory_mode
if use_qfloat8:
convert_model_weight_to_float8(transformer, exclude_module_name=fp8_exclude_module_name, device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
fp32_modules = [m for m in transformer.modules()
if any(p.dtype == torch.float32 for p in m.parameters(recurse=False))]
shard_fn = partial(shard_model, device_id=device, param_dtype=None, cast_dtype=False,
module_to_wrapper=list(transformer.transformer_blocks),
ignored_modules=fp32_modules)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(text_encoder.model.language_model.layers))
pipeline.text_encoder.model = shard_fn(pipeline.text_encoder.model)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
image_start = None if validation_image_start is None else Image.open(validation_image_start)
image_end = None if validation_image_end is None else Image.open(validation_image_end)
pdd_callback = None if pdd_config is None else pdd_step_callback(
transformer, scheduler, audio_scheduler, pdd_config, num_inference_steps
)
with torch.no_grad():
output = pipeline(
prompt=prompt,
image=image_start,
last_image=image_end,
height=None if sample_size is None else sample_size[0],
width=None if sample_size is None else sample_size[1],
num_frames=video_length,
num_inference_steps=num_inference_steps,
flow_shift=flow_shift,
audio_flow_shift=audio_flow_shift,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
callback_on_step_end=pdd_callback,
)
print(f"[{os.environ.get('RANK', '0')}] generation done, decoding", flush=True)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
audio_sample_rate = output.sampling_rate
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=audio_sample_rate)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
# Keep every rank alive until the saving rank finishes; an early exit of one rank makes the elastic launcher
# terminate the others.
dist.barrier()
else:
save_results()
+328
View File
@@ -0,0 +1,328 @@
import os
import sys
import torch
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
MiniMaxH3Transformer3DModel,
Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
Qwen3VLProcessor)
from videox_fun.pipeline import (MiniMaxH3AudioReference,
MiniMaxH3ImageReference,
MiniMaxH3Pipeline,
MiniMaxH3VideoReference)
from videox_fun.utils import (MiniMaxH3Scheduler, register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_with_audio_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus. The Qwen3-VL conditioner is ~62 GB, so with fsdp_dit alone every
# rank still replicates it; fsdp_text_encoder shards it too. Note it must wrap the inner `text_encoder.model`
# (Qwen3VLModel): encode_prompt calls that submodule directly, so a wrap on the top-level module would never fire.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# Load pretrained model if need
# The `ref2va` weights ship in their own subfolder, same architecture as the base transformer. A full finetune goes
# in `transformer_path`, either as the `transformer` folder a training checkpoint writes (config.json included) or
# as a single safetensors file, overriding the `transformer_ref` subfolder. A LoRA goes in `lora_path`: handed to
# `transformer_path` it would match no key at all and load nothing. A PDD LoRA (parallel decoder) goes in
# `pdd_lora_path` and cannot be combined with `lora_path`; use a checkpoint trained with `--train_mode=ref2va`.
transformer_subfolder = "transformer_ref"
transformer_path = None
vae_path = None
lora_path = None
pdd_lora_path = None
# Other params
# MiniMax-H3 generates at a fixed 24 fps, only accepts multiples of 32 as height / width, and snaps video_length up
# to the next 17 * n + 5 the video VAE can decode (the duration has to stay between 5 and 15 seconds). References
# never bind the generated geometry: leaving height / width unset resolves MiniMax-H3's own 16:9 canvas.
sample_size = [1280, 704]
video_length = 124
fps = 24
# The references to condition on, **in the order the model should read them**: the order labels them in the prompt
# presentation and lays them out on the shared rotary clock. One entry per reference, `image=path`, `video=path` or
# `audio=path`; a video's own soundtrack is conditioned on with it. Budgets of the released checkpoint: at most 9
# images, 3 videos, 3 audios and 12 references in total, and an audio reference cannot stand alone.
references = [
"video=asset/ref2va_video.mp4",
"audio=asset/ref2va_audio.wav",
]
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "参考视频中的角色与场景,生成一段动作连贯、镜头流畅的续写视频,环境音与画面同步。"
seed = 43
# Number of denoising steps, i.e. of model evaluations: num_inference_steps = 50 runs 50 of them.
num_inference_steps = 50
# The released `ref2va` checkpoint is guidance-distilled with no unconditional branch, so `references` runs one
# forward pass per step and needs guidance_scale of 1 — the pipeline raises on anything above.
guidance_scale = 1.0
# The exponential sigma shifts of the two schedules. None keeps the ones of the checkpoint (12.0 video, 3.0 audio).
flow_shift = None
audio_flow_shift = None
lora_weight = 0.55
save_path = "samples/minimax-h3-videos-ref2va"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# `model_name` may point either at a converted diffusers layout or at an *original* MiniMax-H3 partition; the
# original shards are converted on the fly while loading, no intermediate copy on disk. The transformer comes from
# the `transformer_ref` subfolder — the released `ref2va` weights, same architecture as the base model.
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
model_name,
subfolder=transformer_subfolder,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if os.path.isdir(transformer_path):
# A training checkpoint's `transformer` folder carries its own config.json, so the loader restores the
# mixed-precision contract of the checkpoint (`_keep_in_fp32_modules`) by itself.
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
transformer_path,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# `strict=False` accepts a file whose keys belong to another model — a LoRA checkpoint, say — by loading
# nothing at all and silently generating with the base weights, so an unexpected key is a hard error.
assert len(u) == 0, (
f"{transformer_path} holds {len(u)} key(s) the transformer does not have, e.g. {u[:3]}. A LoRA "
"checkpoint belongs in `lora_path`, not `transformer_path`."
)
pdd_config = None
if pdd_lora_path is not None:
if lora_path is not None:
raise ValueError("`lora_path` and `pdd_lora_path` cannot be used together.")
from videox_fun.models.minimax_h3_pdd import (load_pdd_lora,
pdd_num_inference_steps,
pdd_step_callback)
pdd_config = load_pdd_lora(transformer, pdd_lora_path)
num_inference_steps = pdd_num_inference_steps(pdd_config, num_inference_steps, teacher_default=50)
# Video VAE. The released weights are float32 and the decode runs under float16 autocast, so the VAE is not
# downcast even when the rest of the pipeline is bfloat16 (this is also how the training scripts load it).
vae = AutoencoderKLMiniMaxH3.from_pretrained(
model_name,
subfolder="vae",
low_cpu_mem_usage=True,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE, waveform in / waveform out: MiniMax-H3 has no separate vocoder. Float32 as released, like the video VAE.
audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained(
model_name,
subfolder="audio_vae",
low_cpu_mem_usage=True,
)
# Get Tokenizer and Processor
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(model_name, "tokenizer"))
processor = Qwen3VLProcessor.from_pretrained(os.path.join(model_name, "processor"))
# Get Text encoder. MiniMax-H3 reads the unnormalized hidden state after the 50th decoder layer of Qwen3-VL.
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Schedulers. MiniMax-H3 steps the video and the audio latents down two schedules inside one transformer call.
scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="scheduler")
audio_scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="audio_scheduler")
pipeline = MiniMaxH3Pipeline(
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
transformer=transformer,
scheduler=scheduler,
audio_scheduler=audio_scheduler,
)
# The float32 modules of the mixed-precision checkpoint stay untouched by the float8 quantization.
fp8_exclude_module_name = [
"proj_in", "audio_proj_in", "context_embedder", "time_embedder", "time_proj",
"token_refiner", "norm_out", "proj_out", "audio_proj_out",
]
use_qfloat8 = "qfloat8" in GPU_memory_mode
if use_qfloat8:
convert_model_weight_to_float8(transformer, exclude_module_name=fp8_exclude_module_name, device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
fp32_modules = [m for m in transformer.modules()
if any(p.dtype == torch.float32 for p in m.parameters(recurse=False))]
shard_fn = partial(shard_model, device_id=device, param_dtype=None, cast_dtype=False,
module_to_wrapper=list(transformer.transformer_blocks),
ignored_modules=fp32_modules)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(text_encoder.model.language_model.layers))
pipeline.text_encoder.model = shard_fn(pipeline.text_encoder.model)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def parse_reference(entry: str):
kind, _, media = entry.partition("=")
kind, media = kind.strip().lower(), media.strip()
if not media:
raise ValueError(f"A reference entry must be `image=path`, `video=path` or `audio=path`, got {entry!r}.")
if kind == "image":
return MiniMaxH3ImageReference.from_file(media)
if kind == "video":
return MiniMaxH3VideoReference.from_file(media)
if kind == "audio":
return MiniMaxH3AudioReference.from_file(media)
raise ValueError(f"A reference entry must start with `image=`, `video=` or `audio=`, got {entry!r}.")
# Decode every reference at the rate its container carries, which the pipeline's setup resamples onto MiniMax-H3's
# own 24 fps and the audio VAE's sample rate.
parsed_references = [parse_reference(entry) for entry in references]
pdd_callback = None if pdd_config is None else pdd_step_callback(
transformer, scheduler, audio_scheduler, pdd_config, num_inference_steps
)
with torch.no_grad():
output = pipeline(
prompt=prompt,
references=parsed_references,
height=None if sample_size is None else sample_size[0],
width=None if sample_size is None else sample_size[1],
num_frames=video_length,
num_inference_steps=num_inference_steps,
flow_shift=flow_shift,
audio_flow_shift=audio_flow_shift,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
callback_on_step_end=pdd_callback,
)
print(f"[{os.environ.get('RANK', '0')}] generation done, decoding", flush=True)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
audio_sample_rate = output.sampling_rate
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=audio_sample_rate)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
# Keep every rank alive until the saving rank finishes; an early exit of one rank makes the elastic launcher
# terminate the others.
dist.barrier()
else:
save_results()
+294
View File
@@ -0,0 +1,294 @@
import os
import sys
import numpy as np
import torch
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
MiniMaxH3Transformer3DModel, Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
Qwen3VLProcessor)
from videox_fun.pipeline import MiniMaxH3Pipeline
from videox_fun.utils import (MiniMaxH3Scheduler, register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_with_audio_grid
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus. The Qwen3-VL conditioner is ~62 GB, so with fsdp_dit alone every
# rank still replicates it; fsdp_text_encoder shards it too. Note it must wrap the inner `text_encoder.model`
# (Qwen3VLModel): encode_prompt calls that submodule directly, so a wrap on the top-level module would never fire.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# Load pretrained model if need
# A full finetune goes in `transformer_path`, either as the `transformer` folder a training checkpoint writes
# (`output_dir_minimax_h3/checkpoint-N/transformer`, config.json included) or as a single safetensors file. A LoRA
# goes in `lora_path`: handed to `transformer_path` it would match no key at all and load nothing. A PDD LoRA
# (parallel decoder) goes in `pdd_lora_path` and cannot be combined with `lora_path`.
transformer_path = None
vae_path = None
lora_path = None
pdd_lora_path = None
# Other params
# MiniMax-H3 generates at a fixed 24 fps, only accepts multiples of 32 as height / width, and snaps video_length up
# to the next 17 * n + 5 the video VAE can decode (the duration has to stay between 5 and 15 seconds).
# Leave sample_size as None to use MiniMax-H3's own 16:9 canvas (768x1344).
sample_size = [704, 1280]
video_length = 124
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A red fox trotting through a snowy pine forest, snow crunching underfoot"
seed = 43
# Number of denoising steps, i.e. of model evaluations: num_inference_steps = 40 runs 40 of them.
num_inference_steps = 40
# The released checkpoint is guidance-distilled: leave guidance_scale at 1 to run one forward pass per step
# with no CFG. A value above 1 enables classifier-free guidance with a negative_prompt, running two passes.
guidance_scale = 1
# The exponential sigma shifts of the two schedules. None keeps the ones of the checkpoint (12.0 video, 3.0 audio).
flow_shift = None
audio_flow_shift = None
lora_weight = 0.55
save_path = "samples/minimax-h3-videos-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# `model_name` may point either at a converted diffusers layout or at an *original* MiniMax-H3 partition (e.g.
# `MiniMax-H3/FL2VA`); the original shards are converted on the fly while loading, no intermediate copy on disk.
# Transformer
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if os.path.isdir(transformer_path):
# A training checkpoint's `transformer` folder carries its own config.json, so the loader restores the
# mixed-precision contract of the checkpoint (`_keep_in_fp32_modules`) by itself.
transformer = MiniMaxH3Transformer3DModel.from_pretrained(
transformer_path,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# `strict=False` accepts a file whose keys belong to another model — a LoRA checkpoint, say — by loading
# nothing at all and silently generating with the base weights, so an unexpected key is a hard error.
assert len(u) == 0, (
f"{transformer_path} holds {len(u)} key(s) the transformer does not have, e.g. {u[:3]}. A LoRA "
"checkpoint belongs in `lora_path`, not `transformer_path`."
)
pdd_config = None
if pdd_lora_path is not None:
if lora_path is not None:
raise ValueError("`lora_path` and `pdd_lora_path` cannot be used together.")
from videox_fun.models.minimax_h3_pdd import (load_pdd_lora,
pdd_num_inference_steps,
pdd_step_callback)
pdd_config = load_pdd_lora(transformer, pdd_lora_path)
num_inference_steps = pdd_num_inference_steps(pdd_config, num_inference_steps, teacher_default=40)
# Video VAE. The released weights are float32 and the decode runs under float16 autocast, so the VAE is not
# downcast even when the rest of the pipeline is bfloat16 (this is also how the training scripts load it).
vae = AutoencoderKLMiniMaxH3.from_pretrained(
model_name,
subfolder="vae",
low_cpu_mem_usage=True,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE, waveform in / waveform out: MiniMax-H3 has no separate vocoder. Float32 as released, like the video VAE.
audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained(
model_name,
subfolder="audio_vae",
low_cpu_mem_usage=True,
)
# Get Tokenizer and Processor
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(model_name, "tokenizer"))
processor = Qwen3VLProcessor.from_pretrained(os.path.join(model_name, "processor"))
# Get Text encoder. MiniMax-H3 reads the unnormalized hidden state after the 50th decoder layer of Qwen3-VL.
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Schedulers. MiniMax-H3 steps the video and the audio latents down two schedules inside one transformer call.
scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="scheduler")
audio_scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="audio_scheduler")
pipeline = MiniMaxH3Pipeline(
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
transformer=transformer,
scheduler=scheduler,
audio_scheduler=audio_scheduler,
)
# The float32 modules of the mixed-precision checkpoint stay untouched by the float8 quantization.
fp8_exclude_module_name = [
"proj_in", "audio_proj_in", "context_embedder", "time_embedder", "time_proj",
"token_refiner", "norm_out", "proj_out", "audio_proj_out",
]
use_qfloat8 = "qfloat8" in GPU_memory_mode
if use_qfloat8:
convert_model_weight_to_float8(transformer, exclude_module_name=fp8_exclude_module_name, device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
fp32_modules = [m for m in transformer.modules()
if any(p.dtype == torch.float32 for p in m.parameters(recurse=False))]
shard_fn = partial(shard_model, device_id=device, param_dtype=None, cast_dtype=False,
module_to_wrapper=list(transformer.transformer_blocks),
ignored_modules=fp32_modules)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(text_encoder.model.language_model.layers))
pipeline.text_encoder.model = shard_fn(pipeline.text_encoder.model)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
pdd_callback = None if pdd_config is None else pdd_step_callback(
transformer, scheduler, audio_scheduler, pdd_config, num_inference_steps
)
with torch.no_grad():
output = pipeline(
prompt=prompt,
height=None if sample_size is None else sample_size[0],
width=None if sample_size is None else sample_size[1],
num_frames=video_length,
num_inference_steps=num_inference_steps,
flow_shift=flow_shift,
audio_flow_shift=audio_flow_shift,
guidance_scale=guidance_scale,
generator=generator,
output_type="pt",
callback_on_step_end=pdd_callback,
)
print(f"[{os.environ.get('RANK', '0')}] generation done, decoding", flush=True)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
audio_sample_rate = output.sampling_rate
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=audio_sample_rate)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
# Keep every rank alive until the saving rank finishes; an early exit of one rank makes the elastic launcher
# terminate the others.
dist.barrier()
else:
save_results()
@@ -0,0 +1,341 @@
import os
import sys
import numpy as np
import torch
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
MiniMaxH3ControlTransformer3DModel,
Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
Qwen3VLProcessor)
from videox_fun.pipeline import MiniMaxH3ControlPipeline
from videox_fun.utils import (MiniMaxH3Scheduler, register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (get_video_to_video_latent,
save_videos_with_audio_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_group_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
# Multi-GPU runs through the xfuser sequence-parallel path and must be launched with torchrun, e.g.
# `torchrun --nproc_per_node=2 examples/minimax_h3_fun/predict_v2v_control.py` for ulysses_degree=2, ring_degree=1.
# It is incompatible with the *cpu_offload* memory modes (accelerate offload hooks own a single device);
# use model_full_load / model_full_load_and_qfloat8 there, with fsdp_dit to save memory.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus. The Qwen3-VL conditioner is ~62 GB, so with fsdp_dit alone every
# rank still replicates it; fsdp_text_encoder shards it too. Note it must wrap the inner `text_encoder.model`
# (Qwen3VLModel): encode_prompt calls that submodule directly, so a wrap on the top-level module would never fire.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# Control branch layout, must match the yaml `train_control.py` ran with: `control_blocks_places` selects the
# layers the control blocks attach to and `control_in_dim` the channels the control rows carry (49 for an
# `--enable_inpaint` checkpoint, whose `control_proj_in` is widened with the mask channels). Leaving it None
# builds the default 24-channel branch, which cannot load an inpaint checkpoint.
config_path = "config/minimax_h3/minimax_h3_control.yaml"
# Load pretrained model if need. The control branch is not part of the released MiniMax-H3 weights, so a base
# `model_name` starts the side branch as an identity (`after_proj` is zero) and the c ontrol video has no effect;
# point `transformer_path` at a control checkpoint trained by `scripts/minimax_h3_fun/train_control.py`.
transformer_path = "models/Diffusion_Transformer/MiniMax-H3-Fun-Controlnet-Union/MiniMax-H3-Fun-Controlnet-Union.safetensors"
vae_path = None
lora_path = None
# Other params
# MiniMax-H3 generates at a fixed 24 fps, only accepts multiples of 32 as height / width, and the generation
# follows the control video's actual length — snapped down to the largest 17 * n + 5 the video VAE can decode so
# a short control video is never padded (the duration has to stay under 15 seconds), capped by video_length.
# Control inference fits the control video onto this canvas with the training's resize + crop geometry, so
# sample_size must be set (it cannot be None).
sample_size = [1280, 704]
video_length = 243
fps = 24
# Scale applied to every control skip before it is added to the main branch. 0.0 switches the control branch off,
# values below 1.0 weaken the guidance of the control video.
control_context_scale = 1.00
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
control_video = "asset/pose.mp4"
# Inpaint inputs, only read by checkpoints trained with `--enable_inpaint` (control_in_dim widened, e.g. 49):
# `inpaint_video` is the source video behind the mask and `inpaint_video_mask` marks the regions to regenerate
# (white = repaint, black = keep). With an inpaint checkpoint but no inpaint inputs given, the pipeline zero-pads
# the mask channels and the run degrades to pure generation; a mask-less checkpoint rejects them outright.
inpaint_video = None
inpaint_video_mask = None
prompt = "视频中,一位年轻女性站在阳光洒满的沙滩上,背景是无垠碧蓝的大海与澄澈如洗的天空,构成一幅充满夏日度假氛围的画面。她身穿一件深海军蓝吊带泳衣,线条简约贴身,凸显健康匀称的身材曲线;外搭一条纯白色背带短裙,裙摆轻盈飘逸,随风微微扬起,增添了几分俏皮与少女感。她的长发柔顺披肩,发梢微卷,在阳光下泛着自然光泽,耳畔垂挂着一对小巧精致的珍珠吊坠耳环,为整体造型注入一丝温柔优雅的气息。她面带甜美笑容,嘴角上扬,露出整齐洁白的牙齿,眼神清澈明亮,直视镜头时流露出真诚与自信,仿佛在与观众分享此刻的快乐。起初,她双臂向两侧张开,手掌舒展,像是在拥抱整个大海与天空;随后手臂缓缓收回并向前挥动,动作节奏轻快而富有韵律,如同在跳舞或做简单的热身操,展现出轻松自在、无忧无虑的状态。她的腿部微微分开站立,姿态稳健又不失灵动,裙摆随着动作轻轻摇曳,与海风形成自然互动。远处海浪轻拍沙滩,发出柔和的“哗哗”声,虽无声但可想象其韵律,与她的动作相得益彰,营造出宁静而愉悦的听觉联想。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
seed = 43
# Number of denoising steps, i.e. of model evaluations: num_inference_steps = 40 runs 40 of them.
num_inference_steps = 40
# The released checkpoint is guidance-distilled: leave guidance_scale at 1 to run one forward pass per step
# with no CFG — the distill checkpoints of train_control_distill.py already bake the teacher's CFG target into
# the weights, so any value above 1 applies guidance twice and degrades the output. A value above 1 enables
# classifier-free guidance with a negative_prompt, running two passes.
guidance_scale = 1.0
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# The exponential sigma shifts of the two schedules. None keeps the ones of the checkpoint (12.0 video, 3.0 audio).
flow_shift = None
audio_flow_shift = None
lora_weight = 0.55
save_path = "samples/minimax-h3-videos-v2v-control"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# The yaml pins the control branch layout exactly as in training (scripts/minimax_h3_fun/train_control.py), where
# `transformer_additional_kwargs` is spread into `from_pretrained` the same way.
transformer_load_kwargs = {}
if config_path is not None:
from omegaconf import OmegaConf
config = OmegaConf.load(config_path)
transformer_load_kwargs.update(
OmegaConf.to_container(config["transformer_additional_kwargs"], resolve=True)
)
# `model_name` may point either at a converted diffusers layout or at an *original* MiniMax-H3 partition (e.g.
# `MiniMax-H3/FL2VA`); the original shards are converted on the fly while loading, no intermediate copy on disk.
# Transformer. `from_pretrained` fills the control branch the released checkpoint does not carry: every control
# block is initialised from the main block it is attached to and `control_proj_in` from `proj_in`, with
# before_proj / after_proj zeroed, so a freshly loaded model is numerically identical to the base MiniMax-H3 model.
transformer = MiniMaxH3ControlTransformer3DModel.from_pretrained(
model_name,
subfolder="transformer",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
**transformer_load_kwargs,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE. The released weights are float32 and the decode runs under float16 autocast, so the VAE is not
# downcast even when the rest of the pipeline is bfloat16.
vae = AutoencoderKLMiniMaxH3.from_pretrained(
model_name,
subfolder="vae",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio VAE, waveform in / waveform out: MiniMax-H3 has no separate vocoder.
audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained(
model_name,
subfolder="audio_vae",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Tokenizer and Processor
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(model_name, "tokenizer"))
processor = Qwen3VLProcessor.from_pretrained(os.path.join(model_name, "processor"))
# Get Text encoder. MiniMax-H3 reads the unnormalized hidden state after the 50th decoder layer of Qwen3-VL.
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(model_name, "text_encoder"),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Schedulers. MiniMax-H3 steps the video and the audio latents down two schedules inside one transformer call.
scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="scheduler")
audio_scheduler = MiniMaxH3Scheduler.from_pretrained(model_name, subfolder="audio_scheduler")
pipeline = MiniMaxH3ControlPipeline(
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
processor=processor,
transformer=transformer,
scheduler=scheduler,
audio_scheduler=audio_scheduler,
)
# The float32 modules of the mixed-precision checkpoint stay untouched by the float8 quantization. The `proj_in`
# entry also covers the control patch projection `control_proj_in`, which shares the video patch projection's dtype.
fp8_exclude_module_name = [
"proj_in", "audio_proj_in", "context_embedder", "time_embedder", "time_proj",
"token_refiner", "norm_out", "proj_out", "audio_proj_out",
]
use_qfloat8 = "qfloat8" in GPU_memory_mode
if use_qfloat8:
convert_model_weight_to_float8(transformer, exclude_module_name=fp8_exclude_module_name, device=device)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
fp32_modules = [m for m in transformer.modules()
if any(p.dtype == torch.float32 for p in m.parameters(recurse=False))]
shard_fn = partial(shard_model, device_id=device, param_dtype=None, cast_dtype=False,
module_to_wrapper=list(transformer.transformer_blocks) + list(transformer.control_blocks),
ignored_modules=fp32_modules)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype,
module_to_wrapper=list(text_encoder.model.language_model.layers))
pipeline.text_encoder.model = shard_fn(pipeline.text_encoder.model)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.transformer_blocks)):
pipeline.transformer.transformer_blocks[i] = torch.compile(pipeline.transformer.transformer_blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def snap_num_frames(actual_num_frames, max_num_frames):
"""
Pick the generation length from the control video instead of padding a short one: the largest `17 * n + 5`
the video VAE can decode that does not exceed the frames actually read (capped by `max_num_frames`), snapping
down so no tail frame is ever repeated. A control video below 5 frames is raised to 5, the smallest count
the video VAE can encode.
"""
num_frames = min(actual_num_frames, max_num_frames)
num_frames = (num_frames - 5) // 17 * 17 + 5
return max(num_frames, 5)
with torch.no_grad():
control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None, keep_aspect_ratio=True)
# Generate at the control video's actual length, never padding; only control videos below the 5 frames the
# video VAE can encode are raised to 5.
num_frames = snap_num_frames(control_video.shape[2], video_length)
if num_frames != video_length:
print(f"[{os.environ.get('RANK', '0')}] control video holds {control_video.shape[2]} frames, generating "
f"{num_frames} instead of {video_length}", flush=True)
mask_video = None
if inpaint_video is not None:
if inpaint_video_mask is None:
raise ValueError("inpaint_video_mask is required when inpaint_video is provided")
inpaint_video, _, _, _ = get_video_to_video_latent(inpaint_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None, keep_aspect_ratio=True)
inpaint_video_mask, _, _, _ = get_video_to_video_latent(inpaint_video_mask, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None, keep_aspect_ratio=True)
# Binarize the grayscale mask onto one channel: 1 marks the regions to regenerate, mirroring the training
# `get_random_mask` convention the visibility map `1 - mask` is built from.
mask_video = (inpaint_video_mask[:, :1] > 0.5).to(inpaint_video_mask.dtype)
output = pipeline(
prompt=prompt,
control_video=control_video,
control_context_scale=control_context_scale,
mask_video=mask_video,
inpaint_video=inpaint_video,
height=None if sample_size is None else sample_size[0],
width=None if sample_size is None else sample_size[1],
num_frames=num_frames,
num_inference_steps=num_inference_steps,
flow_shift=flow_shift,
audio_flow_shift=audio_flow_shift,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt,
generator=generator,
output_type="pt",
)
print(f"[{os.environ.get('RANK', '0')}] generation done, decoding", flush=True)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
sample = output.videos
audio = output.audio
audio_sample_rate = output.sampling_rate
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=audio_sample_rate)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
# Keep every rank alive until the saving rank finishes; an early exit of one rank makes the elastic launcher
# terminate the others.
dist.barrier()
else:
save_results()
+380
View File
@@ -0,0 +1,380 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLMOVAAudio, AutoencoderKLWan,
AutoTokenizer, MOVADualTowerConditionalBridge,
UMT5EncoderModel, WanAudioTransformer3DModel,
WanTransformer3DModel)
from videox_fun.pipeline import MOVAPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import save_videos_with_audio_grid
# 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].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with sequential_cpu_offload.
compile_dit = False
# model path
model_name = "models/Diffusion_Transformer/MOVA-360p"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
boundary_ratio = 0.9
# Load pretrained model if need
# The transformer_path is used for low noise model, the transformer_high_path is used for high noise model.
transformer_path = None
transformer_high_path = None
transformer_audio_path = None
bridge_path = None
vae_path = None
audio_vae_path = None
# Load lora model if need
# The lora_path is used for low noise model, the lora_high_path is used for high noise model.
lora_path = None
lora_high_path = None
# Other params
sample_size = [640, 352]
video_length = 81
fps = 24
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
# Input image for I2V
validation_image = "asset/8.png"
# prompts
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."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指"
guidance_scale = 5.0
seed = 43
num_inference_steps = 50
# The lora_weight is used for low noise model, the lora_high_weight is used for high noise model.
lora_weight = 0.55
lora_high_weight = 0.55
save_path = "samples/mova-videos-i2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
# The from_pretrained method automatically converts WanModel config to WanTransformer3DModel config
print("Loading Video DiT (High Noise) with WanTransformer3DModel...")
transformer = WanTransformer3DModel.from_pretrained(
model_name,
subfolder="video_dit_2",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video DiT 2 (Low Noise) - Using WanTransformer3DModel
print("Loading Video DiT 2 (Low Noise) with WanTransformer3DModel...")
transformer_2 = WanTransformer3DModel.from_pretrained(
model_name,
subfolder="video_dit",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Audio DiT - Using WanAudioTransformer3DModel
print("Loading Audio DiT with WanAudioTransformer3DModel...")
transformer_audio = WanAudioTransformer3DModel.from_pretrained(
model_name,
subfolder="audio_dit",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_audio_path is not None:
print(f"From checkpoint: {transformer_audio_path}")
if transformer_audio_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(transformer_audio_path)
else:
state_dict = torch.load(transformer_audio_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_audio.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Dual Tower Bridge
print("Loading Dual Tower Bridge...")
dual_tower_bridge = MOVADualTowerConditionalBridge.from_pretrained(
model_name,
subfolder="dual_tower_bridge",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if bridge_path is not None:
print(f"From checkpoint: {bridge_path}")
if bridge_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(bridge_path)
else:
state_dict = torch.load(bridge_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = dual_tower_bridge.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Video VAE
print("Loading Video VAE...")
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, "video_vae/diffusion_pytorch_model.safetensors")
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
audio_vae = AutoencoderKLMOVAAudio.from_pretrained(
model_name,
subfolder="audio_vae",
torch_dtype=torch.float32,
)
if audio_vae_path is not None:
print(f"From checkpoint: {audio_vae_path}")
if audio_vae_path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(audio_vae_path)
else:
state_dict = torch.load(audio_vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = audio_vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
print("Loading Tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
model_name,
subfolder="tokenizer",
)
# Get Text Encoder
print("Loading Text Encoder...")
text_encoder = UMT5EncoderModel.from_pretrained(
model_name,
subfolder="text_encoder",
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# Get Scheduler
print("Loading Scheduler...")
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler.from_pretrained(
model_name,
subfolder="scheduler"
)
# Build Pipeline
print("Building MOVAPipeline Pipeline...")
pipeline = MOVAPipeline(
vae=vae,
audio_vae=audio_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
scheduler=scheduler,
transformer=transformer,
transformer_2=transformer_2,
transformer_audio=transformer_audio,
dual_tower_bridge=dual_tower_bridge,
audio_vae_type="dac",
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
# Enable multi-GPU inference for visual transformers
transformer.enable_multi_gpus_inference()
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
# Apply FSDP to visual transformer blocks
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
# Compile MOVAModel blocks
# NOTE: compile_dit is not compatible with fsdp_dit
if fsdp_dit:
print("WARNING: compile_dit is not compatible with fsdp_dit. Disabling compile.")
else:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
for i in range(len(pipeline.transformer_audio.blocks)):
pipeline.transformer_audio.blocks[i] = torch.compile(pipeline.transformer_audio.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(pipeline.transformer, ["modulation",], device=device)
replace_parameters_by_name(pipeline.transformer_2, ["modulation",], device=device)
pipeline.transformer.freqs = pipeline.transformer.freqs.to(device=device)
pipeline.transformer_2.freqs = pipeline.transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(pipeline.transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(pipeline.transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
convert_weight_dtype_wrapper(pipeline.transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(pipeline.transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(pipeline.transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype)
convert_weight_dtype_wrapper(pipeline.transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
# Run inference
print("Running inference...")
with torch.no_grad():
image = Image.open(validation_image).convert("RGB")
output = pipeline(
prompt=prompt,
image=image,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
num_frames=video_length,
frame_rate=fps,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
boundary=boundary_ratio,
)
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
sample = output.videos
audio = output.audio
# Get audio sample rate from pipeline
audio_sample_rate = pipeline.audio_sample_rate
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
sr = getattr(pipeline.audio_vae.config, "output_sampling_rate", audio_sample_rate)
save_videos_with_audio_grid(sample, audio, video_path, fps=fps, audio_sample_rate=sr)
if ulysses_degree > 1 or ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+9 -1
View File
@@ -18,6 +18,8 @@ from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunPhantomPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
@@ -27,7 +29,7 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -38,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -217,6 +222,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+10 -1
View File
@@ -20,12 +20,14 @@ from videox_fun.pipeline import Wan2_2I2VPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +38,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -238,6 +243,10 @@ if GPU_memory_mode == "sequential_cpu_offload":
transformer.freqs = transformer.freqs.to(device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(transformer)
register_auto_device_hook(transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
+10 -1
View File
@@ -23,10 +23,12 @@ from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -37,6 +39,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -192,7 +197,11 @@ if compile_dit:
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+17 -8
View File
@@ -15,18 +15,21 @@ for project_root in project_roots:
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanI2VPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -37,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -218,6 +224,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -238,7 +247,7 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
if cfg_skip_ratio is not None:
if cfg_skip_ratio is not None and cfg_skip_ratio > 0:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+319
View File
@@ -0,0 +1,319 @@
import os
import sys
import argparse
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
from transformers import AutoTokenizer
# 添加项目根目录到系统路径
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanI2VPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid, timer_record)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
def parse_args():
# 解析命令行参数
parser = argparse.ArgumentParser()
# 基础模型与推理参数
parser.add_argument("--GPU_memory_mode", type=str, default="sequential_cpu_offload")
parser.add_argument("--ulysses_degree", type=int, default=1)
parser.add_argument("--ring_degree", type=int, default=1)
parser.add_argument("--fsdp_dit", action="store_true")
parser.add_argument("--fsdp_text_encoder", action="store_true")
parser.add_argument("--compile_dit", action="store_true")
# TeaCache 参数
parser.add_argument("--enable_teacache", action="store_true")
parser.add_argument("--teacache_threshold", type=float, default=0.10)
parser.add_argument("--num_skip_start_steps", type=int, default=5)
parser.add_argument("--teacache_offload", action="store_true")
# CFG Skip 参数
parser.add_argument("--cfg_skip_ratio", type=float, default=0.0)
# Riflex 参数
parser.add_argument("--enable_riflex", action="store_true")
parser.add_argument("--riflex_k", type=int, default=6)
# 模型路径和配置
parser.add_argument("--config_path", type=str, default="config/wan2.1/wan_civitai.yaml")
parser.add_argument("--model_name", type=str, default="models/Diffusion_Transformer/Wan2.1-I2V-14B-480P")
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--vae_path", type=str, default=None)
parser.add_argument("--lora_path", type=str, default=None)
parser.add_argument("--sample_size", nargs='+', type=int, default=[480, 832])
parser.add_argument("--video_length", type=int, default=81)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--weight_dtype", type=str, default="bfloat16")
# 输入图像相关
parser.add_argument("--validation_image_start", type=str, default="asset/1.png")
parser.add_argument("--validation_image_end", type=str, default=None)
# 推理参数
parser.add_argument("--prompt", type=str, default="一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。")
parser.add_argument("--negative_prompt", type=str, default="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走")
parser.add_argument("--guidance_scale", type=float, default=6.0)
parser.add_argument("--seed", type=int, default=43)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--lora_weight", type=float, default=0.55)
parser.add_argument("--save_path", type=str, default="samples/wan-videos-i2v")
# 采样器设置
parser.add_argument("--sampler_name", type=str, choices=["Flow", "Flow_Unipc", "Flow_DPM++"], default="Flow_Unipc")
parser.add_argument("--shift", type=float, default=3.0)
return parser.parse_args()
args = parse_args()
# 获取参数
GPU_memory_mode = args.GPU_memory_mode
ulysses_degree = args.ulysses_degree
ring_degree = args.ring_degree
fsdp_dit = args.fsdp_dit
fsdp_text_encoder = args.fsdp_text_encoder
compile_dit = args.compile_dit
enable_teacache = args.enable_teacache
teacache_threshold = args.teacache_threshold
num_skip_start_steps = args.num_skip_start_steps
teacache_offload = args.teacache_offload
cfg_skip_ratio = args.cfg_skip_ratio
enable_riflex = args.enable_riflex
riflex_k = args.riflex_k
config_path = args.config_path
model_name = args.model_name
transformer_path = args.transformer_path
vae_path = args.vae_path
lora_path = args.lora_path
sample_size = args.sample_size
video_length = args.video_length
fps = args.fps
weight_dtype = torch.bfloat16 if args.weight_dtype == "bfloat16" else torch.float16
prompt = args.prompt
negative_prompt = args.negative_prompt
guidance_scale = args.guidance_scale
seed = args.seed
num_inference_steps = args.num_inference_steps
lora_weight = args.lora_weight
save_path = args.save_path
sampler_name = args.sampler_name
shift = args.shift
validation_image_start = args.validation_image_start
validation_image_end = args.validation_image_end
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# 初始化模型组件
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# 获取VAE
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# 获取分词器
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# 获取文本编码器
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
text_encoder = text_encoder.eval()
# 获取CLIP图像编码器
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')),
).to(weight_dtype)
clip_image_encoder = clip_image_encoder.eval()
# 获取调度器
scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}
Choosen_Scheduler = scheduler_dict[sampler_name]
if sampler_name in ["Flow_Unipc", "Flow_DPM++"]:
config['scheduler_kwargs']['shift'] = 1
scheduler = Choosen_Scheduler(
**filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# 创建Pipeline
pipeline = WanI2VPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
clip_image_encoder=clip_image_encoder
)
# 分布式设置
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
# 编译优化
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
# 内存优化策略
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
for i in range(2):
# TeaCache配置
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
# CFG跳过配置
if cfg_skip_ratio is not None and cfg_skip_ratio > 0:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
# 随机种子
generator = torch.Generator(device=device).manual_seed(seed)
# LoRA加载
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
# 生成视频
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
if enable_riflex:
pipeline.transformer.enable_riflex(k=riflex_k, L_test=latent_frames)
# 输入图像处理
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)
# 执行推理
sample = pipeline(
prompt,
num_frames=video_length,
negative_prompt=negative_prompt,
height=sample_size[0],
width=sample_size[1],
generator=generator,
guidance_scale=guidance_scale,
num_inference_steps=num_inference_steps,
video=input_video,
mask_video=input_video_mask,
clip_image=clip_image,
shift=shift,
).videos
# LoRA卸载
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
# 保存结果
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
# 分布式保存
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+53
View File
@@ -0,0 +1,53 @@
export EXCEL_FILE="./speed.xlsx"
export DIT_EXCEL_COL=0 VAE_EXCEL_COL=1 TOTAL_EXCEL_COL=2
# 14B 720P
export DIT_EXCEL_ROW=1 VAE_EXCEL_ROW=1 TOTAL_EXCEL_ROW=1
python examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --compile_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=2 VAE_EXCEL_ROW=2 TOTAL_EXCEL_ROW=2
torchrun --nproc-per-node=2 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=3 VAE_EXCEL_ROW=3 TOTAL_EXCEL_ROW=3
torchrun --nproc-per-node=4 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=4 VAE_EXCEL_ROW=4 TOTAL_EXCEL_ROW=4
torchrun --nproc-per-node=8 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=8 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
# 14B 480P
export DIT_EXCEL_ROW=5 VAE_EXCEL_ROW=5 TOTAL_EXCEL_ROW=5
python examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --compile_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=6 VAE_EXCEL_ROW=6 TOTAL_EXCEL_ROW=6
torchrun --nproc-per-node=2 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=7 VAE_EXCEL_ROW=7 TOTAL_EXCEL_ROW=7
torchrun --nproc-per-node=4 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=8 VAE_EXCEL_ROW=8 TOTAL_EXCEL_ROW=8
torchrun --nproc-per-node=8 examples/wan2.1/predict_i2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" \
--GPU_memory_mode="model_full_load_and_qfloat8" --ulysses_degree=8 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.30 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
+18 -9
View File
@@ -13,19 +13,22 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, WanT5EncoderModel, AutoTokenizer,
WanTransformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +39,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -205,6 +211,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -225,7 +234,7 @@ if coefficients is not None:
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
if cfg_skip_ratio is not None:
if cfg_skip_ratio is not None and cfg_skip_ratio > 0:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
+303
View File
@@ -0,0 +1,303 @@
import os
import sys
import numpy as np
import torch
import argparse
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, WanT5EncoderModel, AutoTokenizer,
WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid, timer_record)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
def parse_args():
parser = argparse.ArgumentParser(description="Video Generation with Wan2.1-Fun")
parser.add_argument("--GPU_memory_mode", type=str, default="sequential_cpu_offload",
choices=["model_full_load", "model_full_load_and_qfloat8", "model_cpu_offload",
"model_cpu_offload_and_qfloat8", "sequential_cpu_offload"],
help="GPU memory optimization mode.")
parser.add_argument("--ulysses_degree", type=int, default=1,
help="Ulysses parallelism degree.")
parser.add_argument("--ring_degree", type=int, default=1,
help="Ring parallelism degree.")
parser.add_argument("--fsdp_dit", action="store_true",
help="Use FSDP for transformer to save GPU memory.")
parser.add_argument("--fsdp_text_encoder", action="store_true",
help="Use FSDP for text encoder to save GPU memory.")
parser.add_argument("--compile_dit", action="store_true",
help="Compile transformer for fixed resolution speedup.")
parser.add_argument("--enable_teacache", action="store_true",
help="Enable TeaCache optimization.")
parser.add_argument("--teacache_threshold", type=float, default=0.10,
help="TeaCache threshold for step caching.")
parser.add_argument("--num_skip_start_steps", type=int, default=5,
help="Number of steps to skip TeaCache at inference start.")
parser.add_argument("--teacache_offload", action="store_true",
help="Offload TeaCache tensors to CPU.")
parser.add_argument("--cfg_skip_ratio", type=float, default=0.0,
help="CFG skip ratio for inference.")
parser.add_argument("--enable_riflex", action="store_true",
help="Enable Riflex frequency optimization.")
parser.add_argument("--riflex_k", type=int, default=6,
help="Intrinsic frequency index for Riflex.")
parser.add_argument("--config_path", type=str, default="config/wan2.1/wan_civitai.yaml",
help="Path to model config file.")
parser.add_argument("--model_name", type=str, default="models/Diffusion_Transformer/Wan2.1-T2V-1.3B",
help="Path to model directory.")
parser.add_argument("--sampler_name", type=str, default="Flow_Unipc",
choices=["Flow", "Flow_Unipc", "Flow_DPM++"],
help="Sampler type for video generation.")
parser.add_argument("--shift", type=float, default=3.0,
help="Noise schedule shift parameter for Flow_Unipc/Flow_DPM++.")
parser.add_argument("--transformer_path", type=str, default=None,
help="Path to pre-trained transformer checkpoint.")
parser.add_argument("--vae_path", type=str, default=None,
help="Path to pre-trained VAE checkpoint.")
parser.add_argument("--lora_path", type=str, default=None,
help="Path to LoRA weights.")
parser.add_argument("--sample_size", nargs=2, type=int, default=[480, 832],
help="Sample size [height, width].")
parser.add_argument("--video_length", type=int, default=81,
help="Number of frames in the video.")
parser.add_argument("--fps", type=int, default=16,
help="Frames per second for output video.")
parser.add_argument("--weight_dtype", type=str, default="bfloat16",
choices=["float16", "bfloat16"],
help="Weight data type (float16 or bfloat16).")
parser.add_argument("--prompt", type=str, default="一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。",
help="Text prompt for video generation.")
parser.add_argument("--negative_prompt", type=str, default="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
help="Negative prompt for video generation.")
parser.add_argument("--guidance_scale", type=float, default=6.0,
help="Classifier-free guidance scale.")
parser.add_argument("--seed", type=int, default=43,
help="Random seed for reproducibility.")
parser.add_argument("--num_inference_steps", type=int, default=50,
help="Number of inference steps.")
parser.add_argument("--lora_weight", type=float, default=0.55,
help="LoRA weight scaling factor.")
parser.add_argument("--save_path", type=str, default="samples/wan-videos-t2v",
help="Directory to save generated videos.")
return parser.parse_args()
args = parse_args()
# 将 argparse 参数映射到原有变量
GPU_memory_mode = args.GPU_memory_mode
ulysses_degree = args.ulysses_degree
ring_degree = args.ring_degree
fsdp_dit = args.fsdp_dit
fsdp_text_encoder = args.fsdp_text_encoder
compile_dit = args.compile_dit
enable_teacache = args.enable_teacache
teacache_threshold = args.teacache_threshold
num_skip_start_steps = args.num_skip_start_steps
teacache_offload = args.teacache_offload
cfg_skip_ratio = args.cfg_skip_ratio
enable_riflex = args.enable_riflex
riflex_k = args.riflex_k
config_path = args.config_path
model_name = args.model_name
sampler_name = args.sampler_name
shift = args.shift
transformer_path = args.transformer_path
vae_path = args.vae_path
lora_path = args.lora_path
sample_size = args.sample_size
video_length = args.video_length
fps = args.fps
weight_dtype = torch.bfloat16 if args.weight_dtype == "bfloat16" else torch.float16
prompt = args.prompt
negative_prompt = args.negative_prompt
guidance_scale = args.guidance_scale
seed = args.seed
num_inference_steps = args.num_inference_steps
lora_weight = args.lora_weight
save_path = args.save_path
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
transformer = WanTransformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True if not fsdp_dit else False,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Choosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Choosen_Scheduler(
**filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
for i in range(2):
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
if cfg_skip_ratio is not None and cfg_skip_ratio > 0:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+103
View File
@@ -0,0 +1,103 @@
export EXCEL_FILE="./speed.xlsx"
export DIT_EXCEL_COL=0 VAE_EXCEL_COL=1 TOTAL_EXCEL_COL=2
# 1.3B 720P
export DIT_EXCEL_ROW=1 VAE_EXCEL_ROW=1 TOTAL_EXCEL_ROW=1
python examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=2 VAE_EXCEL_ROW=2 TOTAL_EXCEL_ROW=2
torchrun --nproc-per-node=2 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=3 VAE_EXCEL_ROW=3 TOTAL_EXCEL_ROW=3
torchrun --nproc-per-node=4 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=4 VAE_EXCEL_ROW=4 TOTAL_EXCEL_ROW=4
torchrun --nproc-per-node=8 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=2 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
# 1.3B 480P
export DIT_EXCEL_ROW=5 VAE_EXCEL_ROW=5 TOTAL_EXCEL_ROW=5
python examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=6 VAE_EXCEL_ROW=6 TOTAL_EXCEL_ROW=6
torchrun --nproc-per-node=2 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=7 VAE_EXCEL_ROW=7 TOTAL_EXCEL_ROW=7
torchrun --nproc-per-node=4 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=8 VAE_EXCEL_ROW=8 TOTAL_EXCEL_ROW=8
torchrun --nproc-per-node=8 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-1.3B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=2 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
# 14B 720P
export DIT_EXCEL_ROW=9 VAE_EXCEL_ROW=9 TOTAL_EXCEL_ROW=9
python examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=10 VAE_EXCEL_ROW=10 TOTAL_EXCEL_ROW=10
torchrun --nproc-per-node=2 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=11 VAE_EXCEL_ROW=11 TOTAL_EXCEL_ROW=11
torchrun --nproc-per-node=4 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=12 VAE_EXCEL_ROW=12 TOTAL_EXCEL_ROW=12
torchrun --nproc-per-node=8 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=8 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
# 14B 480P
export DIT_EXCEL_ROW=13 VAE_EXCEL_ROW=13 TOTAL_EXCEL_ROW=13
python examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=14 VAE_EXCEL_ROW=14 TOTAL_EXCEL_ROW=14
torchrun --nproc-per-node=2 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=15 VAE_EXCEL_ROW=15 TOTAL_EXCEL_ROW=15
torchrun --nproc-per-node=4 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=16 VAE_EXCEL_ROW=16 TOTAL_EXCEL_ROW=16
torchrun --nproc-per-node=8 examples/wan2.1/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.1-T2V-14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=8 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.15 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
@@ -0,0 +1,328 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++".
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5.0
# `stochastic_sampling`: False: ar
# True : ccd and dmd
stochastic_sampling = True
# Causal-Forcing checkpoint to overlay on top of the Wan2.1 base model.
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-2000/diffusion_pytorch_model.safetensors"
use_ema = False
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Causal-Forcing causal inference config
# `num_frame_per_block`: 3 = chunk-wise:
# 1 = frame-wise:
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = -1
# Others
independent_first_frame = False
context_noise = 0.0
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# -------- Stage selector (uncomment ONE block) --------
# All few-step distilled stages (2, 3) bake CFG into the student weights, so
# inference MUST use guidance_scale=1.0 — CF's CausalInferencePipeline does
# zero CFG (grep "unconditional/cfg/guidance" in pipeline/causal_inference.py
# returns 0). Using gs>1 stacks CFG on top of a CFG-baked model and produces
# over-saturated, AR-unstable outputs.
#
# Stage 1 — AR diffusion (`ar_diffusion.pt`): 50-step UniPC + CFG.
# guidance_scale = 3.0
# num_inference_steps = 50
#
# Stage 2 — CCD (`causal_cd.pt`): 4-step consistency-distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
#
# Stage 3 — DMD (`causal_forcing.pt`): 4-step distribution-matching distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
guidance_scale = 1.0
num_inference_steps = 4
seed = 43
lora_weight = 0.55
save_path = "samples/wan-videos-causal-forcing"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
def _resolve_transformer_path(raw_path: str, prefer_ema: bool) -> str:
"""Resolve a checkpoint path to the actual weights file to load.
File paths are returned unchanged so external `.pt` ckpts (CF official) keep
working. For trainer-output dirs, prefer EMA when asked and available,
otherwise fall back to live `transformer/` weights.
"""
if os.path.isfile(raw_path):
return raw_path
if os.path.isdir(raw_path):
candidates = []
if prefer_ema:
candidates.append(os.path.join(raw_path, "ema_transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "diffusion_pytorch_model.safetensors"))
for c in candidates:
if os.path.isfile(c):
return c
raise FileNotFoundError(
f"transformer_path={raw_path!r} is neither a file nor a checkpoint dir "
f"with a known safetensors layout (transformer/ or ema_transformer/)."
)
_raw_transformer_path = transformer_path
transformer_path = _resolve_transformer_path(transformer_path, prefer_ema=use_ema)
if transformer_path != _raw_transformer_path:
print(f"use_ema={use_ema}: resolved {_raw_transformer_path} -> {transformer_path}")
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
# Causal-Forcing's FSDP-saved ckpts (causal_cd.pt / causal_forcing.pt) keep the
# `model._fsdp_wrapped_module.` prefix; strip it before the generic `model.` strip
# so both kinds of ckpt land at bare parameter names.
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
@@ -0,0 +1,373 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid, StreamVideoSaver,
SegmentVideoSaver)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++".
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5.0
# `stochastic_sampling`: False: ar
# True : ccd and dmd
stochastic_sampling = True
# Causal-Forcing checkpoint to overlay on top of the Wan2.1 base model.
transformer_path = "output_dir_wan2.1_causal_forcing_dmd/checkpoint-2000/diffusion_pytorch_model.safetensors"
use_ema = False
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Causal-Forcing causal inference config
# `num_frame_per_block`: 3 = chunk-wise:
# 1 = frame-wise:
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = -1
# Others
independent_first_frame = False
context_noise = 0.0
# Streaming decode: decode and write each causal block to disk as soon as it is
# generated ("generate a chunk, save a chunk"), instead of holding all latents
# and decoding once at the end. The VAE causal cache is threaded across blocks,
# so the saved video is seam-free and identical to a full decode.
streaming = True
# Streaming save mode (only used when streaming=True):
# "stream" : append every block into one continuous mp4 (finalized on close).
# "segments" : save each decoded block as its own standalone mp4, flushed
# immediately after decoding ("decode a chunk -> save it now").
# Robust to interruption; segments can be concatenated later.
save_mode = "segments"
# Use torch.float16 if GPU does not support torch.bfloat16
# Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# -------- Stage selector (uncomment ONE block) --------
# All few-step distilled stages (2, 3) bake CFG into the student weights, so
# inference MUST use guidance_scale=1.0 — CF's CausalInferencePipeline does
# zero CFG (grep "unconditional/cfg/guidance" in pipeline/causal_inference.py
# returns 0). Using gs>1 stacks CFG on top of a CFG-baked model and produces
# over-saturated, AR-unstable outputs.
#
# Stage 1 — AR diffusion (`ar_diffusion.pt`): 50-step UniPC + CFG.
# guidance_scale = 3.0
# num_inference_steps = 50
#
# Stage 2 — CCD (`causal_cd.pt`): 4-step consistency-distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
#
# Stage 3 — DMD (`causal_forcing.pt`): 4-step distribution-matching distilled.
# guidance_scale = 1.0
# num_inference_steps = 4
guidance_scale = 1.0
num_inference_steps = 4
seed = 43
lora_weight = 0.55
save_path = "samples/wan-videos-causal-forcing"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
def _resolve_transformer_path(raw_path: str, prefer_ema: bool) -> str:
"""Resolve a checkpoint path to the actual weights file to load.
File paths are returned unchanged so external `.pt` ckpts (CF official) keep
working. For trainer-output dirs, prefer EMA when asked and available,
otherwise fall back to live `transformer/` weights.
"""
if os.path.isfile(raw_path):
return raw_path
if os.path.isdir(raw_path):
candidates = []
if prefer_ema:
candidates.append(os.path.join(raw_path, "ema_transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "transformer", "diffusion_pytorch_model.safetensors"))
candidates.append(os.path.join(raw_path, "diffusion_pytorch_model.safetensors"))
for c in candidates:
if os.path.isfile(c):
return c
raise FileNotFoundError(
f"transformer_path={raw_path!r} is neither a file nor a checkpoint dir "
f"with a known safetensors layout (transformer/ or ema_transformer/)."
)
_raw_transformer_path = transformer_path
transformer_path = _resolve_transformer_path(transformer_path, prefer_ema=use_ema)
if transformer_path != _raw_transformer_path:
print(f"use_ema={use_ema}: resolved {_raw_transformer_path} -> {transformer_path}")
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
# Causal-Forcing's FSDP-saved ckpts (causal_cd.pt / causal_forcing.pt) keep the
# `model._fsdp_wrapped_module.` prefix; strip it before the generic `model.` strip
# so both kinds of ckpt land at bare parameter names.
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
# Only the main process (rank 0, or single-GPU) writes files to disk.
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
is_main_process = dist.get_rank() == 0
else:
is_main_process = True
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
# In streaming mode, set up an incremental writer up-front so each block
# can be flushed to disk right after it is decoded.
saver = None
decode_callback = None
if streaming:
if is_main_process:
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if save_mode == "segments":
# One standalone mp4 per decoded block under a per-prompt dir.
saver = SegmentVideoSaver(os.path.join(save_path, prefix), fps)
else:
# One continuous mp4, appended block by block.
saver = StreamVideoSaver(os.path.join(save_path, prefix + ".mp4"), fps)
decode_callback = saver
else:
# Non-main ranks still decode locally (matching non-streaming
# behaviour) but must not write; a no-op avoids accumulation.
decode_callback = lambda video_chunk, block_idx: None
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
streaming = streaming,
decode_callback = decode_callback,
).videos
if saver is not None:
saver.close()
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
# Streaming already wrote the video incrementally above; only the
# non-streaming path needs the one-shot save below.
if not streaming and is_main_process:
save_results()
+16 -7
View File
@@ -15,18 +15,21 @@ for project_root in project_roots:
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
WanTransformer3DModel)
WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunInpaintPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -37,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -219,6 +225,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+16 -7
View File
@@ -15,18 +15,21 @@ for project_root in project_roots:
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel,
WanTransformer3DModel)
WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -37,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -226,6 +232,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+12 -4
View File
@@ -13,12 +13,16 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunControlPipeline, WanPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
@@ -26,10 +30,8 @@ from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +42,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -229,6 +234,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -13,23 +13,26 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunControlPipeline, WanPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +43,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -229,6 +235,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+14 -5
View File
@@ -13,23 +13,26 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel,
WanT5EncoderModel, WanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanFunControlPipeline, WanPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +43,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -229,6 +235,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+274
View File
@@ -0,0 +1,274 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5
stochastic_sampling = True
# Load pretrained model if need
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Self-Forcing causal inference config
# Number of frames to generate per block (1 for standard causal, higher for faster but more memory)
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = -1
# Others
independent_first_frame = False
context_noise = 0.0
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 1.0
seed = 43
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/wan-videos-self-forcing-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
@@ -0,0 +1,337 @@
import os
import sys
import time
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "model_full_load"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
# [NOTE]: Forcing-KV currently supports single GPU only (ulysses_degree = ring_degree = 1).
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5
stochastic_sampling = True
# Load pretrained model if need
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Self-Forcing causal inference config
# Number of frames to generate per block (1 for standard causal, higher for faster but more memory)
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
# For Forcing-KV this caps the rolling KV buffer (memory); grouped heads only
# read sink+~1 hist+current, so 6 (=sink1+hist1+fpb3+margin) suffices.
local_attn_size = 6
# Sink frames always kept at the start of the rolling KV cache
# (official forcing-kv config uses sink_size=1 as the stability anchor)
sink_size = 1
# Others
independent_first_frame = False
context_noise = 0.0
# Forcing-KV (arXiv 2605.09681) hybrid KV cache compression.
# Training-free per-head KV eviction on top of the rolling KV cache,
# aligned with the official zju-jiyicheng/Forcing-KV architecture.
# 1. Run profile_forcing_kv_heads.py first (with the SAME local_attn_size /
# sink_size / num_frame_per_block) to produce the head profile JSON
# (official {"layers": [{"layer_idx", "static_head", "dynamic_head"}]} format).
# 2. Set forcing_kv_enable = True and forcing_kv_head_profile to that JSON.
# When disabled, the original full-window attention path runs bit-identically.
forcing_kv_enable = True
forcing_kv_head_profile = "asset/forcing_kv_head_profile.json"
# Official forcingkv config (defaults match configs/forcing-kv/*.yaml)
forcing_kv_ar_start = 1 # AR step after which compression activates
forcing_kv_spatial_context_length = 1 # history frames for static heads
forcing_kv_temporal_context_length = 1 # recent frames for dynamic heads
forcing_kv_dynamic_context_length = 1 # compressed cache capacity (frames)
forcing_kv_num_frame_patch = 6 # token segments per latent frame
forcing_kv_sim_retention_ratio = 0.33 # fraction of candidate segments kept
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompts = [
"A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
]
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 1.0
seed = 43
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/wan-videos-self-forcing-forcing-kv-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer_additional_kwargs['sink_size'] = sink_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
if forcing_kv_enable:
assert forcing_kv_head_profile is not None, \
"forcing_kv_enable=True requires forcing_kv_head_profile (run profile_forcing_kv_heads.py first)."
print(f"[Forcing-KV] enabled, profile={forcing_kv_head_profile}, "
f"ar_start={forcing_kv_ar_start}, spatial_ctx={forcing_kv_spatial_context_length}, "
f"temporal_ctx={forcing_kv_temporal_context_length}, "
f"dynamic_ctx={forcing_kv_dynamic_context_length}, "
f"num_frame_patch={forcing_kv_num_frame_patch}, "
f"sim_retention={forcing_kv_sim_retention_ratio}")
else:
print("[Forcing-KV] disabled (full-window attention baseline)")
for prompt in prompts:
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
torch.cuda.synchronize()
start_time = time.time()
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
forcing_kv_enable = forcing_kv_enable if forcing_kv_enable else None,
forcing_kv_head_profile = forcing_kv_head_profile,
forcing_kv_ar_start = forcing_kv_ar_start,
forcing_kv_spatial_context_length = forcing_kv_spatial_context_length,
forcing_kv_temporal_context_length = forcing_kv_temporal_context_length,
forcing_kv_dynamic_context_length = forcing_kv_dynamic_context_length,
forcing_kv_num_frame_patch = forcing_kv_num_frame_patch,
forcing_kv_sim_retention_ratio = forcing_kv_sim_retention_ratio,
).videos
torch.cuda.synchronize()
elapsed = time.time() - start_time
print(f"[Timing] {video_length} frames ({latent_frames} latent) in {elapsed:.2f}s "
f"({video_length / elapsed:.2f} frames/s)")
if getattr(pipeline, "kv_cache_pos", None) is not None:
kv_tokens = pipeline.kv_cache_pos[0]["k"].shape[1]
kv_mib = sum(c["k"].numel() + c["v"].numel()
for c in pipeline.kv_cache_pos + pipeline.kv_cache_neg) \
* pipeline.kv_cache_pos[0]["k"].element_size() / (1024 ** 2)
print(f"[KV cache] {kv_tokens} tokens per layer per branch, total {kv_mib:.1f} MiB (pos+neg, all layers)")
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
@@ -0,0 +1,319 @@
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid, StreamVideoSaver,
SegmentVideoSaver)
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory.
#
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
# Multi GPUs config
# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used.
# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4.
# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1.
ulysses_degree = 1
ring_degree = 1
# Use FSDP to save more GPU memory in multi gpus.
fsdp_dit = False
fsdp_text_encoder = True
# Compile will give a speedup in fixed resolution and need a little GPU memory.
# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload.
compile_dit = False
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++"
sampler_name = "Flow"
# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics.
# Used when the sampler is in "Flow_Unipc", "Flow_DPM++".
shift = 5
stochastic_sampling = True
# Load pretrained model if need
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
vae_path = None
lora_path = None
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
# Self-Forcing causal inference config
# Number of frames to generate per block (1 for standard causal, higher for faster but more memory)
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = -1
# Others
independent_first_frame = False
context_noise = 0.0
# Streaming decode: decode and write each causal block to disk as soon as it is
# generated ("generate a chunk, save a chunk"), instead of holding all latents
# and decoding once at the end. The VAE causal cache is threaded across blocks,
# so the saved video is seam-free and identical to a full decode.
streaming = True
# Streaming save mode (only used when streaming=True):
# "stream" : append every block into one continuous mp4 (finalized on close).
# "segments" : save each decoded block as its own standalone mp4, flushed
# immediately after decoding ("decode a chunk -> save it now").
# Robust to interruption; segments can be concatenated later.
save_mode = "segments"
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
weight_dtype = torch.bfloat16
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale = 1.0
seed = 43
num_inference_steps = 4
lora_weight = 0.55
save_path = "samples/wan-videos-self-forcing-t2v"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Chosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Chosen_Scheduler(
**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
# Only the main process (rank 0, or single-GPU) writes files to disk.
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
is_main_process = dist.get_rank() == 0
else:
is_main_process = True
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
# In streaming mode, set up an incremental writer up-front so each block
# can be flushed to disk right after it is decoded.
saver = None
decode_callback = None
if streaming:
if is_main_process:
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if save_mode == "segments":
# One standalone mp4 per decoded block under a per-prompt dir.
saver = SegmentVideoSaver(os.path.join(save_path, prefix), fps)
else:
# One continuous mp4, appended block by block.
saver = StreamVideoSaver(os.path.join(save_path, prefix + ".mp4"), fps)
decode_callback = saver
else:
# Non-main ranks still decode locally (matching non-streaming
# behaviour) but must not write; a no-op avoids accumulation.
decode_callback = lambda video_chunk, block_idx: None
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = stochastic_sampling,
streaming = streaming,
decode_callback = decode_callback,
).videos
if saver is not None:
saver.close()
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
# Streaming already wrote the video incrementally above; only the
# non-streaming path needs the one-shot save below.
if not streaming and is_main_process:
save_results()
@@ -0,0 +1,263 @@
# Offline head profiling for Forcing-KV (arXiv 2605.09681).
#
# Runs the Self-Forcing causal pipeline on one prompt while forward hooks on
# every CasualWanSelfAttention accumulate per-head attention mass by region
# (sink / last-K frames / distant history). A head is classified Static when
# the last-K frames hold >= threshold of the post-sink attention mass (the
# official simplified Eq. 1 criterion from zju-jiyicheng/Forcing-KV
# configs_head/head_profile.py: THRESHOLD=0.8, LAST_K=4, skip sink frames).
# The result is written to forcing_kv_head_profile.json in the official
# {"format": "forcingkv_offline", "layers": [...]} format consumed by
# predict_t2v_forcing_kv.py via forcing_kv_head_profile.
#
# [NOTE]: profile with the SAME local_attn_size / sink_size /
# num_frame_per_block you intend to use at inference, since region boundaries
# depend on them. Single GPU only.
import json
import math
import os
import sys
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel,
WanTransformer3DModel_SelfForcing)
from videox_fun.pipeline import WanSelfForcingPipeline
from videox_fun.utils.utils import filter_kwargs
# Config and model path
config_path = "config/wan2.1/wan_civitai.yaml"
# model path
model_name = "models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
# Load pretrained model if need
transformer_path = "models/Diffusion_Transformer/Self-Forcing/checkpoints/self_forcing_dmd.pt"
# Self-Forcing causal inference config (MUST match the target inference setup)
# Number of frames to generate per block
num_frame_per_block = 3
# Local attention window size (-1 for global attention)
local_attn_size = 6
# Sink frames always kept at the start of the rolling KV cache
sink_size = 1
# Others
independent_first_frame = False
context_noise = 0.0
# Profiling config
# Official criterion (configs_head/head_profile.py): a head is Static when
# last_k frames / post-sink total attention mass >= threshold (default 0.8).
threshold = 0.8
last_k = 4
# Frames to generate while profiling (more frames = better stats, slower)
video_length = 81
sample_size = [480, 832]
shift = 5
guidance_scale = 1.0
num_inference_steps = 4
seed = 43
prompt = "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."
# Output head profile JSON path
output_path = "asset/forcing_kv_head_profile.json"
# Use torch.float16 if GPU does not support torch.bfloat16
weight_dtype = torch.bfloat16
device = set_multi_gpus_devices(1, 1)
config = OmegaConf.load(config_path)
# Load transformer with causal inference support if enabled
transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs'])
transformer_additional_kwargs['local_attn_size'] = local_attn_size
transformer_additional_kwargs['sink_size'] = sink_size
transformer = WanTransformer3DModel_SelfForcing.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=transformer_additional_kwargs,
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
state_dict = state_dict["generator_ema"] if "generator_ema" in state_dict else state_dict
state_dict = state_dict["generator"] if "generator" in state_dict else state_dict
if any("._fsdp_wrapped_module." in k for k in state_dict.keys()):
state_dict = {k.replace("model._fsdp_wrapped_module.", "model.", 1) if k.startswith("model._fsdp_wrapped_module.") else k: v for k, v in state_dict.items()}
if any(k.startswith("model.") for k in state_dict.keys()):
state_dict = {k.replace("model.", "", 1) if k.startswith("model.") else k: v for k, v in state_dict.items()}
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = WanSelfForcingPipeline(
transformer=transformer,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
pipeline.to(device=device)
# Profiling hooks: after every cached self-attn forward, the layer exposes
# kv_cache["_fkv_last_q"] / "_fkv_window_start" / "_fkv_local_end". Recompute
# softmax attention mass in fp32 (row-chunked) and partition it by region.
layer_stats = []
hooks = []
ROW_CHUNK = 128
for layer_idx, block in enumerate(pipeline.transformer.blocks):
attn = block.self_attn
acc = {
"sink": torch.zeros(attn.num_heads, dtype=torch.float64),
"last_k": torch.zeros(attn.num_heads, dtype=torch.float64),
"distant": torch.zeros(attn.num_heads, dtype=torch.float64),
"total": 0.0,
"seen_starts": set(),
}
layer_stats.append(acc)
def make_hook(acc):
def hook(module, inputs, output):
kv_cache = inputs[5] if len(inputs) > 5 else None
if kv_cache is None or "_fkv_last_q" not in kv_cache:
return
current_start = int(inputs[6])
# Only profile the first forward per chunk position (denoise step 0);
# later steps attend over the identical window with noisier keys.
if current_start in acc["seen_starts"]:
return
acc["seen_starts"].add(current_start)
q = kv_cache["_fkv_last_q"] # [B, s, n, d]
window_start = int(kv_cache["_fkv_window_start"])
local_end = int(kv_cache["_fkv_local_end"])
grid_sizes = inputs[2]
frame_seqlen = int(math.prod(grid_sizes[0][1:]))
k_win = kv_cache["k"][:, window_start:local_end] # [B, L, n, d]
sink_tokens = module.sink_size * frame_seqlen
local_start = local_end - q.shape[1]
# Regions cover HISTORY only ([0, local_start)) for sink clamping;
# clamping by local_start avoids overlap with the current chunk on
# early chunks, which would double-count attention mass (score > 1).
sink_end = min(sink_tokens, local_start)
# Last-K frames (inclusive of the current chunk), mirroring the
# official LAST_K criterion; clamped against the sink region.
last_k_start = max(sink_end, local_end - last_k * frame_seqlen)
rel_sink = max(0, sink_end - window_start)
rel_last = max(rel_sink, last_k_start - window_start)
qf = q[0].float() # [s, n, d]
kf = k_win[0].float() # [L, n, d]
scale = kf.shape[-1] ** -0.5
for r0 in range(0, qf.shape[0], ROW_CHUNK):
qc = qf[r0:r0 + ROW_CHUNK] # [c, n, d]
probs = torch.einsum(
"cnd,lnd->ncl", qc, kf).mul_(scale).softmax(dim=-1)
acc["sink"] += probs[:, :, :rel_sink].sum(dim=(1, 2)).double().cpu()
acc["last_k"] += probs[:, :, rel_last:].sum(dim=(1, 2)).double().cpu()
acc["distant"] += probs[:, :, rel_sink:rel_last].sum(dim=(1, 2)).double().cpu()
acc["total"] += float(qc.shape[0])
return hook
hooks.append(attn.register_forward_hook(make_hook(acc)))
generator = torch.Generator(device=device).manual_seed(seed)
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
pipeline(
prompt,
num_frames = video_length,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
shift = shift,
num_frame_per_block = num_frame_per_block,
independent_first_frame = independent_first_frame,
context_noise = context_noise,
stochastic_sampling = True,
output_type = "latent",
)
for h in hooks:
h.remove()
# Classify heads (official criterion: last-K / post-sink >= threshold) and
# dump the profile JSON in the official forcingkv_offline format.
layers_out = []
num_static = 0
print(f"\n{'layer':>5} {'static heads':<40} {'mean score':>10}")
for layer_idx, acc in enumerate(layer_stats):
denom = (torch.clamp(torch.tensor(acc["total"]), min=1e-8) - acc["sink"]).clamp(min=1e-8)
score = acc["last_k"] / denom
static = [h for h in range(score.shape[0]) if score[h].item() >= threshold]
dynamic = [h for h in range(score.shape[0]) if h not in set(static)]
num_heads = score.shape[0]
num_static += len(static)
layers_out.append({
"layer_idx": layer_idx,
"static_head": static,
"dynamic_head": dynamic,
})
print(f"{layer_idx:>5} {str(static):<40} {score.mean().item():>10.4f}")
profile = {
"format": "forcingkv_offline",
"num_layers": len(layer_stats),
"num_heads": num_heads,
"layers": layers_out,
}
with open(output_path, "w") as f:
json.dump(profile, f, indent=2)
total_layers = len(layer_stats)
print(f"\nthreshold={threshold}, last_k={last_k}: {num_static}/{total_layers * num_heads} heads static "
f"({num_static / max(1, total_layers * num_heads) * 100:.1f}%), "
f"dynamic {100 - num_static / max(1, total_layers * num_heads) * 100:.1f}%")
print(f"Profile saved to: {output_path}")
+16 -7
View File
@@ -13,23 +13,26 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel, VaceWanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
VaceWanTransformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanVacePipeline, WanPipeline
from videox_fun.pipeline import WanPipeline, WanVacePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +43,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -220,6 +226,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+16 -7
View File
@@ -13,23 +13,26 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel, VaceWanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
VaceWanTransformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanVacePipeline, WanPipeline
from videox_fun.pipeline import WanPipeline, WanVacePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +43,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -220,6 +226,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+16 -7
View File
@@ -13,23 +13,26 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.data import process_pose_file
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
WanT5EncoderModel, VaceWanTransformer3DModel)
from videox_fun.data.dataset_image_video import process_pose_file
VaceWanTransformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import WanVacePipeline, WanPipeline
from videox_fun.pipeline import WanPipeline, WanVacePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent,
from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -40,6 +43,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -220,6 +226,9 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
+22 -9
View File
@@ -19,6 +19,8 @@ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2AnimatePipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
@@ -30,7 +32,7 @@ from videox_fun.utils.utils import (filter_kwargs, get_image,
get_video_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -41,6 +43,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image,
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -106,9 +111,12 @@ src_bg_path = os.path.join(src_root_path, "src_bg.mp4")
src_mask_path = os.path.join(src_root_path, "src_mask.mp4")
# Other params
sample_size = [480, 832]
video_length = 81
fps = 16
sample_size = [480, 832]
# Total num frames
video_length = 81
# How many frames to generate per clips.
segment_frame_length = 77
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -263,6 +271,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -301,8 +314,8 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
@@ -331,7 +344,7 @@ with torch.no_grad():
sample = pipeline(
prompt,
num_frames = video_length,
segment_frame_length = segment_frame_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
@@ -350,8 +363,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
+64 -38
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel,
WanT5EncoderModel, Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2I2VPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -129,13 +136,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained(
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe":
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
transformer_2 = None
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
@@ -149,17 +158,18 @@ if transformer_path is not None:
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
if transformer_2 is not None:
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
Chosen_AutoencoderKL = {
@@ -221,11 +231,13 @@ pipeline = Wan2_2I2VPipeline(
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
transformer_2.enable_multi_gpus_inference()
if transformer_2 is not None:
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
if transformer_2 is not None:
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
@@ -235,29 +247,38 @@ if ulysses_degree > 1 or ring_degree > 1:
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
if transformer_2 is not None:
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
if transformer_2 is not None:
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
@@ -268,17 +289,20 @@ if coefficients is not None:
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
@@ -287,7 +311,8 @@ with torch.no_grad():
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
if transformer_2 is not None:
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
input_video, input_video_mask, clip_image = get_image_to_video_latent(validation_image_start, None, video_length=video_length, sample_size=sample_size)
@@ -309,6 +334,7 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
+24 -13
View File
@@ -19,6 +19,8 @@ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2S2VPipeline
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
@@ -30,7 +32,7 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
get_video_to_video_latent,
merge_video_audio, save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -41,6 +43,9 @@ from videox_fun.utils.utils import (filter_kwargs, get_image_latent,
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -103,9 +108,10 @@ lora_path = None
lora_high_path = None
# Other params
sample_size = [832, 480]
video_length = 80
fps = 16
sample_size = [832, 480]
# How many frames to generate per clips.
segment_frame_length = 80
fps = 16
# Use torch.float16 if GPU does not support torch.bfloat16
# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16
@@ -269,6 +275,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -307,12 +318,12 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = video_length // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio if video_length != 1 else 1
latent_frames = video_length // vae.config.temporal_compression_ratio
segment_frame_length = segment_frame_length // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio if segment_frame_length != 1 else 1
latent_frames = segment_frame_length // vae.config.temporal_compression_ratio
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
@@ -322,11 +333,11 @@ with torch.no_grad():
if ref_image is not None:
ref_image = get_image_latent(ref_image, sample_size=sample_size)
pose_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None)
pose_video, _, _, _ = get_video_to_video_latent(control_video, video_length=None, sample_size=sample_size, fps=fps, ref_image=None)
sample = pipeline(
prompt,
num_frames = video_length,
segment_frame_length = segment_frame_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
@@ -345,8 +356,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
@@ -354,7 +365,7 @@ def save_results():
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
if sample.size()[2] == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
+64 -38
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, AutoTokenizer,
Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, Wan2_2Transformer3DModel,
WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2Pipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -125,13 +132,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained(
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe":
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
transformer_2 = None
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
@@ -145,17 +154,18 @@ if transformer_path is not None:
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
if transformer_2 is not None:
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
Chosen_AutoencoderKL = {
@@ -216,11 +226,13 @@ pipeline = Wan2_2Pipeline(
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
transformer_2.enable_multi_gpus_inference()
if transformer_2 is not None:
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
if transformer_2 is not None:
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
@@ -230,29 +242,38 @@ if ulysses_degree > 1 or ring_degree > 1:
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
if transformer_2 is not None:
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
if transformer_2 is not None:
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
@@ -263,17 +284,20 @@ if coefficients is not None:
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
@@ -282,7 +306,8 @@ with torch.no_grad():
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
if transformer_2 is not None:
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
sample = pipeline(
prompt,
@@ -299,6 +324,7 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
+352
View File
@@ -0,0 +1,352 @@
import argparse
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
current_file_path = os.path.abspath(__file__)
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)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer,
Wan2_2Transformer3DModel, WanT5EncoderModel,
WanTransformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2Pipeline, WanPipeline
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid, timer_record)
def parse_args():
parser = argparse.ArgumentParser(description="Video Generation with Wan2.2-Fun")
parser.add_argument("--GPU_memory_mode", type=str, default="sequential_cpu_offload",
choices=["model_full_load", "model_full_load_and_qfloat8", "model_cpu_offload",
"model_cpu_offload_and_qfloat8", "sequential_cpu_offload"],
help="GPU memory optimization mode.")
parser.add_argument("--ulysses_degree", type=int, default=1,
help="Ulysses parallelism degree.")
parser.add_argument("--ring_degree", type=int, default=1,
help="Ring parallelism degree.")
parser.add_argument("--fsdp_dit", action="store_true",
help="Use FSDP for transformer to save GPU memory.")
parser.add_argument("--fsdp_text_encoder", action="store_true",
help="Use FSDP for text encoder to save GPU memory.")
parser.add_argument("--compile_dit", action="store_true",
help="Compile transformer for fixed resolution speedup.")
parser.add_argument("--enable_teacache", action="store_true",
help="Enable TeaCache optimization.")
parser.add_argument("--teacache_threshold", type=float, default=0.10,
help="TeaCache threshold for step caching.")
parser.add_argument("--num_skip_start_steps", type=int, default=5,
help="Number of steps to skip TeaCache at inference start.")
parser.add_argument("--teacache_offload", action="store_true",
help="Offload TeaCache tensors to CPU.")
parser.add_argument("--cfg_skip_ratio", type=float, default=0.0,
help="CFG skip ratio for inference.")
parser.add_argument("--enable_riflex", action="store_true",
help="Enable Riflex frequency optimization.")
parser.add_argument("--riflex_k", type=int, default=6,
help="Intrinsic frequency index for Riflex.")
parser.add_argument("--config_path", type=str, default="config/wan2.2/wan_civitai_t2v.yaml",
help="Path to model config file.")
parser.add_argument("--model_name", type=str, default="models/Diffusion_Transformer/Wan2.2-T2V-A14B",
help="Path to model directory.")
parser.add_argument("--sampler_name", type=str, default="Flow_Unipc",
choices=["Flow", "Flow_Unipc", "Flow_DPM++"],
help="Sampler type for video generation.")
parser.add_argument("--shift", type=float, default=3.0,
help="Noise schedule shift parameter for Flow_Unipc/Flow_DPM++.")
parser.add_argument("--transformer_path", type=str, default=None,
help="Path to pre-trained transformer checkpoint.")
parser.add_argument("--transformer_high_path", type=str, default=None,
help="Path to pre-trained high noise transformer checkpoint.")
parser.add_argument("--vae_path", type=str, default=None,
help="Path to pre-trained VAE checkpoint.")
parser.add_argument("--lora_path", type=str, default=None,
help="Path to LoRA weights.")
parser.add_argument("--lora_high_path", type=str, default=None,
help="Path to high noise LoRA weights.")
parser.add_argument("--sample_size", nargs=2, type=int, default=[480, 832],
help="Sample size [height, width].")
parser.add_argument("--video_length", type=int, default=81,
help="Number of frames in the video.")
parser.add_argument("--fps", type=int, default=16,
help="Frames per second for output video.")
parser.add_argument("--weight_dtype", type=str, default="bfloat16",
choices=["float16", "bfloat16"],
help="Weight data type (float16 or bfloat16).")
parser.add_argument("--prompt", type=str, default="一只棕色的狗摇着头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。",
help="Text prompt for video generation.")
parser.add_argument("--negative_prompt", type=str, default="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
help="Negative prompt for video generation.")
parser.add_argument("--guidance_scale", type=float, default=6.0,
help="Classifier-free guidance scale.")
parser.add_argument("--seed", type=int, default=43,
help="Random seed for reproducibility.")
parser.add_argument("--num_inference_steps", type=int, default=50,
help="Number of inference steps.")
parser.add_argument("--lora_weight", type=float, default=0.55,
help="LoRA weight scaling factor.")
parser.add_argument("--lora_high_weight", type=float, default=0.55,
help="High noise LoRA weight scaling factor.")
parser.add_argument("--save_path", type=str, default="samples/wan-videos-t2v",
help="Directory to save generated videos.")
return parser.parse_args()
args = parse_args()
# 将 argparse 参数映射到原有变量
GPU_memory_mode = args.GPU_memory_mode
ulysses_degree = args.ulysses_degree
ring_degree = args.ring_degree
fsdp_dit = args.fsdp_dit
fsdp_text_encoder = args.fsdp_text_encoder
compile_dit = args.compile_dit
enable_teacache = args.enable_teacache
teacache_threshold = args.teacache_threshold
num_skip_start_steps = args.num_skip_start_steps
teacache_offload = args.teacache_offload
cfg_skip_ratio = args.cfg_skip_ratio
enable_riflex = args.enable_riflex
riflex_k = args.riflex_k
config_path = args.config_path
model_name = args.model_name
sampler_name = args.sampler_name
shift = args.shift
transformer_path = args.transformer_path
transformer_high_path = args.transformer_high_path
vae_path = args.vae_path
lora_path = args.lora_path
lora_high_path = args.lora_high_path
sample_size = args.sample_size
video_length = args.video_length
fps = args.fps
weight_dtype = torch.bfloat16 if args.weight_dtype == "bfloat16" else torch.float16
prompt = args.prompt
negative_prompt = args.negative_prompt
guidance_scale = args.guidance_scale
seed = args.seed
num_inference_steps = args.num_inference_steps
lora_weight = args.lora_weight
lora_high_weight = args.lora_high_weight
save_path = args.save_path
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
boundary = config['transformer_additional_kwargs'].get('boundary', 0.875)
transformer = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
if transformer_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_path)
else:
state_dict = torch.load(transformer_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
if vae_path is not None:
print(f"From checkpoint: {vae_path}")
if vae_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(vae_path)
else:
state_dict = torch.load(vae_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = vae.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Tokenizer
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')),
)
# Get Text encoder
text_encoder = WanT5EncoderModel.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
# Get Scheduler
Choosen_Scheduler = scheduler_dict = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++":
config['scheduler_kwargs']['shift'] = 1
scheduler = Choosen_Scheduler(
**filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
# Get Pipeline
pipeline = Wan2_2Pipeline(
transformer=transformer,
transformer_2=transformer_2,
vae=vae,
tokenizer=tokenizer,
text_encoder=text_encoder,
scheduler=scheduler,
)
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.text_encoder = shard_fn(pipeline.text_encoder)
print("Add FSDP TEXT ENCODER")
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
for i in range(2):
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.")
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device)
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
sample = pipeline(
prompt,
num_frames = video_length,
negative_prompt = negative_prompt,
height = sample_size[0],
width = sample_size[1],
generator = generator,
guidance_scale = guidance_scale,
num_inference_steps = num_inference_steps,
boundary = boundary,
shift = shift,
).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device)
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
os.makedirs(save_path, exist_ok=True)
index = len([path for path in os.listdir(save_path)]) + 1
prefix = str(index).zfill(8)
if video_length == 1:
video_path = os.path.join(save_path, prefix + ".png")
image = sample[0, :, 0]
image = image.transpose(0, 1).transpose(1, 2)
image = (image * 255).numpy().astype(np.uint8)
image = Image.fromarray(image)
image.save(video_path)
else:
video_path = os.path.join(save_path, prefix + ".mp4")
save_videos_grid(sample, video_path, fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
+53
View File
@@ -0,0 +1,53 @@
export EXCEL_FILE="./speed.xlsx"
export DIT_EXCEL_COL=0 VAE_EXCEL_COL=1 TOTAL_EXCEL_COL=2
# 14B 720P
export DIT_EXCEL_ROW=1 VAE_EXCEL_ROW=1 TOTAL_EXCEL_ROW=1
python examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=2 VAE_EXCEL_ROW=2 TOTAL_EXCEL_ROW=2
torchrun --nproc-per-node=2 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=3 VAE_EXCEL_ROW=3 TOTAL_EXCEL_ROW=3
torchrun --nproc-per-node=4 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
export DIT_EXCEL_ROW=4 VAE_EXCEL_ROW=4 TOTAL_EXCEL_ROW=4
torchrun --nproc-per-node=8 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=2 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=5 \
--sample_size 720 1280 --num_inference_steps=40
# 14B 480P
export DIT_EXCEL_ROW=5 VAE_EXCEL_ROW=5 TOTAL_EXCEL_ROW=5
python examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=1 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=6 VAE_EXCEL_ROW=6 TOTAL_EXCEL_ROW=6
torchrun --nproc-per-node=2 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=2 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=7 VAE_EXCEL_ROW=7 TOTAL_EXCEL_ROW=7
torchrun --nproc-per-node=4 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=1 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
export DIT_EXCEL_ROW=8 VAE_EXCEL_ROW=8 TOTAL_EXCEL_ROW=8
torchrun --nproc-per-node=8 examples/wan2.2/predict_t2v_speed.py --model_name="models/Diffusion_Transformer/Wan2.2-T2V-A14B" \
--GPU_memory_mode="model_full_load" --ulysses_degree=4 --ring_degree=2 --fsdp_text_encoder --fsdp_dit \
--enable_teacache --teacache_threshold=0.10 --num_skip_start_steps=2 --cfg_skip_ratio=0.25 --shift=3 \
--sample_size 480 832 --num_inference_steps=40
+24 -12
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan3_8, AutoencoderKLWan, WanT5EncoderModel, AutoTokenizer,
Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, Wan2_2Transformer3DModel,
WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2TI2VPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -253,6 +260,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -291,8 +303,8 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
@@ -326,8 +338,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
+25 -13
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel,
WanT5EncoderModel, Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2FunInpaintPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload transfers internal layer groups between CPU/CUDA,
# balancing memory efficiency and speed between full-module and leaf-level offloading methods.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -96,7 +103,7 @@ vae_path = None
# Load lora model if need
# The lora_path is used for low noise model, the lora_high_path is used for high noise model.
lora_path = None
lora_high_path = None
lora_high_path = None
# Other params
sample_size = [480, 832]
@@ -255,6 +262,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -293,8 +305,8 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
@@ -325,8 +337,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
+24 -12
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, AutoencoderKLWan3_8,
WanT5EncoderModel, Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2FunInpaintPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload means that the internal layer groups will be transferred between CPU and CUDA,
# balancing memory efficiency and speed.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -257,6 +264,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -295,8 +307,8 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
@@ -327,8 +339,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):
+64 -38
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, AutoencoderKLWan3_8,
WanT5EncoderModel, Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2FunInpaintPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload means that the internal layer groups will be transferred between CPU and CUDA,
# balancing memory efficiency and speed.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -128,13 +135,15 @@ transformer = Wan2_2Transformer3DModel.from_pretrained(
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
if config['transformer_additional_kwargs'].get('transformer_combination_type', 'single') == "moe":
transformer_2 = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True,
torch_dtype=weight_dtype,
)
else:
transformer_2 = None
if transformer_path is not None:
print(f"From checkpoint: {transformer_path}")
@@ -148,17 +157,18 @@ if transformer_path is not None:
m, u = transformer.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
if transformer_2 is not None:
if transformer_high_path is not None:
print(f"From checkpoint: {transformer_high_path}")
if transformer_high_path.endswith("safetensors"):
from safetensors.torch import load_file, safe_open
state_dict = load_file(transformer_high_path)
else:
state_dict = torch.load(transformer_high_path, map_location="cpu")
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
m, u = transformer_2.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
# Get Vae
Chosen_AutoencoderKL = {
@@ -220,11 +230,13 @@ pipeline = Wan2_2FunInpaintPipeline(
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
transformer_2.enable_multi_gpus_inference()
if transformer_2 is not None:
transformer_2.enable_multi_gpus_inference()
if fsdp_dit:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
pipeline.transformer = shard_fn(pipeline.transformer)
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
if transformer_2 is not None:
pipeline.transformer_2 = shard_fn(pipeline.transformer_2)
print("Add FSDP DIT")
if fsdp_text_encoder:
shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype)
@@ -234,29 +246,38 @@ if ulysses_degree > 1 or ring_degree > 1:
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
if transformer_2 is not None:
for i in range(len(pipeline.transformer_2.blocks)):
pipeline.transformer_2.blocks[i] = torch.compile(pipeline.transformer_2.blocks[i])
print("Add Compile")
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
if transformer_2 is not None:
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
if transformer_2 is not None:
convert_model_weight_to_float8(transformer_2, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer_2, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
@@ -267,17 +288,20 @@ if coefficients is not None:
pipeline.transformer.enable_teacache(
coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload
)
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_teacache(transformer=pipeline.transformer)
if cfg_skip_ratio is not None:
print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.")
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
if transformer_2 is not None:
pipeline.transformer_2.share_cfg_skip(transformer=pipeline.transformer)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
@@ -286,7 +310,8 @@ with torch.no_grad():
if enable_riflex:
pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames)
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
if transformer_2 is not None:
pipeline.transformer_2.enable_riflex(k = riflex_k, L_test = latent_frames)
input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size)
@@ -308,6 +333,7 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
+24 -12
View File
@@ -13,19 +13,23 @@ for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel,
WanT5EncoderModel, Wan2_2Transformer3DModel)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8,
AutoTokenizer, CLIPModel,
Wan2_2Transformer3DModel, WanT5EncoderModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import Wan2_2FunInpaintPipeline
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name,
convert_weight_dtype_wrapper)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
from videox_fun.utils import (register_auto_device_hook,
safe_enable_group_offload)
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# 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].
# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload].
# model_full_load means that the entire model will be moved to the GPU.
#
# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU,
@@ -36,6 +40,9 @@ from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use,
# and the transformer model has been quantized to float8, which can save more GPU memory.
#
# model_group_offload means that the internal layer groups will be transferred between CPU and CUDA,
# balancing memory efficiency and speed.
#
# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use,
# resulting in slower speeds but saving a large amount of GPU memory.
GPU_memory_mode = "sequential_cpu_offload"
@@ -251,6 +258,11 @@ if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer_2, ["modulation",], device=device)
transformer_2.freqs = transformer_2.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
if transformer_2 is not None:
register_auto_device_hook(pipeline.transformer_2)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
@@ -289,8 +301,8 @@ generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = merge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
@@ -321,8 +333,8 @@ with torch.no_grad():
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
if transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
if lora_high_path is not None and transformer_2 is not None:
pipeline = unmerge_lora(pipeline, lora_high_path, lora_high_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2")
def save_results():
if not os.path.exists(save_path):

Some files were not shown because too many files have changed in this diff Show More