Merge branch 'main' into speed
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
---
|
||||
name: integrating-models
|
||||
description: Guides adding, porting, or onboarding a diffusion model (transformer/VAE/encoder, inference pipeline, training script, config) into the VideoX-Fun repository by mirroring the closest existing model family and maximizing reuse of the repository's existing code and shared infrastructure. Use when integrating a new model/architecture, or when creating predict_*.py inference scripts, scripts/*/train*.py training scripts, pipeline_*.py, config/*.yaml, or model definitions under videox_fun/models/.
|
||||
---
|
||||
|
||||
# Integrating Models into VideoX-Fun
|
||||
|
||||
## Core rule: maximize reuse of existing repo code — mirror, extend, never reinvent
|
||||
|
||||
**Prime directive: reuse this repository's existing code to the maximum.** Nearly every building block you need already exists in `videox_fun/` or in a sibling model family. Your job is to **find it, import it, and extend it** — not to write a parallel implementation. A new file should be mostly reused structure plus the genuinely model-specific delta; the less new code you write, the better.
|
||||
|
||||
**Reuse-first protocol — before writing ANY new function / class / util:**
|
||||
1. **Search the repo first.** Grep `videox_fun/` and the closest family for an existing equivalent (weight loader, scheduler, sampler, offload, attention, LoRA, fp8, dataset, dist helper, save/metric util). If one exists → **import and reuse it**. If it is 80% right → **extend / parameterize it**, do not fork it.
|
||||
2. **Only if nothing exists** may you add new code — and then put it in the shared layer (`videox_fun/utils`, `videox_fun/data`, `videox_fun/dist`) so the next model reuses it too, instead of burying it in a family folder.
|
||||
3. **Never copy-paste** a util into a new file (that creates drift); import the single source of truth.
|
||||
|
||||
**Mirror the closest family.** Every model follows the **same layered template**. Integrating a model means finding the closest existing family and mirroring its structure, changing only what genuinely differs:
|
||||
1. Pick the closest existing family by task type (t2v / i2v / v2v-control / s2v / t2i / edit / distill): `wan2.1`, `wan2.1_fun`, `wan2.2`, `qwenimage`, `flux2`, `minimax_h3`, `ltx2`, `longcatvideo`, `cogvideox_fun`, `z_image`, etc.
|
||||
2. Read that family end-to-end across all layers:
|
||||
- `examples/<family>/predict_*.py` (inference entry)
|
||||
- `scripts/<family>/train*.py` + `*.sh` + `README_TRAIN*.md` (training)
|
||||
- `videox_fun/pipeline/pipeline_<family>*.py` (pipeline)
|
||||
- `videox_fun/models/<family>_*.py` (model definitions)
|
||||
- `config/<family>/*.yaml` (config)
|
||||
3. Copy that structure and adapt. Keep names, argument sets, control flow, and reuse points identical in shape.
|
||||
|
||||
Writing a bespoke pipeline, weight loader, trainer, sampler, dataset, or offload scheme from scratch is a **failure mode**. If you are tempted to, **stop** and check the Reuse inventory below first.
|
||||
|
||||
## Repository layout (where each layer lives)
|
||||
|
||||
| Layer | Location | What it is |
|
||||
|-------|----------|------------|
|
||||
| Model definitions | `videox_fun/models/<family>_*.py` | Transformer / VAE / text-audio-image encoders. Diffusers `ModelMixin`+`ConfigMixin`, `@register_to_config`, custom `from_pretrained`. |
|
||||
| Model registry | `videox_fun/models/__init__.py` | Imports every model class. **Must be updated** for a new model. |
|
||||
| Inference pipelines | `videox_fun/pipeline/pipeline_<family>*.py` | `<Family>Pipeline(DiffusionPipeline)` with `__call__`. |
|
||||
| Pipeline registry | `videox_fun/pipeline/__init__.py` | Imports every pipeline + aliases. **Must be updated.** |
|
||||
| Configs (optional) | `config/<family>/*.yaml` | OmegaConf YAML for civitai/custom layouts; a standard diffusers-layout checkpoint can load without one. |
|
||||
| Inference entry scripts | `examples/<family>/predict_*.py` | User-facing, config-block-at-top runnable scripts. |
|
||||
| Inference services | `examples/<family>/{app.py,launch_api.py,post_infer*.py}` | Gradio UI / API server / batch inference. |
|
||||
| Training scripts | `scripts/<family>/train*.py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py`, ... |
|
||||
| Training launchers | `scripts/<family>/train*.sh` | `accelerate launch` / DeepSpeed command with full arg list. |
|
||||
| Training docs | `scripts/<family>/README_TRAIN*.md` | Bilingual pairs: `README_TRAIN.md` + `README_TRAIN_zh-CN.md`. |
|
||||
| Shared: schedulers/utils | `videox_fun/utils/` | `fm_solvers`, `fm_solvers_unipc`, `lora_utils`, `fp8_optimization`, `group_offload`, `utils.py`. |
|
||||
| Shared: distributed | `videox_fun/dist/` | `fsdp.shard_model`, `fuser.set_multi_gpus_devices`, `<family>_xfuser` sequence-parallel attention. |
|
||||
| Shared: data | `videox_fun/data/` | Datasets (`ImageVideoDataset`, `VideoDataset`, ...) + bucket/aspect-ratio samplers. |
|
||||
| Demo / test datasets | `datasets/X-Fun-*-Demo/` | Ready-made smoke-test data, downloaded via `modelscope download --dataset PAI/<name>`; each ships several `metadata*.json` variants. **The only test data to use** (see reference.md §8). |
|
||||
| Preprocessing (data gen) | `scripts/<family>/generate_*.py` / `train_preprocess.py` (+ `.sh`) | Offline multi-GPU generation of cached training data (latents / ODE pairs / embeddings) → per-sample `.safetensors` + `outputs.json`, loaded by `ImageVideoSafetensorsDataset`. |
|
||||
| ComfyUI nodes | `comfyui/<family>/nodes.py` | Optional node integration mirroring the pipeline. |
|
||||
|
||||
## Integration workflow
|
||||
|
||||
Copy this checklist and track progress:
|
||||
|
||||
```
|
||||
Integration Progress:
|
||||
- [ ] Step 0: Choose the closest family to mirror; read it across all layers
|
||||
- [ ] Step 1: Model definitions in videox_fun/models/ + register in models/__init__.py
|
||||
- [ ] Step 2: Pipeline in videox_fun/pipeline/ + register in pipeline/__init__.py
|
||||
- [ ] Step 3: Config YAML in config/<family>/
|
||||
- [ ] Step 4: Inference script(s) in examples/<family>/predict_*.py
|
||||
- [ ] Step 5: Training script(s) in scripts/<family>/train*.py + .sh
|
||||
- [ ] Step 6: Training docs README_TRAIN.md + README_TRAIN_zh-CN.md
|
||||
- [ ] Step 7: Reuse audit + verification (incl. smoke test on the matching demo dataset)
|
||||
```
|
||||
|
||||
**Step 0 — Choose the mirror.** Match by task and architecture. A new control model mirrors an existing `*_fun`/`*_control` family; a new audio/talking model mirrors `minimax_h3`/`longcatvideo`/`infinitetalk`; a new image model mirrors `qwenimage`/`flux2`/`z_image`.
|
||||
|
||||
**Step 1 — Model.** Create `videox_fun/models/<family>_transformer3d.py` (or `2d`), `<family>_vae.py`, encoders as needed. Mirror the class shape: `class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin)`, `_supports_gradient_checkpointing = True`, `@register_to_config __init__`, and a `from_pretrained` that supports `transformer_additional_kwargs`, `dict_mapping`, `low_cpu_mem_usage`, and missing-key init. Add imports to `videox_fun/models/__init__.py`.
|
||||
|
||||
**Step 2 — Pipeline.** Create `videox_fun/pipeline/pipeline_<family>.py`. Mirror `pipeline_wan.py`: module-level `retrieve_timesteps`, a `<Family>PipelineOutput(BaseOutput)` dataclass, `<Family>Pipeline(DiffusionPipeline)` with `model_cpu_offload_seq`, `_callback_tensor_inputs`, `__init__(vae, tokenizer, text_encoder, transformer, scheduler, ...)`, `encode_prompt`, and `__call__`. Add imports/aliases to `videox_fun/pipeline/__init__.py`.
|
||||
|
||||
**Step 3 — Config (optional).** A YAML under `config/<family>/` is **not always required**. It is needed mainly for **civitai-format / custom single-file layouts** — to supply `transformer_additional_kwargs`, `dict_mapping` (civitai key → `__init__` kwarg), component subpaths, and `vae_kwargs`/`text_encoder_kwargs`/`scheduler_kwargs`/`image_encoder_kwargs`. For a **standard diffusers-layout** checkpoint (`model_index.json` + per-subfolder `config.json`), load directly via `from_pretrained(model_name, subfolder=...)` with no YAML — mirror `examples/minimax_h3_fun/predict_v2v_control.py`, which guards `if config_path is not None:`. When you do add a YAML, load it via `OmegaConf.load(config_path)` and spread into `from_pretrained` instead of hardcoding those values.
|
||||
|
||||
**Step 4 — Inference script.** Create `examples/<family>/predict_<task>.py` following the exact template (config block at top → component loading → scheduler dict → pipeline construction → multi-GPU/FSDP/compile → `GPU_memory_mode` branching → TeaCache → LoRA merge → inference → `save_results`). See [examples.md](examples.md).
|
||||
|
||||
**Step 5 — Training script.** Create `scripts/<family>/train.py` (+ `train_lora.py` etc.). Mirror the shared structure: license header, `sys.path` bootstrap, imports from `videox_fun`, `log_validation()` that **reuses the inference Pipeline**, `parse_args()` (reuse the existing shared argument set), `main()`. Add a `train.sh` launcher. Reuse `videox_fun.data` datasets/samplers — do not write a new dataset.
|
||||
|
||||
**Step 6 — Docs.** Write `README_TRAIN.md` and `README_TRAIN_zh-CN.md` as an aligned bilingual pair (same structure, same commands/params, matching section order).
|
||||
|
||||
**Step 7 — Reuse audit + verification.** Confirm you reused shared infra (below), smoke-test the new train/predict path on the **matching official demo dataset** under `datasets/X-Fun-*-Demo/` (pick by task and metadata variant — see reference.md §8), then run the verification checklist. Never invent an ad-hoc test set and never leave `datasets/internal_datasets/` placeholders in shipped scripts/docs.
|
||||
|
||||
## Reuse inventory (use these, do not reimplement)
|
||||
|
||||
**Reuse-first catalog: import from here instead of reimplementing. If a helper you need is not listed, grep `videox_fun/` and the closest family before writing your own.**
|
||||
|
||||
- **Schedulers**: `FlowMatchEulerDiscreteScheduler`, `videox_fun.utils.fm_solvers.FlowDPMSolverMultistepScheduler`, `fm_solvers_unipc.FlowUniPCMultistepScheduler`. Selected via a `sampler_name` dict.
|
||||
- **LoRA**: `videox_fun.utils.lora_utils` — `merge_lora`, `unmerge_lora`, `create_network`, `convert_peft_lora_to_kohya_lora`.
|
||||
- **FP8 / quantization**: `videox_fun.utils.fp8_optimization` — `convert_model_weight_to_float8`, `convert_weight_dtype_wrapper`, `replace_parameters_by_name`.
|
||||
- **Offloading**: `videox_fun.utils.group_offload` — `register_auto_device_hook`, `safe_enable_group_offload`; plus pipeline `enable_sequential_cpu_offload` / `enable_model_cpu_offload` / `.to(device)`.
|
||||
- **Distributed**: `videox_fun.dist` — `set_multi_gpus_devices`, `shard_model` (FSDP), `<family>_xfuser` sequence-parallel attention processors, `enable_multi_gpus_inference()`.
|
||||
- **IO / helpers**: `videox_fun.utils.utils` — `save_videos_grid`, `save_videos_with_audio_grid`, `get_image_to_video_latent`, `get_video_to_video_latent`, `get_image_latent`, `filter_kwargs`, `calculate_dimensions`.
|
||||
- **Data**: `videox_fun.data` — `ImageVideoDataset`, `VideoDataset`, `ImageVideoControlDataset`, `VideoSpeechDataset`, bucket/aspect-ratio samplers, `get_closest_ratio`, `get_random_mask`.
|
||||
- **Caching / speedups**: TeaCache (`models/cache_utils`, `get_teacache_coefficients`, `transformer.enable_teacache`), `enable_cfg_skip`, Riflex (`enable_riflex`), `torch.compile` on `transformer.blocks`.
|
||||
- **Preprocessing (data gen, multi-GPU)**: mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py` — `accelerate launch` + `Accelerator` (interleaved rank sharding), config-driven `from_pretrained` for the teacher/VAE/text-encoder, `safetensors.torch.save_file` per sample + `outputs.json` index, consumed by `videox_fun.data.ImageVideoSafetensorsDataset`. Store as **safetensors only — never LMDB or `.pt`** (see reference.md §10).
|
||||
|
||||
## Non-negotiable conventions
|
||||
|
||||
- **Maximize reuse of existing repo code**: import existing `videox_fun/` helpers and mirror the closest family; never fork or copy-paste a util, and never write a parallel pipeline / loader / scheduler / sampler / offload. Genuinely-new shared code goes in `videox_fun/{utils,data,dist}` (so the next model reuses it), not buried in a family folder.
|
||||
- **`sys.path` bootstrap**: every runnable script starts with the 3-level `project_roots` loop inserting into `sys.path` before importing `videox_fun`.
|
||||
- **Config-driven loading (YAML optional)**: a `config/<family>/*.yaml` is required for civitai-format/custom layouts (it supplies `transformer_additional_kwargs`/`dict_mapping`/subpaths); it is **optional for standard diffusers-layout checkpoints**, which load directly via `from_pretrained(model_name, subfolder=...)`. When a YAML is used, don't hardcode the values it provides.
|
||||
- **`GPU_memory_mode`**: support the standard six modes — `model_full_load`, `model_full_load_and_qfloat8`, `model_cpu_offload`, `model_cpu_offload_and_qfloat8`, `model_group_offload`, `sequential_cpu_offload` — with the exact branching order used in existing `predict_*.py`.
|
||||
- **Naming**: files `<family>_transformer3d.py` / `<family>_vae.py` / `pipeline_<family>.py`; classes `<Family>Transformer3DModel` / `AutoencoderKL<Family>` / `<Family>Pipeline`.
|
||||
- **Resolution args**: drive canvas size with a single square `--video_sample_size` (`type=int`, height = width); never `--video_sample_height` / `--video_sample_width`. For a fixed non-square shape add `--fix_sample_size` (`nargs=2, type=int`, `[height, width]`) that overrides the square size, and derive the effective height/width once in `parse_args()` (see reference.md §5).
|
||||
- **Registries**: a model is not integrated until it is imported in BOTH `videox_fun/models/__init__.py` and `videox_fun/pipeline/__init__.py`.
|
||||
- **Two weight formats**: support `civitai` and `diffusers` via config `format` + `dict_mapping` (maps civitai keys such as `in_dim`→`in_channels`, `dim`→`hidden_size`).
|
||||
- **Bilingual docs**: training READMEs ship as EN + `_zh-CN` pairs with aligned structure and identical commands/params.
|
||||
- **Test data = official demo datasets**: smoke tests, `log_validation` checks, launcher `.sh` defaults, and doc examples all point at `datasets/X-Fun-*-Demo/` (ModelScope `PAI/<name>`), with the metadata variant matching the task — `metadata_add_width_height.json` by default, `_add_objects.json` for VACE/subject-reference, `_add_wav.json` for audio-visual joint models, `metadata_lingbot_video_add_width_height.json` for `lingbot_video`. Selection matrix: reference.md §8.
|
||||
- **Preprocessing = offline data generation, multi-GPU + safetensors**: cached training data (latents / ODE pairs / embeddings) is produced by `accelerate launch` scripts like `generate_ode_pairs.py` (interleaved rank sharding, resume by skipping existing files, `wait_for_everyone`, rank-0 JSON index) and saved with `safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`. **Never single-GPU / `cuda:0`; never LMDB or `.pt`/`torch.save` pickles for preprocessed data.** See reference.md §10.
|
||||
|
||||
## Verification checklist
|
||||
|
||||
- [ ] New model classes imported in `videox_fun/models/__init__.py`
|
||||
- [ ] New pipeline(s) imported in `videox_fun/pipeline/__init__.py`
|
||||
- [ ] Config YAML present **only if** the checkpoint is civitai-format/custom-layout; a diffusers-layout model may load directly via `from_pretrained(model_name, subfolder=...)` with no YAML. When a YAML is used, it drives component loading (no hardcoded kwargs)
|
||||
- [ ] `predict_*.py` mirrors an existing script: `sys.path` bootstrap, config block, scheduler dict, `GPU_memory_mode` branching, LoRA merge, `save_results`
|
||||
- [ ] `train*.py` reuses `videox_fun.data` + shared args, and `log_validation()` reuses the inference Pipeline
|
||||
- [ ] `train*.sh` launcher provided (`accelerate launch` / DeepSpeed)
|
||||
- [ ] Shared infra reused (schedulers / lora_utils / fp8 / group_offload / dist / utils / data) — nothing reimplemented
|
||||
- [ ] Any offline data-generation/preprocessing script runs multi-GPU (`accelerate launch` + `Accelerator`) and saves cached tensors as **safetensors + `outputs.json`** for `ImageVideoSafetensorsDataset` — never LMDB or `.pt`
|
||||
- [ ] `README_TRAIN.md` + `README_TRAIN_zh-CN.md` aligned pair present
|
||||
- [ ] Smoke test / doc examples use the matching `datasets/X-Fun-*-Demo` dataset and the correct `metadata*.json` variant — no `internal_datasets` placeholders (reference.md §8)
|
||||
- [ ] Optional: ComfyUI node in `comfyui/<family>/nodes.py` mirrors the pipeline
|
||||
|
||||
## Additional resources
|
||||
|
||||
- Detailed file-by-file conventions, class/method shapes, and the model-loading internals: [reference.md](reference.md)
|
||||
- **Dataset & sampler selection matrix** (which `videox_fun.data` dataset/loader each training task uses), **demo-dataset / metadata-variant selection matrix** (which `datasets/X-Fun-*-Demo` to smoke-test with), **inference task matrix** (which pipeline each `predict_<task>.py` uses), and **multi-GPU preprocessing patterns**: [reference.md](reference.md) §8–§10
|
||||
- Concrete skeletons (config YAML, `predict_*.py`, pipeline class, training script + DataLoader): [examples.md](examples.md)
|
||||
@@ -0,0 +1,510 @@
|
||||
# VideoX-Fun Integration Skeletons
|
||||
|
||||
Starting templates. **Always open the mirrored family's real file and adapt it** — these skeletons show shape and required reuse points, not full implementations. Replace `<family>` / `<Family>` / `<task>`.
|
||||
|
||||
## Config — `config/<family>/<variant>.yaml` (optional)
|
||||
|
||||
> **Not always required.** Author a YAML only for civitai-format / custom single-file layouts. A standard diffusers-layout checkpoint (`model_index.json` + per-subfolder `config.json`) loads directly via `from_pretrained(model_name, subfolder=...)` with no YAML — set `config_path = None` and guard `if config_path is not None:` (see `examples/minimax_h3_fun/predict_v2v_control.py`).
|
||||
|
||||
```yaml
|
||||
format: civitai
|
||||
pipeline: <Family>
|
||||
transformer_additional_kwargs:
|
||||
transformer_subpath: ./
|
||||
dict_mapping:
|
||||
in_dim: in_channels
|
||||
dim: hidden_size
|
||||
|
||||
vae_kwargs:
|
||||
vae_subpath: <Family>_VAE.pth
|
||||
temporal_compression_ratio: 4
|
||||
spatial_compression_ratio: 8
|
||||
|
||||
text_encoder_kwargs:
|
||||
text_encoder_subpath: <text_encoder>.pth
|
||||
tokenizer_subpath: <tokenizer_id>
|
||||
text_length: 512
|
||||
|
||||
scheduler_kwargs:
|
||||
scheduler_subpath: null
|
||||
num_train_timesteps: 1000
|
||||
shift: 5.0
|
||||
|
||||
# Only for i2v / models with a CLIP image encoder:
|
||||
image_encoder_kwargs:
|
||||
image_encoder_subpath: <image_encoder>.pth
|
||||
```
|
||||
|
||||
## Inference — `examples/<family>/predict_<task>.py`
|
||||
|
||||
```python
|
||||
import os
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
# --- sys.path bootstrap (required, before importing videox_fun) ---
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
from videox_fun.models import (AutoencoderKL<Family>, <Family>TextEncoder,
|
||||
<Family>Transformer3DModel)
|
||||
from videox_fun.models.cache_utils import get_teacache_coefficients
|
||||
from videox_fun.pipeline import <Family>Pipeline
|
||||
from videox_fun.utils import register_auto_device_hook, safe_enable_group_offload
|
||||
from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler
|
||||
from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8,
|
||||
convert_weight_dtype_wrapper,
|
||||
replace_parameters_by_name)
|
||||
from videox_fun.utils.lora_utils import merge_lora, unmerge_lora
|
||||
from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent,
|
||||
save_videos_grid)
|
||||
|
||||
# --- user config block (keep the conventional order + comments) ---
|
||||
GPU_memory_mode = "sequential_cpu_offload"
|
||||
ulysses_degree = 1
|
||||
ring_degree = 1
|
||||
fsdp_dit = False
|
||||
fsdp_text_encoder = True
|
||||
compile_dit = False
|
||||
enable_teacache = True
|
||||
teacache_threshold = 0.10
|
||||
num_skip_start_steps = 5
|
||||
teacache_offload = False
|
||||
cfg_skip_ratio = 0
|
||||
enable_riflex = False
|
||||
riflex_k = 6
|
||||
config_path = "config/<family>/<variant>.yaml"
|
||||
model_name = "models/Diffusion_Transformer/<Family>-Model"
|
||||
sampler_name = "Flow"
|
||||
shift = 3
|
||||
transformer_path = None
|
||||
vae_path = None
|
||||
lora_path = None
|
||||
sample_size = [480, 832]
|
||||
video_length = 81
|
||||
fps = 16
|
||||
weight_dtype = torch.bfloat16
|
||||
prompt = "..."
|
||||
negative_prompt = "..."
|
||||
guidance_scale = 6.0
|
||||
seed = 43
|
||||
num_inference_steps = 50
|
||||
lora_weight = 0.55
|
||||
save_path = "samples/<family>-<task>"
|
||||
|
||||
# --- device + config (config_path may be None for a diffusers-layout checkpoint) ---
|
||||
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
|
||||
config = OmegaConf.load(config_path) # or guard: if config_path is not None: ... (then load components via subfolder=...)
|
||||
|
||||
# --- components (when a YAML is used, paths/kwargs come from config; otherwise pass subfolder=... directly) ---
|
||||
transformer = <Family>Transformer3DModel.from_pretrained(
|
||||
os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
|
||||
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
|
||||
low_cpu_mem_usage=True, torch_dtype=weight_dtype,
|
||||
)
|
||||
# optional transformer_path / vae_path override -> load_state_dict(strict=False) + print missing/unexpected
|
||||
vae = AutoencoderKL<Family>.from_pretrained(
|
||||
os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')),
|
||||
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
|
||||
).to(weight_dtype)
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')))
|
||||
text_encoder = <Family>TextEncoder.from_pretrained(
|
||||
os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')),
|
||||
additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']),
|
||||
low_cpu_mem_usage=True, torch_dtype=weight_dtype).eval()
|
||||
|
||||
# --- scheduler selection dict ---
|
||||
Chosen_Scheduler = {
|
||||
"Flow": FlowMatchEulerDiscreteScheduler,
|
||||
"Flow_Unipc": FlowUniPCMultistepScheduler,
|
||||
"Flow_DPM++": FlowDPMSolverMultistepScheduler,
|
||||
}[sampler_name]
|
||||
scheduler = Chosen_Scheduler(**filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])))
|
||||
|
||||
# --- pipeline ---
|
||||
pipeline = <Family>Pipeline(vae=vae, tokenizer=tokenizer, text_encoder=text_encoder,
|
||||
transformer=transformer, scheduler=scheduler)
|
||||
|
||||
# --- multi-gpu / fsdp / compile ---
|
||||
if ulysses_degree > 1 or ring_degree > 1:
|
||||
from functools import partial
|
||||
transformer.enable_multi_gpus_inference()
|
||||
if fsdp_dit:
|
||||
pipeline.transformer = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.transformer)
|
||||
if fsdp_text_encoder:
|
||||
pipeline.text_encoder = partial(shard_model, device_id=device, param_dtype=weight_dtype)(pipeline.text_encoder)
|
||||
if compile_dit:
|
||||
for i in range(len(pipeline.transformer.blocks)):
|
||||
pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i])
|
||||
|
||||
# --- GPU_memory_mode branching (keep this exact order) ---
|
||||
if GPU_memory_mode == "sequential_cpu_offload":
|
||||
replace_parameters_by_name(transformer, ["modulation",], device=device)
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_group_offload":
|
||||
register_auto_device_hook(pipeline.transformer)
|
||||
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
pipeline.to(device=device)
|
||||
|
||||
# --- teacache / cfg_skip / riflex / lora ---
|
||||
coefficients = get_teacache_coefficients(model_name) if enable_teacache else None
|
||||
if coefficients is not None:
|
||||
pipeline.transformer.enable_teacache(coefficients, num_inference_steps, teacache_threshold,
|
||||
num_skip_start_steps=num_skip_start_steps, offload=teacache_offload)
|
||||
if cfg_skip_ratio is not None:
|
||||
pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps)
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
if lora_path is not None:
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
# --- inference ---
|
||||
with torch.no_grad():
|
||||
video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1
|
||||
if enable_riflex:
|
||||
pipeline.transformer.enable_riflex(k=riflex_k, L_test=(video_length - 1) // vae.config.temporal_compression_ratio + 1)
|
||||
sample = pipeline(prompt, num_frames=video_length, negative_prompt=negative_prompt,
|
||||
height=sample_size[0], width=sample_size[1], generator=generator,
|
||||
guidance_scale=guidance_scale, num_inference_steps=num_inference_steps,
|
||||
shift=shift).videos
|
||||
if lora_path is not None:
|
||||
pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype)
|
||||
|
||||
# --- save (rank 0 only when multi-gpu) ---
|
||||
def save_results():
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
prefix = str(len(os.listdir(save_path)) + 1).zfill(8)
|
||||
if video_length == 1:
|
||||
image = (sample[0, :, 0].transpose(0, 1).transpose(1, 2) * 255).numpy().astype(np.uint8)
|
||||
Image.fromarray(image).save(os.path.join(save_path, prefix + ".png"))
|
||||
else:
|
||||
save_videos_grid(sample, os.path.join(save_path, prefix + ".mp4"), fps=fps)
|
||||
|
||||
if ulysses_degree * ring_degree > 1:
|
||||
import torch.distributed as dist
|
||||
if dist.get_rank() == 0:
|
||||
save_results()
|
||||
else:
|
||||
save_results()
|
||||
```
|
||||
|
||||
For i2v, gate the CLIP image encoder and pass `video`/`mask_video`:
|
||||
```python
|
||||
if transformer.config.in_channels != vae.config.latent_channels:
|
||||
clip_image_encoder = CLIPModel.from_pretrained(
|
||||
os.path.join(model_name, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder'))).to(weight_dtype).eval()
|
||||
input_video, input_video_mask, _ = get_image_to_video_latent(start_image, None, video_length=video_length, sample_size=sample_size)
|
||||
# pipeline = <Family>InpaintPipeline(..., clip_image_encoder=clip_image_encoder)
|
||||
# sample = pipeline(..., video=input_video, mask_video=input_video_mask).videos
|
||||
```
|
||||
|
||||
## Pipeline class — `videox_fun/pipeline/pipeline_<family>.py`
|
||||
|
||||
```python
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Union
|
||||
import torch
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.utils import BaseOutput, logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from ..models import AutoencoderKL<Family>, <Family>Transformer3DModel
|
||||
from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler, get_sampling_sigmas
|
||||
from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
EXAMPLE_DOC_STRING = """Examples:\n```python\npass\n```"""
|
||||
|
||||
# reuse retrieve_timesteps verbatim from pipeline_wan.py
|
||||
|
||||
@dataclass
|
||||
class <Family>PipelineOutput(BaseOutput):
|
||||
videos: torch.Tensor
|
||||
|
||||
class <Family>Pipeline(DiffusionPipeline):
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler):
|
||||
super().__init__()
|
||||
self.register_modules(tokenizer=tokenizer, text_encoder=text_encoder, vae=vae,
|
||||
transformer=transformer, scheduler=scheduler)
|
||||
# video_processor / vae_scale_factor / etc. as in pipeline_wan.py
|
||||
|
||||
def encode_prompt(self, prompt, negative_prompt, device, num_videos_per_prompt=1, ...):
|
||||
... # mirror pipeline_wan.py
|
||||
|
||||
def prepare_latents(self, batch_size, num_channels_latents, height, width, num_frames, dtype, device, generator, latents=None):
|
||||
...
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(self, prompt, negative_prompt=None, height=480, width=832, num_frames=81,
|
||||
num_inference_steps=50, guidance_scale=6.0, generator=None, shift=1.0,
|
||||
callback_on_step_end=None, return_dict=True, **kwargs) -> Union[<Family>PipelineOutput, tuple]:
|
||||
# 1. encode_prompt 2. prepare_latents 3. retrieve_timesteps
|
||||
# 4. denoising loop with guidance 5. vae.decode 6. return <Family>PipelineOutput(videos=...)
|
||||
...
|
||||
```
|
||||
Then register in `videox_fun/pipeline/__init__.py`:
|
||||
```python
|
||||
from .pipeline_<family> import <Family>Pipeline
|
||||
```
|
||||
|
||||
## Model class — `videox_fun/models/<family>_transformer3d.py`
|
||||
|
||||
```python
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from .attention_utils import attention # unified FA/SDPA backend — do not hand-roll SDPA
|
||||
|
||||
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(self, model_type='t2v', in_dim=16, dim=2048, ffn_dim=8192,
|
||||
num_heads=16, num_layers=32, in_channels=16, hidden_size=2048, ...):
|
||||
super().__init__()
|
||||
...
|
||||
|
||||
def _set_gradient_checkpointing(self, *args, **kwargs):
|
||||
self.gradient_checkpointing = True
|
||||
|
||||
def enable_multi_gpus_inference(self): ... # route attn through dist/<family>_xfuser.py
|
||||
def enable_teacache(self, ...): ...
|
||||
def enable_cfg_skip(self, ...): ...
|
||||
|
||||
def forward(self, x, timestep, context, ...): ...
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_path, subfolder=None,
|
||||
transformer_additional_kwargs=None, low_cpu_mem_usage=False,
|
||||
torch_dtype=torch.bfloat16):
|
||||
... # mirror wan_transformer3d.py: config.json -> dict_mapping -> init_empty_weights
|
||||
# -> load .bin/.safetensors -> shape-filter -> initialize missing keys -> load
|
||||
```
|
||||
Then register in `videox_fun/models/__init__.py`:
|
||||
```python
|
||||
from .<family>_transformer3d import <Family>Transformer3DModel
|
||||
from .<family>_vae import AutoencoderKL<Family>
|
||||
```
|
||||
|
||||
## Training — `scripts/<family>/train.py` (key reuse points)
|
||||
|
||||
```python
|
||||
"""Modified from https://github.com/huggingface/diffusers/.../train_text_to_image.py"""
|
||||
import argparse, gc, logging, math, os, sys
|
||||
import accelerate, diffusers, torch, transformers
|
||||
from accelerate import Accelerator
|
||||
from diffusers.optimization import get_scheduler
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
# same sys.path bootstrap as predict scripts
|
||||
from videox_fun.data import (ASPECT_RATIO_512, AspectRatioBatchImageVideoSampler,
|
||||
ImageVideoDataset, ImageVideoSampler, RandomSampler,
|
||||
get_closest_ratio, get_random_mask)
|
||||
from videox_fun.models import AutoencoderKL<Family>, <Family>Transformer3DModel
|
||||
from videox_fun.pipeline import <Family>Pipeline # REUSED for validation
|
||||
from videox_fun.utils.lora_utils import create_network # for train_lora
|
||||
from videox_fun.utils.utils import save_videos_grid, get_image_to_video_latent
|
||||
|
||||
def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config,
|
||||
accelerator, weight_dtype, global_step):
|
||||
# build <Family>Pipeline from accelerator.unwrap_model(transformer3d),
|
||||
# run validation_prompts, save_videos_grid to output_dir/sample/. Reuse the pipeline.
|
||||
...
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(...)
|
||||
# reuse the shared arg surface: --config_path, --pretrained_model_name_or_path,
|
||||
# --train_data_dir, --train_data_meta, --video_sample_n_frames, --train_batch_size,
|
||||
# --gradient_accumulation_steps, --learning_rate, --lr_scheduler, --checkpointing_steps,
|
||||
# --output_dir, --mixed_precision, --gradient_checkpointing, --enable_bucket,
|
||||
# --train_mode, --trainable_modules, --validation_prompts ... (add only what's needed)
|
||||
return parser.parse_args()
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
accelerator = Accelerator(mixed_precision=args.mixed_precision, ...)
|
||||
config = OmegaConf.load(args.config_path)
|
||||
# load transformer/vae/text_encoder via config
|
||||
|
||||
# --- Dataset: pick by task (see reference.md §8) ---
|
||||
# T2V/I2V base + inpaint -> ImageVideoDataset(enable_inpaint = args.train_mode != "normal")
|
||||
# Control -> ImageVideoControlDataset(enable_camera_info = ...)
|
||||
# Image edit -> ImageEditDataset
|
||||
# Speech/audio (S2V) -> VideoSpeechDataset / VideoSpeechControlDataset
|
||||
# Animate -> VideoAnimateDataset
|
||||
# Distill text / GRPO / DPO -> TextDataset
|
||||
# Smoke-test on the matching official demo dataset (reference.md §8), e.g.
|
||||
# datasets/X-Fun-Videos-Demo + metadata_add_width_height.json for T2V/I2V.
|
||||
train_dataset = ImageVideoDataset(
|
||||
args.train_data_meta, args.train_data_dir,
|
||||
video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride,
|
||||
video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat,
|
||||
image_sample_size=args.image_sample_size, enable_bucket=args.enable_bucket,
|
||||
enable_inpaint=True if args.train_mode != "normal" else False)
|
||||
|
||||
# --- Sampler + DataLoader: branch on enable_bucket (see reference.md §8) ---
|
||||
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
|
||||
if args.enable_bucket:
|
||||
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
|
||||
batch_sampler = AspectRatioBatchImageVideoSampler(
|
||||
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
|
||||
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
|
||||
aspect_ratios=aspect_ratio_sample_size)
|
||||
def collate_fn(examples):
|
||||
new_examples = {"pixel_values": [], "text": []}
|
||||
if args.train_mode != "normal":
|
||||
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
|
||||
# get_closest_ratio -> Resize/CenterCrop/Normalize -> stack; masks via get_random_mask
|
||||
return new_examples
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
|
||||
else:
|
||||
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset, batch_sampler=batch_sampler, num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
|
||||
|
||||
# trainable-module filtering or create_network for LoRA
|
||||
# optimizer + get_scheduler; accelerator.prepare; checkpoint hooks
|
||||
# training loop: timestep sampling -> transformer forward -> loss -> backward
|
||||
# periodic log_validation(...); final save weights / LoRA
|
||||
...
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
```
|
||||
|
||||
## Launcher — `scripts/<family>/train.sh`
|
||||
|
||||
```bash
|
||||
export MODEL_NAME="models/Diffusion_Transformer/<Family>-Model"
|
||||
# Test data = the official demo dataset matching the task (reference.md §8). Download once, e.g.:
|
||||
# modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
|
||||
# T2I -> X-Fun-Images-Demo | control -> X-Fun-{Videos,Images}-Controls-Demo
|
||||
# S2V -> X-Fun-Videos-Audios-Demo | image edit -> X-Fun-Images-Edit-Demo
|
||||
export DATASET_NAME="datasets/X-Fun-Videos-Demo/" # = train_data_dir (data_root); media live under train/
|
||||
export DATASET_META_NAME="datasets/X-Fun-Videos-Demo/metadata_add_width_height.json" # = train_data_meta: [{"file_path","text","type","width","height"}] — see reference.md §8
|
||||
# Metadata variants: VACE/subject-ref -> metadata_add_width_height_add_objects.json (X-Fun-Videos-Controls-Demo);
|
||||
# audio-visual joint -> metadata_add_width_height_add_wav.json; lingbot_video -> metadata_lingbot_video_add_width_height.json
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/<family>/train.py \
|
||||
--config_path="config/<family>/<variant>.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--output_dir="output_dir_<family>" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--enable_bucket \
|
||||
--low_vram \
|
||||
--train_mode="normal" \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
## Preprocessing (data gen) — `scripts/<family>/generate_<...>.py`
|
||||
|
||||
Offline generation of cached training data (latents / ODE-trajectory pairs / prompt embeddings). **Always multi-GPU** (`accelerate launch` + `Accelerator`) and **always safetensors** (`safetensors.torch.save_file` + an `outputs.json` index for `ImageVideoSafetensorsDataset`) — never LMDB, never `.pt`. Mirror `scripts/wan2.1_self_forcing/generate_ode_pairs.py`:
|
||||
|
||||
```python
|
||||
# ...license header + sys.path bootstrap...
|
||||
import argparse, json, math, os, torch
|
||||
from accelerate import Accelerator
|
||||
from omegaconf import OmegaConf
|
||||
from safetensors.torch import save_file
|
||||
from tqdm import tqdm
|
||||
from videox_fun.models import AutoencoderKLWan, WanT5EncoderModel, WanTransformer3DModel # reuse repo models
|
||||
from videox_fun.utils.utils import save_videos_grid # reuse repo IO
|
||||
|
||||
def main():
|
||||
args = parse_args() # --pretrained_model_name_or_path --config_path --caption_path --output_folder
|
||||
# --num_inference_steps --guidance_scale --shift --mixed_precision ...
|
||||
accelerator = Accelerator(mixed_precision=args.mixed_precision)
|
||||
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
|
||||
torch.set_grad_enabled(False) # inference-only
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
config = OmegaConf.load(args.config_path) # config-driven loading (Section 3)
|
||||
weight_dtype = {"fp16": torch.float16, "bf16": torch.bfloat16}.get(accelerator.mixed_precision, torch.float32)
|
||||
text_encoder = WanT5EncoderModel.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), torch_dtype=weight_dtype).to(device).eval()
|
||||
vae = AutoencoderKLWan.from_pretrained(..., additional_kwargs=OmegaConf.to_container(config['vae_kwargs'])).to(device, dtype=weight_dtype).eval()
|
||||
transformer = WanTransformer3DModel.from_pretrained(..., transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])).to(device, dtype=weight_dtype).eval()
|
||||
|
||||
prompts = [l.rstrip() for l in open(args.caption_path, encoding="utf-8") if l.strip()]
|
||||
os.makedirs(args.output_folder, exist_ok=True)
|
||||
total_per_rank = math.ceil(len(prompts) / world_size)
|
||||
|
||||
for index in tqdm(range(total_per_rank), disable=rank != 0, desc="Generating"):
|
||||
prompt_index = index * world_size + rank # interleaved multi-GPU shard
|
||||
if prompt_index >= len(prompts):
|
||||
continue
|
||||
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
|
||||
if os.path.exists(out_path): # resume: skip already-done samples
|
||||
continue
|
||||
prompt = prompts[prompt_index]
|
||||
# ... encode prompt, sample noise, run the teacher ODE (CFG), collect latents ...
|
||||
save_file( # safetensors ONLY (no lmdb / no .pt)
|
||||
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
|
||||
out_path, metadata={"prompt": prompt},
|
||||
)
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
if accelerator.is_main_process: # rank-0 writes the JSON index
|
||||
entries = [{"file_path": os.path.join(args.output_folder, f"{i:05d}.safetensors")}
|
||||
for i in range(len(prompts))
|
||||
if os.path.exists(os.path.join(args.output_folder, f"{i:05d}.safetensors"))]
|
||||
json.dump(entries, open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
```
|
||||
|
||||
Launcher (`generate_<...>.sh`) — `accelerate launch` uses every visible GPU:
|
||||
```bash
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B"
|
||||
accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--config_path="config/<family>/*.yaml" \
|
||||
--caption_path="datasets/prompts.txt" \
|
||||
--output_folder="datasets/<family>_ode_pairs" \
|
||||
--num_inference_steps=48 --guidance_scale=6.0 --shift=8.0
|
||||
```
|
||||
|
||||
Training then reads the cache with `ImageVideoSafetensorsDataset(ann_path=".../outputs.json")` (single-file mode `{"file_path": ...}`, or per-tensor mode via `--save_per_tensor`). See reference.md §10.
|
||||
|
||||
> Dataset *curation* (scoring/filtering/captioning under `videox_fun/video_caption/`) is a different activity: also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM tensor-parallel) but writes csv/jsonl metadata, not safetensors. See reference.md §10 “Related but different”.
|
||||
@@ -0,0 +1,414 @@
|
||||
# VideoX-Fun Integration Reference
|
||||
|
||||
Detailed conventions per layer. Read the mirrored family's real files alongside this — the existing code is always the source of truth.
|
||||
|
||||
## 1. Model definitions — `videox_fun/models/<family>_*.py`
|
||||
|
||||
### File naming
|
||||
- Transformer / DiT: `<family>_transformer3d.py` (video) or `<family>_transformer2d.py` (image). Variants append a suffix: `_control`, `_s2v`, `_vace`, `_animate`, `_self_forcing`, `_avatar`.
|
||||
- VAE: `<family>_vae.py` → class `AutoencoderKL<Family>`.
|
||||
- Encoders: `<family>_text_encoder.py`, `<family>_audio_encoder.py`, `<family>_image_encoder.py`.
|
||||
|
||||
### Class shape (mirror `wan_transformer3d.py`)
|
||||
```python
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
class <Family>Transformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(self, model_type='t2v', patch_size=(1,2,2), in_dim=16, dim=2048,
|
||||
ffn_dim=8192, num_heads=16, num_layers=32, in_channels=16,
|
||||
hidden_size=2048, ...):
|
||||
super().__init__()
|
||||
...
|
||||
```
|
||||
- Keep BOTH civitai names (`in_dim`, `dim`, `ffn_dim`) and diffusers aliases (`in_channels`, `hidden_size`) in `__init__` so either format maps cleanly.
|
||||
- Implement `_set_gradient_checkpointing(self, *args, **kwargs)`.
|
||||
- Attention must go through `videox_fun.models.attention_utils.attention` (backend-agnostic), not a hand-rolled `scaled_dot_product_attention`.
|
||||
- Multi-GPU: expose `enable_multi_gpus_inference()` and route attention through the family's `dist/<family>_xfuser.py` processor.
|
||||
- Speedups live on the model: `enable_teacache(...)`, `enable_cfg_skip(...)`, `enable_riflex(...)`.
|
||||
|
||||
### `from_pretrained` internals (do not simplify)
|
||||
The custom classmethod must keep these behaviors (see `wan_transformer3d.py::from_pretrained`):
|
||||
1. Accept `transformer_additional_kwargs`, `subfolder`, `low_cpu_mem_usage`, `torch_dtype`.
|
||||
2. Read `config.json`; auto-convert foreign configs (e.g. diffsynth `has_image_input`) via a `_convert_from_*_config` helper.
|
||||
3. Apply `dict_mapping`: pop it from kwargs, then for each `key: target` set `kwargs[target] = config[key]`.
|
||||
4. Under `low_cpu_mem_usage`, build with `accelerate.init_empty_weights()`, load `.bin`/`.safetensors` (single file or glob all shards), and **filter by exact shape match** before loading.
|
||||
5. Initialize missing keys deliberately: zero-init control/audio projections (`after_proj`, `before_proj`, `processor.k_proj/v_proj`, `audio_injector`, `cond_encoder`, ...), ones for norms, xavier for ≥2D weights, so new branches start as no-ops.
|
||||
|
||||
### Registry — `videox_fun/models/__init__.py`
|
||||
Add an import line for every new public class, grouped with the family. Wrap optional-dependency imports in `try/except` with a helpful upgrade message (see the Qwen2.5-VL / Mistral3 blocks at the top).
|
||||
|
||||
## 2. Pipelines — `videox_fun/pipeline/pipeline_<family>*.py`
|
||||
|
||||
Mirror `pipeline_wan.py`. Required pieces:
|
||||
- Module-level `retrieve_timesteps(scheduler, num_inference_steps, device, timesteps, sigmas, **kwargs)` (copied from diffusers) — reuse verbatim.
|
||||
- `EXAMPLE_DOC_STRING` for the `@replace_example_docstring` decorator.
|
||||
- Output dataclass:
|
||||
```python
|
||||
@dataclass
|
||||
class <Family>PipelineOutput(BaseOutput):
|
||||
videos: torch.Tensor
|
||||
```
|
||||
- Pipeline class:
|
||||
```python
|
||||
class <Family>Pipeline(DiffusionPipeline):
|
||||
_optional_component = [...]
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae" # order matters for offload
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
def __init__(self, tokenizer, text_encoder, vae, transformer, scheduler, ...): ...
|
||||
def encode_prompt(...): ...
|
||||
def prepare_latents(...): ...
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(self, prompt, negative_prompt=..., height=..., width=...,
|
||||
num_frames=..., num_inference_steps=..., guidance_scale=...,
|
||||
generator=None, ..., return_dict=True) -> Union[<Family>PipelineOutput, Tuple]: ...
|
||||
```
|
||||
- Import schedulers from `..utils.fm_solvers` / `..utils.fm_solvers_unipc`, models from `..models`.
|
||||
- Separate pipelines per task: base (`pipeline_<family>.py`), inpaint/i2v (`_inpaint`), control (`_control`), s2v, etc. Register all in `videox_fun/pipeline/__init__.py`, adding convenience aliases (e.g. `WanI2VPipeline = WanFunInpaintPipeline`) where existing code expects them.
|
||||
|
||||
## 3. Config — `config/<family>/<name>.yaml` (optional)
|
||||
|
||||
**The YAML is not mandatory.** Decide by checkpoint layout:
|
||||
- **Required** for civitai-format / custom single-file layouts, where weights and key names are not diffusers-native. The YAML supplies `transformer_additional_kwargs` (incl. `dict_mapping` mapping civitai config keys → model `__init__` kwargs), component `*_subpath`s, and `vae/text_encoder/scheduler/image_encoder` kwargs.
|
||||
- **Optional** for a standard diffusers-layout checkpoint (`model_index.json` + each subfolder carrying its own `config.json`). Load components directly: `<Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", low_cpu_mem_usage=True, torch_dtype=...)`, `AutoencoderKL<Family>.from_pretrained(model_name, subfolder="vae")`, etc. Guard the config path exactly like `examples/minimax_h3_fun/predict_v2v_control.py`:
|
||||
```python
|
||||
transformer_load_kwargs = {}
|
||||
if config_path is not None:
|
||||
from omegaconf import OmegaConf
|
||||
config = OmegaConf.load(config_path)
|
||||
transformer_load_kwargs.update(OmegaConf.to_container(config["transformer_additional_kwargs"], resolve=True))
|
||||
transformer = <Family>Transformer3DModel.from_pretrained(model_name, subfolder="transformer", **transformer_load_kwargs, ...)
|
||||
```
|
||||
|
||||
When you do use a YAML, the canonical schema is below (see `config/wan2.1/wan_civitai.yaml`):
|
||||
```yaml
|
||||
format: civitai # or diffusers — selects weight-key handling
|
||||
pipeline: Wan # family label consumed by API/ComfyUI loaders
|
||||
transformer_additional_kwargs:
|
||||
transformer_subpath: ./ # subfolder under model_name holding the DiT
|
||||
dict_mapping: # civitai config key -> model __init__ kwarg
|
||||
in_dim: in_channels
|
||||
dim: hidden_size
|
||||
vae_kwargs:
|
||||
vae_subpath: Wan2.1_VAE.pth
|
||||
temporal_compression_ratio: 4
|
||||
spatial_compression_ratio: 8
|
||||
text_encoder_kwargs:
|
||||
text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth
|
||||
tokenizer_subpath: google/umt5-xxl
|
||||
text_length: 512
|
||||
...
|
||||
scheduler_kwargs:
|
||||
scheduler_subpath: null
|
||||
num_train_timesteps: 1000
|
||||
shift: 5.0
|
||||
...
|
||||
image_encoder_kwargs: # only for i2v / models with a CLIP image encoder
|
||||
image_encoder_subpath: models_clip_...pth
|
||||
```
|
||||
Every `*_subpath` is joined onto `model_name` in scripts. Load with `OmegaConf.load` and pass `OmegaConf.to_container(config['<section>'])` into `from_pretrained`. Use `filter_kwargs(Cls, OmegaConf.to_container(config['scheduler_kwargs']))` to build schedulers.
|
||||
|
||||
## 4. Inference scripts — `examples/<family>/predict_<task>.py`
|
||||
|
||||
Anatomy, top to bottom (see `examples/wan2.1_fun/predict_t2v.py`):
|
||||
1. **`sys.path` bootstrap** (before importing `videox_fun`):
|
||||
```python
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
sys.path.insert(0, project_root) if project_root not in sys.path else None
|
||||
```
|
||||
2. **User config block** as top-level variables with explanatory comments, in the conventional order: `GPU_memory_mode`, `ulysses_degree`/`ring_degree`, `fsdp_dit`/`fsdp_text_encoder`, `compile_dit`, TeaCache (`enable_teacache`, `teacache_threshold`, `num_skip_start_steps`, `teacache_offload`), `cfg_skip_ratio`, Riflex (`enable_riflex`, `riflex_k`), `config_path`, `model_name`, `sampler_name`, `shift`, `transformer_path`/`vae_path`/`lora_path`, `sample_size`, `video_length`, `fps`, `weight_dtype`, `prompt`/`negative_prompt`, `guidance_scale`, `seed`, `num_inference_steps`, `lora_weight`, `save_path`.
|
||||
3. **Device + config**: `device = set_multi_gpus_devices(ulysses_degree, ring_degree)`; then either `config = OmegaConf.load(config_path)` (civitai/custom layout) **or** guard `if config_path is not None:` and load components directly from a diffusers-layout checkpoint (see §3).
|
||||
4. **Component loading**: transformer (`from_pretrained(..., transformer_additional_kwargs=...)`), optional `transformer_path`/`vae_path` override with `load_state_dict(strict=False)` + missing/unexpected key print, vae, tokenizer, text_encoder, and clip image encoder gated by `transformer.config.in_channels != vae.config.latent_channels`.
|
||||
5. **Scheduler selection dict**: `{"Flow": FlowMatchEulerDiscreteScheduler, "Flow_Unipc": FlowUniPCMultistepScheduler, "Flow_DPM++": FlowDPMSolverMultistepScheduler}[sampler_name]`; build with `filter_kwargs`.
|
||||
6. **Pipeline construction**: choose base vs inpaint/i2v/control pipeline by the model's channel condition.
|
||||
7. **Multi-GPU / FSDP / compile**: if `ulysses_degree>1 or ring_degree>1` call `transformer.enable_multi_gpus_inference()` and optionally `shard_model`; if `compile_dit`, `torch.compile` each `transformer.blocks[i]`.
|
||||
8. **`GPU_memory_mode` branching** — keep this exact order:
|
||||
```python
|
||||
if GPU_memory_mode == "sequential_cpu_offload":
|
||||
replace_parameters_by_name(transformer, ["modulation",], device=device)
|
||||
transformer.freqs = transformer.freqs.to(device=device)
|
||||
pipeline.enable_sequential_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_group_offload":
|
||||
register_auto_device_hook(pipeline.transformer)
|
||||
safe_enable_group_offload(pipeline, onload_device=device, offload_device="cpu", offload_type="leaf_level", use_stream=True)
|
||||
elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_cpu_offload":
|
||||
pipeline.enable_model_cpu_offload(device=device)
|
||||
elif GPU_memory_mode == "model_full_load_and_qfloat8":
|
||||
convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device)
|
||||
convert_weight_dtype_wrapper(transformer, weight_dtype)
|
||||
pipeline.to(device=device)
|
||||
else:
|
||||
pipeline.to(device=device)
|
||||
```
|
||||
9. **TeaCache / cfg_skip / Riflex** enablement, `generator = torch.Generator(device).manual_seed(seed)`, LoRA `merge_lora`.
|
||||
10. **Inference** under `torch.no_grad()`; align `video_length` to `vae.config.temporal_compression_ratio`; pass `video`/`mask_video` for i2v via `get_image_to_video_latent`.
|
||||
11. **`save_results()`**: `save_videos_grid(sample, path, fps=fps)` for video, PIL save for a single frame; only rank 0 saves when multi-GPU. LoRA `unmerge_lora` after.
|
||||
|
||||
Other entry points to mirror when needed: `app.py` (Gradio), `launch_api.py` (API server backed by `videox_fun/api`), `post_infer*.py` (batch/queue inference).
|
||||
|
||||
## 5. Training scripts — `scripts/<family>/train*.py`
|
||||
|
||||
Mirror `scripts/wan2.1_fun/train.py`. Structure:
|
||||
1. Diffusers-derived license header + `"""Modified from ..."""` note.
|
||||
2. Third-party imports, then the **same `sys.path` bootstrap**, then `from videox_fun.data/models/pipeline/utils import ...`.
|
||||
3. Helper funcs: `filter_kwargs`, `resize_mask`, `linear_decay`, `generate_timestep_with_lognorm`.
|
||||
4. **`log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step)`** — builds the **inference Pipeline** from the live (unwrapped) transformer and runs it to produce sample videos under `output_dir/sample/`. Wrapped in try/except; handles DeepSpeed (`transformer3d.config` swap) and restores VAE/text-encoder placement (`low_vram`). **Reuse the pipeline; never write a separate sampler.**
|
||||
5. **`parse_args()`** — reuse the shared argument surface: `--config_path`, `--pretrained_model_name_or_path`, `--train_data_dir`, `--train_data_meta`, `--image_sample_size`/`--video_sample_size`/`--token_sample_size`, `--video_sample_n_frames`, `--video_sample_stride`, `--train_batch_size`, `--gradient_accumulation_steps`, `--learning_rate`, `--lr_scheduler`, `--lr_warmup_steps`, `--checkpointing_steps`, `--output_dir`, `--mixed_precision`, `--gradient_checkpointing`, `--enable_bucket`, `--random_hw_adapt`, `--training_with_video_token_length`, `--uniform_sampling`, `--low_vram`, `--train_mode`, `--trainable_modules`, LoRA args (`--use_lora`, `--rank`, ...), `--validation_prompts`/`--validation_paths`. Add new args only when the family genuinely needs them.
|
||||
6. **`main()`** — Accelerator setup, DeepSpeed/FSDP zero-stage handling (auto-sets `save_state`), model loading via config, dataset + bucket sampler from `videox_fun.data`, trainable-module filtering / LoRA network via `create_network`, optimizer + `get_scheduler`, `accelerator.prepare`, checkpoint save/load hooks, training loop with timestep sampling, loss, `log_validation` at intervals, and final weight/LoRA save.
|
||||
|
||||
### Resolution args — `--video_sample_size` (+ `--fix_sample_size`)
|
||||
Canvas resolution is always driven by a **single square** `--video_sample_size` (`type=int`, height = width) — never by separate `--video_sample_height` / `--video_sample_width`. When a **fixed non-square shape** is required, add `--fix_sample_size` (`nargs=2, type=int, default=None`, `[height, width]`) that overrides the square size; mirror `scripts/wan2.2_fun/train_lora.py`, `scripts/z_image/train_distill.py`. Derive the effective `height` / `width` once in `parse_args()` and reuse them everywhere downstream:
|
||||
```python
|
||||
parser.add_argument("--video_sample_size", type=int, default=1280)
|
||||
parser.add_argument("--fix_sample_size", nargs=2, type=int, default=None,
|
||||
help="Fix Sample size [height, width] to override `--video_sample_size` with a fixed non-square shape.")
|
||||
...
|
||||
if args.fix_sample_size is not None:
|
||||
args.video_sample_height, args.video_sample_width = args.fix_sample_size
|
||||
else:
|
||||
args.video_sample_height = args.video_sample_width = args.video_sample_size
|
||||
```
|
||||
In bucket datasets `--fix_sample_size` also forces `random_hw_adapt=False` / `training_with_video_token_length=False` and bumps `video_sample_size = max(max(fix_sample_size), video_sample_size)`; in data-free scripts (e.g. `scripts/minimax_h3/train_pdd_lora.py`) it simply pins the generation canvas. Always validate the size against the patch/VAE constraint (minimax_h3: `% 32`). The `.sh` launcher passes it space-separated (`nargs=2`): `--fix_sample_size 768 1344`.
|
||||
|
||||
### Launcher — `scripts/<family>/train*.sh`
|
||||
`export MODEL_NAME/DATASET_NAME/DATASET_META_NAME`, then `accelerate launch --mixed_precision="bf16" scripts/<family>/train.py --config_path=... <full arg list>`. Include commented I2V/control variants and DeepSpeed/NCCL notes as the existing scripts do.
|
||||
|
||||
### Docs — `README_TRAIN.md` + `README_TRAIN_zh-CN.md`
|
||||
Aligned bilingual pair: identical section order, identical commands and parameter tables; only the prose language differs. Follow the top-level section order used across existing training READMEs.
|
||||
|
||||
## 6. Shared infrastructure map (reuse, never reimplement)
|
||||
|
||||
| Need | Import from |
|
||||
|------|-------------|
|
||||
| Flow/DPM/UniPC schedulers | `diffusers`, `videox_fun.utils.fm_solvers`, `videox_fun.utils.fm_solvers_unipc` |
|
||||
| LoRA create/merge/unmerge/convert | `videox_fun.utils.lora_utils` |
|
||||
| FP8 quantization | `videox_fun.utils.fp8_optimization` |
|
||||
| Group / leaf offload hooks | `videox_fun.utils.group_offload` |
|
||||
| Multi-GPU device + FSDP shard + seq-parallel attn | `videox_fun.dist` |
|
||||
| Save video/audio, image→video latents, kwarg filter, dimension calc | `videox_fun.utils.utils` |
|
||||
| Datasets + bucket/aspect-ratio samplers + masks | `videox_fun.data` |
|
||||
| TeaCache coefficients | `videox_fun.models.cache_utils` |
|
||||
|
||||
## 7. Naming quick reference
|
||||
|
||||
| Concept | Convention | Example |
|
||||
|---------|-----------|---------|
|
||||
| Model file | `<family>_transformer3d.py` | `wan_transformer3d.py` |
|
||||
| Model class | `<Family>Transformer3DModel` | `WanTransformer3DModel` |
|
||||
| VAE class | `AutoencoderKL<Family>` | `AutoencoderKLWan` |
|
||||
| Pipeline file | `pipeline_<family>.py` | `pipeline_wan.py` |
|
||||
| Pipeline class | `<Family>Pipeline` | `WanPipeline` / `WanFunInpaintPipeline` |
|
||||
| Config | `config/<family>/<variant>.yaml` | `config/wan2.1/wan_civitai.yaml` |
|
||||
| Inference | `examples/<family>/predict_<task>.py` | `predict_t2v.py`, `predict_i2v.py`, `predict_v2v_control.py` |
|
||||
| Training | `scripts/<family>/train[_<variant>].py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py` |
|
||||
|
||||
## 8. Training data pipeline — dataset & sampler selection
|
||||
|
||||
Pick the dataset by **task / `train_mode`**, then the sampler by **`enable_bucket`** and dataset type. All datasets/samplers come from `videox_fun.data` — never write a new one.
|
||||
|
||||
### Annotation format — the `train_data_meta` file (`metadata.json` / `.csv`)
|
||||
Every dataset class reads an annotation file (`args.train_data_meta`) that indexes the media under `args.train_data_dir` (`data_root`). `ImageVideoDataset` accepts **`.json`** (a top-level array of records) or **`.csv`** (`csv.DictReader`; the header row is the field names). Each record for ordinary image/video training:
|
||||
|
||||
| Field | Required | Meaning |
|
||||
|-------|----------|---------|
|
||||
| `file_path` | yes | Media path, resolved **relative to `train_data_dir`** via `os.path.join(data_root, file_path)`. If `data_root is None`, `file_path` is used as-is. |
|
||||
| `text` | yes | Caption / prompt. Dropped to `""` with probability `text_drop_ratio` (default `0.1`) for classifier-free guidance. |
|
||||
| `type` | no | `"video"` or `"image"`; **defaults to `"image"`** when the key is absent (`data_info.get('type', 'image')`). |
|
||||
|
||||
```json
|
||||
[
|
||||
{"file_path": "train/00000000.mp4", "text": "A young woman gently turns her head to the right ...", "type": "video"},
|
||||
{"file_path": "train/00000001.jpg", "text": "a dog running on the beach", "type": "image"}
|
||||
]
|
||||
```
|
||||
The directory layout matches the index — media in a `train/` subdir, the annotation file beside it. Ready-made examples ship in `datasets/X-Fun-Videos-Demo/` (`train/*.mp4` + `metadata.json`) and `datasets/X-Fun-Images-Demo/`. The equivalent `.csv`:
|
||||
```csv
|
||||
file_path,text,type
|
||||
train/00000000.mp4,"A young woman gently turns her head to the right ...",video
|
||||
train/00000001.jpg,"a dog running on the beach",image
|
||||
```
|
||||
|
||||
**Variant datasets append extra fields to this same record shape**, each consumed by its own class (see the table below) — e.g. camera-pose adds `action_path` (`LingbotImageVideoDataset`), object/VACE/S2V variants add object fields (`object_file_path` / `objects`). The demo folders also ship several augmented metadata variants (next subsection). Always read the target class's `get_batch` for the exact fields it consumes.
|
||||
|
||||
### Ready-made demo datasets — the standard test data (never invent a test set)
|
||||
Smoke tests, `log_validation` checks, and doc examples all run on the official demo datasets under `datasets/`, downloaded from ModelScope as `PAI/<name>`:
|
||||
```bash
|
||||
modelscope download --dataset PAI/X-Fun-Videos-Demo --local_dir ./datasets/X-Fun-Videos-Demo
|
||||
```
|
||||
Pick the demo by **task**, matching the dataset class in the table below:
|
||||
|
||||
| Demo dataset (`datasets/...`) | Contents | Extra metadata fields | Task it tests | Dataset class |
|
||||
|-------------------------------|----------|----------------------|---------------|---------------|
|
||||
| `X-Fun-Videos-Demo` | 16 videos (832×480) in `train/` | — | T2V / I2V base + inpaint, distill | `ImageVideoDataset` |
|
||||
| `X-Fun-Videos-Controls-Demo` | 16 videos in `train/` + `canny/` + `object/<video_id>/` + `wav/` | `control_file_path`, `object_file_path` (list), `audio_path` | V2V control, VACE, S2V-with-control | `ImageVideoControlDataset`, `VideoSpeechControlDataset` |
|
||||
| `X-Fun-Videos-Audios-Demo` | 17 video/audio pairs: `train/` (1280×720) + `wav/` (16 kHz mono) + `pose/` | `audio_path`, `control_file_path` | Speech-driven S2V / avatar / talking-head | `VideoSpeechDataset` |
|
||||
| `X-Fun-Images-Demo` | 19 images in `train/` | — | T2I full fine-tune + LoRA (z_image / flux2 / qwenimage / lens / ernie) | `ImageVideoDataset` |
|
||||
| `X-Fun-Images-Controls-Demo` | 19 images in `train/` + `canny/` | `control_file_path` | Image control / ControlNet / i2i inpaint | `ImageVideoControlDataset` |
|
||||
| `X-Fun-Images-Edit-Demo` | 21 records: `source/souce-<id>/` (multi-source supported) → `train/` | `source_file_path` (**list**) | Image edit (Qwen-Image-Edit family) | `ImageEditDataset` |
|
||||
| `X-Fun-Videos-Lingbot-Demo` | video + `intrinsics.npy` / `poses.npy` | camera pose / action | Camera-pose world model (`lingbot_world`) | `LingbotImageVideoDataset` |
|
||||
|
||||
**Which metadata file to point `--train_data_meta` at** (each demo ships several variants beside the media):
|
||||
|
||||
| Metadata file | Use when |
|
||||
|---------------|----------|
|
||||
| `metadata.json` | Base format only (`file_path` / `text` / `type`) — fine for a minimal check |
|
||||
| `metadata_add_width_height.json` | **Default choice.** Adds `width` / `height` so bucketing doesn't decode media (matters on slow storage such as OSS). Used by non-VACE control / S2V training too |
|
||||
| `metadata_add_width_height_add_objects.json` | VACE / subject-reference training (`object_file_path` list → `object/<video_id>/`; shuffled at train time) |
|
||||
| `metadata_add_width_height_add_wav.json` | Audio-visual joint models (e.g. `minimax_h3_fun` control training): `audio_path` → `wav/`. Keep the `.sh` launcher and the README on the same file |
|
||||
| `metadata_lingbot_video_add_width_height.json` | `lingbot_video` — `text` is already a structured JSON caption (lives in `X-Fun-Videos-Demo`) |
|
||||
| `metadata_origin.json` | Pre-processing original kept for reference; not used for training |
|
||||
|
||||
Regenerate the width/height variant with the shipped helper when adding your own media:
|
||||
`python scripts/process_json_add_width_and_height.py --input_file datasets/<Demo>/metadata.json --output_file datasets/<Demo>/metadata_add_width_height.json`.
|
||||
|
||||
`audio_path` optionality differs per class (`videox_fun/data/dataset_video.py`): `VideoSpeechDataset` reads `video_dict['audio_path']` directly, so it is **required**; `VideoSpeechControlDataset` uses `.get('audio_path')` and **falls back to the video file's own audio track** when the field is absent.
|
||||
|
||||
### Dataset by task (all take `train_data_meta, train_data_dir, ...`)
|
||||
| Task / mode | Dataset class | Used by | Key kwargs |
|
||||
|-------------|--------------|---------|-----------|
|
||||
| T2V / I2V base (`normal` + inpaint) | `ImageVideoDataset` | `train.py`, `train_lora.py`, t2i `train.py` | `enable_inpaint = train_mode != "normal"`, `video_sample_size/stride/n_frames`, `image_sample_size`, `video_repeat` |
|
||||
| Image T2I (qwenimage/flux/z_image) | `ImageVideoDataset` | `scripts/<img>/train.py` | `image_sample_size` |
|
||||
| Control (canny/pose/depth/camera) | `ImageVideoControlDataset` | `train_control*.py`, `train_control_distill.py` | `enable_camera_info = train_mode == "control_camera_ref"` |
|
||||
| Image Edit (source→target) | `ImageEditDataset` | `qwenimage/train_edit*.py` | `image_sample_size` |
|
||||
| Speech/audio-driven (S2V, avatar, talking) | `VideoSpeechDataset` | `mova`, `ltx2`, `minimax_h3`, `fantasytalking`, `infinitetalk`, `flashhead`, `longcatvideo/train_avatar*` | audio + video fields |
|
||||
| S2V **with control** | `VideoSpeechControlDataset` | `wan2.2/train_s2v*.py`, `minimax_h3_fun/train_control*` | audio + control |
|
||||
| Motion/pose animate | `VideoAnimateDataset` | `wan2.2/train_animate*.py` | motion/pose driven |
|
||||
| Distill text-only branch, GRPO, DPO | `TextDataset` | `train_distill*.py` (text branch), `z_image/train_grpo_lora.py`, `train_dpo_lora.py` | reads only the `text` field; `text_drop_ratio` |
|
||||
| Precomputed latents (ODE pairs) | `ImageVideoSafetensorsDataset` | `wan2.1_self_forcing/train_ode.py` | `data_root` |
|
||||
| Camera-pose conditioning | `LingbotImageVideoDataset` | `lingbot_world/train.py` | `intrinsics.npy` / `poses.npy` |
|
||||
| Video-only (VAE/TAEHV distill) | `VideoDataset` | `taehv/train_taehv.py` | `sample_size/stride/n_frames`, `enable_inpaint=False` |
|
||||
|
||||
### Sampler by condition
|
||||
| Condition | Sampler | Shape |
|
||||
|-----------|---------|-------|
|
||||
| `enable_bucket=True` (default; image+video) | `AspectRatioBatchImageVideoSampler` | `sampler=RandomSampler(ds, generator=g), dataset=train_dataset.dataset, batch_size, train_folder=args.train_data_dir, drop_last=True, aspect_ratios=aspect_ratio_sample_size` |
|
||||
| `enable_bucket=False` | `ImageVideoSampler` | `ImageVideoSampler(RandomSampler(ds, generator=g), train_dataset, batch_size)` |
|
||||
| `TextDataset` (distill text branch / GRPO / DPO) | `BatchSampler` (plain) | `BatchSampler(RandomSampler(ds, generator=g), batch_size, drop_last=True)`; GRPO adds `k_repeat=args.num_image_per_prompt` |
|
||||
| video-only bucket (available, not used by current scripts) | `AspectRatioBatchSampler` | — |
|
||||
| image-only bucket (available, not used by current scripts) | `AspectRatioBatchImageSampler` | — |
|
||||
|
||||
`aspect_ratio_sample_size` is built from `ASPECT_RATIO_512` scaled by `args.video_sample_size`; `get_closest_ratio` picks the bucket inside `collate_fn`.
|
||||
|
||||
### Universal DataLoader creation pattern
|
||||
```python
|
||||
batch_sampler_generator = torch.Generator().manual_seed(args.seed)
|
||||
if args.enable_bucket:
|
||||
aspect_ratio_sample_size = {k: [x / 512 * args.video_sample_size for x in ASPECT_RATIO_512[k]] for k in ASPECT_RATIO_512}
|
||||
batch_sampler = AspectRatioBatchImageVideoSampler(
|
||||
sampler=RandomSampler(train_dataset, generator=batch_sampler_generator), dataset=train_dataset.dataset,
|
||||
batch_size=args.train_batch_size, train_folder=args.train_data_dir, drop_last=True,
|
||||
aspect_ratios=aspect_ratio_sample_size)
|
||||
def collate_fn(examples):
|
||||
new_examples = {"pixel_values": [], "text": []}
|
||||
if args.train_mode != "normal": # inpaint/i2v adds mask fields
|
||||
new_examples.update({"mask_pixel_values": [], "mask": [], "clip_pixel_values": []})
|
||||
# bucket via get_closest_ratio -> transform (Resize/CenterCrop/Normalize) -> stack
|
||||
# masked branch uses get_random_mask(...)
|
||||
return new_examples
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn,
|
||||
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
|
||||
else:
|
||||
batch_sampler = ImageVideoSampler(RandomSampler(train_dataset, generator=batch_sampler_generator), train_dataset, args.train_batch_size)
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset, batch_sampler=batch_sampler,
|
||||
persistent_workers=args.dataloader_num_workers != 0, num_workers=args.dataloader_num_workers,
|
||||
worker_init_fn=worker_init_fn(args.seed + accelerator.process_index))
|
||||
```
|
||||
`collate_fn` receives the `examples` **list** (not a `batch` dict); build every batch-level field (`text`, `pixel_values`, masks) explicitly from `examples` into `new_examples`. When `--enable_text_encoder_in_dataloader`, encode prompts inside `collate_fn` and emit `encoder_hidden_states` / `encoder_attention_mask`.
|
||||
|
||||
## 9. Inference task matrix — predict script → pipeline → inputs
|
||||
|
||||
Pick the pipeline by **task**; the `predict_<task>.py` name and its inputs follow the same convention across families.
|
||||
|
||||
| Task | `predict_<task>.py` | Pipeline (family example) | Extra `__call__` inputs | Input helper |
|
||||
|------|--------------------|---------------------------|-------------------------|--------------|
|
||||
| Text→Video | `predict_t2v.py` | `WanPipeline`, `Wan2_2Pipeline`, `CogVideoXFunPipeline`, `LongCatVideoPipeline`, `LTX2Pipeline` | `prompt` only | — |
|
||||
| Image→Video | `predict_i2v.py` | `WanI2VPipeline`(=`WanFunInpaintPipeline`), `Wan2_2FunInpaintPipeline`, `Wan2_2I2VPipeline`, `HunyuanVideoI2VPipeline` | `video`, `mask_video` | `get_image_to_video_latent(start_image, end_image, video_length, sample_size)` |
|
||||
| Text+Image→Video (5B) | `predict_ti2v.py` | `Wan2_2TI2VPipeline` | `prompt` (+ optional image) | `get_image_to_video_latent` |
|
||||
| Video→Video Control | `predict_v2v_control.py` | `WanFunControlPipeline`, `Wan2_2FunControlPipeline` | `control_video` | `get_video_to_video_latent(control_video, ...)` |
|
||||
| Control + reference | `predict_v2v_control_ref.py` | `WanFunControlPipeline` | `control_video` + `ref_image` | `get_video_to_video_latent` + `get_image_latent` |
|
||||
| Control + camera | `predict_v2v_control_camera.py` | `WanFunControlPipeline` | `control_video` + camera pose | — |
|
||||
| VACE (control/mask/i2v/s2v) | `predict_v2v_control.py`, `predict_v2v_mask.py`, `predict_s2v.py`, `predict_i2v.py` | `WanVacePipeline`, `Wan2_2VaceFunPipeline` | control/mask/ref | — |
|
||||
| Speech→Video (audio) | `predict_s2v.py` | `Wan2_2S2VPipeline`, `MiniMaxH3Pipeline`, `InfiniteTalkPipeline`, `FantasyTalkingPipeline`, `FlashHeadPipeline`, `MOVAPipeline`, `LongCatVideoAvatarPipeline` | `audio` + reference image | — |
|
||||
| Animate (motion/pose) | `predict_animate.py` | `Wan2_2AnimatePipeline` | motion/pose video + ref | — |
|
||||
| Subject reference | `predict_s2v.py` (phantom) | `WanFunPhantomPipeline` | reference images | — |
|
||||
| Text→Image | `predict_t2i.py` | `QwenImagePipeline`, `Flux2Pipeline`, `ZImagePipeline`, `LensPipeline`, `ErnieImagePipeline` | `prompt` | — |
|
||||
| Image Control (t2i) | `predict_t2i_control.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline`, `QwenImageControlNetPipeline` | `control_image` | — |
|
||||
| Inpaint (i2i) | `predict_i2i_inpaint.py` | `QwenImageControlPipeline`, `ZImageControlPipeline`, `Flux2ControlPipeline` | `image` + `mask` | — |
|
||||
| Image Edit | `predict_t2i_edit.py`, `predict_t2i_edit_plus.py` | `QwenImageEditPipeline`, `QwenImageEditPlusPipeline` | source image + instruction | — |
|
||||
| Layered edit | `predict_i2i_layered.py` | `QwenImageLayeredPipeline` | image | — |
|
||||
| Camera-pose world | `predict_i2v.py` (lingbot_world) | `Wan2_2I2VPipeline`, `WanFunLingbotWorldFastPipeline` | image + camera pose | — |
|
||||
| Latent upsample | `predict_i2v_upsample.py` | `LTX2LatentUpsamplePipeline`, `WanLatentUpsamplePipeline` | low-res latent/video | — |
|
||||
| AR / streaming distill | `predict_t2v_stream.py` | `WanSelfForcingPipeline` | prompt (streamed) | — |
|
||||
|
||||
### Predict-script variant suffixes (same task, different backend/model)
|
||||
| Suffix | Meaning |
|
||||
|--------|---------|
|
||||
| `_tae` | Fast decode via `AutoencoderTinyWan` (TAEHV) instead of the full VAE |
|
||||
| `_2.2vae` | Uses the Wan2.2 VAE (`AutoencoderKLWan3_8`) |
|
||||
| `_5b` | 5B-parameter model variant |
|
||||
| `turbo` / distill | Distilled model, few-step inference (e.g. `predict_turbo_*.py`) |
|
||||
| `_refine` | Two-stage refine pass |
|
||||
| `_ref` / `_camera` | Adds reference-image / camera conditioning |
|
||||
|
||||
All variants keep the identical config block, `GPU_memory_mode` branching, and `save_results()` from Section 4 — only the loaded VAE/transformer and pipeline class change.
|
||||
|
||||
## 10. Preprocessing — offline training-data generation (multi-GPU + safetensors)
|
||||
|
||||
Here "preprocessing" means **generating/caching training data offline** with the teacher / VAE / text-encoder — latents, ODE-trajectory pairs, prompt/text embeddings — so training just reads cached tensors instead of re-encoding every step. Canonical example: `scripts/wan2.1_self_forcing/generate_ode_pairs.py` (+ `generate_ode_pairs.sh`); the loader-side contract is `ImageVideoSafetensorsDataset` in `videox_fun/data/dataset_image_video.py`. Two rules are non-negotiable.
|
||||
|
||||
### Rule 1 — multi-GPU is mandatory
|
||||
Never a single-GPU / hardcoded `cuda:0` loop. Launch with `accelerate launch` and shard work across ranks by interleaving:
|
||||
```python
|
||||
from accelerate import Accelerator
|
||||
accelerator = Accelerator(mixed_precision=args.mixed_precision)
|
||||
device, world_size, rank = accelerator.device, accelerator.num_processes, accelerator.process_index
|
||||
torch.set_grad_enabled(False) # inference-only
|
||||
|
||||
total_per_rank = math.ceil(len(prompts) / world_size)
|
||||
for index in tqdm(range(total_per_rank), disable=rank != 0):
|
||||
prompt_index = index * world_size + rank # interleaved shard
|
||||
if prompt_index >= len(prompts):
|
||||
continue
|
||||
out_path = os.path.join(args.output_folder, f"{prompt_index:05d}.safetensors")
|
||||
if os.path.exists(out_path): # resume-friendly
|
||||
continue
|
||||
... # encode prompt / run teacher ODE / collect latents
|
||||
accelerator.wait_for_everyone()
|
||||
if accelerator.is_main_process: # write the JSON index once, on rank 0
|
||||
json.dump([{"file_path": p} for p in all_safetensor_paths],
|
||||
open(os.path.join(args.output_folder, "outputs.json"), "w"), ensure_ascii=False, indent=4)
|
||||
```
|
||||
Launcher (`.sh`): `accelerate launch --mixed_precision="bf16" scripts/<family>/generate_<...>.py --pretrained_model_name_or_path=... --config_path=config/<family>/*.yaml --output_folder=datasets/<...> ...`. Reuse `videox_fun.models` + config-driven `from_pretrained` (Section 3) and `videox_fun.utils.utils.save_videos_grid` for sample previews — do not write a new loader.
|
||||
|
||||
### Rule 2 — store as safetensors; do NOT use LMDB or `.pt`
|
||||
Save every cached tensor with `safetensors.torch.save_file`, one `.safetensors` per sample (or per tensor), plus a JSON index of `{"file_path": ...}` entries:
|
||||
```python
|
||||
from safetensors.torch import save_file
|
||||
save_file(
|
||||
{"latents": latents.cpu(), "prompt_embeds": text_embeds.cpu(), "prompt_attention_mask": mask.cpu()},
|
||||
out_path, # f"{prompt_index:05d}.safetensors"
|
||||
metadata={"prompt": prompt},
|
||||
)
|
||||
```
|
||||
`ImageVideoSafetensorsDataset(ann_path, data_root=None)` reads that JSON and supports two layouts:
|
||||
- **Single-file (default)**: `{"file_path": "scene.safetensors"}` — whole state dict in one archive.
|
||||
- **Per-tensor (`--save_per_tensor`)**: `{"file_path": "scene_dir", "latents": ".../latents.safetensors", "prompt_embeds": ".../prompt_embeds.safetensors"}` — each key loaded and merged.
|
||||
|
||||
**Do not** cache preprocessed data in **LMDB** or as **`.pt`/`.pth` `torch.save` pickles**. safetensors is the repo-wide standard (also used for LoRA/weight saving), is pickle-free/safe, memory-maps fast, and is exactly what `ImageVideoSafetensorsDataset` loads. (Scope: this governs cached *data tensors*; accelerate optimizer/scheduler/scaler `.pt` states written during training checkpoints are a separate mechanism and unaffected.)
|
||||
|
||||
### Related but different — dataset curation
|
||||
Scoring / filtering / captioning under `videox_fun/video_caption/` (`compute_*.py`, `internvl2_video_recaptioning.py`) is dataset *curation*, not latent caching. It is also multi-GPU (accelerate `PartialState.split_between_processes`/`gather_object`, or vLLM `tensor_parallel_size=device_count()`), but writes csv/jsonl **metadata** (not tensors), so Rule 2 does not apply there.
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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`
|
||||
@@ -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
|
||||
@@ -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]
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user