Merge branch 'main' into speed

This commit is contained in:
bubbliiiing
2026-09-04 21:47:41 +08:00
32 changed files with 8302 additions and 44 deletions
+128
View File
@@ -0,0 +1,128 @@
---
name: integrating-models
description: Guides adding, porting, or onboarding a diffusion model (transformer/VAE/encoder, inference pipeline, training script, config) into the VideoX-Fun repository by mirroring the closest existing model family and maximizing reuse of the repository's existing code and shared infrastructure. Use when integrating a new model/architecture, or when creating predict_*.py inference scripts, scripts/*/train*.py training scripts, pipeline_*.py, config/*.yaml, or model definitions under videox_fun/models/.
---
# Integrating Models into VideoX-Fun
## Core rule: maximize reuse of existing repo code — mirror, extend, never reinvent
**Prime directive: reuse this repository's existing code to the maximum.** Nearly every building block you need already exists in `videox_fun/` or in a sibling model family. Your job is to **find it, import it, and extend it** — not to write a parallel implementation. A new file should be mostly reused structure plus the genuinely model-specific delta; the less new code you write, the better.
**Reuse-first protocol — before writing ANY new function / class / util:**
1. **Search the repo first.** Grep `videox_fun/` and the closest family for an existing equivalent (weight loader, scheduler, sampler, offload, attention, LoRA, fp8, dataset, dist helper, save/metric util). If one exists → **import and reuse it**. If it is 80% right → **extend / parameterize it**, do not fork it.
2. **Only if nothing exists** may you add new code — and then put it in the shared layer (`videox_fun/utils`, `videox_fun/data`, `videox_fun/dist`) so the next model reuses it too, instead of burying it in a family folder.
3. **Never copy-paste** a util into a new file (that creates drift); import the single source of truth.
**Mirror the closest family.** Every model follows the **same layered template**. Integrating a model means finding the closest existing family and mirroring its structure, changing only what genuinely differs:
1. Pick the closest existing family by task type (t2v / i2v / v2v-control / s2v / t2i / edit / distill): `wan2.1`, `wan2.1_fun`, `wan2.2`, `qwenimage`, `flux2`, `minimax_h3`, `ltx2`, `longcatvideo`, `cogvideox_fun`, `z_image`, etc.
2. Read that family end-to-end across all layers:
- `examples/<family>/predict_*.py` (inference entry)
- `scripts/<family>/train*.py` + `*.sh` + `README_TRAIN*.md` (training)
- `videox_fun/pipeline/pipeline_<family>*.py` (pipeline)
- `videox_fun/models/<family>_*.py` (model definitions)
- `config/<family>/*.yaml` (config)
3. Copy that structure and adapt. Keep names, argument sets, control flow, and reuse points identical in shape.
Writing a bespoke pipeline, weight loader, trainer, sampler, dataset, or offload scheme from scratch is a **failure mode**. If you are tempted to, **stop** and check the Reuse inventory below first.
## Repository layout (where each layer lives)
| Layer | Location | What it is |
|-------|----------|------------|
| Model definitions | `videox_fun/models/<family>_*.py` | Transformer / VAE / text-audio-image encoders. Diffusers `ModelMixin`+`ConfigMixin`, `@register_to_config`, custom `from_pretrained`. |
| Model registry | `videox_fun/models/__init__.py` | Imports every model class. **Must be updated** for a new model. |
| Inference pipelines | `videox_fun/pipeline/pipeline_<family>*.py` | `<Family>Pipeline(DiffusionPipeline)` with `__call__`. |
| Pipeline registry | `videox_fun/pipeline/__init__.py` | Imports every pipeline + aliases. **Must be updated.** |
| Configs (optional) | `config/<family>/*.yaml` | OmegaConf YAML for civitai/custom layouts; a standard diffusers-layout checkpoint can load without one. |
| Inference entry scripts | `examples/<family>/predict_*.py` | User-facing, config-block-at-top runnable scripts. |
| Inference services | `examples/<family>/{app.py,launch_api.py,post_infer*.py}` | Gradio UI / API server / batch inference. |
| Training scripts | `scripts/<family>/train*.py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py`, ... |
| Training launchers | `scripts/<family>/train*.sh` | `accelerate launch` / DeepSpeed command with full arg list. |
| Training docs | `scripts/<family>/README_TRAIN*.md` | Bilingual pairs: `README_TRAIN.md` + `README_TRAIN_zh-CN.md`. |
| Shared: schedulers/utils | `videox_fun/utils/` | `fm_solvers`, `fm_solvers_unipc`, `lora_utils`, `fp8_optimization`, `group_offload`, `utils.py`. |
| Shared: distributed | `videox_fun/dist/` | `fsdp.shard_model`, `fuser.set_multi_gpus_devices`, `<family>_xfuser` sequence-parallel attention. |
| Shared: data | `videox_fun/data/` | Datasets (`ImageVideoDataset`, `VideoDataset`, ...) + bucket/aspect-ratio samplers. |
| Demo / test datasets | `datasets/X-Fun-*-Demo/` | Ready-made smoke-test data, downloaded via `modelscope download --dataset PAI/<name>`; each ships several `metadata*.json` variants. **The only test data to use** (see reference.md §8). |
| Preprocessing (data gen) | `scripts/<family>/generate_*.py` / `train_preprocess.py` (+ `.sh`) | Offline multi-GPU generation of cached training data (latents / ODE pairs / embeddings) → per-sample `.safetensors` + `outputs.json`, loaded by `ImageVideoSafetensorsDataset`. |
| ComfyUI nodes | `comfyui/<family>/nodes.py` | Optional node integration mirroring the pipeline. |
## Integration workflow
Copy this checklist and track progress:
```
Integration Progress:
- [ ] Step 0: Choose the closest family to mirror; read it across all layers
- [ ] Step 1: Model definitions in videox_fun/models/ + register in models/__init__.py
- [ ] Step 2: Pipeline in videox_fun/pipeline/ + register in pipeline/__init__.py
- [ ] Step 3: Config YAML in config/<family>/
- [ ] Step 4: Inference script(s) in examples/<family>/predict_*.py
- [ ] Step 5: Training script(s) in scripts/<family>/train*.py + .sh
- [ ] Step 6: Training docs README_TRAIN.md + README_TRAIN_zh-CN.md
- [ ] Step 7: Reuse audit + verification (incl. smoke test on the matching demo dataset)
```
**Step 0 — Choose the mirror.** Match by task and architecture. A new control model mirrors an existing `*_fun`/`*_control` family; a new audio/talking model mirrors `minimax_h3`/`longcatvideo`/`infinitetalk`; a new image model mirrors `qwenimage`/`flux2`/`z_image`.
**Step 1 — Model.** Create `videox_fun/models/<family>_transformer3d.py` (or `2d`), `<family>_vae.py`, encoders as needed. Mirror the class shape: `class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin)`, `_supports_gradient_checkpointing = True`, `@register_to_config __init__`, and a `from_pretrained` that supports `transformer_additional_kwargs`, `dict_mapping`, `low_cpu_mem_usage`, and missing-key init. Add imports to `videox_fun/models/__init__.py`.
**Step 2 — Pipeline.** Create `videox_fun/pipeline/pipeline_<family>.py`. Mirror `pipeline_wan.py`: module-level `retrieve_timesteps`, a `<Family>PipelineOutput(BaseOutput)` dataclass, `<Family>Pipeline(DiffusionPipeline)` with `model_cpu_offload_seq`, `_callback_tensor_inputs`, `__init__(vae, tokenizer, text_encoder, transformer, scheduler, ...)`, `encode_prompt`, and `__call__`. Add imports/aliases to `videox_fun/pipeline/__init__.py`.
**Step 3 — Config (optional).** A YAML under `config/<family>/` is **not always required**. It is needed mainly for **civitai-format / custom single-file layouts** — to supply `transformer_additional_kwargs`, `dict_mapping` (civitai key → `__init__` kwarg), component subpaths, and `vae_kwargs`/`text_encoder_kwargs`/`scheduler_kwargs`/`image_encoder_kwargs`. For a **standard diffusers-layout** checkpoint (`model_index.json` + per-subfolder `config.json`), load directly via `from_pretrained(model_name, subfolder=...)` with no YAML — mirror `examples/minimax_h3_fun/predict_v2v_control.py`, which guards `if config_path is not None:`. When you do add a YAML, load it via `OmegaConf.load(config_path)` and spread into `from_pretrained` instead of hardcoding those values.
**Step 4 — Inference script.** Create `examples/<family>/predict_<task>.py` following the exact template (config block at top → component loading → scheduler dict → pipeline construction → multi-GPU/FSDP/compile → `GPU_memory_mode` branching → TeaCache → LoRA merge → inference → `save_results`). See [examples.md](examples.md).
**Step 5 — Training script.** Create `scripts/<family>/train.py` (+ `train_lora.py` etc.). Mirror the shared structure: license header, `sys.path` bootstrap, imports from `videox_fun`, `log_validation()` that **reuses the inference Pipeline**, `parse_args()` (reuse the existing shared argument set), `main()`. Add a `train.sh` launcher. Reuse `videox_fun.data` datasets/samplers — do not write a new dataset.
**Step 6 — Docs.** Write `README_TRAIN.md` and `README_TRAIN_zh-CN.md` as an aligned bilingual pair (same structure, same commands/params, matching section order).
**Step 7 — Reuse audit + verification.** Confirm you reused shared infra (below), smoke-test the new train/predict path on the **matching official demo dataset** under `datasets/X-Fun-*-Demo/` (pick by task and metadata variant — see reference.md §8), then run the verification checklist. Never invent an ad-hoc test set and never leave `datasets/internal_datasets/` placeholders in shipped scripts/docs.
## Reuse inventory (use these, do not reimplement)
**Reuse-first catalog: import from here instead of reimplementing. If a helper you need is not listed, grep `videox_fun/` and the closest family before writing your own.**
- **Schedulers**: `FlowMatchEulerDiscreteScheduler`, `videox_fun.utils.fm_solvers.FlowDPMSolverMultistepScheduler`, `fm_solvers_unipc.FlowUniPCMultistepScheduler`. Selected via a `sampler_name` dict.
- **LoRA**: `videox_fun.utils.lora_utils` — `merge_lora`, `unmerge_lora`, `create_network`, `convert_peft_lora_to_kohya_lora`.
- **FP8 / quantization**: `videox_fun.utils.fp8_optimization` — `convert_model_weight_to_float8`, `convert_weight_dtype_wrapper`, `replace_parameters_by_name`.
- **Offloading**: `videox_fun.utils.group_offload` — `register_auto_device_hook`, `safe_enable_group_offload`; plus pipeline `enable_sequential_cpu_offload` / `enable_model_cpu_offload` / `.to(device)`.
- **Distributed**: `videox_fun.dist` — `set_multi_gpus_devices`, `shard_model` (FSDP), `<family>_xfuser` sequence-parallel attention processors, `enable_multi_gpus_inference()`.
- **IO / helpers**: `videox_fun.utils.utils` — `save_videos_grid`, `save_videos_with_audio_grid`, `get_image_to_video_latent`, `get_video_to_video_latent`, `get_image_latent`, `filter_kwargs`, `calculate_dimensions`.
- **Data**: `videox_fun.data` — `ImageVideoDataset`, `VideoDataset`, `ImageVideoControlDataset`, `VideoSpeechDataset`, bucket/aspect-ratio samplers, `get_closest_ratio`, `get_random_mask`.
- **Caching / speedups**: TeaCache (`models/cache_utils`, `get_teacache_coefficients`, `transformer.enable_teacache`), `enable_cfg_skip`, Riflex (`enable_riflex`), `torch.compile` on `transformer.blocks`.
- **Preprocessing (data gen, multi-GPU)**: mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py` — `accelerate launch` + `Accelerator` (interleaved rank sharding), config-driven `from_pretrained` for the teacher/VAE/text-encoder, `safetensors.torch.save_file` per sample + `outputs.json` index, consumed by `videox_fun.data.ImageVideoSafetensorsDataset`. Store as **safetensors only — never LMDB or `.pt`** (see reference.md §10).
## Non-negotiable conventions
- **Maximize reuse of existing repo code**: import existing `videox_fun/` helpers and mirror the closest family; never fork or copy-paste a util, and never write a parallel pipeline / loader / scheduler / sampler / offload. Genuinely-new shared code goes in `videox_fun/{utils,data,dist}` (so the next model reuses it), not buried in a family folder.
- **`sys.path` bootstrap**: every runnable script starts with the 3-level `project_roots` loop inserting into `sys.path` before importing `videox_fun`.
- **Config-driven loading (YAML optional)**: a `config/<family>/*.yaml` is required for civitai-format/custom layouts (it supplies `transformer_additional_kwargs`/`dict_mapping`/subpaths); it is **optional for standard diffusers-layout checkpoints**, which load directly via `from_pretrained(model_name, subfolder=...)`. When a YAML is used, don't hardcode the values it provides.
- **`GPU_memory_mode`**: support the standard six modes — `model_full_load`, `model_full_load_and_qfloat8`, `model_cpu_offload`, `model_cpu_offload_and_qfloat8`, `model_group_offload`, `sequential_cpu_offload` — with the exact branching order used in existing `predict_*.py`.
- **Naming**: files `<family>_transformer3d.py` / `<family>_vae.py` / `pipeline_<family>.py`; classes `<Family>Transformer3DModel` / `AutoencoderKL<Family>` / `<Family>Pipeline`.
- **Resolution args**: drive canvas size with a single square `--video_sample_size` (`type=int`, height = width); never `--video_sample_height` / `--video_sample_width`. For a fixed non-square shape add `--fix_sample_size` (`nargs=2, type=int`, `[height, width]`) that overrides the square size, and derive the effective height/width once in `parse_args()` (see reference.md §5).
- **Registries**: a model is not integrated until it is imported in BOTH `videox_fun/models/__init__.py` and `videox_fun/pipeline/__init__.py`.
- **Two weight formats**: support `civitai` and `diffusers` via config `format` + `dict_mapping` (maps civitai keys such as `in_dim`→`in_channels`, `dim`→`hidden_size`).
- **Bilingual docs**: training READMEs ship as EN + `_zh-CN` pairs with aligned structure and identical commands/params.
- **Test data = official demo datasets**: smoke tests, `log_validation` checks, launcher `.sh` defaults, and doc examples all point at `datasets/X-Fun-*-Demo/` (ModelScope `PAI/<name>`), with the metadata variant matching the task — `metadata_add_width_height.json` by default, `_add_objects.json` for VACE/subject-reference, `_add_wav.json` for audio-visual joint models, `metadata_lingbot_video_add_width_height.json` for `lingbot_video`. Selection matrix: reference.md §8.
- **Preprocessing = offline data generation, multi-GPU + safetensors**: cached training data (latents / ODE pairs / embeddings) is produced by `accelerate launch` scripts like `generate_ode_pairs.py` (interleaved rank sharding, resume by skipping existing files, `wait_for_everyone`, rank-0 JSON index) and saved with `safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`. **Never single-GPU / `cuda:0`; never LMDB or `.pt`/`torch.save` pickles for preprocessed data.** See reference.md §10.
## Verification checklist
- [ ] New model classes imported in `videox_fun/models/__init__.py`
- [ ] New pipeline(s) imported in `videox_fun/pipeline/__init__.py`
- [ ] Config YAML present **only if** the checkpoint is civitai-format/custom-layout; a diffusers-layout model may load directly via `from_pretrained(model_name, subfolder=...)` with no YAML. When a YAML is used, it drives component loading (no hardcoded kwargs)
- [ ] `predict_*.py` mirrors an existing script: `sys.path` bootstrap, config block, scheduler dict, `GPU_memory_mode` branching, LoRA merge, `save_results`
- [ ] `train*.py` reuses `videox_fun.data` + shared args, and `log_validation()` reuses the inference Pipeline
- [ ] `train*.sh` launcher provided (`accelerate launch` / DeepSpeed)
- [ ] Shared infra reused (schedulers / lora_utils / fp8 / group_offload / dist / utils / data) — nothing reimplemented
- [ ] Any offline data-generation/preprocessing script runs multi-GPU (`accelerate launch` + `Accelerator`) and saves cached tensors as **safetensors + `outputs.json`** for `ImageVideoSafetensorsDataset` — never LMDB or `.pt`
- [ ] `README_TRAIN.md` + `README_TRAIN_zh-CN.md` aligned pair present
- [ ] Smoke test / doc examples use the matching `datasets/X-Fun-*-Demo` dataset and the correct `metadata*.json` variant — no `internal_datasets` placeholders (reference.md §8)
- [ ] Optional: ComfyUI node in `comfyui/<family>/nodes.py` mirrors the pipeline
## Additional resources
- Detailed file-by-file conventions, class/method shapes, and the model-loading internals: [reference.md](reference.md)
- **Dataset & sampler selection matrix** (which `videox_fun.data` dataset/loader each training task uses), **demo-dataset / metadata-variant selection matrix** (which `datasets/X-Fun-*-Demo` to smoke-test with), **inference task matrix** (which pipeline each `predict_<task>.py` uses), and **multi-GPU preprocessing patterns**: [reference.md](reference.md) §8–§10
- Concrete skeletons (config YAML, `predict_*.py`, pipeline class, training script + DataLoader): [examples.md](examples.md)
+510
View File
@@ -0,0 +1,510 @@
# VideoX-Fun Integration Skeletons
Starting templates. **Always open the mirrored family's real file and adapt it** — these skeletons show shape and required reuse points, not full implementations. Replace `<family>` / `<Family>` / `<task>`.
## Config — `config/<family>/<variant>.yaml` (optional)
> **Not always required.** Author a YAML only for civitai-format / custom single-file layouts. A standard diffusers-layout checkpoint (`model_index.json` + per-subfolder `config.json`) loads directly via `from_pretrained(model_name, subfolder=...)` with no YAML — set `config_path = None` and guard `if config_path is not None:` (see `examples/minimax_h3_fun/predict_v2v_control.py`).
```yaml
format: civitai
pipeline: <Family>
transformer_additional_kwargs:
transformer_subpath: ./
dict_mapping:
in_dim: in_channels
dim: hidden_size
vae_kwargs:
vae_subpath: <Family>_VAE.pth
temporal_compression_ratio: 4
spatial_compression_ratio: 8
text_encoder_kwargs:
text_encoder_subpath: <text_encoder>.pth
tokenizer_subpath: <tokenizer_id>
text_length: 512
scheduler_kwargs:
scheduler_subpath: null
num_train_timesteps: 1000
shift: 5.0
# Only for i2v / models with a CLIP image encoder:
image_encoder_kwargs:
image_encoder_subpath: <image_encoder>.pth
```
## Inference — `examples/<family>/predict_<task>.py`
```python
import os
import sys
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from PIL import Image
from transformers import AutoTokenizer
# --- sys.path bootstrap (required, before importing videox_fun) ---
current_file_path = os.path.abspath(__file__)
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
from videox_fun.dist import set_multi_gpus_devices, shard_model
from videox_fun.models import (AutoencoderKL<Family>, <Family>TextEncoder,
<Family>Transformer3DModel)
from videox_fun.models.cache_utils import get_teacache_coefficients
from videox_fun.pipeline import <Family>Pipeline
from videox_fun.utils import register_auto_device_hook, safe_enable_group_offload
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
convert_weight_dtype_wrapper,
replace_parameters_by_name)
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
save_videos_grid)
# --- user config block (keep the conventional order + comments) ---
GPU_memory_mode = "sequential_cpu_offload"
ulysses_degree = 1
ring_degree = 1
fsdp_dit = False
fsdp_text_encoder = True
compile_dit = False
enable_teacache = True
teacache_threshold = 0.10
num_skip_start_steps = 5
teacache_offload = False
cfg_skip_ratio = 0
enable_riflex = False
riflex_k = 6
config_path = "config/<family>/<variant>.yaml"
model_name = "models/Diffusion_Transformer/<Family>-Model"
sampler_name = "Flow"
shift = 3
transformer_path = None
vae_path = None
lora_path = None
sample_size = [480, 832]
video_length = 81
fps = 16
weight_dtype = torch.bfloat16
prompt = "..."
negative_prompt = "..."
guidance_scale = 6.0
seed = 43
num_inference_steps = 50
lora_weight = 0.55
save_path = "samples/<family>-<task>"
# --- device + config (config_path may be None for a diffusers-layout checkpoint) ---
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path) # or guard: if config_path is not None: ... (then load components via subfolder=...)
# --- components (when a YAML is used, paths/kwargs come from config; otherwise pass subfolder=... directly) ---
transformer = <Family>Transformer3DModel.from_pretrained(
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
low_cpu_mem_usage=True, torch_dtype=weight_dtype,
)
# optional transformer_path / vae_path override -> load_state_dict(strict=False) + print missing/unexpected
vae = AutoencoderKL<Family>.from_pretrained(
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
).to(weight_dtype)
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')))
text_encoder = <Family>TextEncoder.from_pretrained(
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
low_cpu_mem_usage=True, torch_dtype=weight_dtype).eval()
# --- scheduler selection dict ---
Chosen_Scheduler = {
"Flow": FlowMatchEulerDiscreteScheduler,
"Flow_Unipc": FlowUniPCMultistepScheduler,
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
}[sampler_name]
scheduler = Chosen_Scheduler(**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])))
# --- pipeline ---
pipeline = <Family>Pipeline(vae=vae, tokenizer=tokenizer, text_encoder=text_encoder,
transformer=transformer, scheduler=scheduler)
# --- multi-gpu / fsdp / compile ---
if ulysses_degree > 1 or ring_degree > 1:
from functools import partial
transformer.enable_multi_gpus_inference()
if fsdp_dit:
pipeline.transformer = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.transformer)
if fsdp_text_encoder:
pipeline.text_encoder = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.text_encoder)
if compile_dit:
for i in range(len(pipeline.transformer.blocks)):
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
# --- GPU_memory_mode branching (keep this exact order) ---
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
# --- teacache / cfg_skip / riflex / lora ---
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
if coefficients is not None:
pipeline.transformer.enable_teacache(coefficients, num_inference_steps, teacache_threshold,
num_skip_start_steps=num_skip_start_steps, offload=teacache_offload)
if cfg_skip_ratio is not None:
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
generator = torch.Generator(device=device).manual_seed(seed)
if lora_path is not None:
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
# --- inference ---
with torch.no_grad():
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
if enable_riflex:
pipeline.transformer.enable_riflex(k=riflex_k, L_test=(video_length - 1) // vae.config.temporal_compression_ratio + 1)
sample = pipeline(prompt, num_frames=video_length, negative_prompt=negative_prompt,
height=sample_size[0], width=sample_size[1], generator=generator,
guidance_scale=guidance_scale, num_inference_steps=num_inference_steps,
shift=shift).videos
if lora_path is not None:
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
# --- save (rank 0 only when multi-gpu) ---
def save_results():
os.makedirs(save_path, exist_ok=True)
prefix = str(len(os.listdir(save_path)) + 1).zfill(8)
if video_length == 1:
image = (sample[0, :, 0].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
Image.fromarray(image).save(os.path.join(save_path, prefix + ".png"))
else:
save_videos_grid(sample, os.path.join(save_path, prefix + ".mp4"), fps=fps)
if ulysses_degree * ring_degree > 1:
import torch.distributed as dist
if dist.get_rank() == 0:
save_results()
else:
save_results()
```
For i2v, gate the CLIP image encoder and pass `video`/`mask_video`:
```python
if transformer.config.in_channels != vae.config.latent_channels:
clip_image_encoder = CLIPModel.from_pretrained(
os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder'))).to(weight_dtype).eval()
input_video, input_video_mask, _ = get_image_to_video_latent(start_image, None, video_length=video_length, sample_size=sample_size)
# pipeline = <Family>InpaintPipeline(..., clip_image_encoder=clip_image_encoder)
# sample = pipeline(..., video=input_video, mask_video=input_video_mask).videos
```
## Pipeline class — `videox_fun/pipeline/pipeline_<family>.py`
```python
from dataclasses import dataclass
from typing import List, Optional, Union
import torch
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from ..models import AutoencoderKL<Family>, <Family>Transformer3DModel
from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler, get_sampling_sigmas
from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
logger = logging.get_logger(__name__)
EXAMPLE_DOC_STRING = """Examples:\n```python\npass\n```"""
# reuse retrieve_timesteps verbatim from pipeline_wan.py
@dataclass
class <Family>PipelineOutput(BaseOutput):
videos: torch.Tensor
class <Family>Pipeline(DiffusionPipeline):
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler):
super().__init__()
self.register_modules(tokenizer=tokenizer, text_encoder=text_encoder, vae=vae,
transformer=transformer, scheduler=scheduler)
# video_processor / vae_scale_factor / etc. as in pipeline_wan.py
def encode_prompt(self, prompt, negative_prompt, device, num_videos_per_prompt=1, ...):
... # mirror pipeline_wan.py
def prepare_latents(self, batch_size, num_channels_latents, height, width, num_frames, dtype, device, generator, latents=None):
...
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(self, prompt, negative_prompt=None, height=480, width=832, num_frames=81,
num_inference_steps=50, guidance_scale=6.0, generator=None, shift=1.0,
callback_on_step_end=None, return_dict=True, **kwargs) -> Union[<Family>PipelineOutput, tuple]:
# 1. encode_prompt 2. prepare_latents 3. retrieve_timesteps
# 4. denoising loop with guidance 5. vae.decode 6. return <Family>PipelineOutput(videos=...)
...
```
Then register in `videox_fun/pipeline/__init__.py`:
```python
from .pipeline_<family> import <Family>Pipeline
```
## Model class — `videox_fun/models/<family>_transformer3d.py`
```python
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
from .attention_utils import attention # unified FA/SDPA backend — do not hand-roll SDPA
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
_supports_gradient_checkpointing = True
@register_to_config
def __init__(self, model_type='t2v', in_dim=16, dim=2048, ffn_dim=8192,
num_heads=16, num_layers=32, in_channels=16, hidden_size=2048, ...):
super().__init__()
...
def _set_gradient_checkpointing(self, *args, **kwargs):
self.gradient_checkpointing = True
def enable_multi_gpus_inference(self): ... # route attn through dist/<family>_xfuser.py
def enable_teacache(self, ...): ...
def enable_cfg_skip(self, ...): ...
def forward(self, x, timestep, context, ...): ...
@classmethod
def from_pretrained(cls, pretrained_model_path, subfolder=None,
transformer_additional_kwargs=None, low_cpu_mem_usage=False,
torch_dtype=torch.bfloat16):
... # mirror wan_transformer3d.py: config.json -> dict_mapping -> init_empty_weights
# -> load .bin/.safetensors -> shape-filter -> initialize missing keys -> load
```
Then register in `videox_fun/models/__init__.py`:
```python
from .<family>_transformer3d import <Family>Transformer3DModel
from .<family>_vae import AutoencoderKL<Family>
```
## Training — `scripts/<family>/train.py` (key reuse points)
```python
"""Modified from https://github.com/huggingface/diffusers/.../train_text_to_image.py"""
import argparse, gc, logging, math, os, sys
import accelerate, diffusers, torch, transformers
from accelerate import Accelerator
from diffusers.optimization import get_scheduler
from omegaconf import OmegaConf
# same sys.path bootstrap as predict scripts
from videox_fun.data import (ASPECT_RATIO_512, AspectRatioBatchImageVideoSampler,
ImageVideoDataset, ImageVideoSampler, RandomSampler,
get_closest_ratio, get_random_mask)
from videox_fun.models import AutoencoderKL<Family>, <Family>Transformer3DModel
from videox_fun.pipeline import <Family>Pipeline # REUSED for validation
from videox_fun.utils.lora_utils import create_network # for train_lora
from videox_fun.utils.utils import save_videos_grid, get_image_to_video_latent
def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config,
accelerator, weight_dtype, global_step):
# build <Family>Pipeline from accelerator.unwrap_model(transformer3d),
# run validation_prompts, save_videos_grid to output_dir/sample/. Reuse the pipeline.
...
def parse_args():
parser = argparse.ArgumentParser(...)
# reuse the shared arg surface: --config_path, --pretrained_model_name_or_path,
# --train_data_dir, --train_data_meta, --video_sample_n_frames, --train_batch_size,
# --gradient_accumulation_steps, --learning_rate, --lr_scheduler, --checkpointing_steps,
# --output_dir, --mixed_precision, --gradient_checkpointing, --enable_bucket,
# --train_mode, --trainable_modules, --validation_prompts ... (add only what's needed)
return parser.parse_args()
def main():
args = parse_args()
accelerator = Accelerator(mixed_precision=args.mixed_precision, ...)
config = OmegaConf.load(args.config_path)
# load transformer/vae/text_encoder via config
# --- Dataset: pick by task (see reference.md §8) ---
# T2V/I2V base + inpaint -> ImageVideoDataset(enable_inpaint = args.train_mode != "normal")
# Control -> ImageVideoControlDataset(enable_camera_info = ...)
# Image edit -> ImageEditDataset
# Speech/audio (S2V) -> VideoSpeechDataset / VideoSpeechControlDataset
# Animate -> VideoAnimateDataset
# Distill text / GRPO / DPO -> TextDataset
# Smoke-test on the matching official demo dataset (reference.md §8), e.g.
# datasets/X-Fun-Videos-Demo + metadata_add_width_height.json for T2V/I2V.
train_dataset = ImageVideoDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride,
video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat,
image_sample_size=args.image_sample_size, enable_bucket=args.enable_bucket,
enable_inpaint=True if args.train_mode != "normal" else False)
# --- Sampler + DataLoader: branch on enable_bucket (see reference.md §8) ---
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
if args.enable_bucket:
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
batch_sampler = AspectRatioBatchImageVideoSampler(
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
aspect_ratios=aspect_ratio_sample_size)
def collate_fn(examples):
new_examples = {"pixel_values": [], "text": []}
if args.train_mode != "normal":
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
# get_closest_ratio -> Resize/CenterCrop/Normalize -> stack; masks via get_random_mask
return new_examples
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
else:
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
# trainable-module filtering or create_network for LoRA
# optimizer + get_scheduler; accelerator.prepare; checkpoint hooks
# training loop: timestep sampling -> transformer forward -> loss -> backward
# periodic log_validation(...); final save weights / LoRA
...
if __name__ == "__main__":
main()
```
## Launcher — `scripts/<family>/train.sh`
```bash
export MODEL_NAME="models/Diffusion_Transformer/<Family>-Model"
# Test data = the official demo dataset matching the task (reference.md §8). Download once, e.g.:
# modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
# T2I -> X-Fun-Images-Demo | control -> X-Fun-{Videos,Images}-Controls-Demo
# S2V -> X-Fun-Videos-Audios-Demo | image edit -> X-Fun-Images-Edit-Demo
export DATASET_NAME="datasets/X-Fun-Videos-Demo/" # = train_data_dir (data_root); media live under train/
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json" # = train_data_meta: [{"file_path","text","type","width","height"}] — see reference.md §8
# Metadata variants: VACE/subject-ref -> metadata_add_width_height_add_objects.json (X-Fun-Videos-Controls-Demo);
# audio-visual joint -> metadata_add_width_height_add_wav.json; lingbot_video -> metadata_lingbot_video_add_width_height.json
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/<family>/train.py \
--config_path="config/<family>/<variant>.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--checkpointing_steps=50 \
--output_dir="output_dir_<family>" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--enable_bucket \
--low_vram \
--train_mode="normal" \
--trainable_modules "."
```
## Preprocessing (data gen) — `scripts/<family>/generate_<...>.py`
Offline generation of cached training data (latents / ODE-trajectory pairs / prompt embeddings). **Always multi-GPU** (`accelerate launch` + `Accelerator`) and **always safetensors** (`safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`) — never LMDB, never `.pt`. Mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py`:
```python
# ...license header + sys.path bootstrap...
import argparse, json, math, os, torch
from accelerate import Accelerator
from omegaconf import OmegaConf
from safetensors.torch import save_file
from tqdm import tqdm
from videox_fun.models import AutoencoderKLWan, WanT5EncoderModel, WanTransformer3DModel # reuse repo models
from videox_fun.utils.utils import save_videos_grid # reuse repo IO
def main():
args = parse_args() # --pretrained_model_name_or_path --config_path --caption_path --output_folder
# --num_inference_steps --guidance_scale --shift --mixed_precision ...
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
torch.set_grad_enabled(False) # inference-only
torch.backends.cuda.matmul.allow_tf32 = True
config = OmegaConf.load(args.config_path) # config-driven loading (Section 3)
weight_dtype = {"fp16": torch.float16, "bf16": torch.bfloat16}.get(accelerator.mixed_precision, torch.float32)
text_encoder = WanT5EncoderModel.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), torch_dtype=weight_dtype).to(device).eval()
vae = AutoencoderKLWan.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['vae_kwargs'])).to(device, dtype=weight_dtype).eval()
transformer = WanTransformer3DModel.from_pretrained(..., transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])).to(device, dtype=weight_dtype).eval()
prompts = [l.rstrip() for l in open(args.caption_path, encoding="utf-8") if l.strip()]
os.makedirs(args.output_folder, exist_ok=True)
total_per_rank = math.ceil(len(prompts) / world_size)
for index in tqdm(range(total_per_rank), disable=rank != 0, desc="Generating"):
prompt_index = index * world_size + rank # interleaved multi-GPU shard
if prompt_index >= len(prompts):
continue
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(out_path): # resume: skip already-done samples
continue
prompt = prompts[prompt_index]
# ... encode prompt, sample noise, run the teacher ODE (CFG), collect latents ...
save_file( # safetensors ONLY (no lmdb / no .pt)
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
out_path, metadata={"prompt": prompt},
)
accelerator.wait_for_everyone()
if accelerator.is_main_process: # rank-0 writes the JSON index
entries = [{"file_path": os.path.join(args.output_folder, f"{i:05d}.safetensors")}
for i in range(len(prompts))
if os.path.exists(os.path.join(args.output_folder, f"{i:05d}.safetensors"))]
json.dump(entries, open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
if __name__ == "__main__":
main()
```
Launcher (`generate_<...>.sh`) — `accelerate launch` uses every visible GPU:
```bash
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/<family>/*.yaml" \
--caption_path="datasets/prompts.txt" \
--output_folder="datasets/<family>_ode_pairs" \
--num_inference_steps=48 --guidance_scale=6.0 --shift=8.0
```
Training then reads the cache with `ImageVideoSafetensorsDataset(ann_path=".../outputs.json")` (single-file mode `{"file_path": ...}`, or per-tensor mode via `--save_per_tensor`). See reference.md §10.
> Dataset *curation* (scoring/filtering/captioning under `videox_fun/video_caption/`) is a different activity: also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM tensor-parallel) but writes csv/jsonl metadata, not safetensors. See reference.md §10 “Related but different”.
+414
View File
@@ -0,0 +1,414 @@
# VideoX-Fun Integration Reference
Detailed conventions per layer. Read the mirrored family's real files alongside this — the existing code is always the source of truth.
## 1. Model definitions — `videox_fun/models/<family>_*.py`
### File naming
- Transformer / DiT: `<family>_transformer3d.py` (video) or `<family>_transformer2d.py` (image). Variants append a suffix: `_control`, `_s2v`, `_vace`, `_animate`, `_self_forcing`, `_avatar`.
- VAE: `<family>_vae.py` → class `AutoencoderKL<Family>`.
- Encoders: `<family>_text_encoder.py`, `<family>_audio_encoder.py`, `<family>_image_encoder.py`.
### Class shape (mirror `wan_transformer3d.py`)
```python
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders.single_file_model import FromOriginalModelMixin
from diffusers.models.modeling_utils import ModelMixin
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
_supports_gradient_checkpointing = True
@register_to_config
def __init__(self, model_type='t2v', patch_size=(1,2,2), in_dim=16, dim=2048,
ffn_dim=8192, num_heads=16, num_layers=32, in_channels=16,
hidden_size=2048, ...):
super().__init__()
...
```
- Keep BOTH civitai names (`in_dim`, `dim`, `ffn_dim`) and diffusers aliases (`in_channels`, `hidden_size`) in `__init__` so either format maps cleanly.
- Implement `_set_gradient_checkpointing(self, *args, **kwargs)`.
- Attention must go through `videox_fun.models.attention_utils.attention` (backend-agnostic), not a hand-rolled `scaled_dot_product_attention`.
- Multi-GPU: expose `enable_multi_gpus_inference()` and route attention through the family's `dist/<family>_xfuser.py` processor.
- Speedups live on the model: `enable_teacache(...)`, `enable_cfg_skip(...)`, `enable_riflex(...)`.
### `from_pretrained` internals (do not simplify)
The custom classmethod must keep these behaviors (see `wan_transformer3d.py::from_pretrained`):
1. Accept `transformer_additional_kwargs`, `subfolder`, `low_cpu_mem_usage`, `torch_dtype`.
2. Read `config.json`; auto-convert foreign configs (e.g. diffsynth `has_image_input`) via a `_convert_from_*_config` helper.
3. Apply `dict_mapping`: pop it from kwargs, then for each `key: target` set `kwargs[target] = config[key]`.
4. Under `low_cpu_mem_usage`, build with `accelerate.init_empty_weights()`, load `.bin`/`.safetensors` (single file or glob all shards), and **filter by exact shape match** before loading.
5. Initialize missing keys deliberately: zero-init control/audio projections (`after_proj`, `before_proj`, `processor.k_proj/v_proj`, `audio_injector`, `cond_encoder`, ...), ones for norms, xavier for ≥2D weights, so new branches start as no-ops.
### Registry — `videox_fun/models/__init__.py`
Add an import line for every new public class, grouped with the family. Wrap optional-dependency imports in `try/except` with a helpful upgrade message (see the Qwen2.5-VL / Mistral3 blocks at the top).
## 2. Pipelines — `videox_fun/pipeline/pipeline_<family>*.py`
Mirror `pipeline_wan.py`. Required pieces:
- Module-level `retrieve_timesteps(scheduler, num_inference_steps, device, timesteps, sigmas, **kwargs)` (copied from diffusers) — reuse verbatim.
- `EXAMPLE_DOC_STRING` for the `@replace_example_docstring` decorator.
- Output dataclass:
```python
@dataclass
class <Family>PipelineOutput(BaseOutput):
videos: torch.Tensor
```
- Pipeline class:
```python
class <Family>Pipeline(DiffusionPipeline):
_optional_component = [...]
model_cpu_offload_seq = "text_encoder->transformer->vae" # order matters for offload
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler, ...): ...
def encode_prompt(...): ...
def prepare_latents(...): ...
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(self, prompt, negative_prompt=..., height=..., width=...,
num_frames=..., num_inference_steps=..., guidance_scale=...,
generator=None, ..., return_dict=True) -> Union[<Family>PipelineOutput, Tuple]: ...
```
- Import schedulers from `..utils.fm_solvers` / `..utils.fm_solvers_unipc`, models from `..models`.
- Separate pipelines per task: base (`pipeline_<family>.py`), inpaint/i2v (`_inpaint`), control (`_control`), s2v, etc. Register all in `videox_fun/pipeline/__init__.py`, adding convenience aliases (e.g. `WanI2VPipeline = WanFunInpaintPipeline`) where existing code expects them.
## 3. Config — `config/<family>/<name>.yaml` (optional)
**The YAML is not mandatory.** Decide by checkpoint layout:
- **Required** for civitai-format / custom single-file layouts, where weights and key names are not diffusers-native. The YAML supplies `transformer_additional_kwargs` (incl. `dict_mapping` mapping civitai config keys → model `__init__` kwargs), component `*_subpath`s, and `vae/text_encoder/scheduler/image_encoder` kwargs.
- **Optional** for a standard diffusers-layout checkpoint (`model_index.json` + each subfolder carrying its own `config.json`). Load components directly: `<Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", low_cpu_mem_usage=True, torch_dtype=...)`, `AutoencoderKL<Family>.from_pretrained(model_name, subfolder="vae")`, etc. Guard the config path exactly like `examples/minimax_h3_fun/predict_v2v_control.py`:
```python
transformer_load_kwargs = {}
if config_path is not None:
from omegaconf import OmegaConf
config = OmegaConf.load(config_path)
transformer_load_kwargs.update(OmegaConf.to_container(config["transformer_additional_kwargs"], resolve=True))
transformer = <Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", **transformer_load_kwargs, ...)
```
When you do use a YAML, the canonical schema is below (see `config/wan2.1/wan_civitai.yaml`):
```yaml
format: civitai # or diffusers — selects weight-key handling
pipeline: Wan # family label consumed by API/ComfyUI loaders
transformer_additional_kwargs:
transformer_subpath: ./ # subfolder under model_name holding the DiT
dict_mapping: # civitai config key -> model __init__ kwarg
in_dim: in_channels
dim: hidden_size
vae_kwargs:
vae_subpath: Wan2.1_VAE.pth
temporal_compression_ratio: 4
spatial_compression_ratio: 8
text_encoder_kwargs:
text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
tokenizer_subpath: google/umt5-xxl
text_length: 512
...
scheduler_kwargs:
scheduler_subpath: null
num_train_timesteps: 1000
shift: 5.0
...
image_encoder_kwargs: # only for i2v / models with a CLIP image encoder
image_encoder_subpath: models_clip_...pth
```
Every `*_subpath` is joined onto `model_name` in scripts. Load with `OmegaConf.load` and pass `OmegaConf.to_container(config['<section>'])` into `from_pretrained`. Use `filter_kwargs(Cls, OmegaConf.to_container(config['scheduler_kwargs']))` to build schedulers.
## 4. Inference scripts — `examples/<family>/predict_<task>.py`
Anatomy, top to bottom (see `examples/wan2.1_fun/predict_t2v.py`):
1. **`sys.path` bootstrap** (before importing `videox_fun`):
```python
current_file_path = os.path.abspath(__file__)
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
for project_root in project_roots:
sys.path.insert(0, project_root) if project_root not in sys.path else None
```
2. **User config block** as top-level variables with explanatory comments, in the conventional order: `GPU_memory_mode`, `ulysses_degree`/`ring_degree`, `fsdp_dit`/`fsdp_text_encoder`, `compile_dit`, TeaCache (`enable_teacache`, `teacache_threshold`, `num_skip_start_steps`, `teacache_offload`), `cfg_skip_ratio`, Riflex (`enable_riflex`, `riflex_k`), `config_path`, `model_name`, `sampler_name`, `shift`, `transformer_path`/`vae_path`/`lora_path`, `sample_size`, `video_length`, `fps`, `weight_dtype`, `prompt`/`negative_prompt`, `guidance_scale`, `seed`, `num_inference_steps`, `lora_weight`, `save_path`.
3. **Device + config**: `device = set_multi_gpus_devices(ulysses_degree, ring_degree)`; then either `config = OmegaConf.load(config_path)` (civitai/custom layout) **or** guard `if config_path is not None:` and load components directly from a diffusers-layout checkpoint (see §3).
4. **Component loading**: transformer (`from_pretrained(..., transformer_additional_kwargs=...)`), optional `transformer_path`/`vae_path` override with `load_state_dict(strict=False)` + missing/unexpected key print, vae, tokenizer, text_encoder, and clip image encoder gated by `transformer.config.in_channels != vae.config.latent_channels`.
5. **Scheduler selection dict**: `{"Flow": FlowMatchEulerDiscreteScheduler, "Flow_Unipc": FlowUniPCMultistepScheduler, "Flow_DPM++": FlowDPMSolverMultistepScheduler}[sampler_name]`; build with `filter_kwargs`.
6. **Pipeline construction**: choose base vs inpaint/i2v/control pipeline by the model's channel condition.
7. **Multi-GPU / FSDP / compile**: if `ulysses_degree>1 or ring_degree>1` call `transformer.enable_multi_gpus_inference()` and optionally `shard_model`; if `compile_dit`, `torch.compile` each `transformer.blocks[i]`.
8. **`GPU_memory_mode` branching** — keep this exact order:
```python
if GPU_memory_mode == "sequential_cpu_offload":
replace_parameters_by_name(transformer, ["modulation",], device=device)
transformer.freqs = transformer.freqs.to(device=device)
pipeline.enable_sequential_cpu_offload(device=device)
elif GPU_memory_mode == "model_group_offload":
register_auto_device_hook(pipeline.transformer)
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_cpu_offload":
pipeline.enable_model_cpu_offload(device=device)
elif GPU_memory_mode == "model_full_load_and_qfloat8":
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
convert_weight_dtype_wrapper(transformer, weight_dtype)
pipeline.to(device=device)
else:
pipeline.to(device=device)
```
9. **TeaCache / cfg_skip / Riflex** enablement, `generator = torch.Generator(device).manual_seed(seed)`, LoRA `merge_lora`.
10. **Inference** under `torch.no_grad()`; align `video_length` to `vae.config.temporal_compression_ratio`; pass `video`/`mask_video` for i2v via `get_image_to_video_latent`.
11. **`save_results()`**: `save_videos_grid(sample, path, fps=fps)` for video, PIL save for a single frame; only rank 0 saves when multi-GPU. LoRA `unmerge_lora` after.
Other entry points to mirror when needed: `app.py` (Gradio), `launch_api.py` (API server backed by `videox_fun/api`), `post_infer*.py` (batch/queue inference).
## 5. Training scripts — `scripts/<family>/train*.py`
Mirror `scripts/wan2.1_fun/train.py`. Structure:
1. Diffusers-derived license header + `"""Modified from ..."""` note.
2. Third-party imports, then the **same `sys.path` bootstrap**, then `from videox_fun.data/models/pipeline/utils import ...`.
3. Helper funcs: `filter_kwargs`, `resize_mask`, `linear_decay`, `generate_timestep_with_lognorm`.
4. **`log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step)`** — builds the **inference Pipeline** from the live (unwrapped) transformer and runs it to produce sample videos under `output_dir/sample/`. Wrapped in try/except; handles DeepSpeed (`transformer3d.config` swap) and restores VAE/text-encoder placement (`low_vram`). **Reuse the pipeline; never write a separate sampler.**
5. **`parse_args()`** — reuse the shared argument surface: `--config_path`, `--pretrained_model_name_or_path`, `--train_data_dir`, `--train_data_meta`, `--image_sample_size`/`--video_sample_size`/`--token_sample_size`, `--video_sample_n_frames`, `--video_sample_stride`, `--train_batch_size`, `--gradient_accumulation_steps`, `--learning_rate`, `--lr_scheduler`, `--lr_warmup_steps`, `--checkpointing_steps`, `--output_dir`, `--mixed_precision`, `--gradient_checkpointing`, `--enable_bucket`, `--random_hw_adapt`, `--training_with_video_token_length`, `--uniform_sampling`, `--low_vram`, `--train_mode`, `--trainable_modules`, LoRA args (`--use_lora`, `--rank`, ...), `--validation_prompts`/`--validation_paths`. Add new args only when the family genuinely needs them.
6. **`main()`** — Accelerator setup, DeepSpeed/FSDP zero-stage handling (auto-sets `save_state`), model loading via config, dataset + bucket sampler from `videox_fun.data`, trainable-module filtering / LoRA network via `create_network`, optimizer + `get_scheduler`, `accelerator.prepare`, checkpoint save/load hooks, training loop with timestep sampling, loss, `log_validation` at intervals, and final weight/LoRA save.
### Resolution args — `--video_sample_size` (+ `--fix_sample_size`)
Canvas resolution is always driven by a **single square** `--video_sample_size` (`type=int`, height = width) — never by separate `--video_sample_height` / `--video_sample_width`. When a **fixed non-square shape** is required, add `--fix_sample_size` (`nargs=2, type=int, default=None`, `[height, width]`) that overrides the square size; mirror `scripts/wan2.2_fun/train_lora.py`, `scripts/z_image/train_distill.py`. Derive the effective `height` / `width` once in `parse_args()` and reuse them everywhere downstream:
```python
parser.add_argument("--video_sample_size", type=int, default=1280)
parser.add_argument("--fix_sample_size", nargs=2, type=int, default=None,
help="Fix Sample size [height, width] to override `--video_sample_size` with a fixed non-square shape.")
...
if args.fix_sample_size is not None:
args.video_sample_height, args.video_sample_width = args.fix_sample_size
else:
args.video_sample_height = args.video_sample_width = args.video_sample_size
```
In bucket datasets `--fix_sample_size` also forces `random_hw_adapt=False` / `training_with_video_token_length=False` and bumps `video_sample_size = max(max(fix_sample_size), video_sample_size)`; in data-free scripts (e.g. `scripts/minimax_h3/train_pdd_lora.py`) it simply pins the generation canvas. Always validate the size against the patch/VAE constraint (minimax_h3: `% 32`). The `.sh` launcher passes it space-separated (`nargs=2`): `--fix_sample_size 768 1344`.
### Launcher — `scripts/<family>/train*.sh`
`export MODEL_NAME/DATASET_NAME/DATASET_META_NAME`, then `accelerate launch --mixed_precision="bf16" scripts/<family>/train.py --config_path=... <full arg list>`. Include commented I2V/control variants and DeepSpeed/NCCL notes as the existing scripts do.
### Docs — `README_TRAIN.md` + `README_TRAIN_zh-CN.md`
Aligned bilingual pair: identical section order, identical commands and parameter tables; only the prose language differs. Follow the top-level section order used across existing training READMEs.
## 6. Shared infrastructure map (reuse, never reimplement)
| Need | Import from |
|------|-------------|
| Flow/DPM/UniPC schedulers | `diffusers`, `videox_fun.utils.fm_solvers`, `videox_fun.utils.fm_solvers_unipc` |
| LoRA create/merge/unmerge/convert | `videox_fun.utils.lora_utils` |
| FP8 quantization | `videox_fun.utils.fp8_optimization` |
| Group / leaf offload hooks | `videox_fun.utils.group_offload` |
| Multi-GPU device + FSDP shard + seq-parallel attn | `videox_fun.dist` |
| Save video/audio, image→video latents, kwarg filter, dimension calc | `videox_fun.utils.utils` |
| Datasets + bucket/aspect-ratio samplers + masks | `videox_fun.data` |
| TeaCache coefficients | `videox_fun.models.cache_utils` |
## 7. Naming quick reference
| Concept | Convention | Example |
|---------|-----------|---------|
| Model file | `<family>_transformer3d.py` | `wan_transformer3d.py` |
| Model class | `<Family>Transformer3DModel` | `WanTransformer3DModel` |
| VAE class | `AutoencoderKL<Family>` | `AutoencoderKLWan` |
| Pipeline file | `pipeline_<family>.py` | `pipeline_wan.py` |
| Pipeline class | `<Family>Pipeline` | `WanPipeline` / `WanFunInpaintPipeline` |
| Config | `config/<family>/<variant>.yaml` | `config/wan2.1/wan_civitai.yaml` |
| Inference | `examples/<family>/predict_<task>.py` | `predict_t2v.py`, `predict_i2v.py`, `predict_v2v_control.py` |
| Training | `scripts/<family>/train[_<variant>].py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py` |
## 8. Training data pipeline — dataset & sampler selection
Pick the dataset by **task / `train_mode`**, then the sampler by **`enable_bucket`** and dataset type. All datasets/samplers come from `videox_fun.data` — never write a new one.
### Annotation format — the `train_data_meta` file (`metadata.json` / `.csv`)
Every dataset class reads an annotation file (`args.train_data_meta`) that indexes the media under `args.train_data_dir` (`data_root`). `ImageVideoDataset` accepts **`.json`** (a top-level array of records) or **`.csv`** (`csv.DictReader`; the header row is the field names). Each record for ordinary image/video training:
| Field | Required | Meaning |
|-------|----------|---------|
| `file_path` | yes | Media path, resolved **relative to `train_data_dir`** via `os.path.join(data_root, file_path)`. If `data_root is None`, `file_path` is used as-is. |
| `text` | yes | Caption / prompt. Dropped to `""` with probability `text_drop_ratio` (default `0.1`) for classifier-free guidance. |
| `type` | no | `"video"` or `"image"`; **defaults to `"image"`** when the key is absent (`data_info.get('type', 'image')`). |
```json
[
{"file_path": "train/00000000.mp4", "text": "A young woman gently turns her head to the right ...", "type": "video"},
{"file_path": "train/00000001.jpg", "text": "a dog running on the beach", "type": "image"}
]
```
The directory layout matches the index — media in a `train/` subdir, the annotation file beside it. Ready-made examples ship in `datasets/X-Fun-Videos-Demo/` (`train/*.mp4` + `metadata.json`) and `datasets/X-Fun-Images-Demo/`. The equivalent `.csv`:
```csv
file_path,text,type
train/00000000.mp4,"A young woman gently turns her head to the right ...",video
train/00000001.jpg,"a dog running on the beach",image
```
**Variant datasets append extra fields to this same record shape**, each consumed by its own class (see the table below) — e.g. camera-pose adds `action_path` (`LingbotImageVideoDataset`), object/VACE/S2V variants add object fields (`object_file_path` / `objects`). The demo folders also ship several augmented metadata variants (next subsection). Always read the target class's `get_batch` for the exact fields it consumes.
### Ready-made demo datasets — the standard test data (never invent a test set)
Smoke tests, `log_validation` checks, and doc examples all run on the official demo datasets under `datasets/`, downloaded from ModelScope as `PAI/<name>`:
```bash
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
```
Pick the demo by **task**, matching the dataset class in the table below:
| Demo dataset (`datasets/...`) | Contents | Extra metadata fields | Task it tests | Dataset class |
|-------------------------------|----------|----------------------|---------------|---------------|
| `X-Fun-Videos-Demo` | 16 videos (832×480) in `train/` | — | T2V / I2V base + inpaint, distill | `ImageVideoDataset` |
| `X-Fun-Videos-Controls-Demo` | 16 videos in `train/` + `canny/` + `object/<video_id>/` + `wav/` | `control_file_path`, `object_file_path` (list), `audio_path` | V2V control, VACE, S2V-with-control | `ImageVideoControlDataset`, `VideoSpeechControlDataset` |
| `X-Fun-Videos-Audios-Demo` | 17 video/audio pairs: `train/` (1280×720) + `wav/` (16 kHz mono) + `pose/` | `audio_path`, `control_file_path` | Speech-driven S2V / avatar / talking-head | `VideoSpeechDataset` |
| `X-Fun-Images-Demo` | 19 images in `train/` | — | T2I full fine-tune + LoRA (z_image / flux2 / qwenimage / lens / ernie) | `ImageVideoDataset` |
| `X-Fun-Images-Controls-Demo` | 19 images in `train/` + `canny/` | `control_file_path` | Image control / ControlNet / i2i inpaint | `ImageVideoControlDataset` |
| `X-Fun-Images-Edit-Demo` | 21 records: `source/souce-<id>/` (multi-source supported) → `train/` | `source_file_path` (**list**) | Image edit (Qwen-Image-Edit family) | `ImageEditDataset` |
| `X-Fun-Videos-Lingbot-Demo` | video + `intrinsics.npy` / `poses.npy` | camera pose / action | Camera-pose world model (`lingbot_world`) | `LingbotImageVideoDataset` |
**Which metadata file to point `--train_data_meta` at** (each demo ships several variants beside the media):
| Metadata file | Use when |
|---------------|----------|
| `metadata.json` | Base format only (`file_path` / `text` / `type`) — fine for a minimal check |
| `metadata_add_width_height.json` | **Default choice.** Adds `width` / `height` so bucketing doesn't decode media (matters on slow storage such as OSS). Used by non-VACE control / S2V training too |
| `metadata_add_width_height_add_objects.json` | VACE / subject-reference training (`object_file_path` list → `object/<video_id>/`; shuffled at train time) |
| `metadata_add_width_height_add_wav.json` | Audio-visual joint models (e.g. `minimax_h3_fun` control training): `audio_path` → `wav/`. Keep the `.sh` launcher and the README on the same file |
| `metadata_lingbot_video_add_width_height.json` | `lingbot_video` — `text` is already a structured JSON caption (lives in `X-Fun-Videos-Demo`) |
| `metadata_origin.json` | Pre-processing original kept for reference; not used for training |
Regenerate the width/height variant with the shipped helper when adding your own media:
`python scripts/process_json_add_width_and_height.py --input_file datasets/<Demo>/metadata.json --output_file datasets/<Demo>/metadata_add_width_height.json`.
`audio_path` optionality differs per class (`videox_fun/data/dataset_video.py`): `VideoSpeechDataset` reads `video_dict['audio_path']` directly, so it is **required**; `VideoSpeechControlDataset` uses `.get('audio_path')` and **falls back to the video file's own audio track** when the field is absent.
### Dataset by task (all take `train_data_meta, train_data_dir, ...`)
| Task / mode | Dataset class | Used by | Key kwargs |
|-------------|--------------|---------|-----------|
| T2V / I2V base (`normal` + inpaint) | `ImageVideoDataset` | `train.py`, `train_lora.py`, t2i `train.py` | `enable_inpaint = train_mode != "normal"`, `video_sample_size/stride/n_frames`, `image_sample_size`, `video_repeat` |
| Image T2I (qwenimage/flux/z_image) | `ImageVideoDataset` | `scripts/<img>/train.py` | `image_sample_size` |
| Control (canny/pose/depth/camera) | `ImageVideoControlDataset` | `train_control*.py`, `train_control_distill.py` | `enable_camera_info = train_mode == "control_camera_ref"` |
| Image Edit (source→target) | `ImageEditDataset` | `qwenimage/train_edit*.py` | `image_sample_size` |
| Speech/audio-driven (S2V, avatar, talking) | `VideoSpeechDataset` | `mova`, `ltx2`, `minimax_h3`, `fantasytalking`, `infinitetalk`, `flashhead`, `longcatvideo/train_avatar*` | audio + video fields |
| S2V **with control** | `VideoSpeechControlDataset` | `wan2.2/train_s2v*.py`, `minimax_h3_fun/train_control*` | audio + control |
| Motion/pose animate | `VideoAnimateDataset` | `wan2.2/train_animate*.py` | motion/pose driven |
| Distill text-only branch, GRPO, DPO | `TextDataset` | `train_distill*.py` (text branch), `z_image/train_grpo_lora.py`, `train_dpo_lora.py` | reads only the `text` field; `text_drop_ratio` |
| Precomputed latents (ODE pairs) | `ImageVideoSafetensorsDataset` | `wan2.1_self_forcing/train_ode.py` | `data_root` |
| Camera-pose conditioning | `LingbotImageVideoDataset` | `lingbot_world/train.py` | `intrinsics.npy` / `poses.npy` |
| Video-only (VAE/TAEHV distill) | `VideoDataset` | `taehv/train_taehv.py` | `sample_size/stride/n_frames`, `enable_inpaint=False` |
### Sampler by condition
| Condition | Sampler | Shape |
|-----------|---------|-------|
| `enable_bucket=True` (default; image+video) | `AspectRatioBatchImageVideoSampler` | `sampler=RandomSampler(ds, generator=g), dataset=train_dataset.dataset, batch_size, train_folder=args.train_data_dir, drop_last=True, aspect_ratios=aspect_ratio_sample_size` |
| `enable_bucket=False` | `ImageVideoSampler` | `ImageVideoSampler(RandomSampler(ds, generator=g), train_dataset, batch_size)` |
| `TextDataset` (distill text branch / GRPO / DPO) | `BatchSampler` (plain) | `BatchSampler(RandomSampler(ds, generator=g), batch_size, drop_last=True)`; GRPO adds `k_repeat=args.num_image_per_prompt` |
| video-only bucket (available, not used by current scripts) | `AspectRatioBatchSampler` | — |
| image-only bucket (available, not used by current scripts) | `AspectRatioBatchImageSampler` | — |
`aspect_ratio_sample_size` is built from `ASPECT_RATIO_512` scaled by `args.video_sample_size`; `get_closest_ratio` picks the bucket inside `collate_fn`.
### Universal DataLoader creation pattern
```python
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
if args.enable_bucket:
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
batch_sampler = AspectRatioBatchImageVideoSampler(
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
aspect_ratios=aspect_ratio_sample_size)
def collate_fn(examples):
new_examples = {"pixel_values": [], "text": []}
if args.train_mode != "normal": # inpaint/i2v adds mask fields
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
# bucket via get_closest_ratio -> transform (Resize/CenterCrop/Normalize) -> stack
# masked branch uses get_random_mask(...)
return new_examples
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
else:
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
train_dataloader = torch.utils.data.DataLoader(
train_dataset, batch_sampler=batch_sampler,
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
```
`collate_fn` receives the `examples` **list** (not a `batch` dict); build every batch-level field (`text`, `pixel_values`, masks) explicitly from `examples` into `new_examples`. When `--enable_text_encoder_in_dataloader`, encode prompts inside `collate_fn` and emit `encoder_hidden_states` / `encoder_attention_mask`.
## 9. Inference task matrix — predict script → pipeline → inputs
Pick the pipeline by **task**; the `predict_<task>.py` name and its inputs follow the same convention across families.
| Task | `predict_<task>.py` | Pipeline (family example) | Extra `__call__` inputs | Input helper |
|------|--------------------|---------------------------|-------------------------|--------------|
| Text→Video | `predict_t2v.py` | `WanPipeline`, `Wan2_2Pipeline`, `CogVideoXFunPipeline`, `LongCatVideoPipeline`, `LTX2Pipeline` | `prompt` only | — |
| Image→Video | `predict_i2v.py` | `WanI2VPipeline`(=`WanFunInpaintPipeline`), `Wan2_2FunInpaintPipeline`, `Wan2_2I2VPipeline`, `HunyuanVideoI2VPipeline` | `video`, `mask_video` | `get_image_to_video_latent(start_image, end_image, video_length, sample_size)` |
| Text+Image→Video (5B) | `predict_ti2v.py` | `Wan2_2TI2VPipeline` | `prompt` (+ optional image) | `get_image_to_video_latent` |
| Video→Video Control | `predict_v2v_control.py` | `WanFunControlPipeline`, `Wan2_2FunControlPipeline` | `control_video` | `get_video_to_video_latent(control_video, ...)` |
| Control + reference | `predict_v2v_control_ref.py` | `WanFunControlPipeline` | `control_video` + `ref_image` | `get_video_to_video_latent` + `get_image_latent` |
| Control + camera | `predict_v2v_control_camera.py` | `WanFunControlPipeline` | `control_video` + camera pose | — |
| VACE (control/mask/i2v/s2v) | `predict_v2v_control.py`, `predict_v2v_mask.py`, `predict_s2v.py`, `predict_i2v.py` | `WanVacePipeline`, `Wan2_2VaceFunPipeline` | control/mask/ref | — |
| Speech→Video (audio) | `predict_s2v.py` | `Wan2_2S2VPipeline`, `MiniMaxH3Pipeline`, `InfiniteTalkPipeline`, `FantasyTalkingPipeline`, `FlashHeadPipeline`, `MOVAPipeline`, `LongCatVideoAvatarPipeline` | `audio` + reference image | — |
| Animate (motion/pose) | `predict_animate.py` | `Wan2_2AnimatePipeline` | motion/pose video + ref | — |
| Subject reference | `predict_s2v.py` (phantom) | `WanFunPhantomPipeline` | reference images | — |
| Text→Image | `predict_t2i.py` | `QwenImagePipeline`, `Flux2Pipeline`, `ZImagePipeline`, `LensPipeline`, `ErnieImagePipeline` | `prompt` | — |
| Image Control (t2i) | `predict_t2i_control.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline`, `QwenImageControlNetPipeline` | `control_image` | — |
| Inpaint (i2i) | `predict_i2i_inpaint.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline` | `image` + `mask` | — |
| Image Edit | `predict_t2i_edit.py`, `predict_t2i_edit_plus.py` | `QwenImageEditPipeline`, `QwenImageEditPlusPipeline` | source image + instruction | — |
| Layered edit | `predict_i2i_layered.py` | `QwenImageLayeredPipeline` | image | — |
| Camera-pose world | `predict_i2v.py` (lingbot_world) | `Wan2_2I2VPipeline`, `WanFunLingbotWorldFastPipeline` | image + camera pose | — |
| Latent upsample | `predict_i2v_upsample.py` | `LTX2LatentUpsamplePipeline`, `WanLatentUpsamplePipeline` | low-res latent/video | — |
| AR / streaming distill | `predict_t2v_stream.py` | `WanSelfForcingPipeline` | prompt (streamed) | — |
### Predict-script variant suffixes (same task, different backend/model)
| Suffix | Meaning |
|--------|---------|
| `_tae` | Fast decode via `AutoencoderTinyWan` (TAEHV) instead of the full VAE |
| `_2.2vae` | Uses the Wan2.2 VAE (`AutoencoderKLWan3_8`) |
| `_5b` | 5B-parameter model variant |
| `turbo` / distill | Distilled model, few-step inference (e.g. `predict_turbo_*.py`) |
| `_refine` | Two-stage refine pass |
| `_ref` / `_camera` | Adds reference-image / camera conditioning |
All variants keep the identical config block, `GPU_memory_mode` branching, and `save_results()` from Section 4 — only the loaded VAE/transformer and pipeline class change.
## 10. Preprocessing — offline training-data generation (multi-GPU + safetensors)
Here "preprocessing" means **generating/caching training data offline** with the teacher / VAE / text-encoder — latents, ODE-trajectory pairs, prompt/text embeddings — so training just reads cached tensors instead of re-encoding every step. Canonical example: `scripts/wan2.1_self_forcing/generate_ode_pairs.py` (+ `generate_ode_pairs.sh`); the loader-side contract is `ImageVideoSafetensorsDataset` in `videox_fun/data/dataset_image_video.py`. Two rules are non-negotiable.
### Rule 1 — multi-GPU is mandatory
Never a single-GPU / hardcoded `cuda:0` loop. Launch with `accelerate launch` and shard work across ranks by interleaving:
```python
from accelerate import Accelerator
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
torch.set_grad_enabled(False) # inference-only
total_per_rank = math.ceil(len(prompts) / world_size)
for index in tqdm(range(total_per_rank), disable=rank != 0):
prompt_index = index * world_size + rank # interleaved shard
if prompt_index >= len(prompts):
continue
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(out_path): # resume-friendly
continue
... # encode prompt / run teacher ODE / collect latents
accelerator.wait_for_everyone()
if accelerator.is_main_process: # write the JSON index once, on rank 0
json.dump([{"file_path": p} for p in all_safetensor_paths],
open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
```
Launcher (`.sh`): `accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py --pretrained_model_name_or_path=... --config_path=config/<family>/*.yaml --output_folder=datasets/<...> ...`. Reuse `videox_fun.models` + config-driven `from_pretrained` (Section 3) and `videox_fun.utils.utils.save_videos_grid` for sample previews — do not write a new loader.
### Rule 2 — store as safetensors; do NOT use LMDB or `.pt`
Save every cached tensor with `safetensors.torch.save_file`, one `.safetensors` per sample (or per tensor), plus a JSON index of `{"file_path": ...}` entries:
```python
from safetensors.torch import save_file
save_file(
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
out_path, # f"{prompt_index:05d}.safetensors"
metadata={"prompt": prompt},
)
```
`ImageVideoSafetensorsDataset(ann_path, data_root=None)` reads that JSON and supports two layouts:
- **Single-file (default)**: `{"file_path": "scene.safetensors"}` — whole state dict in one archive.
- **Per-tensor (`--save_per_tensor`)**: `{"file_path": "scene_dir", "latents": ".../latents.safetensors", "prompt_embeds": ".../prompt_embeds.safetensors"}` — each key loaded and merged.
**Do not** cache preprocessed data in **LMDB** or as **`.pt`/`.pth` `torch.save` pickles**. safetensors is the repo-wide standard (also used for LoRA/weight saving), is pickle-free/safe, memory-maps fast, and is exactly what `ImageVideoSafetensorsDataset` loads. (Scope: this governs cached *data tensors*; accelerate optimizer/scheduler/scaler `.pt` states written during training checkpoints are a separate mechanism and unaffected.)
### Related but different — dataset curation
Scoring / filtering / captioning under `videox_fun/video_caption/` (`compute_*.py`, `internvl2_video_recaptioning.py`) is dataset *curation*, not latent caching. It is also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM `tensor_parallel_size=device_count()`), but writes csv/jsonl **metadata** (not tensors), so Rule 2 does not apply there.
+18 -1
View File
@@ -62,10 +62,12 @@ 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.
# 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
@@ -136,6 +138,16 @@ if transformer_path is not None:
"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(
@@ -245,6 +257,10 @@ if lora_path is not None:
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,
@@ -259,6 +275,7 @@ with torch.no_grad():
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)
+18 -1
View File
@@ -65,11 +65,13 @@ model_name = "models/Diffusion_Transformer/MiniMax-H3"
# 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.
# `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
@@ -143,6 +145,16 @@ if transformer_path is not None:
"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(
@@ -268,6 +280,10 @@ def parse_reference(entry: str):
# 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,
@@ -281,6 +297,7 @@ with torch.no_grad():
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)
+18 -1
View File
@@ -61,10 +61,12 @@ 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.
# 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
@@ -129,6 +131,16 @@ if transformer_path is not None:
"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(
@@ -235,6 +247,10 @@ 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,
@@ -247,6 +263,7 @@ with torch.no_grad():
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)
@@ -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,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}")
@@ -3,8 +3,10 @@ import sys
import numpy as np
import torch
import transformers
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from packaging.version import Version
from PIL import Image
current_file_path = os.path.abspath(__file__)
@@ -137,8 +139,14 @@ if vae_path is not None:
tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
# `Qwen3ForCausalLM.from_pretrained` renamed `torch_dtype` -> `dtype` in transformers
# 4.56 (PR #39782); pick whichever keyword the installed version accepts.
def _dtype_kwargs(dtype):
return {"dtype": dtype} if Version(transformers.__version__) >= Version("4.56") else {"torch_dtype": dtype}
text_encoder = Qwen3ForCausalLM.from_pretrained(
model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
model_name, subfolder="text_encoder", **_dtype_kwargs(weight_dtype),
low_cpu_mem_usage=True,
)
@@ -3,8 +3,10 @@ import sys
import numpy as np
import torch
import transformers
from diffusers import FlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from packaging.version import Version
from PIL import Image
current_file_path = os.path.abspath(__file__)
@@ -137,8 +139,14 @@ if vae_path is not None:
tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
# `Qwen3ForCausalLM.from_pretrained` renamed `torch_dtype` -> `dtype` in transformers
# 4.56 (PR #39782); pick whichever keyword the installed version accepts.
def _dtype_kwargs(dtype):
return {"dtype": dtype} if Version(transformers.__version__) >= Version("4.56") else {"torch_dtype": dtype}
text_encoder = Qwen3ForCausalLM.from_pretrained(
model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
model_name, subfolder="text_encoder", **_dtype_kwargs(weight_dtype),
low_cpu_mem_usage=True,
)
@@ -247,4 +255,4 @@ if ulysses_degree * ring_degree > 1:
if dist.get_rank() == 0:
save_results()
else:
save_results()
save_results()
@@ -28,6 +28,15 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent,
get_image_to_video_latent,
get_video_to_video_latent,
save_videos_grid)
import transformers
from packaging.version import Version
def _dtype_kwargs(dtype):
"""`dtype` keyword of `from_pretrained` exists since transformers 4.56 (PR #39782);
older versions use `torch_dtype`."""
if Version(transformers.__version__) >= Version("4.56"):
return {"dtype": dtype}
return {"torch_dtype": dtype}
# 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.
@@ -138,7 +147,7 @@ tokenizer = AutoTokenizer.from_pretrained(
model_name, subfolder="tokenizer"
)
text_encoder = Qwen3ForCausalLM.from_pretrained(
model_name, subfolder="text_encoder", torch_dtype=weight_dtype,
model_name, subfolder="text_encoder", **_dtype_kwargs(weight_dtype),
low_cpu_mem_usage=True,
)
@@ -247,4 +256,4 @@ if ulysses_degree * ring_degree > 1:
if dist.get_rank() == 0:
save_results()
else:
save_results()
save_results()
+2
View File
@@ -256,6 +256,7 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR
| `--train_data_meta` | Training data metadata file | `datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json` |
| `--train_batch_size` | Samples per batch | 1 |
| `--video_sample_size` | Maximum video resolution for training | 960 |
| `--fix_sample_size` | Fixed `[height, width]` canvas (both multiples of 32) overriding `--video_sample_size`; requires `--enable_bucket` and turns off `--random_hw_adapt` / `--training_with_video_token_length` | None |
| `--token_sample_size` | Token length sampling size | 960 |
| `--video_sample_stride` | Frame sampling stride (MiniMax-H3 is 24 fps) | 1 |
| `--video_sample_n_frames` | Number of frames to sample, must follow the `17*n+5` form of the video VAE (duration stays between 5 and 15 seconds) | 124 |
@@ -606,5 +607,6 @@ torchrun --nproc_per_node=2 examples/minimax_h3/predict_t2v.py
## 5. Additional Resources
- **MiniMax-H3 PDD LoRA training**: `scripts/minimax_h3/README_TRAIN_PDD_LORA.md`
- **MiniMax-H3 Official GitHub**: https://github.com/MiniMax-AI/MiniMax-H3
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
@@ -256,6 +256,7 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR
| `--train_data_meta` | 训练数据元信息文件 | `datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json` |
| `--train_batch_size` | 每批样本数 | 1 |
| `--video_sample_size` | 训练的最大视频分辨率 | 960 |
| `--fix_sample_size` | 固定的 `[高, 宽]` 画布(都为 32 的倍数),覆盖 `--video_sample_size`;需配合 `--enable_bucket`,并关闭 `--random_hw_adapt` / `--training_with_video_token_length` | None |
| `--token_sample_size` | token 长度采样尺寸 | 960 |
| `--video_sample_stride` | 抽帧步长(MiniMax-H3 为 24 fps) | 1 |
| `--video_sample_n_frames` | 采样帧数,须满足视频 VAE 的 `17*n+5` 形式(时长保持在 5 到 15 秒之间) | 124 |
@@ -606,5 +607,6 @@ torchrun --nproc_per_node=2 examples/minimax_h3/predict_t2v.py
## 五、更多资源
- **MiniMax-H3 PDD LoRA 训练**:`scripts/minimax_h3/README_TRAIN_PDD_LORA_zh-CN.md`
- **MiniMax-H3 官方 GitHub**:https://github.com/MiniMax-AI/MiniMax-H3
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
+588
View File
@@ -0,0 +1,588 @@
# MiniMax-H3 PDD LoRA Training Guide
This document provides a complete workflow for Parallel Decoding Distillation (PDD, [arXiv 2607.26004](https://arxiv.org/abs/2607.26004)) LoRA training of MiniMax-H3, including environment configuration, conditioning-cache preparation, distributed training, and inference testing.
> **Note**: MiniMax-H3 is an audio-visual generative video model that can simultaneously generate video and corresponding audio. PDD training is **data-free**: it never reads target videos. Each rank carries one trajectory, rolls it forward with the student's own predictions, and is supervised by a frozen teacher on the same backbone. Only cached Qwen3-VL conditioning is needed, which keeps the ~62 GB text encoder out of the training run.
PDD turns the pre-trained flow model into a *parallel decoder*. The sampling interval is discretized into `N` intervals grouped into blocks of size `L`; one network evaluation predicts the mean velocity of every interval of the next block, so generation advances `L` intervals per evaluation (`NFE = N / L`). The default recipe is `N = 32`, `L = 4` (8 NFE). The student is the teacher's own transformer with the two final heads (`proj_out` / `audio_proj_out`) repeated `N` times; switching LoRA off is still the teacher, so there is no second copy of the 33 B backbone.
---
## Table of Contents
- [1. Environment Configuration](#1-environment-configuration)
- [2. Data Preparation](#2-data-preparation)
- [2.1 Data-free Conditioning](#21-data-free-conditioning)
- [2.2 Cache Structure](#22-cache-structure)
- [2.3 Generating the Cache](#23-generating-the-cache)
- [2.4 Annotation JSON Format](#24-annotation-json-format)
- [2.5 Ref2VA Request Cache](#25-ref2va-request-cache)
- [3. PDD LoRA Training](#3-pdd-lora-training)
- [3.1 Download Pretrained Model](#31-download-pretrained-model)
- [3.2 Quick Start (FSDP)](#32-quick-start-fsdp)
- [3.2.1 Ref2VA Training](#321-ref2va-training)
- [3.3 PDD Training Parameters](#33-pdd-training-parameters)
- [3.4 Training Validation](#34-training-validation)
- [3.5 Checkpoint Layout](#35-checkpoint-layout)
- [3.6 Training with DeepSpeed-Zero-2](#36-training-with-deepspeed-zero-2)
- [3.7 Training Without DeepSpeed or FSDP](#37-training-without-deepspeed-or-fsdp)
- [3.8 Multi-Machine Distributed Training](#38-multi-machine-distributed-training)
- [4. Inference Testing](#4-inference-testing)
- [4.1 Inference Parameters](#41-inference-parameters)
- [4.2 Single GPU Inference](#42-single-gpu-inference)
- [4.3 Multi-GPU Parallel Inference](#43-multi-gpu-parallel-inference)
- [5. Additional Resources](#5-additional-resources)
---
## 1. Environment Configuration
**Method 1: Using requirements.txt**
```bash
pip install -r requirements.txt
```
**Method 2: Manual Dependency Installation**
```bash
pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image
pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime
pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2"
pip install yunchang xfuser modelscope openpyxl
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**Method 3: Using Docker**
When using Docker, please ensure that the GPU driver and CUDA environment are correctly installed on your machine, then execute the following commands:
```
# pull image
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
# enter image
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
```
---
## 2. Data Preparation
PDD is **data-free**: the student trajectory is sampled from noise, so no target video/audio media is ever read for the loss. Training consumes only *conditioning* — the prompt for `fl2va`, and the prompt plus reference media for `ref2va`. Both `--train_mode`s support both conditioning routes, selected by `--enable_preprocess_training`:
| `--train_mode` | Conditions on | **Cache route** (`--enable_preprocess_training`) | **Direct-load route** (flag omitted) |
|----------------|---------------|--------------------------------------------------|--------------------------------------|
| `fl2va` | prompt | Pre-encode with `generate_prompt_cache.py`, read the `outputs.json` | Read a `{"text": ...}` annotation (`TextDataset`), encode on the fly |
| `ref2va` | prompt + reference media | Pre-encode with `generate_ref2va_request_cache.py`, read the `outputs.json` | Read a request annotation (`load_requests`), encode the prompt + reference latents on the fly |
- **Cache route** (recommended for long / repeated runs): the Qwen3-VL embeddings — and, for `ref2va`, the VAE-encoded reference latents — are pre-encoded to safetensors **once**, so the ~62 GB conditioner never loads during training.
- **Direct-load route** (recommended to start, or for a small request set): there is no separate preprocessing step; the run loads the ~62 GB conditioner (and, for `ref2va`, the two VAEs) and encodes each entry on the fly. Keep `--low_vram` on so they move onto the GPU only while encoding.
> 💡 Either route feeds **both** the train trajectories and the validation renders — data-free PDD has no train/val split, so `--train_data_meta` and `--val_data_meta` usually point at the same annotation (e.g. the official `datasets/X-Fun-Videos-Audios-Demo`, the standard test data — never an ad-hoc prompt set).
### 2.1 Data-free Conditioning
`--train_mode=fl2va` (the default recipe, FL2VA / t2va packed layout) needs the prompt conditioning; `--train_mode=ref2va` additionally needs reference media (images / videos / audio) and loads `transformer_ref` by default. On the **cache route**, generate the cache **once** with `scripts/minimax_h3/generate_prompt_cache.sh` (`fl2va`) or `scripts/minimax_h3/generate_ref2va_request_cache.sh` (`ref2va`); both run multi-GPU under `accelerate launch`, and the ~62 GB Qwen3-VL conditioner is then not loaded during PDD training. On the **direct-load route** there is no preprocessing step: point `--train_data_meta` straight at the annotation and the run encodes it on the fly (see **3.2.1** for both `ref2va` routes end to end).
### 2.2 Cache Structure
```
📦 datasets/
├── 📂 X-Fun-Videos-Audios-Demo/ # official demo dataset (source of the `text` captions)
│ └── 📄 metadata_add_width_height.json
└── 📂 minimax_h3_pdd_prompt_cache/ # generated once; feeds both train and validation
├── 📄 outputs.json
├── 📄 00000.safetensors
├── 📄 00001.safetensors
└── 📄 ...
```
`outputs.json` is a list of `{"file_path": ".../00000.safetensors"}` records that `ImageVideoSafetensorsDataset` reads. `--train_data_dir` is the optional root prepended to each `file_path`; leave it empty when `outputs.json` already stores repo-relative (or absolute) paths, as the generators do. Each `fl2va` `.safetensors` holds:
| Field | Description |
|-------|-------------|
| `prompt_embeds` | Qwen3-VL hidden states at the MiniMax-H3 text-encoder layer (bfloat16) |
| `text_token_tags` | Per-token tags for the packed sequence (int64) |
### 2.3 Generating the Cache
This step is only for the **cache route**; the direct-load route (**3.2.1 Route A**) reads the annotation directly and skips it.
```bash
# fl2va: cache the prompt conditioning of the official demo dataset (multi-GPU)
bash scripts/minimax_h3/generate_prompt_cache.sh
# ref2va: cache the request conditioning (prompt embeds + reference latents)
bash scripts/minimax_h3/generate_ref2va_request_cache.sh
```
Each launcher runs `accelerate launch ... generate_*_cache.py` once: every rank walks an interleaved slice of the annotation, `.safetensors` that already exist are skipped (resume), and rank0 finally writes `outputs.json`. Edit the `MODEL_NAME` / `DATASET_META` / `CACHE_ROOT` variables at the top of each `.sh` first.
> 💡 `--pretrained_model_name_or_path` may be either the converted diffusers layout or an original MiniMax-H3 partition (e.g. `MiniMax-H3/FL2VA`). The tokenizer is read from `tokenizer/`, the processor from `processor/`, and the conditioner from `text_encoder/`.
### 2.4 Annotation JSON Format
The `fl2va` generator reads the official demo dataset's annotation JSON — a list of records whose `text` caption is the only field PDD uses (the `file_path` / `audio_path` / `control_file_path` / `width` / `height` fields the audio-visual dataset carries are ignored). A bare list of strings, or `{"prompt": ...}` / `{"examples": [...]}` jobs, are accepted too:
```json
[
{"file_path": "train/00000001.mp4", "text": "A young woman gently turns her head to the right ...", "audio_path": "wav/00000001.wav", "type": "video"}
]
```
The `ref2va` route — both `generate_ref2va_request_cache.py` (cache) and direct-load training — reads a list of request records instead. Either an explicit `{"prompt": ..., "references": [...]}` request, where each reference is an `"image=..."` / `"video=..."` / `"audio=..."` string (the `predict_ref2va.py` schema, in the order the model reads them):
```json
[
{"prompt": "The character turns and waves", "references": ["image=ref/face.png", "audio=ref/voice.wav"]}
]
```
or — the default — the *same* audio-visual demo annotation as `fl2va` above: `load_requests` derives the `video=<file_path>` + `audio=<audio_path>` references from each record (relative media paths resolve against the annotation's directory) and uses its `text` as the prompt.
Only the conditioning is encoded. Resolution and duration are training/inference flags (`--video_sample_size` / `--fix_sample_size` / `--video_sample_n_frames`), not cache fields.
### 2.5 Ref2VA Request Cache
`--train_mode=ref2va` loads `transformer_ref` by default. On the **cache route** it reads the request cache written by `generate_ref2va_request_cache.py`: besides `prompt_embeds` / `text_token_tags`, each request `.safetensors` also carries the reference latents for the Ref2VA packed layout, flattened into tensors — `reference_kind_ids` / `reference_has_audio` (per-reference kind and has-audio flag), `num_condition_latents` / `num_audio_condition_latents`, and the indexed `condition_latents_{i}` / `audio_condition_latents_{i}`. On the **direct-load route** the very same latents are produced on the fly instead — `train_pdd_lora.py` reads the request annotation with `load_requests`, encodes the prompt with Qwen3-VL and the references with the two VAEs — so no `.safetensors` cache is needed (`--train_mode=ref2va` without `--enable_preprocess_training`).
---
## 3. PDD LoRA Training
### 3.1 Download Pretrained Model
```bash
# Create model directory
mkdir -p models/Diffusion_Transformer
# Download MiniMax-H3 official weights
hf download MiniMax-AI/MiniMax-H3 --local-dir models/Diffusion_Transformer/MiniMax-H3
```
> 💡 The loader accepts either the converted diffusers layout above or an *original* MiniMax-H3 partition (e.g. `MiniMax-H3/FL2VA`); the original shards are converted on the fly while loading, with no intermediate copy on disk.
### 3.2 Quick Start (FSDP)
If you have generated a conditioning cache as per **2.3** and downloaded the weights as per **3.1**, you can copy and run the command below. `scripts/minimax_h3/train_pdd_lora.sh` is the same launch.
FSDP is recommended: even though PDD does not load Qwen3-VL, the frozen transformer is still about 62 GB in bfloat16 and must be sharded — which FSDP (`FULL_SHARD`) does but DeepSpeed-Zero-2 does not.
`--mixed_precision=no` is required. The released checkpoint already pins `proj_out` / `audio_proj_out` in float32 (`_keep_in_fp32_modules`); the parallel heads built from them stay float32 master weights over a bfloat16 backbone, and the run does not use autocast.
**fl2va (from a prompt cache — the default recipe):**
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
--video_sample_n_frames=124 \
--fix_sample_size 768 1344 \
--train_batch_size=1 \
--max_train_steps=3000 \
--checkpointing_steps=200 \
--learning_rate=1e-5 \
--lora_learning_rate=1e-4 \
--seed=43 \
--output_dir="output_dir_minimax_h3_pdd_lora" \
--gradient_checkpointing \
--gradient_checkpointing_save_on_cpu \
--mixed_precision="no" \
--adam_weight_decay=0.0 \
--max_grad_norm=1.0 \
--rank=64 \
--network_alpha=64 \
--low_vram \
--target_name="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear" \
--train_mode="fl2va" \
--pdd_num_steps=32 \
--pdd_block_size=4 \
--validation_steps=200 \
--resume_from_checkpoint=latest
```
#### 3.2.1 Ref2VA Training
`--train_mode=ref2va` conditions on reference media (images / videos / audio) in addition to the prompt, and loads `transformer_ref` by default. It runs on either conditioning route; both commands below are the FSDP Quick Start with only the mode / data flags changed.
**Route A — direct load (no request cache).** Point `--train_data_meta` straight at a request annotation: either the explicit `{"prompt", "references"}` schema or the official `X-Fun-Videos-Audios-Demo` annotation, whose own video + audio become the references (see **2.4**). The run loads the Qwen3-VL conditioner and the two VAEs and encodes every request on the fly, so keep `--low_vram` — they are onloaded only while encoding.
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export REQUEST_META="datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_mode="ref2va" \
--train_data_meta=$REQUEST_META \
--val_data_meta=$REQUEST_META \
--video_sample_n_frames=124 \
--fix_sample_size 768 1344 \
--train_batch_size=1 \
--max_train_steps=3000 \
--checkpointing_steps=200 \
--learning_rate=1e-5 \
--lora_learning_rate=1e-4 \
--seed=43 \
--output_dir="output_dir_minimax_h3_pdd_ref2va_lora" \
--gradient_checkpointing \
--gradient_checkpointing_save_on_cpu \
--mixed_precision="no" \
--adam_weight_decay=0.0 \
--max_grad_norm=1.0 \
--rank=64 \
--network_alpha=64 \
--low_vram \
--target_name="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear" \
--pdd_num_steps=32 \
--pdd_block_size=4 \
--validation_steps=200 \
--resume_from_checkpoint=latest
```
> ⚠️ Route A keeps the ~62 GB Qwen3-VL conditioner resident in the run (sharded by FSDP alongside the transformer) and VAE-encodes each request's references on every trajectory reset. For long or repeated runs, pre-encode once with Route B so training never loads the conditioner.
**Route B — request cache (recommended for long / repeated runs).** Pre-encode the requests once with `generate_ref2va_request_cache.sh` (**2.3**), then train with `--enable_preprocess_training`; the ~62 GB conditioner stays out of the run:
```bash
# 1) Pre-encode the ref2va requests once (multi-GPU): prompt embeds + reference latents → safetensors
bash scripts/minimax_h3/generate_ref2va_request_cache.sh
# 2) Train off the cache — Route A with these flags changed
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export REQUEST_CACHE_META="datasets/minimax_h3_pdd_request_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_mode="ref2va" \
--enable_preprocess_training \
--train_data_meta=$REQUEST_CACHE_META \
--val_data_meta=$REQUEST_CACHE_META \
... # the remaining arguments identical to Route A
```
### 3.3 PDD Training Parameters
**PDD / LoRA parameters**:
| Parameter | Description | Example Value |
|-----------|-------------|----------------|
| `--pdd_num_steps` | Grid size `N` | 32 |
| `--pdd_block_size` | `L_min`: intervals the carried state advances by (`NFE = N / L`) | 4 |
| `--pdd_max_block_size` | `L_max`: widest block a loss target is drawn from. Defaults to `--pdd_block_size` | 4 |
| `--pdd_solver` | Runge-Kutta method for the teacher's mean velocity: `euler` or `midpoint` | `midpoint` |
| `--pdd_num_targets` | How many intra-block indices `k` one student evaluation is supervised at | 2 |
| `--rank` | Dimension of LoRA update matrices | 64 |
| `--network_alpha` | Scale of LoRA update matrices | 64 |
| `--target_name` | Modules to apply LoRA (comma-separated) | `to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear` |
| `--learning_rate` | Learning rate of the parallel heads | 1e-5 |
| `--lora_learning_rate` | Learning rate of the low-rank updates | 1e-4 |
| `--use_ema` | Keep an EMA of the trainable set; validation and `pdd_ema.safetensors` use it | off |
| `--ema_decay` | EMA decay | 0.99 |
**Common training parameters**:
| Parameter | Description | Example Value |
|-----------|-------------|----------------|
| `--pretrained_model_name_or_path` | Path to pretrained model | `models/Diffusion_Transformer/MiniMax-H3` |
| `--enable_preprocess_training` | Train on the pre-processed safetensors cache instead of encoding the conditioning on the fly; both `fl2va` and `ref2va` support either route | on |
| `--train_data_dir` | Optional root prepended to each `file_path` of `--train_data_meta`; empty when it stores repo-relative/absolute paths | `""` |
| `--train_data_meta` | With the flag: the cache `outputs.json`. Without it: the on-the-fly annotation (`fl2va`: `{"text": ...}`; `ref2va`: the request list) | `datasets/minimax_h3_pdd_prompt_cache/outputs.json` |
| `--val_data_meta` | Mirrors `--train_data_meta`: with the flag the cache `outputs.json`, without it the on-the-fly annotation (`fl2va`: `{"text": ...}`; `ref2va`: the request list). Skipped when empty | `datasets/minimax_h3_pdd_prompt_cache/outputs.json` |
| `--train_mode` | `fl2va` (t2va layout) or `ref2va` (`transformer_ref` + reference media); both run on the cache or the direct-load route | `fl2va` |
| `--transformer_subfolder` | Transformer subfolder. Default: `transformer_ref` for `ref2va`, else `transformer` | None |
| `--train_batch_size` | Must be 1: each rank carries one trajectory | 1 |
| `--num_train_epochs` | Training epochs when `--max_train_steps` is omitted. One epoch is one pass through the conditioning set | 100 |
| `--max_train_steps` | Total optimization steps. If set, overrides `--num_train_epochs` | 3000 |
| `--video_sample_n_frames` | Number of frames, must follow the `17*n+5` form of the video VAE (duration stays between 5 and 15 seconds) | 124 |
| `--video_sample_size` | Square canvas size (height = width); must be a multiple of 32 | 1280 |
| `--fix_sample_size` | Fixed `[height, width]` overriding `--video_sample_size` for a non-square canvas; both must be multiples of 32 | 768 1344 |
| `--gradient_accumulation_steps` | Gradient accumulation steps | 1 |
| `--checkpointing_steps` | Save checkpoint every N steps | 200 |
| `--seed` | Random seed | 43 |
| `--output_dir` | Output directory | `output_dir_minimax_h3_pdd_lora` |
| `--gradient_checkpointing` | Enable activation checkpointing | - |
| `--gradient_checkpointing_save_on_cpu` | Offload the activations saved for backward of the transformer blocks to CPU memory | - |
| `--mixed_precision` | Use `no`. The parallel heads stay float32 over a bfloat16 backbone | `no` |
| `--adam_weight_decay` | AdamW weight decay | 0.0 |
| `--max_grad_norm` | Gradient clipping threshold | 1.0 |
| `--low_vram` | Keep the VAEs on CPU; they move to GPU only inside validation decode | - |
| `--resume_from_checkpoint` | Resume from a checkpoint path, or `"latest"` to auto-select | `latest` |
| `--validation_steps` | Run validation every N steps | 200 |
| `--validation_nfe` | Student NFE during validation; must divide `--pdd_num_steps` | 8 |
| `--video_loss_weight` / `--audio_loss_weight` | Weights of the joint video + audio MSE | 0.5 / 0.5 |
### 3.4 Training Validation
Validation does **not** take `--validation_prompts`. It generates every entry of `--val_data_meta`, sharded across ranks, at `--validation_nfe`: the cache `outputs.json` under `--enable_preprocess_training`, or the on-the-fly annotation without it — mirroring the training route for both `fl2va` (prompts) and `ref2va` (requests). Validation is skipped when `--val_data_meta` is empty.
| Parameter | Description | Recommended Value |
|-----------|-------------|-------------------|
| `--validation_steps` | Validate every N steps | 200 |
| `--validation_nfe` | Student evaluations per clip (`N / NFE` must be an integer) | 8 |
Videos are saved under `output_dir/sample/` as `sample-{step}-prompt{index}-{train_mode}-nfe{nfe}.mp4` (with audio).
### 3.5 Checkpoint Layout
Each `checkpoint-{step}/` holds:
| File | Role |
|------|------|
| `pdd.safetensors` | Live (non-EMA) gathered trainable tensors (parallel heads + LoRA). Used to resume DDP |
| `pdd_ema.safetensors` | EMA export when `--use_ema` is on; this is the inference file |
| `pdd_config.json` | Rank / alpha / targets / grid, read by `examples/minimax_h3/predict_t2v.py` |
| `optimizer.pt` / `scheduler.pt` / `scaler.pt` / `ema.pt` | DDP trainer state (`optimizer.bin` / `scheduler.bin` from Accelerate `--save_state` are also accepted on resume) |
FSDP stage 3 / ZeRO-3 also write `accelerator.save_state` into the same folder (auto `--save_state`) and still export a gathered `pdd.safetensors` (live) plus `pdd_ema.safetensors` when EMA is on. Checkpoints written before the rename stored live weights in `pdd_live.safetensors`; DDP resume still loads that file when it is present.
### 3.6 Training with DeepSpeed-Zero-2
> ⚠️ **Warning**: DeepSpeed-Zero-2 only partitions optimizer states and gradients, **not the model weights**. The MiniMax-H3 transformer is about 62 GB, so each GPU still holds a full weight replica and this setup usually runs out of memory. Prefer FSDP (**3.2**); the command below is provided for reference only.
```sh
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # the same train_pdd_lora.py arguments as the Quick Start
```
### 3.7 Training Without DeepSpeed or FSDP
**This approach is not recommended on 80 GB cards**: every GPU keeps a full ~62 GB transformer replica. PDD does not load Qwen3-VL, so DDP is lighter than `scripts/minimax_h3/train_lora.py`, but FSDP (**3.2**) is still the default. Drop `--use_fsdp` and the FSDP wrap flags from the Quick Start command; DDP resume reads `pdd.safetensors` plus `optimizer.pt` / `optimizer.bin`.
```sh
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # the same train_pdd_lora.py arguments as the Quick Start
```
### 3.8 Multi-Machine Distributed Training
**Suitable for**: more GPUs, faster training
#### 3.8.1 Environment Configuration
Assuming 2 machines with 8 GPUs each:
**Machine 0 (Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export MASTER_ADDR="192.168.1.100" # Master machine IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # Total number of machines
export NUM_PROCESS=16 # Total processes = machines × 8
export RANK=0 # Current machine rank (0 or 1)
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # the same train_pdd_lora.py arguments as the Quick Start
```
**Machine 1 (Worker)**:
```bash
export RANK=1 # Note this is 1
# Other environment variables identical to Machine 0
# Use the same accelerate launch command as Machine 0
```
#### 3.8.2 Multi-Machine Training Notes
- **Network Requirements**:
- RDMA/InfiniBand recommended (high performance)
- Without RDMA, add environment variables:
```bash
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
```
- **Data Synchronization**: All machines must be able to access the same conditioning cache and model paths (NFS/shared storage)
## 4. Inference Testing
PDD inference attaches the parallel heads and LoRA from `pdd_ema.safetensors` (falling back to `pdd.safetensors` when EMA was not saved), then samples at `num_inference_steps` NFE. Use `examples/minimax_h3/predict_t2v.py`; do not set `lora_path` at the same time (`lora_path` and `pdd_lora_path` cannot be used together).
The default recipe (`N = 32`, `L = 4`) runs at **8** inference steps. `num_inference_steps` must divide `pdd_num_steps` from `pdd_config.json`. If it is left at the teacher default of 40, the script snaps it to `N / L`.
### 4.1 Inference Parameters
**Key Parameter Descriptions**:
| Parameter | Description | Example Value |
|------|------|-------|
| `GPU_memory_mode` | GPU memory mode, see table below for options | `model_cpu_offload` |
| `ulysses_degree` | Head dimension parallelization degree, 1 for single GPU | 1 |
| `ring_degree` | Sequence dimension parallelization degree, 1 for single GPU | 1 |
| `fsdp_dit` | Use FSDP for Transformer in multi-GPU inference to save VRAM | `False` |
| `fsdp_text_encoder` | Use FSDP for the Qwen3-VL text encoder in multi-GPU inference to save VRAM | `False` |
| `compile_dit` | Compile Transformer to accelerate inference (effective at fixed resolution) | `False` |
| `model_name` | Model path | `models/Diffusion_Transformer/MiniMax-H3` |
| `transformer_path` | Path to trained Transformer weights | `None` |
| `vae_path` | Path to trained VAE weights | `None` |
| `pdd_lora_path` | PDD checkpoint directory (loads `pdd_ema.safetensors` if present, else `pdd.safetensors`, plus `pdd_config.json`) or a `.safetensors` file | `output_dir_minimax_h3_pdd_lora/checkpoint-3000` |
| `lora_path` | Turbo/PEFT LoRA; **cannot** be combined with `pdd_lora_path` | `None` |
| `sample_size` | Generated video resolution `[height, width]`; height/width must be multiples of 32. `None` uses MiniMax-H3's own 16:9 canvas (768x1344) | `[768, 1344]` |
| `video_length` | Number of frames to generate, snapped up to the next `17*n+5` the video VAE can decode (duration stays between 5 and 15 seconds) | 124 |
| `fps` | Frames per second (MiniMax-H3 generates at a fixed 24 fps) | 24 |
| `weight_dtype` | Model weight precision, use `torch.float16` for GPUs without bf16 support | `torch.bfloat16` |
| `prompt` | Positive prompt describing the content to generate | `"A red fox trotting..."` |
| `seed` | Random seed for reproducibility | 43 |
| `num_inference_steps` | Student NFE. Default PDD recipe uses 8 (not the teacher's 40) | 8 |
| `guidance_scale` | Guidance strength. The released checkpoint is guidance-distilled: keep it at 1 to run one forward pass per step with no CFG | 1 |
| `flow_shift` | Exponential sigma shift of the video schedule, `None` keeps the one of the checkpoint (12.0) | `None` |
| `audio_flow_shift` | Exponential sigma shift of the audio schedule, `None` keeps the one of the checkpoint (3.0) | `None` |
| `save_path` | Generated video save path | `samples/minimax-h3-videos-t2v` |
**GPU Memory Mode Description**:
| Mode | Description | VRAM Usage |
|------|------|---------|
| `model_full_load` | Load entire model to GPU | Highest |
| `model_full_load_and_qfloat8` | Full load + FP8 quantization | High |
| `model_cpu_offload` | Offload model to CPU after use | Medium |
| `model_cpu_offload_and_qfloat8` | CPU offload + FP8 quantization | Medium-Low |
| `model_group_offload` | Layer group offload between CPU/CUDA | Low |
| `sequential_cpu_offload` | Offload each layer individually (slowest) | Lowest |
> 💡 The transformer alone is 61.7 GB in bfloat16 and the Qwen3-VL conditioner is another 62.1 GB, so a single 80 GB card needs `model_cpu_offload` or `model_group_offload`. Inference *does* load the text encoder; training does not.
### 4.2 Single GPU Inference
Run single GPU inference with:
```bash
python examples/minimax_h3/predict_t2v.py
```
Edit `examples/minimax_h3/predict_t2v.py` according to your needs. For PDD inference, focus on these parameters:
```python
# Choose based on your GPU VRAM
GPU_memory_mode = "model_cpu_offload"
# Your actual model path
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# PDD checkpoint directory or weights file; a directory prefers pdd_ema.safetensors. Rank / alpha / targets / grid are read from pdd_config.json
pdd_lora_path = "output_dir_minimax_h3_pdd_lora/checkpoint-3000"
# Must stay None when pdd_lora_path is set
lora_path = None
# Student NFE; 8 for the default N=32 / L=4 recipe. Left at 40, the script snaps to N / L
num_inference_steps = 8
# Write based on content to generate
prompt = "A red fox trotting through a snowy pine forest, snow crunching underfoot"
# ...
```
Image-to-video and Ref2VA use the same `pdd_lora_path` field in `examples/minimax_h3/predict_i2v.py` and `examples/minimax_h3/predict_ref2va.py`. Ref2VA needs a checkpoint trained with `--train_mode=ref2va`.
### 4.3 Multi-GPU Parallel Inference
**Suitable for**: High-resolution generation, accelerated inference
#### Install Parallel Inference Dependencies
```bash
pip install xfuser yunchang
```
#### Configure Parallel Strategy
Edit `examples/minimax_h3/predict_t2v.py`:
```python
# Ensure ulysses_degree × ring_degree = number of GPUs
# For example, using 2 GPUs:
ulysses_degree = 2 # Head dimension parallelization
ring_degree = 1 # Sequence dimension parallelization
```
**Configuration Principles**:
- `ulysses_degree` must evenly divide the model's number of heads
- `ring_degree` splits on sequence dimension, affecting communication overhead; avoid using it when heads can be divided
- Multi-GPU runs through the xfuser sequence-parallel path and is **incompatible with the `*cpu_offload*` memory modes** (accelerate offload hooks own a single device); use `model_full_load` / `model_full_load_and_qfloat8` across GPUs there, with `fsdp_dit` / `fsdp_text_encoder` to save memory
- Sequence parallel needs a working FlashAttention. Without it, run independent single-GPU jobs (`CUDA_VISIBLE_DEVICES=i`, `ulysses_degree = 1`, `ring_degree = 1`) instead
**Example Configurations**:
| GPU Count | ulysses_degree | ring_degree | Description |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | Single GPU |
| 4 | 4 | 1 | Head parallelization |
| 8 | 2 | 4 | Hybrid parallelization |
| 8 | 8 | 1 | Head parallelization |
#### Run Multi-GPU Inference
```bash
torchrun --nproc_per_node=2 examples/minimax_h3/predict_t2v.py
```
## 5. Additional Resources
- **PDD paper**: https://arxiv.org/abs/2607.26004
- **MiniMax-H3 Official GitHub**: https://github.com/MiniMax-AI/MiniMax-H3
- **Official GitHub**: https://github.com/aigc-apps/VideoX-Fun
- **Base MiniMax-H3 LoRA training**: `scripts/minimax_h3/README_TRAIN_LORA.md`
- **MiniMax-H3 Fun control training**: `scripts/minimax_h3_fun/README_TRAIN.md`
@@ -0,0 +1,588 @@
# MiniMax-H3 PDD LoRA 训练指南
本文档提供 MiniMax-H3 的 Parallel Decoding Distillation(PDD,[arXiv 2607.26004](https://arxiv.org/abs/2607.26004))LoRA 训练完整工作流,包括环境配置、条件 cache 准备、分布式训练和推理测试。
> **注意**:MiniMax-H3 是一个音视频生成模型,可以同时生成视频和对应音频。PDD 训练是 **data-free** 的:从不读取目标视频。每个 rank 携带一条轨迹,用学生自己的预测向前滚动,并由同一骨干上的冻结教师监督。训练只需要缓存好的 Qwen3-VL 条件,因此约 62 GB 的文本编码器不会进入训练进程。
PDD 把预训练 flow 模型变成 *parallel decoder*。采样区间被离散成 `N` 个 interval,再按大小 `L` 分块;一次网络前向预测下一块中每个 interval 的平均速度,因此生成每步前进 `L` 个 interval(`NFE = N / L`)。默认配方是 `N = 32`、`L = 4`(8 NFE)。学生就是教师自己的 transformer,两个最终头(`proj_out` / `audio_proj_out`)各重复 `N` 次;关掉 LoRA 仍是教师,不需要第二份 33 B 骨干。
---
## 目录
- [一、环境配置](#一环境配置)
- [二、数据准备](#二数据准备)
- [2.1 Data-free 条件](#21-data-free-条件)
- [2.2 Cache 结构](#22-cache-结构)
- [2.3 生成 Cache](#23-生成-cache)
- [2.4 标注 JSON 格式](#24-标注-json-格式)
- [2.5 Ref2VA Request Cache](#25-ref2va-request-cache)
- [三、PDD LoRA 训练](#三pdd-lora-训练)
- [3.1 下载预训练模型](#31-下载预训练模型)
- [3.2 快速开始(FSDP)](#32-快速开始fsdp)
- [3.2.1 Ref2VA 训练](#321-ref2va-训练)
- [3.3 PDD 训练参数](#33-pdd-训练参数)
- [3.4 训练验证](#34-训练验证)
- [3.5 Checkpoint 布局](#35-checkpoint-布局)
- [3.6 使用 DeepSpeed-Zero-2 训练](#36-使用-deepspeed-zero-2-训练)
- [3.7 不使用 DeepSpeed 或 FSDP 训练](#37-不使用-deepspeed-或-fsdp-训练)
- [3.8 多机分布式训练](#38-多机分布式训练)
- [四、推理测试](#四推理测试)
- [4.1 推理参数](#41-推理参数)
- [4.2 单 GPU 推理](#42-单-gpu-推理)
- [4.3 多 GPU 并行推理](#43-多-gpu-并行推理)
- [五、更多资源](#五更多资源)
---
## 一、环境配置
**方式一:使用 requirements.txt**
```bash
pip install -r requirements.txt
```
**方式二:手动安装依赖**
```bash
pip install Pillow einops safetensors timm tomesd librosa "torch>=2.1.2" torchdiffeq torchsde decord datasets numpy scikit-image
pip install omegaconf SentencePiece imageio[ffmpeg] imageio[pyav] tensorboard beautifulsoup4 ftfy func_timeout onnxruntime
pip install "peft>=0.17.0" "accelerate>=0.25.0" "gradio>=3.41.2" "diffusers>=0.30.1" "transformers>=4.46.2"
pip install yunchang xfuser modelscope openpyxl
pip uninstall opencv-python opencv-contrib-python opencv-python-headless -y
pip install opencv-python-headless
pip install deepspeed==0.17.0 numpy==1.26.4
```
**方式三:使用 Docker**
使用 Docker 时,请先确保本机已正确安装 GPU 驱动和 CUDA 环境,然后执行以下命令:
```
# 拉取镜像
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
# 进入镜像
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
```
---
## 二、数据准备
PDD 是 **data-free** 的:学生轨迹从噪声采样,因此完全不读取用于 loss 的目标视频/音频媒体。训练只读取 *条件*——`fl2va` 是 prompt,`ref2va` 是 prompt 加参考媒体。两种 `--train_mode` 都支持两条条件路径,由 `--enable_preprocess_training` 选择:
| `--train_mode` | 条件 | **Cache 路径**(`--enable_preprocess_training`) | **直接加载路径**(不加该参数) |
|----------------|------|--------------------------------------------------|-------------------------------|
| `fl2va` | prompt | 用 `generate_prompt_cache.py` 预编码,读取 `outputs.json` | 读取 `{"text": ...}` 标注(`TextDataset`),现场编码 |
| `ref2va` | prompt + 参考媒体 | 用 `generate_ref2va_request_cache.py` 预编码,读取 `outputs.json` | 读取 request 标注(`load_requests`),现场编码 prompt + 参考 latent |
- **Cache 路径**(适合长时间 / 反复训练):Qwen3-VL embedding——`ref2va` 还有 VAE 编码的参考 latent——**只**预编码一次写入 safetensors,训练时约 62 GB 的条件器完全不加载。
- **直接加载路径**(适合快速起步或 request 集较小):没有单独的预处理步骤;训练进程加载约 62 GB 条件器(`ref2va` 还有两个 VAE),逐条现场编码。请开启 `--low_vram`,使其仅在编码时搬上 GPU。
> 💡 两条路径都同时供训练轨迹和验证渲染使用——data-free PDD 没有 train/val 之分,因此 `--train_data_meta` 与 `--val_data_meta` 通常指向同一份标注(例如官方 `datasets/X-Fun-Videos-Audios-Demo`,标准测试数据,绝不用临时拼凑的 prompt 集)。
### 2.1 Data-free 条件
`--train_mode=fl2va`(默认配方,FL2VA / t2va packed 布局)需要 prompt 条件;`--train_mode=ref2va` 还需要参考媒体(图像 / 视频 / 音频),并默认加载 `transformer_ref`。走 **cache 路径** 时,用 `scripts/minimax_h3/generate_prompt_cache.sh`(`fl2va`)或 `scripts/minimax_h3/generate_ref2va_request_cache.sh`(`ref2va`)**只生成一次** cache;两者都在 `accelerate launch` 下多卡运行,此后 PDD 训练过程中不会加载约 62 GB 的 Qwen3-VL 条件器。走 **直接加载路径** 时没有预处理步骤:把 `--train_data_meta` 直接指向标注,训练进程现场编码(两条 `ref2va` 路径的完整流程见 **3.2.1**)。
### 2.2 Cache 结构
```
📦 datasets/
├── 📂 X-Fun-Videos-Audios-Demo/ # 官方 demo 数据集(`text` 字幕的来源)
│ └── 📄 metadata_add_width_height.json
└── 📂 minimax_h3_pdd_prompt_cache/ # 只生成一次;同时供训练与验证使用
├── 📄 outputs.json
├── 📄 00000.safetensors
├── 📄 00001.safetensors
└── 📄 ...
```
`outputs.json` 是 `ImageVideoSafetensorsDataset` 读取的 `{"file_path": ".../00000.safetensors"}` 记录列表。`--train_data_dir` 是拼接到每个 `file_path` 之前的可选根目录;当 `outputs.json` 已使用仓库相对(或绝对)路径时(生成器即如此)可留空。每个 `fl2va` `.safetensors` 包含:
| 字段 | 说明 |
|------|------|
| `prompt_embeds` | MiniMax-H3 文本编码器层上的 Qwen3-VL hidden states(bfloat16) |
| `text_token_tags` | packed sequence 用的逐 token 标签(int64) |
### 2.3 生成 Cache
该步骤仅 **cache 路径** 需要;直接加载路径(**3.2.1 Route A**)直接读取标注,跳过此步。
```bash
# fl2va:多卡缓存官方 demo 数据集的 prompt 条件
bash scripts/minimax_h3/generate_prompt_cache.sh
# ref2va:缓存 request 条件(prompt embedding + 参考 latent)
bash scripts/minimax_h3/generate_ref2va_request_cache.sh
```
每个启动器只运行一次 `accelerate launch ... generate_*_cache.py`:每个 rank 走标注的交错分片,已存在的 `.safetensors` 会被跳过(resume),最后由 rank0 写出 `outputs.json`。请先编辑各 `.sh` 顶部的 `MODEL_NAME` / `DATASET_META` / `CACHE_ROOT` 变量。
> 💡 `--pretrained_model_name_or_path` 可以是转换后的 diffusers 布局,也可以是原始 MiniMax-H3 分区(例如 `MiniMax-H3/FL2VA`)。分词器从 `tokenizer/` 读取,processor 从 `processor/` 读取,条件器从 `text_encoder/` 读取。
### 2.4 标注 JSON 格式
`fl2va` 生成器读取官方 demo 数据集的标注 JSON——记录列表中 PDD 只用到 `text` 字幕(音视频数据集携带的 `file_path` / `audio_path` / `control_file_path` / `width` / `height` 字段在此被忽略)。也接受纯字符串列表,或 `{"prompt": ...}` / `{"examples": [...]}` jobs:
```json
[
{"file_path": "train/00000001.mp4", "text": "A young woman gently turns her head to the right ...", "audio_path": "wav/00000001.wav", "type": "video"}
]
```
`ref2va` 路径——`generate_ref2va_request_cache.py`(cache)与直接加载训练皆然——则读取 request 记录的列表。既可以是显式的 `{"prompt": ..., "references": [...]}` request,其中每个 reference 是 `"image=..."` / `"video=..."` / `"audio=..."` 字符串(`predict_ref2va.py` 的 schema,按模型读取顺序排列):
```json
[
{"prompt": "The character turns and waves", "references": ["image=ref/face.png", "audio=ref/voice.wav"]}
]
```
也可以是——默认方式,遵循官方 demo 约定——与上面 `fl2va` 完全相同的音视频 demo 标注:`load_requests` 会从每条记录自身的 `file_path`(视频)+ `audio_path`(音频)推导出 `video=`/`audio=` reference(相对媒体路径按标注文件所在目录解析),并用其 `text` 作为 prompt。
只编码条件。分辨率和时长由训练/推理参数决定(`--video_sample_size` / `--fix_sample_size` / `--video_sample_n_frames`),不是 cache 字段。
### 2.5 Ref2VA Request Cache
`--train_mode=ref2va` 默认加载 `transformer_ref`。走 **cache 路径** 时读取 `generate_ref2va_request_cache.py` 写出的 request cache:除 `prompt_embeds` / `text_token_tags` 外,每条 request `.safetensors` 还带有 Ref2VA packed 布局所需的参考 latent,并扁平化为张量——`reference_kind_ids` / `reference_has_audio`(每个 reference 的类别与是否含音频)、`num_condition_latents` / `num_audio_condition_latents`,以及带下标的 `condition_latents_{i}` / `audio_condition_latents_{i}`。走 **直接加载路径** 时,同样的 latent 改为现场生成——`train_pdd_lora.py` 用 `load_requests` 读取 request 标注,用 Qwen3-VL 编码 prompt、用两个 VAE 编码参考——因此不需要 `.safetensors` cache(即 `--train_mode=ref2va` 且不加 `--enable_preprocess_training`)。
---
## 三、PDD LoRA 训练
### 3.1 下载预训练模型
```bash
# 创建模型目录
mkdir -p models/Diffusion_Transformer
# 下载 MiniMax-H3 官方权重
hf download MiniMax-AI/MiniMax-H3 --local-dir models/Diffusion_Transformer/MiniMax-H3
```
> 💡 加载器既接受上面转换后的 diffusers 布局,也接受 *原始* MiniMax-H3 分区(例如 `MiniMax-H3/FL2VA`);原始分片在加载时即时转换,磁盘上不会留下中间副本。
### 3.2 快速开始(FSDP)
若已按 **2.3** 生成条件 cache、按 **3.1** 下载权重,可直接复制运行下面的命令。`scripts/minimax_h3/train_pdd_lora.sh` 是同一套 launch。
推荐使用 FSDP:虽然 PDD 不加载 Qwen3-VL,冻结的 transformer 在 bfloat16 下仍约 62 GB,必须跨 GPU 切分——FSDP(`FULL_SHARD`)会切分权重,DeepSpeed-Zero-2 不会。
必须使用 `--mixed_precision=no`。发布权重已将 `proj_out` / `audio_proj_out` 钉在 float32(`_keep_in_fp32_modules`);由此构建的 parallel head 作为 float32 master 权重叠在 bfloat16 骨干上,训练过程不走 autocast。
**fl2va(读取 prompt cache——默认配方):**
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
# 无 RDMA 的多机环境可设置 NCCL_IB_DISABLE=1 和 NCCL_P2P_DISABLE=1。
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
--video_sample_n_frames=124 \
--fix_sample_size 768 1344 \
--train_batch_size=1 \
--max_train_steps=3000 \
--checkpointing_steps=200 \
--learning_rate=1e-5 \
--lora_learning_rate=1e-4 \
--seed=43 \
--output_dir="output_dir_minimax_h3_pdd_lora" \
--gradient_checkpointing \
--gradient_checkpointing_save_on_cpu \
--mixed_precision="no" \
--adam_weight_decay=0.0 \
--max_grad_norm=1.0 \
--rank=64 \
--network_alpha=64 \
--low_vram \
--target_name="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear" \
--train_mode="fl2va" \
--pdd_num_steps=32 \
--pdd_block_size=4 \
--validation_steps=200 \
--resume_from_checkpoint=latest
```
#### 3.2.1 Ref2VA 训练
`--train_mode=ref2va` 除 prompt 外还对参考媒体(图像 / 视频 / 音频)加以条件,并默认加载 `transformer_ref`。它可走任一条条件路径;下面两条命令都是在 FSDP 快速开始的基础上,只改动 mode / data 相关参数。
**Route A —— 直接加载(无 request cache)。** 把 `--train_data_meta` 直接指向 request 标注:既可以是显式的 `{"prompt", "references"}` schema,也可以是官方 `X-Fun-Videos-Audios-Demo` 标注(其自身的视频 + 音频会被当作 reference,见 **2.4**)。训练进程会加载 Qwen3-VL 条件器与两个 VAE,对每条 request 现场编码,因此请保留 `--low_vram`——它们仅在编码时搬上 GPU。
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export REQUEST_META="datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_mode="ref2va" \
--train_data_meta=$REQUEST_META \
--val_data_meta=$REQUEST_META \
--video_sample_n_frames=124 \
--fix_sample_size 768 1344 \
--train_batch_size=1 \
--max_train_steps=3000 \
--checkpointing_steps=200 \
--learning_rate=1e-5 \
--lora_learning_rate=1e-4 \
--seed=43 \
--output_dir="output_dir_minimax_h3_pdd_ref2va_lora" \
--gradient_checkpointing \
--gradient_checkpointing_save_on_cpu \
--mixed_precision="no" \
--adam_weight_decay=0.0 \
--max_grad_norm=1.0 \
--rank=64 \
--network_alpha=64 \
--low_vram \
--target_name="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear" \
--pdd_num_steps=32 \
--pdd_block_size=4 \
--validation_steps=200 \
--resume_from_checkpoint=latest
```
> ⚠️ Route A 会把约 62 GB 的 Qwen3-VL 条件器常驻在训练进程中(由 FSDP 与 transformer 一起切分),并在每次轨迹 reset 时用 VAE 编码该 request 的参考。若长时间或反复训练,请用 Route B 预编码一次,训练时就不再加载条件器。
**Route B —— request cache(适合长时间 / 反复训练)。** 先用 `generate_ref2va_request_cache.sh`(**2.3**)预编码一次,再加 `--enable_preprocess_training` 训练;约 62 GB 条件器不会进入训练进程:
```bash
# 1) 预编码 ref2va request 一次(多卡):prompt embedding + 参考 latent → safetensors
bash scripts/minimax_h3/generate_ref2va_request_cache.sh
# 2) 从 cache 训练——即 Route A 改动以下参数
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export REQUEST_CACHE_META="datasets/minimax_h3_pdd_request_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_mode="ref2va" \
--enable_preprocess_training \
--train_data_meta=$REQUEST_CACHE_META \
--val_data_meta=$REQUEST_CACHE_META \
... # 其余参数与 Route A 相同
```
### 3.3 PDD 训练参数
**PDD / LoRA 参数**:
| 参数 | 说明 | 示例值 |
|------|------|--------|
| `--pdd_num_steps` | 网格大小 `N` | 32 |
| `--pdd_block_size` | `L_min`:携带状态每次前进的 interval 数(`NFE = N / L`) | 4 |
| `--pdd_max_block_size` | `L_max`:抽 loss 目标时最宽的块。默认等于 `--pdd_block_size` | 4 |
| `--pdd_solver` | 估计教师平均速度的 Runge-Kutta 方法:`euler` 或 `midpoint` | `midpoint` |
| `--pdd_num_targets` | 一次学生前向在块内监督的下标 `k` 个数 | 2 |
| `--rank` | LoRA 更新矩阵的维度 | 64 |
| `--network_alpha` | LoRA 更新矩阵的缩放 | 64 |
| `--target_name` | 施加 LoRA 的模块(逗号分隔) | `to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear` |
| `--learning_rate` | parallel head 的学习率 | 1e-5 |
| `--lora_learning_rate` | 低秩更新的学习率 | 1e-4 |
| `--use_ema` | 对可训练参数做 EMA;验证和 `pdd_ema.safetensors` 都用 EMA | 默认关闭 |
| `--ema_decay` | EMA 衰减 | 0.99 |
**通用训练参数**:
| 参数 | 说明 | 示例值 |
|------|------|--------|
| `--pretrained_model_name_or_path` | 预训练模型路径 | `models/Diffusion_Transformer/MiniMax-H3` |
| `--enable_preprocess_training` | 使用预处理好的 safetensors cache 训练,而非现场编码条件;`fl2va` 与 `ref2va` 都支持两条路径 | 开启 |
| `--train_data_dir` | 拼接到 `--train_data_meta` 每个 `file_path` 之前的可选根目录;当其存的是仓库相对/绝对路径时留空 | `""` |
| `--train_data_meta` | 加该参数时:cache 的 `outputs.json`。不加时:现场编码的标注(`fl2va`:`{"text": ...}`;`ref2va`:request 列表) | `datasets/minimax_h3_pdd_prompt_cache/outputs.json` |
| `--val_data_meta` | 与 `--train_data_meta` 一致:加该参数时用 cache 的 `outputs.json`,不加时用现场编码的标注(`fl2va`:`{"text": ...}`;`ref2va`:request 列表)。留空则跳过验证 | `datasets/minimax_h3_pdd_prompt_cache/outputs.json` |
| `--train_mode` | `fl2va`(t2va 布局)或 `ref2va`(`transformer_ref` + 参考媒体);两者都可走 cache 或直接加载路径 | `fl2va` |
| `--transformer_subfolder` | Transformer 子目录。默认:`ref2va` 用 `transformer_ref`,否则 `transformer` | None |
| `--train_batch_size` | 必须为 1:每个 rank 只携带一条轨迹 | 1 |
| `--num_train_epochs` | 未指定 `--max_train_steps` 时的训练轮数。一轮对应条件集走一遍 | 100 |
| `--max_train_steps` | 总优化步数。若设置则覆盖 `--num_train_epochs` | 3000 |
| `--video_sample_n_frames` | 采样帧数,须符合视频 VAE 的 `17*n+5`(时长保持在 5 到 15 秒) | 124 |
| `--video_sample_size` | 正方形画布尺寸(高 = 宽);必须是 32 的倍数 | 1280 |
| `--fix_sample_size` | 固定的 `[高, 宽]`,用于非正方形画布并覆盖 `--video_sample_size`;都必须是 32 的倍数 | 768 1344 |
| `--gradient_accumulation_steps` | 梯度累积步数 | 1 |
| `--checkpointing_steps` | 每 N 步保存一次 checkpoint | 200 |
| `--seed` | 随机种子 | 43 |
| `--output_dir` | 输出目录 | `output_dir_minimax_h3_pdd_lora` |
| `--gradient_checkpointing` | 启用 activation checkpointing | - |
| `--gradient_checkpointing_save_on_cpu` | 将 transformer block 反向所需的激活卸载到 CPU | - |
| `--mixed_precision` | 使用 `no`。parallel head 保持 float32,骨干为 bfloat16 | `no` |
| `--adam_weight_decay` | AdamW weight decay | 0.0 |
| `--max_grad_norm` | 梯度裁剪阈值 | 1.0 |
| `--low_vram` | VAE 放在 CPU,仅在验证解码时搬到 GPU | - |
| `--resume_from_checkpoint` | 从 checkpoint 恢复,`"latest"` 自动选最新 | `latest` |
| `--validation_steps` | 每 N 步做一次验证 | 200 |
| `--validation_nfe` | 验证时学生的 NFE;必须整除 `--pdd_num_steps` | 8 |
| `--video_loss_weight` / `--audio_loss_weight` | 视频 + 音频联合 MSE 的权重 | 0.5 / 0.5 |
### 3.4 训练验证
验证 **不使用** `--validation_prompts`。它会按 rank 分片、以 `--validation_nfe` 生成 `--val_data_meta` 的每一条:加 `--enable_preprocess_training` 时是 cache 的 `outputs.json`,不加时是现场编码的标注——与训练路径一致,`fl2va`(prompt)与 `ref2va`(request)皆然。当 `--val_data_meta` 留空时跳过验证。
| 参数 | 说明 | 推荐值 |
|------|------|--------|
| `--validation_steps` | 每 N 步验证一次 | 200 |
| `--validation_nfe` | 每条 clip 的学生前向次数(`N / NFE` 必须为整数) | 8 |
视频保存在 `output_dir/sample/`,文件名为 `sample-{step}-prompt{index}-{train_mode}-nfe{nfe}.mp4`(带音频)。
### 3.5 Checkpoint 布局
每个 `checkpoint-{step}/` 包含:
| 文件 | 作用 |
|------|------|
| `pdd.safetensors` | 现场(非 EMA)收集后的可训练张量(parallel head + LoRA),供 DDP resume |
| `pdd_ema.safetensors` | 开启 `--use_ema` 时的 EMA 导出;这是推理用文件 |
| `pdd_config.json` | rank / alpha / targets / 网格,供 `examples/minimax_h3/predict_t2v.py` 读取 |
| `optimizer.pt` / `scheduler.pt` / `scaler.pt` / `ema.pt` | DDP 训练器状态(Accelerate `--save_state` 写出的 `optimizer.bin` / `scheduler.bin` 在 resume 时同样接受) |
FSDP stage 3 / ZeRO-3 还会把 `accelerator.save_state` 写进同一目录(自动 `--save_state`),并仍然导出一份收集后的 `pdd.safetensors`(现场权重);开启 EMA 时另写 `pdd_ema.safetensors`。更早的 checkpoint 把现场权重存在 `pdd_live.safetensors`;DDP resume 在该文件存在时仍会读取它。
### 3.6 使用 DeepSpeed-Zero-2 训练
> ⚠️ **警告**:DeepSpeed-Zero-2 只切分优化器状态和梯度,**不切分模型权重**。MiniMax-H3 transformer 约 62 GB,每张 GPU 仍会持有完整权重副本,通常会显存不足。请优先使用 FSDP(**3.2**);下面命令仅供参考。
```sh
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # 与快速开始相同的 train_pdd_lora.py 参数
```
### 3.7 不使用 DeepSpeed 或 FSDP 训练
**不建议在 80 GB 卡上使用**:每张 GPU 仍会保留完整约 62 GB 的 transformer 副本。PDD 不加载 Qwen3-VL,因此 DDP 比 `scripts/minimax_h3/train_lora.py` 更轻,但默认仍应使用 FSDP(**3.2**)。从快速开始命令中去掉 `--use_fsdp` 和 FSDP wrap 参数即可;DDP resume 读取 `pdd.safetensors` 以及 `optimizer.pt` / `optimizer.bin`。
```sh
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # 与快速开始相同的 train_pdd_lora.py 参数
```
### 3.8 多机分布式训练
**适用场景**:更多 GPU、更快训练
#### 3.8.1 环境配置
假设 2 台机器,每台 8 卡:
**机器 0(Master)**:
```bash
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export MASTER_ADDR="192.168.1.100" # Master 机器 IP
export MASTER_PORT=10086
export WORLD_SIZE=2 # 机器总数
export NUM_PROCESS=16 # 总进程数 = 机器数 × 8
export RANK=0 # 当前机器 rank(0 或 1)
# 无 RDMA 的多机环境可设置 NCCL_IB_DISABLE=1 和 NCCL_P2P_DISABLE=1。
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
... # 与快速开始相同的 train_pdd_lora.py 参数
```
**机器 1(Worker)**:
```bash
export RANK=1 # 注意这里是 1
# 其余环境变量与机器 0 相同
# 使用与机器 0 相同的 accelerate launch 命令
```
#### 3.8.2 多机训练注意事项
- **网络要求**:
- 推荐 RDMA/InfiniBand(高性能)
- 无 RDMA 时,添加环境变量:
```bash
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
```
- **数据同步**:所有机器必须能访问相同的条件 cache 和模型路径(NFS/共享存储)
## 四、推理测试
PDD 推理从 `pdd_ema.safetensors` 挂上 parallel head 和 LoRA(没有 EMA 文件时回退到 `pdd.safetensors`),再按 `num_inference_steps` NFE 采样。使用 `examples/minimax_h3/predict_t2v.py`;不要同时设置 `lora_path`(`lora_path` 与 `pdd_lora_path` 不能一起用)。
默认配方(`N = 32`、`L = 4`)以 **8** 步推理。`num_inference_steps` 必须整除 `pdd_config.json` 中的 `pdd_num_steps`。若仍为教师默认的 40,脚本会自动改成 `N / L`。
### 4.1 推理参数
**关键参数说明**:
| 参数 | 说明 | 示例值 |
|------|------|-------|
| `GPU_memory_mode` | GPU 显存模式,可选项见下表 | `model_cpu_offload` |
| `ulysses_degree` | 头维度并行度,单卡为 1 | 1 |
| `ring_degree` | 序列维度并行度,单卡为 1 | 1 |
| `fsdp_dit` | 多卡推理时对 Transformer 使用 FSDP 以节省显存 | `False` |
| `fsdp_text_encoder` | 多卡推理时对 Qwen3-VL 文本编码器使用 FSDP 以节省显存 | `False` |
| `compile_dit` | 编译 Transformer 以加速推理(固定分辨率下有效) | `False` |
| `model_name` | 模型路径 | `models/Diffusion_Transformer/MiniMax-H3` |
| `transformer_path` | 训练好的 Transformer 权重路径 | `None` |
| `vae_path` | 训练好的 VAE 权重路径 | `None` |
| `pdd_lora_path` | PDD checkpoint 目录(优先加载 `pdd_ema.safetensors`,否则 `pdd.safetensors`,外加 `pdd_config.json`)或 `.safetensors` 文件 | `output_dir_minimax_h3_pdd_lora/checkpoint-3000` |
| `lora_path` | Turbo/PEFT LoRA;**不能**与 `pdd_lora_path` 同时使用 | `None` |
| `sample_size` | 生成视频分辨率 `[height, width]`;宽高必须是 32 的倍数。设为 `None` 时使用 MiniMax-H3 自带的 16:9 画布(768x1344) | `[768, 1344]` |
| `video_length` | 生成帧数,会向上取整到视频 VAE 可解码的下一个 `17*n+5`(时长保持在 5 到 15 秒) | 124 |
| `fps` | 每秒帧数(MiniMax-H3 固定以 24 fps 生成) | 24 |
| `weight_dtype` | 模型权重精度,不支持 bf16 的 GPU 请使用 `torch.float16` | `torch.bfloat16` |
| `prompt` | 描述生成内容的正向提示词 | `"A red fox trotting..."` |
| `seed` | 用于复现的随机种子 | 43 |
| `num_inference_steps` | 学生 NFE。默认 PDD 配方用 8(不是教师的 40) | 8 |
| `guidance_scale` | 引导强度。发布权重已做 guidance 蒸馏:保持 1 时每步只做一次前向、不走 CFG | 1 |
| `flow_shift` | 视频调度的指数 sigma shift,`None` 时沿用权重自带值(12.0) | `None` |
| `audio_flow_shift` | 音频调度的指数 sigma shift,`None` 时沿用权重自带值(3.0) | `None` |
| `save_path` | 生成视频保存路径 | `samples/minimax-h3-videos-t2v` |
**GPU 显存模式说明**:
| 模式 | 说明 | 显存占用 |
|------|------|---------|
| `model_full_load` | 整个模型加载到 GPU | 最高 |
| `model_full_load_and_qfloat8` | 全量加载 + FP8 量化 | 高 |
| `model_cpu_offload` | 模型用完后卸载到 CPU | 中 |
| `model_cpu_offload_and_qfloat8` | CPU 卸载 + FP8 量化 | 中低 |
| `model_group_offload` | 层级分组在 CPU/CUDA 间换入换出 | 低 |
| `sequential_cpu_offload` | 逐层卸载(最慢) | 最低 |
> 💡 transformer 在 bfloat16 下有 61.7 GB,Qwen3-VL 条件器还有 62.1 GB,因此单张 80 GB 卡需要使用 `model_cpu_offload` 或 `model_group_offload`。推理会加载文本编码器;训练不会。
### 4.2 单 GPU 推理
运行单卡推理:
```bash
python examples/minimax_h3/predict_t2v.py
```
按需编辑 `examples/minimax_h3/predict_t2v.py`。PDD 推理请重点关注以下参数:
```python
# 根据 GPU 显存选择
GPU_memory_mode = "model_cpu_offload"
# 您的实际模型路径
model_name = "models/Diffusion_Transformer/MiniMax-H3"
# PDD checkpoint 目录或权重文件;目录优先加载 pdd_ema.safetensors。rank / alpha / targets / 网格从 pdd_config.json 读取
pdd_lora_path = "output_dir_minimax_h3_pdd_lora/checkpoint-3000"
# 设置 pdd_lora_path 时必须保持 None
lora_path = None
# 学生 NFE;默认 N=32 / L=4 配方为 8。留在 40 时脚本会改成 N / L
num_inference_steps = 8
# 按要生成的内容填写
prompt = "A red fox trotting through a snowy pine forest, snow crunching underfoot"
# ...
```
图生视频和 Ref2VA 在 `examples/minimax_h3/predict_i2v.py`、`examples/minimax_h3/predict_ref2va.py` 里使用同样的 `pdd_lora_path` 字段。Ref2VA 需要 `--train_mode=ref2va` 训出的 checkpoint。
### 4.3 多 GPU 并行推理
**适用场景**:高分辨率生成、推理加速
#### 安装并行推理依赖
```bash
pip install xfuser yunchang
```
#### 配置并行策略
编辑 `examples/minimax_h3/predict_t2v.py`:
```python
# 保证 ulysses_degree × ring_degree = GPU 数量
# 例如使用 2 张 GPU:
ulysses_degree = 2 # 头维度并行
ring_degree = 1 # 序列维度并行
```
**配置原则**:
- `ulysses_degree` 必须能整除模型的注意力头数
- `ring_degree` 在序列维切分,影响通信开销;头数能均分时尽量不用
- 多卡走 xfuser 序列并行路径,与 `*cpu_offload*` 显存模式 **不兼容**(accelerate offload hook 占用单一 device);此时用 `model_full_load` / `model_full_load_and_qfloat8`,并用 `fsdp_dit` / `fsdp_text_encoder` 省显存
- 序列并行需要可用的 FlashAttention。没有时改为独立的单卡任务(`CUDA_VISIBLE_DEVICES=i`,`ulysses_degree = 1`,`ring_degree = 1`)
**配置示例**:
| GPU 数量 | ulysses_degree | ring_degree | 说明 |
|---------|---------------|-------------|------|
| 1 | 1 | 1 | 单卡 |
| 4 | 4 | 1 | 头并行 |
| 8 | 2 | 4 | 混合并行 |
| 8 | 8 | 1 | 头并行 |
#### 运行多卡推理
```bash
torchrun --nproc_per_node=2 examples/minimax_h3/predict_t2v.py
```
## 五、更多资源
- **PDD 论文**:https://arxiv.org/abs/2607.26004
- **MiniMax-H3 官方 GitHub**:https://github.com/MiniMax-AI/MiniMax-H3
- **官方 GitHub**:https://github.com/aigc-apps/VideoX-Fun
- **基础 MiniMax-H3 LoRA 训练**:`scripts/minimax_h3/README_TRAIN_LORA.md`
- **MiniMax-H3 Fun 控制训练**:`scripts/minimax_h3_fun/README_TRAIN.md`
+178
View File
@@ -0,0 +1,178 @@
r"""Cache Qwen3-VL conditioning for data-free PDD (`fl2va`), multi-GPU, in safetensors.
Rewritten from the single-GPU `encode_prompts.py` to follow the repository's preprocessing convention
(mirrored on `scripts/wan2.1_self_forcing/generate_ode_pairs.py`):
`accelerate launch` over every rank, interleaved rank sharding, resume on files that already exist, per-entry
`.safetensors` (never `.pt` / LMDB), `wait_for_everyone`, then rank0 writes an `outputs.json` that
`ImageVideoSafetensorsDataset` consumes. Each entry holds the exact keys `FL2VATrajectory.reset` reads —
`prompt_embeds` (`hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER]`) and `text_token_tags` — so the ~62 GB conditioner
stays out of the PDD training run under `--enable_preprocess_training`.
accelerate launch --mixed_precision="bf16" scripts/minimax_h3/generate_prompt_cache.py \
--pretrained_model_name_or_path=models/Diffusion_Transformer/MiniMax-H3 \
--train_data_meta=datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json \
--output_folder=datasets/minimax_h3_pdd_prompt_cache
`--train_data_meta` is the official demo dataset's annotation JSON (never an ad-hoc prompt set). PDD is
data-free, so only its `text` captions are read — the `file_path` / `audio_path` / `control_file_path` / `width` /
`height` fields the audio-visual dataset carries are ignored here. `load_prompts` also accepts a bare list of prompt
strings or `{"prompt": ...}` records. Data-free PDD has no train/val split: cache once and point both
`--train_data_meta` and `--val_data_meta` of `train_pdd_lora.py` at the same `outputs.json`.
"""
import argparse
import json
import math
import os
import sys
import torch
from accelerate import Accelerator
from safetensors.torch import save_file
from tqdm import tqdm
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 Qwen2TokenizerFast, Qwen3VLForConditionalGeneration
# Reuse the canonical MiniMax-H3 conditioner recipe rather than re-deriving it: `train_lora.encode_prompt` builds the
# presentation and reads `hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER]` exactly as the pipeline does. For a text-only
# `fl2va` request it never touches `processor`, so `None` is passed.
from train_lora import encode_prompt
def load_prompts(path):
r"""Read the annotation JSON into a list of prompt strings.
Accepts the repository annotation format (a list of `{"text": ...}`), a bare list of strings, and the legacy
`{"prompt": ...}` job records / `{"examples": [...]}` wrapper the old `encode_prompts.py` took.
"""
with open(path, encoding="utf-8") as handle:
document = json.load(handle)
if isinstance(document, dict):
document = document.get("examples", document)
if not isinstance(document, list) or not document:
raise ValueError(f"{path} must be a non-empty list of prompts or `{{'text': ...}}` records.")
prompts = []
for index, entry in enumerate(document, start=1):
if isinstance(entry, str):
prompt = entry.strip()
elif isinstance(entry, dict) and isinstance(entry.get("text"), str):
prompt = entry["text"].strip()
elif isinstance(entry, dict) and isinstance(entry.get("prompt"), str):
prompt = entry["prompt"].strip()
else:
raise ValueError(f"Entry {index} in {path} is not a prompt string or a `{{'text'/'prompt': ...}}` record.")
if not prompt:
raise ValueError(f"Entry {index} in {path} is empty.")
prompts.append(prompt)
return prompts
def parse_args():
parser = argparse.ArgumentParser(description="Cache Qwen3-VL `fl2va` conditioning for data-free PDD (multi-GPU).")
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
required=True,
help="Path to the MiniMax-H3 partition; its `tokenizer/` and `text_encoder/` subfolders are read.",
)
parser.add_argument(
"--train_data_meta",
type=str,
required=True,
help="Annotation JSON of prompts: a list of `{\"text\": ...}` records (bare strings / `{\"prompt\": ...}` also accepted).",
)
parser.add_argument(
"--output_folder",
type=str,
required=True,
help="Directory the per-prompt `.safetensors` and the rank0 `outputs.json` are written to (one split).",
)
parser.add_argument(
"--mixed_precision",
type=str,
default="bf16",
choices=["no", "fp16", "bf16"],
help="Mixed precision the conditioner runs at. MiniMax-H3 conditions in bfloat16. Default: bf16.",
)
parser.add_argument("--local_rank", type=int, default=-1, help="For distributed preprocessing: local_rank.")
return parser.parse_args()
def main():
args = parse_args()
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device = accelerator.device
world_size = accelerator.num_processes
rank = accelerator.process_index
torch.set_grad_enabled(False)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
weight_dtype = torch.float32
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
prompts = load_prompts(args.train_data_meta)
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(args.pretrained_model_name_or_path, "tokenizer"))
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, "text_encoder"), low_cpu_mem_usage=True, torch_dtype=weight_dtype,
).to(device).eval()
text_encoder.requires_grad_(False)
os.makedirs(args.output_folder, exist_ok=True)
total_per_rank = int(math.ceil(len(prompts) / world_size))
# Each rank walks an interleaved slice of the prompt list; a file that already exists is skipped so a killed run
# resumes instead of recomputing.
for index in tqdm(range(total_per_rank), disable=rank != 0, desc="Caching fl2va prompts"):
prompt_index = index * world_size + rank
if prompt_index >= len(prompts):
continue
prompt = prompts[prompt_index]
output_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(output_path):
continue
prompt_embeds, text_token_tags = encode_prompt(
text_encoder, tokenizer, None, prompt, device=device, dtype=weight_dtype,
)
save_file(
{
"prompt_embeds": prompt_embeds.to(torch.bfloat16).cpu().contiguous(),
"text_token_tags": text_token_tags.cpu().contiguous().long(),
},
output_path,
metadata={"format": "pt", "prompt": prompt},
)
accelerator.wait_for_everyone()
# rank0 lists every generated file so `ImageVideoSafetensorsDataset` can load the split.
if accelerator.is_main_process:
records = []
for prompt_index in range(len(prompts)):
safetensor_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
if os.path.exists(safetensor_path):
records.append({"file_path": safetensor_path})
json_path = os.path.join(args.output_folder, "outputs.json")
with open(json_path, "w", encoding="utf-8") as handle:
json.dump(records, handle, ensure_ascii=False, indent=4)
print(f"Done. Cached {len(records)} fl2va prompts to {args.output_folder}")
print(f"Annotation JSON: {json_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,12 @@
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATASET_META="datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json"
export CACHE_ROOT="datasets/minimax_h3_pdd_prompt_cache"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/minimax_h3/generate_prompt_cache.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_meta=$DATASET_META \
--output_folder=$CACHE_ROOT
@@ -0,0 +1,293 @@
r"""Cache `ref2va` request conditioning for data-free PDD, multi-GPU, in safetensors.
The `ref2va` counterpart of `generate_prompt_cache.py`: a `ref2va` request conditions on reference media, so
besides the Qwen3-VL `prompt_embeds` / `text_token_tags` it also needs the VAE-encoded reference latents
`Ref2VATrajectory.reset` lays in front of the generated rows. This cache is the *optional* pre-encode route
(README_TRAIN_PDD_LORA.md §3.2.1 Route B): `train_pdd_lora.py --train_mode=ref2va` can instead encode those
latents on the fly from a request annotation (Route A, launched without `--enable_preprocess_training`), but
pre-encoding once keeps the ~62 GB conditioner and both VAEs out of long or repeated training runs.
Follows the same preprocessing convention as `generate_prompt_cache.py` (mirrored on
`scripts/wan2.1_self_forcing/generate_ode_pairs.py`): `accelerate launch`, interleaved rank
sharding, resume on existing files, per-request `.safetensors`, `wait_for_everyone`, then rank0 `outputs.json`.
accelerate launch --mixed_precision="bf16" scripts/minimax_h3/generate_ref2va_request_cache.py \
--pretrained_model_name_or_path=models/Diffusion_Transformer/MiniMax-H3 \
--train_data_meta=datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json \
--output_folder=datasets/minimax_h3_pdd_request_cache \
--transformer_subfolder=transformer_ref \
--video_sample_n_frames=124
`--train_data_meta` is a JSON list of requests. Either an explicit `{"prompt": ..., "references": ["image=...",
"video=...", "audio=..."]}` record (the `predict_ref2va.py` schema, references in the order the model reads them), or
the official audio-visual demo record `{"text", "file_path", "audio_path"}` (never an ad-hoc request set),
whose own video + audio become the `video=`/`audio=` references — see `load_requests`.
"""
import argparse
import json
import os
import sys
import torch
from accelerate import Accelerator
from safetensors.torch import save_file
from tqdm import tqdm
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 (AutoencoderKLMiniMaxH3,
AutoencoderKLMiniMaxH3Audio,
MiniMaxH3Transformer3DModel, Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
Qwen3VLProcessor)
from videox_fun.pipeline.pipeline_minimax_h3 import (
MiniMaxH3AudioReference, MiniMaxH3ImageReference, MiniMaxH3VideoReference,
align_num_frames, check_ref2va_references, normalize_ref2va_references)
# Reuse the canonical recipes rather than re-deriving them: `encode_prompt(references=...)` builds the ref2va
# presentation and `encode_reference_latents_for_training` mirrors `MiniMaxH3Pipeline.encode_reference_latents`
# without needing a pipeline instance.
from train_lora import encode_prompt, encode_reference_latents_for_training
_REFERENCE_KIND_IDS = {"image": 0, "video": 1, "audio": 2}
def parse_reference(entry: str):
r"""Decode one `image=path` / `video=path` / `audio=path` entry into its `MiniMaxH3Reference` (the
`predict_ref2va.py` schema)."""
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}.")
def load_requests(path):
r"""Read the request annotation JSON into a list of `{"prompt": str, "references": [str, ...]}`.
Two record shapes are accepted. An explicit `ref2va` request carries its own `references` (the `predict_ref2va.py`
schema). The official audio-visual demo record (`{"text", "file_path", "audio_path"}`) has none, so its
own video + audio become the references — `video=<file_path>` and `audio=<audio_path>` — which lets `ref2va` be
driven straight from `datasets/X-Fun-Videos-Audios-Demo` exactly as the `fl2va` prompt cache is. A demo media path
is relative to the dataset, so it resolves against the annotation file's own directory.
"""
with open(path, encoding="utf-8") as handle:
document = json.load(handle)
if isinstance(document, dict):
document = document.get("examples", document)
if not isinstance(document, list) or not document:
raise ValueError(f"{path} must be a non-empty list of ref2va requests.")
root = os.path.dirname(os.path.abspath(path))
requests = []
for index, entry in enumerate(document, start=1):
if not isinstance(entry, dict):
raise ValueError(f"Entry {index} in {path} is not a request record.")
prompt = entry.get("prompt", entry.get("text"))
if not isinstance(prompt, str) or not prompt.strip():
raise ValueError(f"Entry {index} in {path} needs a non-empty `prompt`/`text`.")
references = entry.get("references")
if references is None:
# Official audio-visual demo record: derive the references from its own video + audio (an audio reference
# cannot stand alone, and every demo entry ships both), so no ad-hoc request file is needed.
references = []
for key, kind in (("file_path", "video"), ("audio_path", "audio")):
media = entry.get(key)
if isinstance(media, str) and media.strip():
media = media.strip()
references.append(f"{kind}={media if os.path.isabs(media) else os.path.join(root, media)}")
if not isinstance(references, list) or not references or not all(isinstance(r, str) for r in references):
raise ValueError(f"Entry {index} in {path} needs a non-empty `references` list of strings.")
requests.append({"prompt": prompt.strip(), "references": list(references)})
return requests
def parse_args():
parser = argparse.ArgumentParser(description="Cache `ref2va` request conditioning for data-free PDD (multi-GPU).")
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
required=True,
help="Path to the MiniMax-H3 `ref2va` partition; its `tokenizer/`, `processor/`, `text_encoder/`, `vae/`, "
"`audio_vae/` and `transformer_ref/` subfolders are read.",
)
parser.add_argument(
"--train_data_meta",
type=str,
required=True,
help="Annotation JSON of requests: a list of `{\"prompt\": ..., \"references\": [\"image=...\", ...]}`.",
)
parser.add_argument(
"--output_folder",
type=str,
required=True,
help="Directory the per-request `.safetensors` and the rank0 `outputs.json` are written to (one split).",
)
parser.add_argument(
"--transformer_subfolder",
type=str,
default="transformer_ref",
help="Subfolder the `patch_size` / `audio_in_channels` are read from (config only, no weights). Default: transformer_ref.",
)
parser.add_argument(
"--video_sample_n_frames",
type=int,
default=124,
help="Generated frame count (form 17 * n + 5) the references are normalized onto. Default: 124.",
)
parser.add_argument(
"--mixed_precision",
type=str,
default="bf16",
choices=["no", "fp16", "bf16"],
help="Mixed precision the conditioner runs at. MiniMax-H3 conditions in bfloat16. Default: bf16.",
)
parser.add_argument("--local_rank", type=int, default=-1, help="For distributed preprocessing: local_rank.")
return parser.parse_args()
def main():
args = parse_args()
accelerator = Accelerator(mixed_precision=args.mixed_precision)
device = accelerator.device
world_size = accelerator.num_processes
rank = accelerator.process_index
torch.set_grad_enabled(False)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
weight_dtype = torch.float32
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
requests = load_requests(args.train_data_meta)
num_frames = align_num_frames(int(args.video_sample_n_frames))
# `patch_size` / `audio_in_channels` / `sampling_rate` are read from config files alone — neither the 33 B
# transformer nor the audio VAE weights are loaded just to read them.
transformer_config = MiniMaxH3Transformer3DModel.load_config(
args.pretrained_model_name_or_path, subfolder=args.transformer_subfolder
)
patch_size = tuple(transformer_config["patch_size"])
audio_channels = transformer_config["audio_in_channels"]
audio_sr = AutoencoderKLMiniMaxH3Audio.load_config(
args.pretrained_model_name_or_path, subfolder="audio_vae"
).get("sampling_rate", 32000)
os.makedirs(args.output_folder, exist_ok=True)
# Interleaved rank sharding with resume: `pending` holds this rank's not-yet-cached request indices. Both passes
# walk the same list, so every rank stays independent and no collective depends on the request count dividing
# evenly across ranks — which is why this generator does not shard the conditioner with FSDP the way inference does.
pending = [
request_index
for request_index in range(rank, len(requests), world_size)
if not os.path.exists(os.path.join(args.output_folder, f"{request_index:05d}.safetensors"))
]
def normalized_references(request):
# Deterministic, so pass 2 reproduces exactly the references pass 1 encoded the prompt against.
references = [parse_reference(entry) for entry in request["references"]]
references = check_ref2va_references(references)
return normalize_ref2va_references(references, num_frames, audio_sr)
# ---- Pass 1/2: the ~62 GB Qwen3-VL conditioner is resident; encode every prompt, then release it. A 124-frame
# reference video does not VAE-encode alongside the conditioner within 80 GB, so the two big models never overlap.
tokenizer = Qwen2TokenizerFast.from_pretrained(os.path.join(args.pretrained_model_name_or_path, "tokenizer"))
processor = Qwen3VLProcessor.from_pretrained(os.path.join(args.pretrained_model_name_or_path, "processor"))
text_encoder = Qwen3VLForConditionalGeneration.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, "text_encoder"), low_cpu_mem_usage=True, torch_dtype=weight_dtype,
).to(device).eval()
text_encoder.requires_grad_(False)
prompt_cache = {}
for request_index in tqdm(pending, disable=rank != 0, desc="Caching ref2va prompts (pass 1/2)"):
request = requests[request_index]
references = normalized_references(request)
prompt_embeds, text_token_tags = encode_prompt(
text_encoder, tokenizer, processor, request["prompt"], references=references, device=device, dtype=weight_dtype,
)
prompt_cache[request_index] = (
prompt_embeds.to(torch.bfloat16).cpu().contiguous(),
text_token_tags.cpu().contiguous().long(),
)
del references, prompt_embeds, text_token_tags
del text_encoder, tokenizer, processor
torch.cuda.empty_cache()
# ---- Pass 2/2: only the two VAEs are resident; encode the reference latents and write each request. The VAEs stay
# float32 as released (the encode recipe is float16 autocast over float32 weights).
vae = AutoencoderKLMiniMaxH3.from_pretrained(
args.pretrained_model_name_or_path, subfolder="vae", low_cpu_mem_usage=True,
).to(device).eval()
audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained(
args.pretrained_model_name_or_path, subfolder="audio_vae", low_cpu_mem_usage=True,
).to(device).eval()
vae.requires_grad_(False)
audio_vae.requires_grad_(False)
for request_index in tqdm(pending, disable=rank != 0, desc="Caching ref2va latents (pass 2/2)"):
request = requests[request_index]
references = normalized_references(request)
condition_latents, audio_condition_latents = encode_reference_latents_for_training(
vae, audio_vae, references, patch_size, device, audio_latent_channels=audio_channels,
)
reference_kinds = [(reference.kind, bool(reference.has_audio)) for reference in references]
prompt_embeds, text_token_tags = prompt_cache.pop(request_index)
tensors = {
"prompt_embeds": prompt_embeds,
"text_token_tags": text_token_tags,
# safetensors holds tensors only, so the ragged reference structure is flattened: the kind / has-audio
# pairs become two int vectors and the per-reference latents become indexed tensors under a count.
"reference_kind_ids": torch.tensor([_REFERENCE_KIND_IDS[kind] for kind, _ in reference_kinds], dtype=torch.long),
"reference_has_audio": torch.tensor([int(has_audio) for _, has_audio in reference_kinds], dtype=torch.long),
"num_condition_latents": torch.tensor(len(condition_latents), dtype=torch.long),
"num_audio_condition_latents": torch.tensor(len(audio_condition_latents), dtype=torch.long),
}
for position, latent in enumerate(condition_latents):
tensors[f"condition_latents_{position}"] = latent.cpu().contiguous().float()
for position, latent in enumerate(audio_condition_latents):
tensors[f"audio_condition_latents_{position}"] = latent.cpu().contiguous().float()
save_file(
tensors,
os.path.join(args.output_folder, f"{request_index:05d}.safetensors"),
metadata={"format": "pt", "prompt": request["prompt"], "references": json.dumps(request["references"])},
)
del references, condition_latents, audio_condition_latents
accelerator.wait_for_everyone()
if accelerator.is_main_process:
records = []
for request_index in range(len(requests)):
safetensor_path = os.path.join(args.output_folder, f"{request_index:05d}.safetensors")
if os.path.exists(safetensor_path):
records.append({"file_path": safetensor_path})
json_path = os.path.join(args.output_folder, "outputs.json")
with open(json_path, "w", encoding="utf-8") as handle:
json.dump(records, handle, ensure_ascii=False, indent=4)
print(f"Done. Cached {len(records)} ref2va requests to {args.output_folder}")
print(f"Annotation JSON: {json_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,15 @@
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export REQUEST_META="datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json"
export CACHE_ROOT="datasets/minimax_h3_pdd_request_cache"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" \
scripts/minimax_h3/generate_ref2va_request_cache.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_meta=$REQUEST_META \
--output_folder=$CACHE_ROOT \
--transformer_subfolder=transformer_ref \
--video_sample_n_frames=124
+23 -2
View File
@@ -753,6 +753,11 @@ def parse_args():
default=512,
help="Sample size of the video.",
)
parser.add_argument(
"--fix_sample_size",
nargs=2, type=int, default=None,
help="Fix Sample size [height, width] when using bucket and collate_fn.",
)
parser.add_argument(
"--video_sample_stride",
type=int,
@@ -835,6 +840,11 @@ def main():
f"`video_sample_size` {args.video_sample_size} must be a multiple of 32: the canvas is patched "
"2x2 into the transformer and its RoPE grid keys off that."
)
if args.fix_sample_size is not None and (args.fix_sample_size[0] % 32 or args.fix_sample_size[1] % 32):
raise ValueError(
f"`fix_sample_size` {args.fix_sample_size} must be multiples of 32: the canvas is patched "
"2x2 into the transformer and its RoPE grid keys off that."
)
aligned_frames = align_num_frames(int(args.video_sample_n_frames))
if aligned_frames != int(args.video_sample_n_frames):
raise ValueError(
@@ -1154,6 +1164,14 @@ def main():
# once with the pipeline's torchaudio pass onto the audio VAE's sample rate (32 kHz, 40 latents/s), stereo kept
# as released, over the `num_frames / fps` span the audio latent grid keys off.
audio_sr = getattr(audio_vae.config, "sampling_rate", 32000)
# A fixed canvas overrides the bucket's aspect-ratio search: pin every sample to `fix_sample_size`, so the
# dataset's resize ceiling must cover it and the random-resolution paths turn off.
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
args.training_with_video_token_length = False
args.random_hw_adapt = False
train_dataset = VideoSpeechDataset(
args.train_data_meta, args.train_data_dir,
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride,
@@ -1281,8 +1299,11 @@ def main():
aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 32) * 32 for x in closest_size]
if args.fix_sample_size is not None:
closest_size = [int(x / 32) * 32 for x in args.fix_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 32) * 32 for x in closest_size]
min_example_length = min(
[example["pixel_values"].shape[0] for example in examples]
+4
View File
@@ -3,6 +3,10 @@ export DATASET_NAME="datasets/X-Fun-Videos-Audios-Demo/"
export DATASET_META_NAME="datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json"
NCCL_DEBUG=INFO
# Resolution: --video_sample_size is the square bucket ceiling used with --random_hw_adapt / --enable_bucket.
# To pin every sample to a fixed non-square canvas instead, add "--fix_sample_size <height> <width>" (e.g. 768 1344);
# it overrides --video_sample_size and turns off --random_hw_adapt / --training_with_video_token_length.
accelerate launch --mixed_precision="bf16" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
File diff suppressed because it is too large Load Diff
+40
View File
@@ -0,0 +1,40 @@
# fl2va off a pre-encoded prompt cache (README_TRAIN_PDD_LORA.md §3.2, `--enable_preprocess_training`).
# For ref2va — direct load (Route A) or request cache (Route B) — see §3.2.1 and generate_ref2va_request_cache.sh.
export MODEL_NAME="models/Diffusion_Transformer/MiniMax-H3"
export DATA_DIR=""
export PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
export VAL_PROMPT_CACHE_META="datasets/minimax_h3_pdd_prompt_cache/outputs.json"
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="no" --use_fsdp \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=MiniMaxH3TransformerBlock \
--fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT \
--fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False \
scripts/minimax_h3/train_pdd_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--enable_preprocess_training \
--train_data_dir=$DATA_DIR \
--train_data_meta=$PROMPT_CACHE_META \
--val_data_meta=$VAL_PROMPT_CACHE_META \
--video_sample_n_frames=124 \
--fix_sample_size 768 1344 \
--train_batch_size=1 \
--max_train_steps=3000 \
--checkpointing_steps=200 \
--learning_rate=1e-5 \
--lora_learning_rate=1e-4 \
--seed=43 \
--output_dir="output_dir_minimax_h3_pdd_lora" \
--gradient_checkpointing \
--gradient_checkpointing_save_on_cpu \
--mixed_precision="no" \
--adam_weight_decay=0.0 \
--max_grad_norm=1.0 \
--rank=64 \
--network_alpha=64 \
--low_vram \
--target_name="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear" \
--train_mode="fl2va" \
--pdd_num_steps=32 \
--pdd_block_size=4 \
--validation_steps=200
+12
View File
@@ -1,4 +1,16 @@
import importlib.util
import os
if importlib.util.find_spec("paifuser") is not None:
import paifuser
# Imported conditionally rather than unconditionally-then-bailing the way `videox_fun.pipeline` does it, because
# this runs on `import videox_fun` itself: with the variable unset that import does not load `perf_metrics` at all,
# and costs exactly what it did before. (`videox_fun.utils` exports from it and so loads it either way, which buys
# nothing but a bytecode load -- the module imports only the standard library and `torch`.) Here rather than in
# `videox_fun.pipeline` because the training scripts are what this measures, and three of them never import that
# package -- but all 113 of them import `videox_fun.models`, which this reaches through.
if os.environ.get("VIDEOX_PERF", "0").strip() not in ("", "0"):
from .utils.perf_metrics import install_training as _install_perf_training
_install_perf_training()
+148
View File
@@ -0,0 +1,148 @@
r"""Parallel Decoding Distillation (PDD) parts for MiniMax-H3.
The model-agnostic pieces — the plan math, [`PDDParallelHead`], [`PDDLoRALinear`], [`pdd_teacher_mode`], the checkpoint
resolution — live in `videox_fun/utils/lora_utils_pdd.py` and are re-exported below, so every existing
`from videox_fun.models.minimax_h3_pdd import ...` site keeps working. This file keeps only the MiniMax-H3 glue:
which layers become parallel heads, how the packed-sequence teacher forward is called, and how the pipeline's step
callback arms the plans.
PDD (arXiv 2607.26004) turns a pre-trained flow model into a *parallel decoder*: the sampling interval is discretized
into `N` intervals grouped into blocks of size `L`, and one network evaluation predicts the **mean velocity of every
interval of the next block** instead of the single instantaneous velocity. Generation then advances `L` intervals per
evaluation, i.e. `NFE = N / L`.
MiniMax-H3 has two final heads — `proj_out` for the video rows and `audio_proj_out` for the audio rows — and both are
repeated, as the paper does for the two towers of LTX-2.3. Video and audio ride the same block structure on two
schedules (`shift=12.0` / `3.0`), so both modalities take the stage together and each advances by its own step size.
The frozen teacher is the same module under [`pdd_teacher_mode`]: the low-rank updates of the backbone are switched off and
both heads fall back to the pre-trained weights they were built from. There is no second copy of the 33 B backbone.
"""
import os
import torch
from videox_fun.utils.lora_utils_pdd import (
PDD_EMA_WEIGHTS_NAME, PDD_LEGACY_LIVE_WEIGHTS_NAME, PDD_WEIGHTS_NAME, PDDLoRALinear, PDDParallelHead, add_pdd_lora,
load_pdd_config, merge_pdd_lora, pdd_num_inference_steps, pdd_sampling_plan, pdd_state_dict, pdd_time_grid,
pdd_training_plan, pdd_teacher_mode, resolve_pdd_lora_path, shifted_sigma,
)
# The generic class was `MiniMaxH3ParallelHead` before the model-agnostic parts moved to
# `videox_fun/utils/lora_utils_pdd.py`; the alias is the same class object, so old imports and `isinstance` checks
# keep working.
MiniMaxH3ParallelHead = PDDParallelHead
def attach_parallel_decoder(transformer, num_steps: int) -> None:
r"""
Turn a `MiniMaxH3Transformer3DModel` into a PDD parallel decoder, in place.
Both final heads are replaced by [`PDDParallelHead`]s of `num_steps` heads each, initialized from the weights they
replace. Nothing else about the model changes: the two heads keep the names `proj_out` and `audio_proj_out`, so
the float32 pinning of the mixed-precision checkpoint (`_keep_in_fp32_modules`) and the forward that reads
`self.proj_out.weight.dtype` both still apply.
Args:
transformer (`MiniMaxH3Transformer3DModel`): The model to convert.
num_steps (`int`): The PDD grid size `N`.
"""
transformer.proj_out = PDDParallelHead(transformer.proj_out, num_steps)
transformer.audio_proj_out = PDDParallelHead(transformer.audio_proj_out, num_steps)
def set_parallel_plan(transformer, video_plan: torch.Tensor, audio_plan: torch.Tensor) -> None:
r"""Set the plans of both parallel heads for the next forward pass."""
transformer.proj_out.set_plan(video_plan)
transformer.audio_proj_out.set_plan(audio_plan)
def pdd_teacher_mean_velocity(teacher, forward_kwargs, video, audio, index, grids, solver: str):
r"""
A Runge-Kutta estimate of the teacher's mean velocity over interval `index` of the grid (eq. 5 / eq. 6).
Video and audio ride the same block structure on two schedules, so both modalities take the stage together and
each advances by its own step size. The caller must already have put the model in [`pdd_teacher_mode`].
Args:
teacher: The transformer, under [`pdd_teacher_mode`].
forward_kwargs (`Callable[[float, float], dict]`):
Builds everything but the two latent streams for a forward at a given `(video_time, audio_time)` — the
conditioning, the row timesteps and the packed layout, all of which are the caller's business.
video (`torch.Tensor`), audio (`torch.Tensor`): The state the mean velocity is estimated at.
index (`int`): The grid interval.
grids (`tuple`): `(video_grid, audio_grid, video_step_sizes, audio_step_sizes)`.
solver (`str`): `"euler"` for one evaluation, `"midpoint"` for two.
Returns:
`tuple[torch.Tensor, torch.Tensor]`: the video and audio mean velocities, in float32.
"""
video_grid, audio_grid, video_steps, audio_steps = grids
video_time, audio_time = float(video_grid[index]), float(audio_grid[index])
velocity = teacher(
hidden_states=video[None], audio_hidden_states=audio[None], **forward_kwargs(video_time, audio_time)
)
if solver == "euler":
return velocity[0][0].float(), velocity[1][0].float()
half_video, half_audio = 0.5 * float(video_steps[index]), 0.5 * float(audio_steps[index])
mid_video = video + half_video * velocity[0][0].float()
mid_audio = audio + half_audio * velocity[1][0].float()
velocity = teacher(
hidden_states=mid_video[None],
audio_hidden_states=mid_audio[None],
**forward_kwargs(video_time + half_video, audio_time + half_audio),
)
return velocity[0][0].float(), velocity[1][0].float()
def load_pdd_lora(transformer, pdd_lora_path):
r"""
Attach the parallel heads and LoRA, then load the resolved PDD weights into `transformer`.
A checkpoint directory loads `pdd_ema.safetensors` when present (EMA inference export) and otherwise
`pdd.safetensors`. Returns the config the predict scripts need to arm the heads and pick NFE.
"""
path = resolve_pdd_lora_path(pdd_lora_path)
config = load_pdd_config(path)
add_pdd_lora(
transformer,
config["lora_targets"].split(","),
int(config["lora_rank"]),
float(config["lora_alpha"]),
)
attach_parallel_decoder(transformer, int(config["pdd_num_steps"]))
if path.endswith("safetensors"):
from safetensors.torch import load_file
state_dict = load_file(path)
else:
state_dict = torch.load(path, map_location="cpu")
_, unexpected = transformer.load_state_dict(state_dict, strict=False)
print(f"From PDD checkpoint: {path} ({len(state_dict)} tensors, unexpected keys: {len(unexpected)})", flush=True)
assert not unexpected, f"{path} holds keys the parallel decoder does not have, e.g. {unexpected[:3]}."
return config
def pdd_step_callback(transformer, scheduler, audio_scheduler, config, num_inference_steps):
r"""Arm the fused block-mean plan before each pipeline step. Call this, then pass the return value as `callback_on_step_end`."""
video_steps = pdd_time_grid(scheduler.shift, int(config["pdd_num_steps"])).diff()
audio_steps = pdd_time_grid(audio_scheduler.shift, int(config["pdd_num_steps"])).diff()
block_size = int(config["pdd_num_steps"]) // int(num_inference_steps)
def arm(step_index):
start = step_index * block_size
set_parallel_plan(
transformer,
pdd_sampling_plan(video_steps, start, block_size).float(),
pdd_sampling_plan(audio_steps, start, block_size).float(),
)
arm(0)
def callback(pipe, step_index, timestep, callback_kwargs):
if step_index + 1 < int(num_inference_steps):
arm(step_index + 1)
return {}
return callback
@@ -112,6 +112,23 @@ class CasualWanSelfAttention(nn.Module):
self.local_attn_size = local_attn_size
self.sink_size = sink_size
# Forcing-KV (arXiv 2605.09681) hybrid KV cache compression state,
# aligned with the official zju-jiyicheng/Forcing-KV architecture.
# Configured by WanSelfForcingPipeline.set_forcing_kv_config; defaults
# keep the original full-window attention path bit-identical.
self.forcing_kv_enable = False
self.forcing_kv_static_heads = None
self.forcing_kv_dynamic_heads = None
self.forcing_kv_layer_idx = -1
self.forcing_kv_ar_start = 1
self.forcing_kv_spatial_context_length = 1
self.forcing_kv_temporal_context_length = 1
self.forcing_kv_dynamic_context_length = 1
self.forcing_kv_num_frame_patch = 6
self.forcing_kv_sim_retention_ratio = 0.33
# Lazily-built, device-cached head index tensors (see _fkv_head_indices).
self._fkv_idx_cache = None
# Layers
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
@@ -131,7 +148,8 @@ class CasualWanSelfAttention(nn.Module):
current_start=0,
cache_start=None,
dtype=torch.bfloat16,
t=0
t=0,
forcing_kv_state=None,
):
r"""
Args:
@@ -143,6 +161,9 @@ class CasualWanSelfAttention(nn.Module):
kv_cache: KV cache for causal self-attention
current_start: Current starting position in token sequence
cache_start: Cache starting position
forcing_kv_state: Dict shared across blocks within one forward call,
used by Forcing-KV to pass per-chunk metadata (e.g. frames per
chunk) from the model down to the self-attention layers.
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
if cache_start is None:
@@ -259,6 +280,35 @@ class CasualWanSelfAttention(nn.Module):
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
# Forcing-KV: per-cache compression state, lazily created on first
# use. Before the switch step (ar_start) every layer attends over
# [sink + recent history + current chunk]; after the switch, layer 0
# keeps that path while layers >= 1 split heads into static/dynamic
# groups with separate history budgets (official architecture).
forcing_kv_active = bool(
self.forcing_kv_enable and self.forcing_kv_static_heads is not None)
is_appending = current_end > kv_cache["global_end_index"].item()
clean_pass = bool(forcing_kv_state is not None and forcing_kv_state.get("clean_pass", False))
fkv = None
fkv_switched = False
if forcing_kv_active:
if "forcing_kv" not in kv_cache:
kv_cache["forcing_kv"] = {
"switched_at": None,
"dyn_k": None,
"dyn_v": None,
"dyn_valid": 0,
}
fkv = kv_cache["forcing_kv"]
fpb = max(1, int(self.num_frame_per_block))
ar_step = current_start // (frame_seqlen * fpb)
if clean_pass and ar_step >= self.forcing_kv_ar_start and fkv["switched_at"] is None:
# Mirror the official switch: it only takes effect from the
# next chunk on, so the switching chunk itself stays
# ungrouped.
fkv["switched_at"] = ar_step
fkv_switched = fkv["switched_at"] is not None and ar_step > fkv["switched_at"]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
@@ -283,17 +333,31 @@ class CasualWanSelfAttention(nn.Module):
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
# Compute attention with local window
if self.local_attn_size == -1:
max_attention_size = local_end_index
else:
max_attention_size = self.local_attn_size * frame_seqlen
x = attention(
roped_query,
kv_cache["k"][:, max(0, local_end_index - max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - max_attention_size):local_end_index]
)
window_start = max(0, local_end_index - max_attention_size)
if forcing_kv_active:
x = self._forcing_kv_grouped_attention(
roped_query, kv_cache, fkv, fkv_switched,
local_start_index, local_end_index, frame_seqlen)
if clean_pass and fkv_switched and self.forcing_kv_layer_idx != 0:
self._forcing_kv_dynamic_update(
kv_cache, fkv, local_start_index, local_end_index,
frame_seqlen, forcing_kv_state=forcing_kv_state)
else:
x = attention(
roped_query,
kv_cache["k"][:, window_start:local_end_index],
kv_cache["v"][:, window_start:local_end_index]
)
# Expose the attention window for offline head profiling
kv_cache["_fkv_last_q"] = roped_query
kv_cache["_fkv_window_start"] = window_start
kv_cache["_fkv_local_end"] = local_end_index
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
@@ -302,6 +366,204 @@ class CasualWanSelfAttention(nn.Module):
x = self.o(x)
return x
def _fkv_head_indices(self, device):
"""Lazily build and cache the static/dynamic head index tensors.
Constant for a given head profile; cached on the module keyed by device
so the hot path (grouped attention / dynamic update) never rebuilds the
tensors or pays their host->device copies. Reset by set_forcing_kv_config.
"""
cache = self._fkv_idx_cache
if cache is not None and cache[0] == device:
return cache[1], cache[2]
s_idx = torch.tensor(
self.forcing_kv_static_heads or [], dtype=torch.long, device=device)
d_idx = torch.tensor(
self.forcing_kv_dynamic_heads or [], dtype=torch.long, device=device)
self._fkv_idx_cache = (device, s_idx, d_idx)
return s_idx, d_idx
# ------------------------------------------------------------------
# Forcing-KV: Hybrid KV Cache Compression for Efficient Autoregressive
# Video Diffusion Models (arXiv 2605.09681), aligned with the official
# zju-jiyicheng/Forcing-KV implementation. Static ("spatial") heads
# attend over [sink + last spatial_context_length frames + current
# chunk]; dynamic ("temporal") heads attend over [sink + compressed
# segment cache + last temporal_context_length frames + current chunk].
# Layer 0 and pre-switch chunks use the ungrouped path over
# [sink + max(spatial, temporal) history frames + current chunk].
# ------------------------------------------------------------------
def _forcing_kv_grouped_attention(self, roped_query, kv_cache, fkv,
fkv_switched, local_start_index,
local_end_index, frame_seqlen):
"""Per-head-group attention over the compacted KV cache."""
k_full = kv_cache["k"]
v_full = kv_cache["v"]
sink_tokens = self.sink_size * frame_seqlen
# Early chunks may hold fewer tokens than the sink budget
sink_end = min(sink_tokens, local_start_index)
cur = slice(local_start_index, local_end_index)
def hist_slice(n_frames):
start = max(sink_end, local_start_index - n_frames * frame_seqlen)
return slice(start, local_start_index)
layer_idx = self.forcing_kv_layer_idx
if not fkv_switched or layer_idx == 0:
# Ungrouped path shared by layer 0 and pre-switch chunks
n_hist = max(self.forcing_kv_spatial_context_length,
self.forcing_kv_temporal_context_length)
hist = hist_slice(n_hist)
return attention(
roped_query,
torch.cat([k_full[:, :sink_end], k_full[:, hist], k_full[:, cur]], dim=1),
torch.cat([v_full[:, :sink_end], v_full[:, hist], v_full[:, cur]], dim=1),
)
# Head index tensors are built lazily once and cached on the module
# (constant for a given head profile); see _fkv_head_indices.
s_idx, d_idx = self._fkv_head_indices(roped_query.device)
x = roped_query.new_empty(roped_query.shape)
if s_idx.numel() > 0:
hist_s = hist_slice(self.forcing_kv_spatial_context_length)
k_s = torch.cat([
k_full[:, :sink_end].index_select(2, s_idx),
k_full[:, hist_s].index_select(2, s_idx),
k_full[:, cur].index_select(2, s_idx),
], dim=1)
v_s = torch.cat([
v_full[:, :sink_end].index_select(2, s_idx),
v_full[:, hist_s].index_select(2, s_idx),
v_full[:, cur].index_select(2, s_idx),
], dim=1)
x.index_copy_(2, s_idx, attention(
roped_query.index_select(2, s_idx), k_s, v_s))
if d_idx.numel() > 0:
hist_t = hist_slice(self.forcing_kv_temporal_context_length)
k_parts = [k_full[:, :sink_end].index_select(2, d_idx)]
v_parts = [v_full[:, :sink_end].index_select(2, d_idx)]
dyn_valid = int(fkv["dyn_valid"])
if dyn_valid > 0 and fkv["dyn_k"] is not None:
k_parts.append(fkv["dyn_k"][:, :dyn_valid])
v_parts.append(fkv["dyn_v"][:, :dyn_valid])
k_parts.append(k_full[:, hist_t].index_select(2, d_idx))
k_parts.append(k_full[:, cur].index_select(2, d_idx))
v_parts.append(v_full[:, hist_t].index_select(2, d_idx))
v_parts.append(v_full[:, cur].index_select(2, d_idx))
x.index_copy_(2, d_idx, attention(
roped_query.index_select(2, d_idx),
torch.cat(k_parts, dim=1), torch.cat(v_parts, dim=1)))
return x
def _forcing_kv_dynamic_update(self, kv_cache, fkv, local_start_index,
local_end_index, frame_seqlen,
forcing_kv_state=None):
"""Refresh the dynamic compressed cache (official dynamic_compression).
Candidate frames are the last temporal_context_length history frames
plus the current chunk; the final chunk frame acts as the boundary.
Adjacent frame pairs are scored segment-wise with fp32 cosine
similarity computed per token (heads x dim vectors) and averaged over
each segment's tokens (official token_scores.mean(dim=3)); the least
similar segments are kept. Layer 1 computes the keep decision once
and every other layer of the same forward call reuses it through the
per-forward forcing_kv_state (the official build shares via a class
attribute on a single cache; here the state dict is fresh per forward
so the two CFG caches stay independent).
"""
layer_idx = self.forcing_kv_layer_idx
d_idx = self._fkv_head_indices(kv_cache["k"].device)[1]
temporal = max(0, int(self.forcing_kv_temporal_context_length))
n_patch = int(self.forcing_kv_num_frame_patch)
retention = float(self.forcing_kv_sim_retention_ratio)
if d_idx.numel() == 0 or temporal == 0 or n_patch <= 0:
fkv["dyn_valid"] = 0
return
if frame_seqlen % n_patch != 0:
raise ValueError(
f"Frame token count {frame_seqlen} must be divisible by "
f"forcing_kv_num_frame_patch {n_patch}")
k_full = kv_cache["k"]
v_full = kv_cache["v"]
hist_start = max(self.sink_size * frame_seqlen,
local_start_index - temporal * frame_seqlen)
old_k = k_full[:, hist_start:local_start_index]
new_k = k_full[:, local_start_index:local_end_index]
if old_k.shape[1] < temporal * frame_seqlen or new_k.shape[1] < frame_seqlen:
# Not enough history yet; skip the dynamic cache for this chunk
fkv["dyn_valid"] = 0
return
old_v = v_full[:, hist_start:local_start_index]
new_v = v_full[:, local_start_index:local_end_index]
chain_k = torch.cat([old_k, new_k], dim=1)
chain_v = torch.cat([old_v, new_v], dim=1)
num_chain_frames = chain_k.shape[1] // frame_seqlen
pairs = num_chain_frames - 1
chunk_tokens = frame_seqlen // n_patch
total_chunks = pairs * n_patch
device = chain_k.device
capacity_chunks = int(self.forcing_kv_dynamic_context_length) * n_patch
keep_chunk_count = max(0, min(
min(capacity_chunks, total_chunks),
int(round(total_chunks * retention))))
keep_indices = (forcing_kv_state or {}).get("fkv_keep_indices")
if keep_indices is None:
# Official scoring: cosine similarity per token (heads x dim)
# between adjacent frames, averaged over the segment's tokens.
chain = chain_k.index_select(2, d_idx).float()
frames = chain.reshape(
chain.shape[0], num_chain_frames, n_patch, chunk_tokens, -1)
token_scores = torch.nn.functional.cosine_similarity(
frames[:, :-1], frames[:, 1:], dim=-1) # [B, pairs, n_patch, chunk_tokens]
patch_scores = token_scores.mean(dim=3)[0] # [pairs, n_patch]
if keep_chunk_count > 0:
keep = torch.topk(
patch_scores.reshape(-1), k=keep_chunk_count, largest=False).indices
keep_indices = torch.sort(keep).values
else:
keep_indices = torch.empty(
(0,), device=device, dtype=torch.long)
if forcing_kv_state is not None and layer_idx == 1:
# First grouped layer: publish the decision for this forward
forcing_kv_state["fkv_keep_indices"] = keep_indices
if keep_indices.numel() == 0:
fkv["dyn_valid"] = 0
return
keep_chunk_count = keep_indices.numel()
# Fill this layer's dynamic cache with the selected candidate
# segments (all chain frames except the boundary frame)
cand_k = chain_k[:, : pairs * frame_seqlen].reshape(
chain_k.shape[0], total_chunks, chunk_tokens,
chain_k.shape[2], chain_k.shape[3]).index_select(1, keep_indices)
cand_v = chain_v[:, : pairs * frame_seqlen].reshape(
chain_v.shape[0], total_chunks, chunk_tokens,
chain_v.shape[2], chain_v.shape[3]).index_select(1, keep_indices)
n_dyn = d_idx.numel()
head_dim = chain_k.shape[-1]
sel_k = cand_k.index_select(3, d_idx).reshape(
cand_k.shape[0], keep_chunk_count * chunk_tokens, n_dyn, head_dim)
sel_v = cand_v.index_select(3, d_idx).reshape(
cand_v.shape[0], keep_chunk_count * chunk_tokens, n_dyn, head_dim)
capacity_tokens = int(self.forcing_kv_dynamic_context_length) * frame_seqlen
if fkv["dyn_k"] is None or fkv["dyn_k"].shape[1] < sel_k.shape[1] \
or fkv["dyn_k"].shape[2] != n_dyn:
# Allocate once per layer; reused (overwritten) every chunk
fkv["dyn_k"] = sel_k.new_zeros(
(sel_k.shape[0], max(capacity_tokens, sel_k.shape[1]), n_dyn, head_dim))
fkv["dyn_v"] = sel_v.new_zeros(
(sel_v.shape[0], max(capacity_tokens, sel_v.shape[1]), n_dyn, head_dim))
fkv["dyn_k"][:, : sel_k.shape[1]] = sel_k
fkv["dyn_v"][:, : sel_v.shape[1]] = sel_v
fkv["dyn_valid"] = sel_k.shape[1]
class CasualWanT2VCrossAttention(CasualWanSelfAttention):
"""Text-to-video cross-attention layer."""
@@ -482,6 +744,7 @@ class CasualWanAttentionBlock(nn.Module):
block_mask=None,
dtype=torch.bfloat16,
t=0,
forcing_kv_state=None,
):
r"""
Args:
@@ -497,15 +760,23 @@ class CasualWanAttentionBlock(nn.Module):
current_start: Current starting position in token sequence
cache_start: Cache starting position
block_mask: Block mask for flex attention
forcing_kv_state: Dict shared across blocks within one forward call,
forwarded to the self-attention layer for Forcing-KV.
"""
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
# Self-attention with modulation
attn_kwargs = {}
if forcing_kv_state is not None:
# Only forwarded when set: the USP-replaced self-attn forward used
# in multi-GPU inference does not accept this keyword.
attn_kwargs["forcing_kv_state"] = forcing_kv_state
y = self.self_attn(
(self.norm1(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2),
seq_lens, grid_sizes,
freqs, block_mask, kv_cache, current_start, cache_start)
freqs, block_mask, kv_cache, current_start, cache_start,
**attn_kwargs)
# Residual connection with modulation
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2)
@@ -601,6 +872,17 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
# Self-Forcing causal inference parameters
local_attn_size=-1,
sink_size=0,
# Forcing-KV hybrid KV cache compression (arXiv 2605.09681),
# aligned with the official zju-jiyicheng/Forcing-KV architecture
forcing_kv_enable=False,
forcing_kv_head_profile=None,
forcing_kv_ar_start=1,
forcing_kv_spatial_context_length=1,
forcing_kv_temporal_context_length=1,
forcing_kv_dynamic_context_length=1,
forcing_kv_num_frame_patch=6,
forcing_kv_sim_retention_ratio=0.33,
):
r"""
Initialize the diffusion model backbone.
@@ -697,6 +979,17 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
block.self_attn.layer_idx = layer_idx
block.self_attn.num_layers = self.num_layers
# Forcing-KV: propagate per-head compression config to all self-attn
# layers (training-free; disabled by default keeps behavior unchanged).
self.set_forcing_kv_config(
forcing_kv_enable, forcing_kv_head_profile,
ar_start=forcing_kv_ar_start,
spatial_context_length=forcing_kv_spatial_context_length,
temporal_context_length=forcing_kv_temporal_context_length,
dynamic_context_length=forcing_kv_dynamic_context_length,
num_frame_patch=forcing_kv_num_frame_patch,
sim_retention_ratio=forcing_kv_sim_retention_ratio)
# Head
self.head = CausalHead(dim, out_dim, patch_size, eps)
@@ -714,6 +1007,72 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
self.sp_world_rank = 0
self.init_weights()
def set_forcing_kv_config(self, enable, head_profile=None, ar_start=1,
spatial_context_length=1, temporal_context_length=1,
dynamic_context_length=1, num_frame_patch=6,
sim_retention_ratio=0.33):
"""Propagate Forcing-KV settings to every self-attention layer.
Args:
enable: Master switch. When False the original full-window
attention path is used and results are bit-identical.
head_profile: Offline head classification. Supports the official
format {"layers": [{"layer_idx", "static_head", "dynamic_head"}]}
as well as a plain {layer: {"static": [...]}} mapping.
ar_start: First AR step (chunk index) after which grouped
head compression activates.
spatial_context_length: History frames kept for static heads.
temporal_context_length: Recent frames kept for dynamic heads.
dynamic_context_length: Compressed segment cache capacity (frames).
num_frame_patch: Token segments per latent frame.
sim_retention_ratio: Fraction of candidate segments kept.
"""
head_map = {}
if head_profile is not None:
layers = head_profile.get("layers") if isinstance(head_profile, dict) else None
if layers is not None:
# Official configs_head/*.json format
for entry in layers:
key = entry.get("layer_idx", entry.get("layer"))
static = entry.get("static_head", entry.get("static", []))
head_map[int(key)] = static
else:
for key, entry in head_profile.items():
try:
layer = int(key)
except (TypeError, ValueError):
continue
if isinstance(entry, dict):
entry = entry.get("static", [])
head_map[layer] = entry
for layer_idx, block in enumerate(self.blocks):
attn = block.self_attn
attn.forcing_kv_enable = bool(enable)
attn.forcing_kv_layer_idx = layer_idx
attn.forcing_kv_ar_start = int(ar_start)
attn.forcing_kv_spatial_context_length = int(spatial_context_length)
attn.forcing_kv_temporal_context_length = int(temporal_context_length)
attn.forcing_kv_dynamic_context_length = int(dynamic_context_length)
attn.forcing_kv_num_frame_patch = int(num_frame_patch)
attn.forcing_kv_sim_retention_ratio = float(sim_retention_ratio)
if head_profile is None:
attn.forcing_kv_static_heads = None
static_list = []
else:
static = head_map.get(layer_idx)
static_list = [int(h) for h in static] if static else []
attn.forcing_kv_static_heads = static_list
# Precompute the head partition once (constant for a given profile).
# Index tensors are built lazily on first forward via
# _fkv_head_indices() so they land on the real (post-materialization)
# device and get cached there - this avoids per-forward rebuilds and
# any meta-device buffer copy errors.
static_set = set(static_list)
attn.forcing_kv_dynamic_heads = [
h for h in range(attn.num_heads) if h not in static_set]
attn._fkv_idx_cache = None
def enable_multi_gpus_inference(self):
"""Enable multi-GPU inference with sequence parallelism for KV cache mode."""
self.sp_world_size = get_sequence_parallel_world_size()
@@ -922,6 +1281,7 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
cache_start: int = 0,
clean_x=None,
aug_t=None,
forcing_kv_state: dict = None,
):
r"""
Run the diffusion model with kv caching.
@@ -1088,13 +1448,22 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
return module(*inputs, **kwargs)
return custom_forward
# Forcing-KV: one fresh state dict per forward call, shared across all
# blocks so each self-attn layer can detect the clean context pass.
if forcing_kv_state is None and kv_cache is not None:
forcing_kv_state = {}
if kv_cache is not None:
for block in self.blocks:
block.self_attn.num_frame_per_block = self.num_frame_per_block
for block_index, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
kwargs.update(
{
"kv_cache": kv_cache[block_index] if kv_cache else None,
"current_start": current_start,
"cache_start": cache_start
"cache_start": cache_start,
"forcing_kv_state": forcing_kv_state,
}
)
x = torch.utils.checkpoint.checkpoint(
@@ -1108,7 +1477,8 @@ class WanTransformer3DModel_SelfForcing(WanTransformer3DModel):
"kv_cache": kv_cache[block_index] if kv_cache else None,
"crossattn_cache": crossattn_cache[block_index] if crossattn_cache else None,
"current_start": current_start,
"cache_start": cache_start
"cache_start": cache_start,
"forcing_kv_state": forcing_kv_state,
}
)
x = block(x, **kwargs)
+7 -1
View File
@@ -52,4 +52,10 @@ WanFunPipeline = WanPipeline
WanI2VPipeline = WanFunInpaintPipeline
Wan2_2FunPipeline = Wan2_2Pipeline
Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline
Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline
# Must stay last: `install` measures the pipeline classes present in this namespace, so it has to see all of them.
# It returns on its first line unless `VIDEOX_PERF` is set, so a default import is untouched.
from ..utils.perf_metrics import install as _install_perf_metrics
_install_perf_metrics()
+57 -23
View File
@@ -1709,8 +1709,10 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
def attention_kwargs(self):
return self._attention_kwargs
def check_inputs(self, prompt, height, width, num_frames, num_inference_steps):
if not isinstance(prompt, str):
def check_inputs(self, prompt, height, width, num_frames, num_inference_steps, prompt_embeds=None, text_token_tags=None):
if (prompt_embeds is None) != (text_token_tags is None):
raise ValueError("`prompt_embeds` and `text_token_tags` have to be passed together.")
if prompt_embeds is None and not isinstance(prompt, str):
raise ValueError(
f"MiniMax-H3 packs one request into one sequence, so `prompt` must be a single string, got "
f"{type(prompt)}."
@@ -2269,6 +2271,11 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
audio_latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
text_token_tags: Optional[torch.Tensor] = None,
normalized_references: Optional[List[Any]] = None,
condition_latents: Optional[Union[torch.Tensor, List[torch.Tensor]]] = None,
audio_condition_latents: Optional[List[torch.Tensor]] = None,
output_type: str = "pt",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
@@ -2330,6 +2337,22 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
instead of the draw.
audio_latents (`torch.Tensor`, *optional*):
Pre-generated audio noise of shape `(2, 32, num_audio_latents)`.
prompt_embeds (`torch.Tensor`, *optional*):
A pre-computed conditioning of shape `(1, num_text_tokens, 5120)`, used instead of running the
conditioner. Passed together with `text_token_tags`, it lets a caller that has cached its prompts keep
the 62 GB Qwen3-VL conditioner out of the run entirely.
text_token_tags (`torch.Tensor`, *optional*):
The `(num_text_tokens,)` per-row modality tags that go with `prompt_embeds`.
normalized_references (`list`, *optional*):
References already put through [`normalize_ref2va_references`], used instead of normalizing
`references` again. Together with `condition_latents` the only fields read off them are `kind` and
`has_audio`, so a caller working from a cache may pass lightweight stand-ins.
condition_latents (`torch.Tensor` or `list[torch.Tensor]`, *optional*):
Pre-encoded visual conditioning. A list is the `ref2va` reference latents already produced by
[`~MiniMaxH3Pipeline.encode_reference_latents`]; a tensor is the `fl2va` keyframe rows.
audio_condition_latents (`list[torch.Tensor]`, *optional*):
Pre-encoded `ref2va` soundtrack latents, used together with `condition_latents` so the video VAE
encoder stays out of the request.
output_type (`str`, defaults to `"pt"`):
Output format: `"pil"`, `"np"`, `"pt"`, or `"latent"` for the raw latents.
return_dict (`bool`, defaults to `True`):
@@ -2346,13 +2369,14 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
The generated video, the stereo soundtrack of shape `(1, 2, num_samples)` and its sample rate. Muxing
the two into one file is left to the caller, e.g. with `save_videos_with_audio_grid`.
"""
self.check_inputs(prompt, height, width, num_frames, num_inference_steps)
self.check_inputs(prompt, height, width, num_frames, num_inference_steps, prompt_embeds, text_token_tags)
self._attention_kwargs = attention_kwargs
device = self._execution_device
# `ref2va` is a task of its own: the keyframes of `fl2va` are mutually exclusive with it, and the released
# `ref2va` checkpoint is guidance-distilled with no unconditional branch, so there is no CFG to run.
do_ref2va = bool(references)
# Cached requests pass `normalized_references` (and usually pre-encoded latents) instead of raw media.
do_ref2va = bool(references) or normalized_references is not None
if do_ref2va:
if image is not None or last_image is not None:
raise ValueError(
@@ -2364,7 +2388,8 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
"The `ref2va` checkpoint is guidance-distilled and has no unconditional branch, so `references` "
f"needs `guidance_scale <= 1`, got {guidance_scale}."
)
references = check_ref2va_references(list(references))
if normalized_references is None:
references = check_ref2va_references(list(references))
# 1. Resolve the plan: the canvas, the frame count the video VAE can decode, the latent geometry every later
# step keys off, and the keyframes put onto that canvas.
@@ -2393,8 +2418,11 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
num_audio_latents = audio_latent_num_frames(num_frames)
if do_ref2va:
# The references never bind the generated geometry: they are normalized onto their own resolutions, with
# soundtracks truncated to the resolved duration.
references = normalize_ref2va_references(references, num_frames, self.audio_sampling_rate)
# soundtracks truncated to the resolved duration. A cached request already carries normalized stand-ins.
if normalized_references is None:
references = normalize_ref2va_references(references, num_frames, self.audio_sampling_rate)
else:
references = normalized_references
else:
keyframes = [
prepare_keyframe_image(keyframe, height, width, stretch=index == 0)
@@ -2403,28 +2431,34 @@ class MiniMaxH3Pipeline(DiffusionPipeline):
# 2. Encode MiniMax-H3's presentation of the request. The released checkpoint is guidance-distilled, so the
# default guidance_scale of 1 runs one forward pass per step with no CFG; a guidance_scale above 1 enables
# classifier-free guidance with a negative prompt.
# classifier-free guidance with a negative prompt. Cached `prompt_embeds` skip the 62 GB conditioner.
do_cfg = guidance_scale > 1.0
if do_ref2va:
prompt_embeds, text_token_tags = self.encode_prompt(
prompt, references=references, device=device, dtype=self.transformer.dtype
)
else:
prompt_embeds, text_token_tags = self.encode_prompt(
prompt, keyframes, device=device, dtype=self.transformer.dtype
)
if do_cfg:
negative_prompt = negative_prompt if negative_prompt is not None else ""
negative_prompt_embeds, negative_text_token_tags = self.encode_prompt(
negative_prompt, keyframes, device=device, dtype=self.transformer.dtype
if prompt_embeds is None:
if do_ref2va:
prompt_embeds, text_token_tags = self.encode_prompt(
prompt, references=references, device=device, dtype=self.transformer.dtype
)
else:
prompt_embeds, text_token_tags = self.encode_prompt(
prompt, keyframes, device=device, dtype=self.transformer.dtype
)
else:
prompt_embeds = prompt_embeds.to(device=device, dtype=self.transformer.dtype)
if do_cfg:
if do_ref2va:
raise ValueError("The `ref2va` checkpoint has no unconditional branch, so CFG cannot run.")
negative_prompt = negative_prompt if negative_prompt is not None else ""
negative_prompt_embeds, negative_text_token_tags = self.encode_prompt(
negative_prompt, keyframes, device=device, dtype=self.transformer.dtype
)
# 3. Encode the conditioning and noise it to MiniMax-H3's conditioning level. The anchors are the whole
# denoising loop's invariant: the loop only ever writes the generated rows.
audio_condition_latents = []
condition_latents = None
if do_ref2va:
condition_latents, audio_condition_latents = self.encode_reference_latents(references, device=device)
if condition_latents is None:
condition_latents, audio_condition_latents = self.encode_reference_latents(references, device=device)
elif audio_condition_latents is None:
audio_condition_latents = []
elif keyframes:
condition_latents = self.encode_keyframes(keyframes, device=device)
noise = keyframe_condition_noise(
@@ -1,5 +1,6 @@
# Modified from https://github.com/guandeh17/Self-Forcing/blob/main/pipeline/causal_diffusion_inference.py
import inspect
import json
import math
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
@@ -418,6 +419,79 @@ class WanSelfForcingPipeline(DiffusionPipeline):
def interrupt(self):
return self._interrupt
def _unwrap_transformer(self):
"""Return the bare transformer, undoing FSDP/DDP/torch.compile wraps."""
transformer = self.transformer
while hasattr(transformer, "module"):
transformer = transformer.module
if hasattr(transformer, "_orig_mod"):
transformer = transformer._orig_mod
return transformer
def set_forcing_kv_config(self, enable, head_profile=None, ar_start=1,
spatial_context_length=1, temporal_context_length=1,
dynamic_context_length=1, num_frame_patch=6,
sim_retention_ratio=0.33):
"""Enable/disable Forcing-KV hybrid KV cache compression (arXiv 2605.09681).
Args:
enable: Master switch. False restores the original full-window
attention path bit-identically.
head_profile: Offline head classification - the official
{"layers": [{"layer_idx", "static_head", "dynamic_head"}]}
format, a plain {layer: {"static": [...]}} mapping, or a path
to a JSON file.
ar_start: First AR step (chunk index) after which grouped head
compression activates.
spatial_context_length: History frames kept for static heads.
temporal_context_length: Recent frames kept for dynamic heads.
dynamic_context_length: Compressed segment cache capacity (frames).
num_frame_patch: Token segments per latent frame.
sim_retention_ratio: Fraction of candidate segments kept.
"""
if isinstance(head_profile, str):
with open(head_profile, "r") as f:
head_profile = json.load(f)
self._unwrap_transformer().set_forcing_kv_config(
enable, head_profile,
ar_start=ar_start,
spatial_context_length=spatial_context_length,
temporal_context_length=temporal_context_length,
dynamic_context_length=dynamic_context_length,
num_frame_patch=num_frame_patch,
sim_retention_ratio=sim_retention_ratio)
def _forcing_kv_enabled(self):
"""True when Forcing-KV is active (enabled and a head profile is set)."""
try:
attn = self._unwrap_transformer().blocks[0].self_attn
except (AttributeError, IndexError):
return False
return bool(getattr(attn, "forcing_kv_enable", False)
and getattr(attn, "forcing_kv_static_heads", None) is not None)
def _forcing_kv_cache_tokens(self, frame_seq_length):
"""Per-layer rolling KV budget (tokens) under Forcing-KV.
The shared rolling buffer only needs to cover the largest grouped
read: sink + max(spatial, temporal) history frames + the current
chunk. The dynamic heads' compressed segment cache is allocated
lazily per layer and lives outside this buffer.
"""
cfg = self.transformer.config
local_attn_size = getattr(cfg, 'local_attn_size', -1)
if local_attn_size == -1:
# No rolling window: keep the full buffer; the speed-up then
# comes purely from the grouped short-KV reads.
return None
sink_size = getattr(cfg, 'sink_size', 0)
attn = self._unwrap_transformer().blocks[0].self_attn
spatial_ctx = int(getattr(attn, "forcing_kv_spatial_context_length", 1))
temporal_ctx = int(getattr(attn, "forcing_kv_temporal_context_length", 1))
fpb = max(1, int(getattr(self._unwrap_transformer(), "num_frame_per_block", 1)))
budget_frames = sink_size + max(spatial_ctx, temporal_ctx) + fpb
return budget_frames * frame_seq_length
def _initialize_kv_cache(self, batch_size, dtype, device, frame_seq_length, num_latent_frames):
"""
Initialize KV cache for causal self-attention.
@@ -428,6 +502,11 @@ class WanSelfForcingPipeline(DiffusionPipeline):
local_attn_size = getattr(self.transformer.config, 'local_attn_size', -1)
if local_attn_size != -1:
kv_cache_size = local_attn_size * frame_seq_length
if self._forcing_kv_enabled():
# Forcing-KV: shrink the rolling window to the per-head budget
fkv_cache_size = self._forcing_kv_cache_tokens(frame_seq_length)
if fkv_cache_size is not None:
kv_cache_size = fkv_cache_size
else:
kv_cache_size = num_latent_frames * frame_seq_length
@@ -512,6 +591,14 @@ class WanSelfForcingPipeline(DiffusionPipeline):
stochastic_sampling: bool = True,
streaming: bool = False,
decode_callback: Optional[Callable[[torch.Tensor, int], None]] = None,
forcing_kv_enable: Optional[bool] = None,
forcing_kv_head_profile: Optional[Union[str, Dict[str, Any]]] = None,
forcing_kv_ar_start: Optional[int] = None,
forcing_kv_spatial_context_length: Optional[int] = None,
forcing_kv_temporal_context_length: Optional[int] = None,
forcing_kv_dynamic_context_length: Optional[int] = None,
forcing_kv_num_frame_patch: Optional[int] = None,
forcing_kv_sim_retention_ratio: Optional[float] = None,
) -> Union[WanSelfForcingPipelineOutput, Tuple]:
r"""
Function invoked when calling the pipeline for Self-Forcing causal generation.
@@ -534,6 +621,24 @@ class WanSelfForcingPipeline(DiffusionPipeline):
tensor of shape [B, C, F_pixels, H, W] in [0, 1]. When provided, chunks are NOT
accumulated in memory and the returned `videos` is an empty tensor (the caller is
responsible for consuming/saving each chunk). Only used when `streaming` is True.
forcing_kv_enable: Optional Forcing-KV (arXiv 2605.09681) master switch.
None keeps the transformer's current configuration untouched.
forcing_kv_head_profile: Offline head profile - the official
{"layers": [{"layer_idx", "static_head", "dynamic_head"}]}
format, a plain mapping, or a path to a JSON file. Required
when enabling Forcing-KV.
forcing_kv_ar_start: AR step (chunk index) after which grouped
head compression activates (default 1).
forcing_kv_spatial_context_length: History frames for static heads
(default 1).
forcing_kv_temporal_context_length: Recent frames for dynamic
heads (default 1).
forcing_kv_dynamic_context_length: Compressed cache capacity in
frames (default 1).
forcing_kv_num_frame_patch: Token segments per latent frame
(default 6).
forcing_kv_sim_retention_ratio: Fraction of candidate segments
kept (default 0.33).
Examples:
```python
@@ -688,9 +793,37 @@ class WanSelfForcingPipeline(DiffusionPipeline):
current_start_frame = start_frame_index
cache_start_frame = 0
# Forcing-KV (arXiv 2605.09681): apply requested configuration before
# the KV cache is allocated so the rolling window can be budget-sized.
if forcing_kv_enable is not None:
fkv_kwargs = {}
if forcing_kv_ar_start is not None:
fkv_kwargs["ar_start"] = forcing_kv_ar_start
if forcing_kv_spatial_context_length is not None:
fkv_kwargs["spatial_context_length"] = forcing_kv_spatial_context_length
if forcing_kv_temporal_context_length is not None:
fkv_kwargs["temporal_context_length"] = forcing_kv_temporal_context_length
if forcing_kv_dynamic_context_length is not None:
fkv_kwargs["dynamic_context_length"] = forcing_kv_dynamic_context_length
if forcing_kv_num_frame_patch is not None:
fkv_kwargs["num_frame_patch"] = forcing_kv_num_frame_patch
if forcing_kv_sim_retention_ratio is not None:
fkv_kwargs["sim_retention_ratio"] = forcing_kv_sim_retention_ratio
self.set_forcing_kv_config(
forcing_kv_enable,
head_profile=forcing_kv_head_profile,
**fkv_kwargs)
# Anchor frame window in the attention layers follows the causal block size
self._unwrap_transformer().num_frame_per_block = num_frame_per_block
# 8. Initialize KV cache and cross-attention cache
# Reset caches if they exist (for multiple inference calls)
required_kv_size = num_latent_frames * frame_seq_length
if (self._forcing_kv_enabled()
and getattr(self.transformer.config, 'local_attn_size', -1) != -1):
fkv_cache_size = self._forcing_kv_cache_tokens(frame_seq_length)
if fkv_cache_size is not None:
required_kv_size = fkv_cache_size
if self.kv_cache_pos is not None and self.kv_cache_pos[0]["k"].shape[1] >= required_kv_size:
for block_index in range(len(self.kv_cache_pos)):
self.kv_cache_pos[block_index]["global_end_index"] = torch.tensor(
@@ -701,6 +834,10 @@ class WanSelfForcingPipeline(DiffusionPipeline):
[0], dtype=torch.long, device=device)
self.kv_cache_neg[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=device)
# Drop per-cache Forcing-KV state so the keep-mask restarts clean
for cache in (self.kv_cache_pos[block_index], self.kv_cache_neg[block_index]):
for key in ("forcing_kv", "_fkv_last_q", "_fkv_window_start", "_fkv_local_end"):
cache.pop(key, None)
for block_index in range(len(self.crossattn_cache_pos)):
self.crossattn_cache_pos[block_index]["is_init"] = False
self.crossattn_cache_neg[block_index]["is_init"] = False
@@ -863,6 +1000,7 @@ class WanSelfForcingPipeline(DiffusionPipeline):
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
forcing_kv_state={"clean_pass": True},
)
self.transformer(
x=denoised_pred,
@@ -873,6 +1011,7 @@ class WanSelfForcingPipeline(DiffusionPipeline):
crossattn_cache=self.crossattn_cache_neg,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
forcing_kv_state={"clean_pass": True},
)
else:
with torch.cuda.amp.autocast(dtype=weight_dtype):
@@ -885,6 +1024,7 @@ class WanSelfForcingPipeline(DiffusionPipeline):
crossattn_cache=self.crossattn_cache_pos,
current_start=current_start_frame * frame_seq_length,
cache_start=None,
forcing_kv_state={"clean_pass": True},
)
current_start_frame += current_num_frames
+3
View File
@@ -12,6 +12,9 @@ from .group_offload import (register_auto_device_hook,
safe_remove_group_offloading)
from .lora_utils import (convert_peft_lora_to_kohya_lora, create_network,
merge_lora, unmerge_lora)
from .perf_metrics import install as install_perf_metrics
from .perf_metrics import install_training as install_perf_training
from .perf_metrics import instrument_pipeline
from .sd3_sde_with_logprob import sde_step_with_logprob
from .trigflow_sampler import (RectifiedFlow_TrigFlowWrapper,
sample_trigflow_timesteps)
+425
View File
@@ -0,0 +1,425 @@
r"""Model-agnostic parts of Parallel Decoding Distillation (PDD).
PDD (arXiv 2607.26004) turns a pre-trained flow model into a *parallel decoder*: the sampling interval is discretized
into `N` intervals grouped into blocks of size `L`, and one network evaluation predicts the **mean velocity of every
interval of the next block** instead of the single instantaneous velocity. Generation then advances `L` intervals per
evaluation, i.e. `NFE = N / L`.
Everything here works on plain `nn.Linear`s and tensors, with no reference to any particular transformer: the math of
the time grid and the head plans, the [`PDDParallelHead`] / [`PDDLoRALinear`] modules, the teacher switch, and the checkpoint
resolution. Porting a model supplies only the glue that knows where the final linear layers live and how a forward is
called — `videox_fun/models/minimax_h3_pdd.py` is the MiniMax-H3 reference: an attach step that swaps the final heads
for [`PDDParallelHead`]s, a teacher mean-velocity estimate in the model's calling convention, a step callback in the
pipeline's callback protocol, and a `load_pdd_lora` that ties them together.
"""
import contextlib
import json
import os
from typing import Optional, Sequence
import torch
import torch.nn as nn
import torch.nn.functional as F
def shifted_sigma(shift: float, sigma: torch.Tensor) -> torch.Tensor:
r"""The exponential sigma shift of a rectified-flow schedule, `sigma' = s*sigma / (1 + (s-1)*sigma)`."""
return shift * sigma / (1 + (shift - 1) * sigma)
def pdd_time_grid(shift: float, num_steps: int) -> torch.Tensor:
r"""
The PDD time discretization `0 = t_0 < ... < t_N = 1` of a rectified-flow schedule with an exponential sigma shift.
The paper's shift reparameterization (eq. 16), `t_n = shift_s(n/N)` with `shift_s(t) = (t/s) / (1 + (1/s - 1) t)`,
is algebraically the same grid as `t = 1 - sigma'` over a uniform sigma grid — the convention of MiniMax-H3's
scheduler, where `t = 1` is clean. A consequence worth relying on when the model's scheduler is a plain Euler
rectified-flow one: the block boundaries of this grid, taken every `L` indices, are exactly the grid
`set_timesteps(N / L + 1)` builds, so PDD generation reuses such a released scheduler unchanged.
Args:
shift (`float`): The exponential shift of the schedule (`12.0` video / `3.0` audio for MiniMax-H3; `1.0` is
a uniform grid).
num_steps (`int`): The grid size `N`.
Returns:
`torch.Tensor` of shape `(num_steps + 1,)`, float64: the grid, ascending from `0` to `1`.
"""
sigma = torch.linspace(1.0, 0.0, num_steps + 1, dtype=torch.float64)
return 1.0 - shifted_sigma(shift, sigma)
def pdd_training_plan(step_sizes: torch.Tensor, start: int, targets: Sequence[int], advance: int) -> torch.Tensor:
r"""
Every direction one PDD training step needs, from a single backbone evaluation.
Args:
step_sizes (`torch.Tensor` of shape `(N,)`): The grid step sizes `h_l = t_{l+1} - t_l`.
start (`int`): The block start `n`, i.e. the index the state is currently at.
targets (`Sequence[int]`): The intra-block indices `k` the loss is evaluated at, each `n <= k < N`.
advance (`int`): How many intervals the carried state moves after the step, i.e. `L_min`.
Returns:
`torch.Tensor` of shape `(2 * len(targets) + 1, N)`: for every target, the displacement from `X_n` to `X_k`
followed by the row that selects `u_k`; then, last, the displacement from `X_n` to `X_{n+L_min}`.
"""
plan = torch.zeros(2 * len(targets) + 1, step_sizes.shape[0], dtype=step_sizes.dtype, device=step_sizes.device)
for position, target in enumerate(targets):
plan[2 * position, start:target] = step_sizes[start:target]
plan[2 * position + 1, target] = 1.0
plan[-1, start : start + advance] = step_sizes[start : start + advance]
return plan
def pdd_sampling_plan(step_sizes: torch.Tensor, start: int, block_size: int) -> torch.Tensor:
r"""
The single direction a PDD generation step needs: the *mean* velocity of the whole block.
Normalizing the fused displacement by the block span turns it into the block's average velocity, which is what an
ordinary Euler step over the block boundaries consumes — so a plain rectified-flow scheduler drives PDD
generation unchanged.
Args:
step_sizes (`torch.Tensor` of shape `(N,)`): The grid step sizes `h_l = t_{l+1} - t_l`.
start (`int`): The block start `n`.
block_size (`int`): The block size `L`.
Returns:
`torch.Tensor` of shape `(1, N)`: the plan.
"""
plan = torch.zeros(1, step_sizes.shape[0], dtype=step_sizes.dtype, device=step_sizes.device)
span = step_sizes[start : start + block_size].sum()
plan[0, start : start + block_size] = step_sizes[start : start + block_size] / span
return plan
class PDDParallelHead(nn.Module):
r"""
The `N` per-interval output heads of a PDD parallel decoder, in place of one final linear layer.
The heads are held as a single `(num_steps, out_features, in_features)` parameter, every slice initialized from the
pre-trained layer this replaces — so at initialization every interval predicts exactly the teacher's velocity and
the parallel decoder starts as the teacher. That pre-trained layer is also kept as a frozen buffer pair, which is
what [`pdd_teacher_mode`] switches to: the teacher's instantaneous velocity stays available from the same module after
the heads have moved.
`forward` does not evaluate the heads one by one: it fuses them into the `num_directions` linear maps of the
current `plan` and applies those, which is the paper's layer fusion (§3.1) and keeps the head's cost independent of
`num_steps`.
Args:
source (`nn.Linear`): The pre-trained final layer to repeat.
num_steps (`int`): The grid size `N`, i.e. how many heads to hold.
"""
def __init__(self, source: nn.Linear, num_steps: int):
super().__init__()
self.num_steps = num_steps
self.in_features = source.in_features
self.out_features = source.out_features
self.weight = nn.Parameter(source.weight.detach()[None].repeat(num_steps, 1, 1).clone())
self.bias = (
None if source.bias is None else nn.Parameter(source.bias.detach()[None].repeat(num_steps, 1).clone())
)
self.register_buffer("teacher_weight", source.weight.detach().clone(), persistent=False)
self.register_buffer(
"teacher_bias", None if source.bias is None else source.bias.detach().clone(), persistent=False
)
self.teacher = False
# `plan` is a plain attribute rather than a buffer: it is per-step control flow, not model state to
# serialize or shard. The default reproduces the source layer, so an unplanned head is the teacher's head.
self.plan = torch.zeros(1, num_steps)
self.plan[0, 0] = 1.0
def set_plan(self, plan: torch.Tensor) -> None:
r"""
Set the `(num_directions, num_steps)` coefficient matrix the next forward fuses the heads with.
Args:
plan (`torch.Tensor`): The plan. Row `p` weights the `N` heads into the `p`-th output direction.
"""
if plan.ndim != 2 or plan.shape[1] != self.num_steps:
raise ValueError(
f"A PDD plan must be a `(num_directions, {self.num_steps})` matrix, got {list(plan.shape)}."
)
self.plan = plan
@property
def num_directions(self) -> int:
return 1 if self.teacher else self.plan.shape[0]
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
r"""
Args:
hidden_states (`torch.Tensor` of shape `(..., in_features)`): The backbone's final hidden state.
Returns:
`torch.Tensor` of shape `(..., num_directions * out_features)`: the planned directions, stacked on the
channel axis in plan-row order. Under [`pdd_teacher_mode`] this is the single pre-trained direction.
"""
if self.teacher:
return F.linear(hidden_states, self.teacher_weight, self.teacher_bias)
plan = self.plan.to(device=self.weight.device, dtype=self.weight.dtype)
weight = torch.einsum("pn,noi->poi", plan, self.weight).flatten(0, 1)
bias = None if self.bias is None else torch.einsum("pn,no->po", plan, self.bias).flatten()
return F.linear(hidden_states, weight, bias)
class PDDLoRALinear(nn.Module):
r"""
A frozen `nn.Linear` with a trainable low-rank update, `y = W x + b + (alpha / rank) * B A x`.
The PDD counterpart of `lora_utils.py`'s `LoRAModule`, and deliberately not built on it: the adapter is a node in
the model tree (so FSDP shards it and the base layer stays visible for dtype pinning), not an out-of-tree
forward patch, and it must collapse to exactly the frozen layer when disabled.
The adapter parameters are held in float32 and cast to the activation dtype inside `forward`, so the optimizer sees
float32 master weights while the matmuls stay at the backbone's precision. `B` starts at zero, so the wrapped
module is exactly the frozen layer at initialization — and is again exactly the frozen layer whenever `enabled` is
false, which is how [`pdd_teacher_mode`] recovers the teacher without a second copy of the backbone.
Args:
base (`nn.Linear`): The layer to wrap. It is frozen here.
rank (`int`): The rank of the update.
alpha (`float`): The scaling numerator; `alpha == rank` means a unit-scaled update.
"""
def __init__(self, base: nn.Linear, rank: int, alpha: float):
super().__init__()
self.base = base
self.base.requires_grad_(False)
self.scaling = alpha / rank
self.enabled = True
self.lora_down = nn.Parameter(torch.empty(rank, base.in_features, dtype=torch.float32))
self.lora_up = nn.Parameter(torch.zeros(base.out_features, rank, dtype=torch.float32))
nn.init.kaiming_uniform_(self.lora_down, a=5**0.5)
# Models may read `linear.weight.dtype` off their projections to align activations with a mixed-precision
# checkpoint (MiniMax-H3 does), so the wrapper has to present the wrapped layer's own tensors under the usual
# names.
@property
def weight(self) -> torch.Tensor:
return self.base.weight
@property
def bias(self) -> Optional[torch.Tensor]:
return self.base.bias
@property
def in_features(self) -> int:
return self.base.in_features
@property
def out_features(self) -> int:
return self.base.out_features
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
out = self.base(hidden_states)
if not self.enabled:
return out
update = F.linear(
F.linear(hidden_states, self.lora_down.to(hidden_states.dtype)),
self.lora_up.to(hidden_states.dtype),
)
return out + self.scaling * update.to(out.dtype)
def add_pdd_lora(module: nn.Module, target_names: Sequence[str], rank: int, alpha: float) -> int:
r"""
Wrap every `nn.Linear` whose qualified name ends in one of `target_names` with a [`PDDLoRALinear`], in place.
Args:
module (`nn.Module`): The root to walk.
target_names (`Sequence[str]`): Qualified-name suffixes to match, e.g. `("to_q", "ff.net.2")`.
rank (`int`): The rank of every adapter.
alpha (`float`): The scaling numerator of every adapter.
Returns:
`int`: The number of layers wrapped.
"""
targets = [
(name, child)
for name, child in module.named_modules()
if isinstance(child, nn.Linear) and any(name.endswith(suffix) for suffix in target_names)
]
for name, child in targets:
parent_name, _, attribute = name.rpartition(".")
parent = module.get_submodule(parent_name) if parent_name else module
setattr(parent, attribute, PDDLoRALinear(child, rank, alpha))
return len(targets)
def merge_pdd_lora(module: nn.Module) -> int:
r"""
Fold every [`PDDLoRALinear`] into its frozen base layer and unwrap it, in place.
An inference-only optimization: afterwards each wrapped layer is a plain `nn.Linear` whose weight already carries
`scaling * B A`, so a forward costs one matmul instead of three. The [`PDDParallelHead`]s are left untouched — their
effective weight changes every step with the `plan`, so they cannot be folded into a static layer.
Two things to know before calling this. Merging overwrites the base weight, so it destroys the `enabled=False`
fallback [`pdd_teacher_mode`] relies on: only merge on a pure student inference path. And the update is accumulated
in float32 (the adapters' storage dtype) then cast back to the base weight's dtype, so the result is not
bit-for-bit the un-merged forward — it rounds the low-rank delta into the backbone precision.
Call this before any device offload / quantization / FSDP wrap is registered on the model, so no hook has to be
rebuilt and the delta lands on the unquantized weight.
Args:
module (`nn.Module`): The root to walk, e.g. the transformer returned by `load_pdd_lora`.
Returns:
`int`: The number of adapters merged.
"""
# Collect first, then mutate: replacing a child while walking `module.modules()` would perturb the traversal.
adapters = [
(parent, attribute, child)
for parent in module.modules()
for attribute, child in parent.named_children()
if isinstance(child, PDDLoRALinear)
]
for parent, attribute, adapter in adapters:
base = adapter.base
with torch.no_grad():
weight = base.weight
delta = adapter.scaling * (
adapter.lora_up.to(weight.device) @ adapter.lora_down.to(weight.device)
)
weight.data = (weight.data.to(delta.dtype) + delta).to(weight.dtype)
setattr(parent, attribute, base)
return len(adapters)
@contextlib.contextmanager
def pdd_teacher_mode(transformer):
r"""
Run `transformer` as the frozen pre-trained teacher.
The low-rank updates of the backbone are switched off and every [`PDDParallelHead`] falls back to the weights it was
built from, so the forward is bit-for-bit the released model's instantaneous velocity — with a single output
direction rather than the planned ones.
"""
heads = [module for module in transformer.modules() if isinstance(module, PDDParallelHead)]
adapters = [module for module in transformer.modules() if isinstance(module, PDDLoRALinear)]
for head in heads:
head.teacher = True
for adapter in adapters:
adapter.enabled = False
try:
yield transformer
finally:
for head in heads:
head.teacher = False
for adapter in adapters:
adapter.enabled = True
def _strip_fsdp_wrapper(name: str) -> str:
r"""Drop the `_fsdp_wrapped_module` path segments FSDP injects around each separately-wrapped child unit."""
return ".".join(part for part in name.split(".") if part != "_fsdp_wrapped_module")
def pdd_state_dict(transformer, state_dict: Optional[dict] = None) -> dict:
r"""
The trainable PDD state of a parallel decoder: the enlarged heads and every low-rank update.
The frozen backbone is not included, so a checkpoint is a few gigabytes rather than the full size of the base
model.
Args:
transformer (`nn.Module`): The parallel decoder; its module tree decides which keys are trainable.
state_dict (`dict`, optional): The weights to filter, defaulting to `transformer.state_dict()`. Under FSDP,
pass an already-gathered `FULL_STATE_DICT` here, since the live module views are sharded.
FSDP wraps every child unit in a `_fsdp_wrapped_module`, so a wrapped module's `named_modules` path
(`blocks.0._fsdp_wrapped_module.attn.to_q`) never matches the clean keys of a gathered `FULL_STATE_DICT`
(`blocks.0.attn.to_q.lora_down`) and only the root wrap unit would survive the filter. Both sides are normalized
through [`_strip_fsdp_wrapper`] so the block LoRA and the parallel heads are kept too.
"""
if state_dict is None:
state_dict = transformer.state_dict()
trainable = {
_strip_fsdp_wrapper(name)
for name, module in transformer.named_modules()
if isinstance(module, (PDDParallelHead, PDDLoRALinear))
}
return {
_strip_fsdp_wrapper(name): value.detach().cpu()
for name, value in state_dict.items()
if any(_strip_fsdp_wrapper(name).startswith(f"{prefix}.") for prefix in trainable) and ".base." not in name
}
PDD_WEIGHTS_NAME = "pdd.safetensors"
PDD_EMA_WEIGHTS_NAME = "pdd_ema.safetensors"
# Pre-rename checkpoints stored live weights here; resume still accepts it.
PDD_LEGACY_LIVE_WEIGHTS_NAME = "pdd_live.safetensors"
# The released MiniMax-H3 recipe. Every field is meant to be overridden by the `pdd_config.json` a training run
# writes next to its weights; a port to another model passes its own `defaults` to [`load_pdd_config`] instead.
PDD_DEFAULT_CONFIG = {
"pdd_num_steps": 32,
"pdd_block_size": 4,
"lora_rank": 64,
"lora_alpha": 64.0,
"lora_targets": "to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear",
}
def resolve_pdd_lora_path(path):
r"""
A checkpoint directory or a weights file.
A directory prefers `pdd_ema.safetensors` (the EMA inference export) and falls back to `pdd.safetensors`
(live weights, or the EMA file on checkpoints written before the rename).
"""
if path is None:
return None
path = os.path.abspath(os.path.expanduser(path))
if os.path.isdir(path):
ema = os.path.join(path, PDD_EMA_WEIGHTS_NAME)
live = os.path.join(path, PDD_WEIGHTS_NAME)
if os.path.isfile(ema):
path = ema
elif os.path.isfile(live):
path = live
else:
raise FileNotFoundError(
f"PDD checkpoint directory {path} has neither {PDD_EMA_WEIGHTS_NAME} nor {PDD_WEIGHTS_NAME}."
)
if not os.path.isfile(path):
raise FileNotFoundError(f"PDD checkpoint does not exist: {path}")
return path
def load_pdd_config(weights_path, defaults=None):
r"""Rank / alpha / targets / grid next to the weights file (`pdd_config.json`), over `defaults`."""
config = dict(PDD_DEFAULT_CONFIG if defaults is None else defaults)
config_path = os.path.join(os.path.dirname(weights_path), "pdd_config.json")
if os.path.isfile(config_path):
with open(config_path, encoding="utf-8") as handle:
saved = json.load(handle)
aliases = {"lora_rank": "rank", "lora_alpha": "network_alpha", "lora_targets": "target_name"}
for key in config:
if key in saved:
config[key] = saved[key]
elif aliases.get(key) in saved:
config[key] = saved[aliases[key]]
if not isinstance(config["lora_targets"], str):
config["lora_targets"] = ",".join(config["lora_targets"])
return config
def pdd_num_inference_steps(config, num_inference_steps, teacher_default=None):
r"""Keep `num_inference_steps` when it divides `N`; otherwise snap a leftover teacher default to `N / L`."""
grid = int(config["pdd_num_steps"])
steps = int(num_inference_steps)
if grid % steps == 0:
return steps
block = int(config["pdd_block_size"])
if teacher_default is not None and steps == int(teacher_default) and block > 0 and grid % block == 0:
nfe = grid // block
print(f"PDD checkpoint: using num_inference_steps {nfe} (grid {grid}, block {block})", flush=True)
return nfe
raise ValueError(f"num_inference_steps {steps} must divide PDD grid size {grid}.")
File diff suppressed because it is too large Load Diff