Compare commits
48
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1c581bb17 | ||
|
|
968f0e2192 | ||
|
|
43739895a1 | ||
|
|
d31de92872 | ||
|
|
6f3fb60dad | ||
|
|
35bfe679dc | ||
|
|
b0acf916c2 | ||
|
|
6787dc8ed4 | ||
|
|
d555ea056b | ||
|
|
b6e5a32b3e | ||
|
|
248ab0ac0e | ||
|
|
403f1f7b78 | ||
|
|
1fd9ed9208 | ||
|
|
2b5596b8e6 | ||
|
|
804a4258e2 | ||
|
|
8fb0bb165c | ||
|
|
0266bab98b | ||
|
|
199a43b544 | ||
|
|
34036517a8 | ||
|
|
54bf97ab66 | ||
|
|
4a86483cc2 | ||
|
|
ad72867c0f | ||
|
|
5202421e7c | ||
|
|
745acc1f47 | ||
|
|
a4b60f40bf | ||
|
|
7a085c9535 | ||
|
|
1ca4162447 | ||
|
|
ee44fbc950 | ||
|
|
67ab552885 | ||
|
|
413f831ced | ||
|
|
f4ffd1c64b | ||
|
|
ed42563dd3 | ||
|
|
8ee48e420c | ||
|
|
b757b8edae | ||
|
|
924dd8528a | ||
|
|
e258d4158b | ||
|
|
48c8288323 | ||
|
|
1aacfe6bca | ||
|
|
288e88eb46 | ||
|
|
dfb8ca04a7 | ||
|
|
499eef12c3 | ||
|
|
fa8623b22b | ||
|
|
61fb59833f | ||
|
|
4b0b009fd3 | ||
|
|
518c2bdc7e | ||
|
|
6188e66fc4 | ||
|
|
7c9d655822 | ||
|
|
873f622dae |
+12
-10
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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”.
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
@@ -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
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.
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
+1
-1
@@ -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"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
+1
-1
@@ -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"
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
+1
-1
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
Executable
+303
@@ -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()
|
||||
Executable
+103
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}")
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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():
|
||||
|
||||
Executable
+352
@@ -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()
|
||||
Executable
+53
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user