diff --git a/.skills/integrating-models/SKILL.md b/.skills/integrating-models/SKILL.md new file mode 100644 index 0000000..c0f5b7f --- /dev/null +++ b/.skills/integrating-models/SKILL.md @@ -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//predict_*.py` (inference entry) + - `scripts//train*.py` + `*.sh` + `README_TRAIN*.md` (training) + - `videox_fun/pipeline/pipeline_*.py` (pipeline) + - `videox_fun/models/_*.py` (model definitions) + - `config//*.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/_*.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_*.py` | `Pipeline(DiffusionPipeline)` with `__call__`. | +| Pipeline registry | `videox_fun/pipeline/__init__.py` | Imports every pipeline + aliases. **Must be updated.** | +| Configs (optional) | `config//*.yaml` | OmegaConf YAML for civitai/custom layouts; a standard diffusers-layout checkpoint can load without one. | +| Inference entry scripts | `examples//predict_*.py` | User-facing, config-block-at-top runnable scripts. | +| Inference services | `examples//{app.py,launch_api.py,post_infer*.py}` | Gradio UI / API server / batch inference. | +| Training scripts | `scripts//train*.py` | `train.py`, `train_lora.py`, `train_control.py`, `train_distill.py`, ... | +| Training launchers | `scripts//train*.sh` | `accelerate launch` / DeepSpeed command with full arg list. | +| Training docs | `scripts//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`, `_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/`; each ships several `metadata*.json` variants. **The only test data to use** (see reference.md §8). | +| Preprocessing (data gen) | `scripts//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//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// +- [ ] Step 4: Inference script(s) in examples//predict_*.py +- [ ] Step 5: Training script(s) in scripts//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/_transformer3d.py` (or `2d`), `_vae.py`, encoders as needed. Mirror the class shape: `class 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_.py`. Mirror `pipeline_wan.py`: module-level `retrieve_timesteps`, a `PipelineOutput(BaseOutput)` dataclass, `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//` 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//predict_.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//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), `_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//*.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 `_transformer3d.py` / `_vae.py` / `pipeline_.py`; classes `Transformer3DModel` / `AutoencoderKL` / `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/`), 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//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_.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) diff --git a/.skills/integrating-models/examples.md b/.skills/integrating-models/examples.md new file mode 100644 index 0000000..b5e9ea4 --- /dev/null +++ b/.skills/integrating-models/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 `` / `` / ``. + +## Config — `config//.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: +transformer_additional_kwargs: + transformer_subpath: ./ + dict_mapping: + in_dim: in_channels + dim: hidden_size + +vae_kwargs: + vae_subpath: _VAE.pth + temporal_compression_ratio: 4 + spatial_compression_ratio: 8 + +text_encoder_kwargs: + text_encoder_subpath: .pth + tokenizer_subpath: + 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: .pth +``` + +## Inference — `examples//predict_.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, TextEncoder, + Transformer3DModel) +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import 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//.yaml" +model_name = "models/Diffusion_Transformer/-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/-" + +# --- 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 = 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.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 = 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 = 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 = InpaintPipeline(..., clip_image_encoder=clip_image_encoder) + # sample = pipeline(..., video=input_video, mask_video=input_video_mask).videos +``` + +## Pipeline class — `videox_fun/pipeline/pipeline_.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, 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 PipelineOutput(BaseOutput): + videos: torch.Tensor + +class 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[PipelineOutput, tuple]: + # 1. encode_prompt 2. prepare_latents 3. retrieve_timesteps + # 4. denoising loop with guidance 5. vae.decode 6. return PipelineOutput(videos=...) + ... +``` +Then register in `videox_fun/pipeline/__init__.py`: +```python +from .pipeline_ import Pipeline +``` + +## Model class — `videox_fun/models/_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 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/_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 ._transformer3d import Transformer3DModel +from ._vae import AutoencoderKL +``` + +## Training — `scripts//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, Transformer3DModel +from videox_fun.pipeline import 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 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//train.sh` + +```bash +export MODEL_NAME="models/Diffusion_Transformer/-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//train.py \ + --config_path="config//.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_" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --enable_bucket \ + --low_vram \ + --train_mode="normal" \ + --trainable_modules "." +``` + +## Preprocessing (data gen) — `scripts//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//generate_<...>.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --config_path="config//*.yaml" \ + --caption_path="datasets/prompts.txt" \ + --output_folder="datasets/_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”. diff --git a/.skills/integrating-models/reference.md b/.skills/integrating-models/reference.md new file mode 100644 index 0000000..c05b8e7 --- /dev/null +++ b/.skills/integrating-models/reference.md @@ -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/_*.py` + +### File naming +- Transformer / DiT: `_transformer3d.py` (video) or `_transformer2d.py` (image). Variants append a suffix: `_control`, `_s2v`, `_vace`, `_animate`, `_self_forcing`, `_avatar`. +- VAE: `_vae.py` → class `AutoencoderKL`. +- Encoders: `_text_encoder.py`, `_audio_encoder.py`, `_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 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/_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_*.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 PipelineOutput(BaseOutput): + videos: torch.Tensor + ``` +- Pipeline class: + ```python + class 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[PipelineOutput, Tuple]: ... + ``` +- Import schedulers from `..utils.fm_solvers` / `..utils.fm_solvers_unipc`, models from `..models`. +- Separate pipelines per task: base (`pipeline_.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//.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: `Transformer3DModel.from_pretrained(model_name, subfolder="transformer", low_cpu_mem_usage=True, torch_dtype=...)`, `AutoencoderKL.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 = 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['
'])` into `from_pretrained`. Use `filter_kwargs(Cls, OmegaConf.to_container(config['scheduler_kwargs']))` to build schedulers. + +## 4. Inference scripts — `examples//predict_.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//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//train*.sh` +`export MODEL_NAME/DATASET_NAME/DATASET_META_NAME`, then `accelerate launch --mixed_precision="bf16" scripts//train.py --config_path=... `. 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 | `_transformer3d.py` | `wan_transformer3d.py` | +| Model class | `Transformer3DModel` | `WanTransformer3DModel` | +| VAE class | `AutoencoderKL` | `AutoencoderKLWan` | +| Pipeline file | `pipeline_.py` | `pipeline_wan.py` | +| Pipeline class | `Pipeline` | `WanPipeline` / `WanFunInpaintPipeline` | +| Config | `config//.yaml` | `config/wan2.1/wan_civitai.yaml` | +| Inference | `examples//predict_.py` | `predict_t2v.py`, `predict_i2v.py`, `predict_v2v_control.py` | +| Training | `scripts//train[_].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/`: +```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//` + `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-/` (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//`; 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//metadata.json --output_file datasets//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//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_.py` name and its inputs follow the same convention across families. + +| Task | `predict_.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//generate_<...>.py --pretrained_model_name_or_path=... --config_path=config//*.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. diff --git a/examples/minimax_h3/predict_i2v.py b/examples/minimax_h3/predict_i2v.py index bb2147f..59571e1 100644 --- a/examples/minimax_h3/predict_i2v.py +++ b/examples/minimax_h3/predict_i2v.py @@ -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) diff --git a/examples/minimax_h3/predict_ref2va.py b/examples/minimax_h3/predict_ref2va.py index 3277ede..6121bb7 100644 --- a/examples/minimax_h3/predict_ref2va.py +++ b/examples/minimax_h3/predict_ref2va.py @@ -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) diff --git a/examples/minimax_h3/predict_t2v.py b/examples/minimax_h3/predict_t2v.py index ac5fde5..ede6535 100644 --- a/examples/minimax_h3/predict_t2v.py +++ b/examples/minimax_h3/predict_t2v.py @@ -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) diff --git a/examples/wan2.1_self_forcing/predict_t2v_forcing_kv.py b/examples/wan2.1_self_forcing/predict_t2v_forcing_kv.py new file mode 100644 index 0000000..7662a41 --- /dev/null +++ b/examples/wan2.1_self_forcing/predict_t2v_forcing_kv.py @@ -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() diff --git a/examples/wan2.1_self_forcing/profile_forcing_kv_heads.py b/examples/wan2.1_self_forcing/profile_forcing_kv_heads.py new file mode 100644 index 0000000..8bdee43 --- /dev/null +++ b/examples/wan2.1_self_forcing/profile_forcing_kv_heads.py @@ -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}") diff --git a/examples/z_image_fun/predict_t2i_control_2.1.py b/examples/z_image_fun/predict_t2i_control_2.1.py index b31bff7..a707129 100644 --- a/examples/z_image_fun/predict_t2i_control_2.1.py +++ b/examples/z_image_fun/predict_t2i_control_2.1.py @@ -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, ) diff --git a/examples/z_image_fun/predict_t2i_control_2.1_lite.py b/examples/z_image_fun/predict_t2i_control_2.1_lite.py index 6defa3c..a4c348f 100644 --- a/examples/z_image_fun/predict_t2i_control_2.1_lite.py +++ b/examples/z_image_fun/predict_t2i_control_2.1_lite.py @@ -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() \ No newline at end of file + save_results() diff --git a/examples/z_image_fun/predict_turbo_t2i_control_2.1.py b/examples/z_image_fun/predict_turbo_t2i_control_2.1.py index 7b38bcf..dd833eb 100644 --- a/examples/z_image_fun/predict_turbo_t2i_control_2.1.py +++ b/examples/z_image_fun/predict_turbo_t2i_control_2.1.py @@ -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() \ No newline at end of file + save_results() diff --git a/scripts/minimax_h3/README_TRAIN_LORA.md b/scripts/minimax_h3/README_TRAIN_LORA.md index 1db3072..5615a04 100644 --- a/scripts/minimax_h3/README_TRAIN_LORA.md +++ b/scripts/minimax_h3/README_TRAIN_LORA.md @@ -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 diff --git a/scripts/minimax_h3/README_TRAIN_LORA_zh-CN.md b/scripts/minimax_h3/README_TRAIN_LORA_zh-CN.md index d72e1a0..7ddeb8b 100644 --- a/scripts/minimax_h3/README_TRAIN_LORA_zh-CN.md +++ b/scripts/minimax_h3/README_TRAIN_LORA_zh-CN.md @@ -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 diff --git a/scripts/minimax_h3/README_TRAIN_PDD_LORA.md b/scripts/minimax_h3/README_TRAIN_PDD_LORA.md new file mode 100644 index 0000000..3d0d0f4 --- /dev/null +++ b/scripts/minimax_h3/README_TRAIN_PDD_LORA.md @@ -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=` + `audio=` 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` diff --git a/scripts/minimax_h3/README_TRAIN_PDD_LORA_zh-CN.md b/scripts/minimax_h3/README_TRAIN_PDD_LORA_zh-CN.md new file mode 100644 index 0000000..893641c --- /dev/null +++ b/scripts/minimax_h3/README_TRAIN_PDD_LORA_zh-CN.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` diff --git a/scripts/minimax_h3/generate_prompt_cache.py b/scripts/minimax_h3/generate_prompt_cache.py new file mode 100644 index 0000000..47b344a --- /dev/null +++ b/scripts/minimax_h3/generate_prompt_cache.py @@ -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() diff --git a/scripts/minimax_h3/generate_prompt_cache.sh b/scripts/minimax_h3/generate_prompt_cache.sh new file mode 100644 index 0000000..49aaaea --- /dev/null +++ b/scripts/minimax_h3/generate_prompt_cache.sh @@ -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 diff --git a/scripts/minimax_h3/generate_ref2va_request_cache.py b/scripts/minimax_h3/generate_ref2va_request_cache.py new file mode 100644 index 0000000..1650627 --- /dev/null +++ b/scripts/minimax_h3/generate_ref2va_request_cache.py @@ -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=` and `audio=` — 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() diff --git a/scripts/minimax_h3/generate_ref2va_request_cache.sh b/scripts/minimax_h3/generate_ref2va_request_cache.sh new file mode 100644 index 0000000..069f59a --- /dev/null +++ b/scripts/minimax_h3/generate_ref2va_request_cache.sh @@ -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 diff --git a/scripts/minimax_h3/train_lora.py b/scripts/minimax_h3/train_lora.py index 0c33142..ebd6c28 100644 --- a/scripts/minimax_h3/train_lora.py +++ b/scripts/minimax_h3/train_lora.py @@ -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] diff --git a/scripts/minimax_h3/train_lora.sh b/scripts/minimax_h3/train_lora.sh index 7865215..4bd4023 100644 --- a/scripts/minimax_h3/train_lora.sh +++ b/scripts/minimax_h3/train_lora.sh @@ -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 " (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 \ diff --git a/scripts/minimax_h3/train_pdd_lora.py b/scripts/minimax_h3/train_pdd_lora.py new file mode 100644 index 0000000..c116b55 --- /dev/null +++ b/scripts/minimax_h3/train_pdd_lora.py @@ -0,0 +1,1786 @@ +# Modified from scripts/minimax_h3/train_lora.py for Parallel Decoding Distillation (PDD, arXiv 2607.26004). +# Scaffold (parameter set, resume, checkpointing) follows `scripts/minimax_h3/train_lora.py`. +# +# Data-free PDD LoRA of the packed-sequence transformer, covering `fl2va` (FL2VA / t2va layout) and `ref2va`. +# PDD trains a *parallel decoder*: the sampling interval is discretized into `N` intervals grouped into blocks of +# `L`, and 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 student is the teacher's own backbone with the two +# final heads repeated `N` times (`videox_fun/models/minimax_h3_pdd.py`); the loss is a plain MSE onto a +# Runge-Kutta estimate of the teacher's mean velocity — no VSD, no adversarial term, no JVP. +# +# Training is *data-free* (Algorithm 3 of the paper): no target video is ever read. Each rank carries one +# trajectory, rolls it forward with the student's own predictions, and resets to fresh noise and a fresh prompt +# when it reaches the end of the grid. The conditioning is either pre-encoded (`--enable_preprocess_training`, +# which keeps the 62 GB conditioner out of the run) or encoded on the fly from an annotation JSON; `ref2va` +# supports both, its on-the-fly route additionally VAE-encoding each request's reference latents. +# FSDP / DeepSpeed follow `scripts/minimax_h3/train_lora.py`: the plugin is read off Accelerator, ZeRO-3 +# skips `zero.Init` on the frozen VAEs, FSDP stage 3 / ZeRO-3 resume through `accelerator.save_state`, and the +# student / teacher forwards always go through the prepared wrapper so a sharded 33 B backbone all-gathers. +# +# MiniMax-H3's rectified-flow convention is the *opposite* of Wan's and is reproduced here from +# `MiniMaxH3Scheduler.scale_noise` / `MiniMaxH3Scheduler.step`, the single source of truth: +# * noising: `x_t = t * x0 + (1 - t) * noise` with `t = 1` clean, `t = 1 - sigma`, +# * the sigma grid is exponentially shifted, `sigma' = s * sigma / (1 + (s - 1) * sigma)`, `s = 12.0` for video and `3.0` for audio, +# * the transformer predicts a data-ward velocity, so the regression target is `x0 - noise`. +# +# The checkpoint is guidance-distilled: one forward per step, no unconditional branch. +# +# Usage: +# # fl2va off a pre-encoded prompt cache (`--enable_preprocess_training`): +# accelerate launch --mixed_precision no scripts/minimax_h3/train_pdd_lora.py \ +# --pretrained_model_name_or_path=models/Diffusion_Transformer/MiniMax-H3 \ +# --train_mode=fl2va --enable_preprocess_training \ +# --train_data_meta=datasets/minimax_h3_pdd_prompt_cache/outputs.json \ +# --output_dir=output_dir_minimax_h3_pdd_lora --gradient_checkpointing --resume_from_checkpoint=latest +# +# # ref2va loaded directly from a request annotation (no request cache; encodes on the fly): +# accelerate launch --mixed_precision no scripts/minimax_h3/train_pdd_lora.py \ +# --pretrained_model_name_or_path=models/Diffusion_Transformer/MiniMax-H3 \ +# --train_mode=ref2va \ +# --train_data_meta=datasets/X-Fun-Videos-Audios-Demo/metadata_add_width_height.json \ +# --output_dir=output_dir_minimax_h3_pdd_ref2va_lora --gradient_checkpointing --resume_from_checkpoint=latest + +import argparse +import gc +import json +import logging +import math +import os +import shutil +import sys +import time +import warnings +from types import SimpleNamespace + +import accelerate +import diffusers +import numpy as np +import torch +import torch.nn.functional as F +import transformers +from torch.utils.data import DataLoader, Dataset +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers.optimization import get_scheduler +from diffusers.training_utils import EMAModel +from diffusers.utils.torch_utils import is_compiled_module +from packaging import version +from tqdm.auto import tqdm +from transformers.utils import ContextManagers + +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.data import (BatchSampler, ImageVideoSafetensorsDataset, + RandomSampler, TextDataset) +from videox_fun.models import (AutoencoderKLMiniMaxH3, + AutoencoderKLMiniMaxH3Audio, + MiniMaxH3Transformer3DModel, Qwen2TokenizerFast, + Qwen3VLForConditionalGeneration, + Qwen3VLProcessor) +from videox_fun.models.minimax_h3_pdd import (attach_parallel_decoder, + pdd_teacher_mean_velocity, + set_parallel_plan) +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, + pdd_sampling_plan, + pdd_state_dict, + pdd_teacher_mode, + pdd_time_grid, + pdd_training_plan) +from videox_fun.pipeline import MiniMaxH3Pipeline +from videox_fun.pipeline.pipeline_minimax_h3 import ( + MINIMAX_H3_KEYFRAME_NOISE_AUG, align_num_frames, audio_latent_num_frames, + build_packed_sequence, build_ref2va_packed_sequence, build_row_timesteps, + check_ref2va_references, normalize_ref2va_references, + patchify_video_latents, video_latent_num_frames) +from videox_fun.utils import MiniMaxH3Scheduler +from videox_fun.utils.utils import save_videos_with_audio_grid + +# The on-the-fly route (without `--enable_preprocess_training`) encodes conditioning with the canonical MiniMax-H3 +# recipes rather than re-deriving them; these modules sit in the same directory, already on `sys.path`. +# `encode_prompt` builds the `fl2va` / `ref2va` presentation, `encode_reference_latents_for_training` the `ref2va` +# reference latents, and `load_requests` / `parse_reference` read a `ref2va` request annotation exactly as +# `generate_ref2va_request_cache.py` does. +from train_lora import encode_prompt, encode_reference_latents_for_training +from generate_ref2va_request_cache import load_requests, parse_reference + +# Silences diffusers' `randn_tensor` notice about CPU generators producing CUDA tensors (the tensor is created +# on CPU and moved to GPU; harmless, only a marginal speed note). +warnings.filterwarnings("ignore", message="The passed generator was created on") + + +def linear_decay(initial_value, final_value, total_steps, current_step): + if current_step >= total_steps: + return final_value + current_step = max(0, current_step) + step_size = (final_value - initial_value) / total_steps + return initial_value + step_size * current_step + + +# The `ref2va` request cache flattens its ragged reference structure into tensors (safetensors holds tensors only): +# the per-reference kind / has-audio pair becomes two int vectors and the per-reference latents become indexed +# tensors under a count. This maps the kind id back to the string `Ref2VATrajectory` and `log_validation` expect. +_REFERENCE_KINDS = ("image", "video", "audio") + + +def reconstruct_cache_entry(state_dict, train_mode): + r"""Rebuild one normalized conditioning entry from a flat safetensors `state_dict`. + + `ImageVideoSafetensorsDataset.__getitem__` returns the raw tensor dict written by `generate_prompt_cache.py` + (`fl2va`) or `generate_ref2va_request_cache.py` (`ref2va`); this restores the `{prompt_embeds, text_token_tags}` + (plus, for `ref2va`, `{reference_kinds, condition_latents, audio_condition_latents}`) shape the trajectories and + validation consume, so both are agnostic to how the cache was serialized. + """ + entry = { + "prompt_embeds": state_dict["prompt_embeds"], + "text_token_tags": state_dict["text_token_tags"], + } + if train_mode == "ref2va": + kind_ids = state_dict["reference_kind_ids"].tolist() + has_audio = state_dict["reference_has_audio"].tolist() + entry["reference_kinds"] = [ + (_REFERENCE_KINDS[int(kind)], bool(int(flag))) for kind, flag in zip(kind_ids, has_audio) + ] + num_condition = int(state_dict["num_condition_latents"]) + entry["condition_latents"] = [state_dict[f"condition_latents_{i}"] for i in range(num_condition)] + num_audio_condition = int(state_dict["num_audio_condition_latents"]) + entry["audio_condition_latents"] = [ + state_dict[f"audio_condition_latents_{i}"] for i in range(num_audio_condition) + ] + return entry + + +class _RequestDataset(Dataset): + r"""Map-style wrapper over the in-memory `ref2va` request list `load_requests` returns. + + The on-the-fly `ref2va` route has no pre-encoded safetensors to read, so the requests (each a + `{"prompt": str, "references": [str, ...]}` record) are carried in memory and flow through the same + accelerate-sharded DataLoader as the other conditioning routes. The reference media are parsed and encoded in the + conditioning iterator (main process), never in the collate / DataLoader workers. + """ + + def __init__(self, requests): + self.requests = list(requests) + + def __len__(self): + return len(self.requests) + + def __getitem__(self, index): + return self.requests[index] + + +class FL2VATrajectory: + r""" + One rank's carried trajectory of the data-free PDD algorithm on the FL2VA / t2va layout. + + The state is a partially denoised sample plus the grid index it sits at. A step reads it, rolls it forward by + `L_min` intervals with the student's own prediction, and the trajectory is thrown away and re-drawn from noise + (with a new prompt) once it reaches the end of the grid. + """ + + def __init__(self, geometry, patch_size, latent_channels, audio_channels, condition_iter, device): + self.geometry = geometry + self.patch_size = patch_size + self.latent_channels = latent_channels + self.audio_channels = audio_channels + self.condition_iter = condition_iter + self.device = device + self.index = None + + def reset(self): + num_latent_frames, latent_height, latent_width, num_audio_latents = self.geometry + # Draw the next conditioning entry from the (accelerate-sharded, cycling) dataloader iterator instead of + # picking at random from a whole in-memory cache: each rank now sees a distinct slice of the data. + cached = next(self.condition_iter) + self.prompt_embeds = cached["prompt_embeds"].to(self.device) + text_token_tags = cached["text_token_tags"] + if not torch.is_tensor(text_token_tags): + text_token_tags = torch.tensor(text_token_tags, dtype=torch.long) + else: + # `accelerator.prepare` moves the cached conditioning onto the GPU, but `build_packed_sequence` assembles + # the layout on CPU (its outputs are moved to `self.device` just below), so keep the tags on CPU to match. + text_token_tags = text_token_tags.cpu() + self.layout = build_packed_sequence( + text_token_tags, + num_latent_frames, + latent_height, + latent_width, + num_audio_latents, + self.patch_size, + ) + self.indices = { + name: getattr(self.layout, name).to(self.device) + for name in ("token_tags", "position_ids", "video_indices", "audio_indices", "text_indices") + } + rows_per_frame = (latent_height // self.patch_size[1]) * (latent_width // self.patch_size[2]) + video_patch_dim = self.latent_channels * math.prod(self.patch_size) + self.video = torch.randn( + num_latent_frames * rows_per_frame, video_patch_dim, device=self.device, dtype=torch.float32 + ) + self.audio = torch.randn( + num_audio_latents * 2, self.audio_channels, device=self.device, dtype=torch.float32 + ) + self.index = 0 + + def forward_kwargs(self, video_time, audio_time): + unique_timesteps, timestep_indices = build_row_timesteps( + self.layout, float(video_time), float(audio_time), float(video_time), 1.0 + ) + return dict( + encoder_hidden_states=self.prompt_embeds, + timestep=unique_timesteps.to(self.device), + timestep_indices=timestep_indices.to(self.device), + return_dict=False, + **self.indices, + ) + + def generated(self, video, audio): + return video, audio + + def with_generated(self, video_tail, audio_tail): + return video_tail, audio_tail + + +class Ref2VATrajectory: + r""" + One rank's carried trajectory of the data-free PDD algorithm, on a `ref2va` layout. + + The state is the two packed streams — each of them the request's fixed conditioning rows followed by the + partially denoised generated rows — plus the grid index they sit at. The conditioning rows are drawn once per + trajectory and never move; only the generated tail is rolled forward and supervised. + """ + + def __init__(self, geometry, patch_size, latent_channels, audio_channels, condition_iter, scheduler, device): + self.geometry = geometry + self.patch_size = patch_size + self.latent_channels = latent_channels + self.audio_channels = audio_channels + self.condition_iter = condition_iter + self.scheduler = scheduler + self.device = device + self.index = None + + def reset(self): + num_latent_frames, latent_height, latent_width, num_audio_latents = self.geometry + # Draw the next cached request from the (accelerate-sharded, cycling) dataloader iterator. + request = next(self.condition_iter) + self.prompt_embeds = request["prompt_embeds"].to(self.device) + references = [SimpleNamespace(kind=kind, has_audio=has_audio) for kind, has_audio in request["reference_kinds"]] + condition_latents = request["condition_latents"] + audio_condition_latents = request["audio_condition_latents"] + text_token_tags = request["text_token_tags"] + if not torch.is_tensor(text_token_tags): + text_token_tags = torch.tensor(text_token_tags, dtype=torch.long) + else: + # `accelerator.prepare` moves the cached conditioning onto the GPU, but `build_ref2va_packed_sequence` + # assembles the layout on CPU (its outputs are moved to `self.device` just below), so keep the tags there. + text_token_tags = text_token_tags.cpu() + + self.layout = build_ref2va_packed_sequence( + text_token_tags, + references, + condition_latents, + audio_condition_latents, + num_latent_frames, + latent_height, + latent_width, + num_audio_latents, + self.patch_size, + ) + self.indices = { + name: getattr(self.layout, name).to(self.device) + for name in ("token_tags", "position_ids", "video_indices", "audio_indices", "text_indices") + } + self.num_condition_video_rows = self.layout.num_condition_video_rows + self.num_condition_audio_rows = self.layout.num_condition_audio_rows + + condition_rows = [ + patchify_video_latents( + self.scheduler.scale_noise( + condition.to(self.device), + MINIMAX_H3_KEYFRAME_NOISE_AUG, + torch.randn(condition.shape, device=self.device, dtype=torch.float32), + ), + self.patch_size, + ) + for condition in condition_latents + ] + + rows_per_frame = (latent_height // self.patch_size[1]) * (latent_width // self.patch_size[2]) + video_patch_dim = self.latent_channels * math.prod(self.patch_size) + video = torch.randn( + num_latent_frames * rows_per_frame, video_patch_dim, device=self.device, dtype=torch.float32 + ) + audio = torch.randn(num_audio_latents * 2, self.audio_channels, device=self.device, dtype=torch.float32) + self.video = torch.cat(condition_rows + [video]) if condition_rows else video + self.audio = ( + torch.cat([rows.to(self.device) for rows in audio_condition_latents] + [audio]) + if audio_condition_latents + else audio + ) + self.index = 0 + + def forward_kwargs(self, video_time, audio_time): + unique_timesteps, timestep_indices = build_row_timesteps( + self.layout, + float(video_time), + float(audio_time), + max(float(video_time), MINIMAX_H3_KEYFRAME_NOISE_AUG), + 1.0, + ) + return dict( + encoder_hidden_states=self.prompt_embeds, + timestep=unique_timesteps.to(self.device), + timestep_indices=timestep_indices.to(self.device), + return_dict=False, + **self.indices, + ) + + def generated(self, video, audio): + return video[self.num_condition_video_rows :], audio[self.num_condition_audio_rows :] + + def with_generated(self, video_tail, audio_tail): + video = torch.cat([self.video[: self.num_condition_video_rows], video_tail]) + audio = torch.cat([self.audio[: self.num_condition_audio_rows], audio_tail]) + return video, audio + + +logger = get_logger(__name__, log_level="INFO") + + +def log_validation( + vae, audio_vae, transformer, scheduler, audio_scheduler, args, accelerator, val_cache, grids, global_step, +): + r""" + Render every validation cache entry with the student at `--validation_nfe` and save the mp4 (video + audio) + next to the run. Each entry is already-cached Qwen3-VL conditioning, so no text encoder is loaded here. + + PDD generation is an ordinary Euler loop over the *block boundaries* of the grid: those boundaries are exactly + the schedule `MiniMaxH3Scheduler` builds for `NFE` steps, so the released pipeline drives the student unchanged + and the only PDD-specific work is arming the heads before each step. + """ + sharded = ( + getattr(accelerator.state, "fsdp_plugin", None) is not None + or getattr(accelerator.state, "deepspeed_plugin", None) is not None + ) + if sharded: + # Under FSDP / DeepSpeed every forward all-gathers the sharded params, so it is a *collective* that all ranks + # must enter the same number of times. Splitting the entries by `index % num_processes` gives ranks different + # counts whenever `len(val_cache)` is not a multiple of `num_processes` (2 entries over 4 ranks leaves ranks 2 + # and 3 with nothing to render); a rank that returns early never joins the all-gather the others block on, and + # NCCL times out. So every rank walks the same `ceil(len / ranks)` rounds, cycling the entry list, and writes + # a file only for the entries it owns (`index < len`) — the extra cycles are collective ballast that keep the + # ranks in lockstep, not duplicate renders. + num_rounds = math.ceil(len(val_cache) / accelerator.num_processes) + assigned = [] + for round_index in range(num_rounds): + index = round_index * accelerator.num_processes + accelerator.process_index + assigned.append((index, val_cache[index % len(val_cache)], index < len(val_cache))) + else: + assigned = [ + (index, entry, True) + for index, entry in enumerate(val_cache) + if index % accelerator.num_processes == accelerator.process_index + ] + if not assigned: + return + + try: + with torch.no_grad(): + logger.info("Running validation... ") + _, _, video_steps, audio_steps = grids + num_steps = video_steps.shape[0] + student = accelerator.unwrap_model(transformer) + + pipeline = MiniMaxH3Pipeline( + vae=vae, + audio_vae=audio_vae, + text_encoder=None, + tokenizer=None, + processor=None, + transformer=transformer, + scheduler=scheduler, + audio_scheduler=audio_scheduler, + ) + # Avoid `.to()` on an FSDP / DeepSpeed wrapper: it rematerializes FlatParameters on every rank. + if not sharded: + pipeline = pipeline.to(accelerator.device) + block_size = num_steps // args.validation_nfe + + def arm(step_index): + start = step_index * block_size + set_parallel_plan( + student, + pdd_sampling_plan(video_steps, start, block_size).float(), + pdd_sampling_plan(audio_steps, start, block_size).float(), + ) + + def callback(pipe, step_index, timestep, callback_kwargs): + if step_index + 1 < args.validation_nfe: + arm(step_index + 1) + return {} + + vae.to(accelerator.device) + audio_vae.to(accelerator.device) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + + for index, entry, owner in assigned: + arm(0) + if args.seed is None: + generator = None + else: + prompt_seed = args.seed + index + generator = torch.Generator(device=accelerator.device).manual_seed(prompt_seed) + logger.info(f"Rank {accelerator.process_index} prompt {index} using seed: {prompt_seed}") + + call_kwargs = dict( + prompt=None, + prompt_embeds=entry["prompt_embeds"], + text_token_tags=entry["text_token_tags"], + height=args.video_sample_height, + width=args.video_sample_width, + num_frames=args.video_sample_n_frames, + num_inference_steps=args.validation_nfe, + generator=generator, + output_type="pt", + callback_on_step_end=callback, + ) + if args.train_mode == "ref2va": + call_kwargs.update( + normalized_references=[ + SimpleNamespace(kind=kind, has_audio=has_audio) for kind, has_audio in entry["reference_kinds"] + ], + condition_latents=entry["condition_latents"], + audio_condition_latents=entry["audio_condition_latents"], + ) + + output = pipeline(**call_kwargs) + if owner: + save_videos_with_audio_grid( + output.videos, + output.audio, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-prompt{index}-{args.train_mode}-nfe{args.validation_nfe}.mp4", + ), + fps=24, + audio_sample_rate=output.sampling_rate, + ) + del output + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu") + audio_vae.to(accelerator.device if not args.low_vram else "cpu") + except Exception as e: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu") + audio_vae.to(accelerator.device if not args.low_vram else "cpu") + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Parallel Decoding Distillation LoRA of MiniMax-H3 (FL2VA / Ref2VA, video + audio)." + ) + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--output_dir", + type=str, + default="output_dir_minimax_h3_pdd_lora", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument("--seed", type=int, default=43, help="A seed for reproducible training.") + parser.add_argument( + "--train_batch_size", type=int, default=1, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--gradient_checkpointing_save_on_cpu", + action="store_true", + help="Offload the activations saved for backward of the transformer blocks to CPU memory.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-5, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=0, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training." + ), + ) + parser.add_argument("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=0.0, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10 and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=50, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument("--save_state", action="store_true", help="Whether or not to save state.") + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--rank", + type=int, + default=64, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--network_alpha", + type=int, + default=64, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--target_name", + type=str, + default="to_q,to_k,to_v,to_out.0,ff.net.0.proj,ff.net.2,adaln_proj.linear", + help=("The module is trained in loras."), + ) + # MiniMax-H3 specific + parser.add_argument( + "--train_mode", + type=str, + default="fl2va", + choices=["fl2va", "ref2va"], + help="fl2va: FL2VA / t2va packed layout. ref2va: Ref2VA layout, which also consumes cached reference latents.", + ) + parser.add_argument( + "--video_loss_weight", + type=float, + default=0.5, + help="Weight of the video flow-matching loss in the joint video + audio loss.", + ) + parser.add_argument( + "--audio_loss_weight", + type=float, + default=0.5, + help="Weight of the audio flow-matching loss in the joint video + audio loss.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=124, + help="Number of frames (form 17*n+5).", + ) + parser.add_argument( + "--low_vram", + action="store_true", + help="Keep VAE and conditioner on CPU, move to GPU only while encoding.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="minimax_h3_pdd_lora", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + parser.add_argument( + "--validation_steps", + type=int, + default=50, + help="Run validation every X steps.", + ) + # PDD + parser.add_argument( + "--enable_preprocess_training", + action="store_true", + help=( + "Train on the pre-processed (cached) conditioning instead of encoding it on the fly. When set, read the " + "pre-encoded safetensors via `ImageVideoSafetensorsDataset(--train_data_meta=outputs.json, " + "data_root=--train_data_dir)` — the ~62 GB conditioner stays out of training. When unset, load a Qwen3-VL " + "conditioner in the run and encode the conditioning on the fly: `fl2va` reads prompts via " + "`TextDataset(--train_data_meta)`, and `ref2va` reads a request annotation via `load_requests` and also " + "VAE-encodes each request's reference latents." + ), + ) + parser.add_argument( + "--train_data_dir", + type=str, + default=None, + help="Optional root prepended to each `file_path` of `--train_data_meta` (used with " + "`--enable_preprocess_training`); leave empty when `outputs.json` already stores absolute paths.", + ) + parser.add_argument( + "--train_data_meta", + type=str, + default=None, + help="Annotation JSON of the training conditioning. With `--enable_preprocess_training`: the `outputs.json` " + "written by `generate_prompt_cache.py` / `generate_ref2va_request_cache.py`. Without it: `fl2va` reads a " + "list of `{\"text\": ...}` records (`TextDataset`), and `ref2va` reads a list of requests " + "(`{\"prompt\": ..., \"references\": [...]}`, or the demo record `load_requests` derives them from).", + ) + parser.add_argument( + "--val_data_meta", + type=str, + default=None, + help="Optional annotation JSON of the held-out conditioning used by validation, mirroring `--train_data_meta`: " + "with `--enable_preprocess_training` the `outputs.json` written by `generate_prompt_cache.py` / " + "`generate_ref2va_request_cache.py`; without it the on-the-fly annotation (`fl2va`: `{\"text\": ...}` " + "records; `ref2va`: the request list). Validation is skipped when it is left empty.", + ) + parser.add_argument( + "--transformer_subfolder", + type=str, + default=None, + help="Transformer subfolder. Default: `transformer_ref` for `--train_mode=ref2va`, else `transformer`.", + ) + parser.add_argument( + "--pdd_num_steps", + type=int, + default=32, + help="The grid size `N`. The paper uses 128 for video models with the midpoint solver, 256 with Euler.", + ) + parser.add_argument( + "--pdd_block_size", + type=int, + default=4, + help="`L_min`: the block the carried state advances by, so the student is trained for `N / L_min` NFE.", + ) + parser.add_argument( + "--pdd_max_block_size", + type=int, + default=None, + help="`L_max`: the widest block a loss target is drawn from. Defaults to `--pdd_block_size`.", + ) + parser.add_argument( + "--pdd_solver", + type=str, + default="midpoint", + choices=["euler", "midpoint"], + help="Runge-Kutta method the teacher's mean velocity is estimated with.", + ) + parser.add_argument( + "--pdd_num_targets", + type=int, + default=2, + help="How many intra-block indices `k` one student evaluation is supervised at.", + ) + parser.add_argument("--lora_learning_rate", type=float, default=1e-4, help="Learning rate of the low-rank updates.") + parser.add_argument( + "--use_ema", + action="store_true", + help="Keep an exponential moving average of the trainable set. Validation and `pdd_ema.safetensors` use it.", + ) + parser.add_argument("--ema_decay", type=float, default=0.99) + parser.add_argument("--abnormal_norm_clip_start", type=int, default=1000) + parser.add_argument("--initial_grad_norm_ratio", type=int, default=5) + parser.add_argument( + "--video_sample_size", + type=int, + default=1280, + help="Square canvas size (height = width) for validation and latent geometry; must be a multiple of 32.", + ) + 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.", + ) + parser.add_argument("--validation_nfe", type=int, default=8) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + if args.pdd_max_block_size is None: + args.pdd_max_block_size = args.pdd_block_size + 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 + return args + + +def gather_full_state_dict(model, accelerator): + r"""Consolidated `state_dict` of a (possibly sharded) model, gathered to rank-0 CPU. + + Under `--fsdp_state_dict_type=SHARDED_STATE_DICT`, `accelerator.get_state_dict(..., unwrap=True)` hands back + *sharded* tensors, so `pdd_state_dict`'s `.detach().cpu()` only materializes the params of the root wrap + unit (the token refiner) and silently drops every separately-wrapped child unit — the `MiniMaxH3TransformerBlock` + LoRA and both `PDDParallelHead`s — leaving a 13 MB stub instead of the full ~1.4 GB trainable set. Temporarily + switching the FSDP root to `FULL_STATE_DICT` (offloaded to CPU, rank-0 only) reconstructs every original-named + param across all wrap units. DeepSpeed ZeRO-3 and DDP already consolidate in `accelerator.get_state_dict`, so + they keep using it. Collective: call on every rank; only rank 0 receives a non-empty dict. + """ + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + # `accelerator.get_state_dict(..., unwrap=True)` unwraps the FSDP root and reads `.state_dict()` off the *sharded* + # original-param views, so it silently returns only the root wrap unit's tensors (the token refiner) and drops + # every separately-wrapped child unit. Detect the FSDP wrapper directly — not via `fsdp_plugin`, which is not + # reliably plumbed this far — and gather a FULL_STATE_DICT on the wrapped root instead. + if isinstance(model, FSDP): + from torch.distributed.fsdp import FullStateDictConfig, StateDictType + with FSDP.state_dict_type( + model, + StateDictType.FULL_STATE_DICT, + FullStateDictConfig(offload_to_cpu=True, rank0_only=True), + ): + state_dict = model.state_dict() + if accelerator.is_main_process: + logger.info( + "gather_full_state_dict[FSDP FULL_STATE_DICT]: %d tensors (blocks.*=%d, proj_out*=%d).", + len(state_dict), + sum("blocks." in key for key in state_dict), + sum("proj_out" in key for key in state_dict), + ) + return state_dict + state_dict = accelerator.get_state_dict(model, unwrap=True) + if accelerator.is_main_process: + logger.info( + "gather_full_state_dict[get_state_dict fallback]: %d tensors (type(model)=%s).", + len(state_dict) if state_dict else 0, + type(model).__name__, + ) + return state_dict + + +def save_pdd_weights(path, state_dict): + from safetensors.torch import save_file + save_file( + {name: tensor.detach().contiguous().cpu() for name, tensor in state_dict.items()}, + path, + metadata={"format": "pt"}, + ) + + +def dump_pdd_config(args, save_path): + r"""Write `pdd_config.json` with both this script's LoRA flags and the inference aliases `predict_t2v.py` reads.""" + config = dict(vars(args)) + config["lora_rank"] = args.rank + config["lora_alpha"] = args.network_alpha + config["lora_targets"] = args.target_name + with open(os.path.join(save_path, "pdd_config.json"), "w") as handle: + json.dump(config, handle, indent=1) + + +def save_resume_state(save_path, student, optimizer, lr_scheduler, ema, accelerator): + r"""Optimizer / scheduler / live `pdd.safetensors` / EMA shadow — the pieces `pdd_ema.safetensors` does not hold.""" + os.makedirs(save_path, exist_ok=True) + torch.save(optimizer.state_dict(), os.path.join(save_path, "optimizer.pt")) + torch.save(lr_scheduler.state_dict(), os.path.join(save_path, "scheduler.pt")) + save_pdd_weights(os.path.join(save_path, PDD_WEIGHTS_NAME), pdd_state_dict(student)) + if ema is not None: + torch.save(ema.state_dict(), os.path.join(save_path, "ema.pt")) + if getattr(accelerator, "scaler", None) is not None: + torch.save(accelerator.scaler.state_dict(), os.path.join(save_path, "scaler.pt")) + + +def load_resume_state(save_path, student, optimizer, lr_scheduler, ema, trainable_params, accelerator): + r"""Load the trainer state written by [`save_resume_state`]. Prefers legacy `pdd_live.safetensors` when present.""" + from safetensors.torch import load_file + + legacy_live = os.path.join(save_path, PDD_LEGACY_LIVE_WEIGHTS_NAME) + weights_path = legacy_live if os.path.isfile(legacy_live) else os.path.join(save_path, PDD_WEIGHTS_NAME) + state_dict = load_file(weights_path, device="cpu") + m, u = student.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + print(f"Loaded {len(state_dict)} PDD tensors from {weights_path}") + + device = accelerator.device + optimizer_file_pt = os.path.join(save_path, "optimizer.pt") + optimizer_file_bin = os.path.join(save_path, "optimizer.bin") + optimizer_file_to_load = None + if os.path.exists(optimizer_file_pt): + optimizer_file_to_load = optimizer_file_pt + elif os.path.exists(optimizer_file_bin): + optimizer_file_to_load = optimizer_file_bin + if optimizer_file_to_load: + try: + accelerator.print(f"Loading optimizer state from {optimizer_file_to_load}") + optimizer.load_state_dict(torch.load(optimizer_file_to_load, map_location=device)) + accelerator.print("Optimizer state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load optimizer state from {optimizer_file_to_load}: {e}") + + scheduler_file_pt = os.path.join(save_path, "scheduler.pt") + scheduler_file_bin = os.path.join(save_path, "scheduler.bin") + scheduler_file_to_load = None + if os.path.exists(scheduler_file_pt): + scheduler_file_to_load = scheduler_file_pt + elif os.path.exists(scheduler_file_bin): + scheduler_file_to_load = scheduler_file_bin + if scheduler_file_to_load: + try: + accelerator.print(f"Loading scheduler state from {scheduler_file_to_load}") + lr_scheduler.load_state_dict(torch.load(scheduler_file_to_load, map_location=device)) + accelerator.print("Scheduler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load scheduler state from {scheduler_file_to_load}: {e}") + + if getattr(accelerator, "scaler", None) is not None: + scaler_file = os.path.join(save_path, "scaler.pt") + if os.path.exists(scaler_file): + try: + accelerator.print(f"Loading GradScaler state from {scaler_file}") + accelerator.scaler.load_state_dict(torch.load(scaler_file, map_location=device)) + accelerator.print("GradScaler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load GradScaler state: {e}") + + if ema is None: + return + ema_path = os.path.join(save_path, "ema.pt") + if os.path.exists(ema_path): + try: + print(f"Loading EMA state from {ema_path}") + ema.load_state_dict(torch.load(ema_path, map_location="cpu")) + print("EMA state loaded successfully.") + return + except Exception as e: + print(f"Failed to load EMA state from {ema_path}: {e}") + ema.shadow_params = [parameter.detach().clone() for parameter in trainable_params] + print(f"No ema.pt under {save_path}; EMA is re-seeded from the loaded weights.") + + +def main(): + args = parse_args() + + if args.train_mode not in ("fl2va", "ref2va"): + raise ValueError(f"`train_mode` must be 'fl2va' or 'ref2va', got {args.train_mode!r}.") + aligned_frames = align_num_frames(int(args.video_sample_n_frames)) + if aligned_frames != int(args.video_sample_n_frames): + raise ValueError( + f"`video_sample_n_frames` has to be of the form 17 * n + 5 the video VAE encodes, got " + f"{args.video_sample_n_frames} (nearest is {aligned_frames})." + ) + if args.video_sample_height % 32 or args.video_sample_width % 32: + raise ValueError( + f"`video_sample_size` / `fix_sample_size` ({args.video_sample_height}x{args.video_sample_width}) " + "must be multiples of 32: the canvas is patched 2x2 into the transformer and its RoPE grid keys off that." + ) + if args.pdd_num_steps % args.pdd_block_size: + raise ValueError( + f"The grid size {args.pdd_num_steps} must be a multiple of the block size {args.pdd_block_size}: the " + "block starts of the data-free algorithm are the multiples of `L_min` and the last one has to be the end " + "of the grid." + ) + if args.pdd_num_steps % args.validation_nfe: + raise ValueError( + f"`--validation_nfe` {args.validation_nfe} must divide the grid size {args.pdd_num_steps}: generation " + "advances `N / NFE` intervals per evaluation." + ) + if not args.train_data_meta: + raise ValueError( + "`--train_data_meta` is required: the `outputs.json` of the cached conditioning (with " + "`--enable_preprocess_training`) or the on-the-fly annotation JSON without it (`fl2va`: `{\"text\": ...}` " + "records; `ref2va`: the request list `load_requests` reads)." + ) + if args.train_batch_size != 1: + raise ValueError("Data-free PDD carries one trajectory per rank and requires --train_batch_size=1.") + + logging_dir = os.path.join(args.output_dir, args.logging_dir) + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + ) + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None + if deepspeed_plugin is not None: + zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + args.use_deepspeed = True + if zero_stage == 3: + print("Auto set save_state to True because zero_stage == 3") + args.save_state = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + else: + fsdp_stage = 0 + print(f"Using FSDP stage: {fsdp_stage}") + args.use_fsdp = True + if fsdp_stage == 3: + print("Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed/FSDP is not enabled.") + + logging.basicConfig(format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", datefmt="%m/%d/%Y %H:%M:%S", level=logging.INFO) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + # Per-rank seeding: the ranks of one global batch have to roll out *different* noise, otherwise the trajectories + # of a step differ only in their prompt. The conditioning order is drawn by the dataloader's own seeded sampler. + if args.seed is not None: + set_seed(args.seed + accelerator.process_index) + print(f"Init seed {args.seed + accelerator.process_index}. Process_index is {accelerator.process_index}") + else: + print(f"Init without fixed seed. Process_index is {accelerator.process_index}") + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast non-trainable weights to half-precision. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + # PDD: the released checkpoint already pins `proj_out` / `audio_proj_out` in float32 + # (`_keep_in_fp32_modules`), so the parallel heads built from them are float32 master weights over a bfloat16 + # backbone. The model casts every input to its projection's dtype itself, so the run needs no autocast. + weight_dtype = torch.bfloat16 + + # ------------------------------------------------------------------ models + # `pretrained_model_name_or_path` may point at a converted diffusers layout or at an *original* MiniMax-H3 + # partition; every component's `from_pretrained` auto-detects the layout and stream-converts the original + # shards on the fly, so the caller never branches on the format itself. + transformer_subfolder = args.transformer_subfolder or ( + "transformer_ref" if args.train_mode == "ref2va" else "transformer" + ) + print(f"Loading transformer from subfolder `{transformer_subfolder}` (train_mode={args.train_mode}).") + transformer = MiniMaxH3Transformer3DModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder=transformer_subfolder, low_cpu_mem_usage=True, torch_dtype=weight_dtype, + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained`. So the two VAEs will not enjoy the parameter sharding across multiple gpus + # and only the transformer will get ZeRO sharded. The 62 GB conditioner is loaded only when conditioning is + # encoded on the fly (without `--enable_preprocess_training`); the cached route keeps it out of the run entirely. + uses_text_encoder = not args.enable_preprocess_training + tokenizer = processor = text_encoder = None + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + # The two VAEs stay float32 as released (the encode/decode recipe is float16 autocast over float32 + # weights), so they are loaded without `torch_dtype`; the mixed-precision loader mixin restores the + # pinned fp32 modules anyway. PDD validation is the only consumer. + vae = AutoencoderKLMiniMaxH3.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", low_cpu_mem_usage=True, + ) + audio_vae = AutoencoderKLMiniMaxH3Audio.from_pretrained( + args.pretrained_model_name_or_path, subfolder="audio_vae", low_cpu_mem_usage=True, + ) + if uses_text_encoder: + # On-the-fly conditioning (`fl2va` prompts, `ref2va` requests): the same Qwen3-VL components + # `train_lora.py` loads, so `train_lora.encode_prompt` runs unchanged. Mirrors `scripts/minimax_h3/train_lora.py`. + 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, + ) + text_encoder = text_encoder.eval() + scheduler = MiniMaxH3Scheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + audio_scheduler = MiniMaxH3Scheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="audio_scheduler") + + # Freeze everything; the LoRA modules and parallel heads created below are the only trainable parameters. + transformer.requires_grad_(False) + vae.requires_grad_(False) + audio_vae.requires_grad_(False) + if uses_text_encoder: + text_encoder.requires_grad_(False) + + # ------------------------------------------------------------------ LoRA + num_adapters = add_pdd_lora(transformer, args.target_name.split(","), args.rank, args.network_alpha) + attach_parallel_decoder(transformer, args.pdd_num_steps) + transformer.train() + # FSDP flattens one dtype per wrap unit. DDP keeps float32 LoRA master weights; under FSDP the adapters match + # the Linear they wrap (bf16) so each `MiniMaxH3TransformerBlock` is uniform. The parallel heads stay float32 + # and are wrapped as their own units. Frozen `_keep_in_fp32_modules` embeddings are ignored so they are not + # mixed into the bf16 root flatten (`--mixed_precision=no` does not install an FSDP MixedPrecision policy). + if fsdp_plugin is not None: + for module in transformer.modules(): + if isinstance(module, PDDLoRALinear): + dtype = module.base.weight.dtype + module.lora_down.data = module.lora_down.data.to(dtype) + module.lora_up.data = module.lora_up.data.to(dtype) + wrap_names = list(fsdp_plugin.transformer_cls_names_to_wrap or []) + if "PDDParallelHead" not in wrap_names: + wrap_names.append("PDDParallelHead") + fsdp_plugin.transformer_cls_names_to_wrap = wrap_names + ignored = [] + for name in ("proj_in", "audio_proj_in", "time_embedder", "rope"): + module = getattr(transformer, name, None) + if isinstance(module, torch.nn.Module): + # `sync_module_states=True` rejects CPU params on ignored modules; FSDP's `device_id` only + # moves the flattened units. + module.to(accelerator.device) + ignored.append(module) + fsdp_plugin.ignored_modules = ignored + logger.info( + "FSDP: LoRA adapters cast to the backbone dtype; wrap %s; ignored_modules=%s.", + wrap_names, + [module.__class__.__name__ for module in ignored], + ) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # ------------------------------------------------------------------ save / load hooks + # `accelerate` 0.16.0+ supports custom saving hooks. Under FSDP / ZeRO-3 the hook writes + # live `pdd.safetensors` from the gathered trainable tensors so DDP `--save_state` resume can reload + # the current step; popping `weights` on the DDP path keeps `save_state` from serializing the frozen + # backbone. The EMA inference export `pdd_ema.safetensors` is written after `ema.copy_to`. + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + if fsdp_stage != 0 or zero_stage == 3: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = gather_full_state_dict(models[-1], accelerator) + if accelerator.is_main_process and accelerate_state_dict is not None: + os.makedirs(output_dir, exist_ok=True) + save_pdd_weights( + os.path.join(output_dir, PDD_WEIGHTS_NAME), + pdd_state_dict(unwrap_model(models[-1]), accelerate_state_dict), + ) + dump_pdd_config(args, output_dir) + + def load_model_hook(models, input_dir): + return + + else: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = gather_full_state_dict(models[-1], accelerator) + if accelerator.is_main_process and accelerate_state_dict is not None: + os.makedirs(output_dir, exist_ok=True) + save_pdd_weights( + os.path.join(output_dir, PDD_WEIGHTS_NAME), + pdd_state_dict(unwrap_model(models[-1]), accelerate_state_dict), + ) + dump_pdd_config(args, output_dir) + if not args.use_deepspeed: + for _ in range(len(weights)): + weights.pop() + + def load_model_hook(models, input_dir): + return + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # see https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + lr_scale = args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + args.learning_rate = args.learning_rate * lr_scale + args.lora_learning_rate = args.lora_learning_rate * lr_scale + + head_params, lora_params = [], [] + for name, parameter in transformer.named_parameters(): + if not parameter.requires_grad: + continue + (head_params if "proj_out" in name else lora_params).append(parameter) + trainable_params = head_params + lora_params + logger.info( + f"LoRA created: {num_adapters} adapters, {sum(p.numel() for p in lora_params) / 1e6:.2f} M parameters; " + f"{len(head_params)} parallel head tensors, {sum(p.numel() for p in head_params) / 1e6:.2f} M parameters." + ) + + # ------------------------------------------------------------------ optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError("Please install bitsandbytes to use 8-bit Adam.") + optimizer_cls = bnb.optim.AdamW8bit + else: + optimizer_cls = torch.optim.AdamW + optimizer = optimizer_cls( + [ + {"params": head_params, "lr": args.learning_rate}, + {"params": lora_params, "lr": args.lora_learning_rate}, + ], + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # ------------------------------------------------------------------ data + # Data-free PDD never reads a target video: each rank carries one trajectory and needs only the conditioning — + # either pre-encoded safetensors (`--enable_preprocess_training`) or an annotation encoded on the fly (without it). + # All three routes go through a DataLoader so `accelerator.prepare` shards the entries across ranks, replacing the + # old whole-cache random pick. + if args.enable_preprocess_training: + train_dataset = ImageVideoSafetensorsDataset(args.train_data_meta, data_root=args.train_data_dir) + + def collate_fn(examples): + return reconstruct_cache_entry(examples[0], args.train_mode) + elif args.train_mode == "ref2va": + # On-the-fly `ref2va`: read the request annotation with `load_requests` (the same reader the cache generator + # uses) and encode each request's prompt + reference latents in the conditioning iterator below. + train_dataset = _RequestDataset(load_requests(args.train_data_meta)) + + def collate_fn(examples): + # `--train_batch_size=1`: one request per trajectory reset, so hand over the single record; the Qwen3-VL / + # VAE encode runs in the conditioning iterator below (main process), not in the collate / workers. + return {"prompt": examples[0]["prompt"], "references": list(examples[0]["references"])} + else: + train_dataset = TextDataset(args.train_data_meta) + + def collate_fn(examples): + # `--train_batch_size=1`: one conditioning entry per trajectory reset, so hand over the single record; the + # Qwen3-VL encode runs in the conditioning iterator below (main process), not in the collate / workers. + return {"text": examples[0]["text"]} + + batch_sampler_generator = torch.Generator().manual_seed(args.seed) + batch_sampler = BatchSampler( + RandomSampler(train_dataset, generator=batch_sampler_generator), + batch_size=args.train_batch_size, + drop_last=True, + ) + train_dataloader = DataLoader(train_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn) + + # The held-out validation conditioning mirrors the training route (a cache *or* on-the-fly prompts) instead of + # forcing a pre-processed cache, so it is built further down — once the conditioner is sharded / on-device — right + # before the trajectory setup. Validation is skipped when `--val_data_meta` is empty. + + # Scheduler and math around the number of training steps. One epoch is one pass through the conditioning set + # (batch size is 1). `--max_train_steps` overrides `--num_train_epochs`, matching `scripts/minimax_h3/train_lora.py`. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + transformer.gradient_checkpointing_save_on_cpu = args.gradient_checkpointing_save_on_cpu + transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer, optimizer, train_dataloader, lr_scheduler + ) + + # Shard the frozen text encoder *after* prepare (mirrors `train_lora.py`): the Qwen3-VL conditioner (~62 GB) is + # wrapped per decoder layer so the per-step unshard footprint stays small, and a post-prepare shard keeps it out of + # the trainable FSDP unit. Only the on-the-fly text route loads it (`uses_text_encoder`). + sharded_text_encoder = uses_text_encoder and (fsdp_stage != 0 or zero_stage != 0) + if sharded_text_encoder: + from videox_fun.dist import shard_model + text_encoder.model = shard_model( + text_encoder.model, + device_id=accelerator.device, + param_dtype=weight_dtype, + module_to_wrapper=list(text_encoder.model.language_model.layers), + ) + + device = accelerator.device + # The two VAEs stay float32 (mirrors the pipeline: float32 weights, float16 autocast only at the + # encode/decode call site), so they are moved without a dtype cast. + vae.to(device if not args.low_vram else "cpu") + audio_vae.to(device if not args.low_vram else "cpu") + if fsdp_stage == 0 and zero_stage == 0: + transformer.to(device) + # An FSDP/ZeRO-sharded text encoder is already on-device and dtype-pinned by `shard_model`; otherwise move it to + # the GPU, or keep it on CPU under `--low_vram` (the conditioning iterator moves it up only for each encode). + if uses_text_encoder and not sharded_text_encoder: + text_encoder.to(device if not args.low_vram else "cpu", dtype=weight_dtype) + + trainable_params = [parameter for parameter in transformer.parameters() if parameter.requires_grad] + ema = ( + EMAModel(trainable_params, decay=args.ema_decay, use_ema_warmup=False, foreach=True) + if args.use_ema + else None + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. One + # epoch is one pass through the conditioning set (batch size 1), keeping `--num_train_epochs` / `--max_train_steps` + # consistent with `train_lora.py`. + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + if accelerator.is_main_process: + master_dtypes = {parameter.dtype for parameter in transformer.parameters()} + num_local_params = sum(parameter.numel() for parameter in transformer.parameters()) + logger.info( + f"Master parameter dtype(s): {master_dtypes}, {num_local_params / 1e9:.2f} B parameters per rank " + f"over {accelerator.num_processes} process(es)." + ) + + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + tracker_config = {k: v for k, v in tracker_config.items() if not isinstance(v, list)} + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # ------------------------------------------------------------------ constants + # Read the transformer config through the unwrap so it works under FSDP (the prepared `transformer` is a + # sharded wrapper) as well as single-process. + student = unwrap_model(transformer) + patch_size = tuple(student.config.patch_size) + latent_channels = student.config.in_channels + audio_channels = student.config.audio_in_channels + geometry = ( + video_latent_num_frames(args.video_sample_n_frames), + args.video_sample_height // vae.spatial_compression_ratio, + args.video_sample_width // vae.spatial_compression_ratio, + audio_latent_num_frames(args.video_sample_n_frames), + ) + video_grid = pdd_time_grid(scheduler.shift, args.pdd_num_steps) + audio_grid = pdd_time_grid(audio_scheduler.shift, args.pdd_num_steps) + grids = (video_grid, audio_grid, video_grid.diff(), audio_grid.diff()) + logger.info( + f"Grid N={args.pdd_num_steps}, L_min={args.pdd_block_size}, L_max={args.pdd_max_block_size}: block starts at " + f"t = {[round(float(video_grid[i]), 4) for i in range(0, args.pdd_num_steps + 1, args.pdd_block_size)]}" + ) + + # On-the-fly `ref2va` conditioning, shared by the training iterator and validation: parse the request's + # references, then encode the prompt (Qwen3-VL) and the reference latents (the two VAEs) with the same recipes + # `generate_ref2va_request_cache.py` uses. Under `--low_vram` the two encodes are serialized — conditioner up then + # down, VAEs up then down — so the ~62 GB text encoder and a 124-frame reference video never share the GPU. + ref2va_num_frames = align_num_frames(int(args.video_sample_n_frames)) + ref2va_audio_sr = getattr(audio_vae.config, "sampling_rate", 32000) + + def encode_request_on_the_fly(request): + references = [parse_reference(entry) for entry in request["references"]] + references = check_ref2va_references(references) + references = normalize_ref2va_references(references, ref2va_num_frames, ref2va_audio_sr) + + if args.low_vram and not sharded_text_encoder: + text_encoder.to(device) + with torch.no_grad(): + prompt_embeds, text_token_tags = encode_prompt( + text_encoder, tokenizer, processor, request["prompt"], + references=references, device=device, dtype=weight_dtype, + ) + if args.low_vram and not sharded_text_encoder: + text_encoder.to("cpu") + torch.cuda.empty_cache() + + if args.low_vram: + vae.to(device) + audio_vae.to(device) + with torch.no_grad(): + condition_latents, audio_condition_latents = encode_reference_latents_for_training( + vae, audio_vae, references, patch_size, device, audio_latent_channels=audio_channels, + ) + if args.low_vram: + vae.to("cpu") + audio_vae.to("cpu") + torch.cuda.empty_cache() + + return { + "prompt_embeds": prompt_embeds, + "text_token_tags": text_token_tags, + "reference_kinds": [(reference.kind, bool(reference.has_audio)) for reference in references], + "condition_latents": condition_latents, + "audio_condition_latents": audio_condition_latents, + } + + # Validation conditioning, built now that the conditioner is sharded / on-device so it can mirror the training + # route: with `--enable_preprocess_training` a `generate_*_cache.py` cache (`outputs.json` + safetensors); without + # it, the on-the-fly annotation encoded here (`fl2va` prompts via `TextDataset`, `ref2va` requests via + # `encode_request_on_the_fly`). Every rank builds the full list (exactly like the cache route); `log_validation` + # shards it at render time. + val_cache = [] + if args.val_data_meta: + if args.enable_preprocess_training: + val_dataset = ImageVideoSafetensorsDataset(args.val_data_meta, data_root=args.train_data_dir) + val_cache = [reconstruct_cache_entry(val_dataset[i], args.train_mode) for i in range(len(val_dataset))] + elif args.train_mode == "ref2va": + val_dataset = _RequestDataset(load_requests(args.val_data_meta)) + val_cache = [encode_request_on_the_fly(val_dataset[i]) for i in range(len(val_dataset))] + else: + val_dataset = TextDataset(args.val_data_meta) + if args.low_vram and not sharded_text_encoder: + text_encoder.to(device) + with torch.no_grad(): + for i in range(len(val_dataset)): + prompt_embeds, text_token_tags = encode_prompt( + text_encoder, tokenizer, processor, val_dataset[i]["text"], device=device, dtype=weight_dtype, + ) + val_cache.append({"prompt_embeds": prompt_embeds, "text_token_tags": text_token_tags}) + if args.low_vram and not sharded_text_encoder: + text_encoder.to("cpu") + torch.cuda.empty_cache() + + # The conditioning iterator turns the (accelerate-sharded, cycling) dataloader into the normalized entries the + # trajectories pull on `reset()`. The pre-processed route already yields `{prompt_embeds, text_token_tags}` (plus + # the ref2va reference tensors); the on-the-fly route encodes each entry here in the main process — `fl2va` prompts + # via `encode_prompt`, `ref2va` requests via `encode_request_on_the_fly` — moving the conditioner / VAEs up only + # for the encode under `--low_vram`. + def conditioning_iterator(): + while True: + for batch in train_dataloader: + if not uses_text_encoder: + yield batch + elif args.train_mode == "ref2va": + yield encode_request_on_the_fly(batch) + else: + if args.low_vram and not sharded_text_encoder: + text_encoder.to(device) + with torch.no_grad(): + prompt_embeds, text_token_tags = encode_prompt( + text_encoder, tokenizer, processor, batch["text"], device=device, dtype=weight_dtype, + ) + if args.low_vram and not sharded_text_encoder: + text_encoder.to("cpu") + torch.cuda.empty_cache() + yield {"prompt_embeds": prompt_embeds, "text_token_tags": text_token_tags} + + condition_iter = conditioning_iterator() + + if args.train_mode == "ref2va": + trajectory = Ref2VATrajectory( + geometry, patch_size, latent_channels, audio_channels, condition_iter, scheduler, device, + ) + else: + trajectory = FL2VATrajectory( + geometry, patch_size, latent_channels, audio_channels, condition_iter, device, + ) + target_seed = (args.seed if args.seed is not None else 0) + 1000 + accelerator.process_index + target_rng = np.random.default_rng(np.random.PCG64(target_seed)) + + # ------------------------------------------------------------------ train loop + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + logger.info(f" Video / audio loss weights = {args.video_loss_weight} / {args.audio_loss_weight}") + + global_step = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) if os.path.isdir(args.output_dir) else [] + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + if args.resume_from_checkpoint != "latest" and os.path.isdir(args.resume_from_checkpoint): + checkpoint_folder_path = args.resume_from_checkpoint + else: + checkpoint_folder_path = os.path.join(args.output_dir, path) + if zero_stage != 3 and not args.use_fsdp: + load_resume_state( + checkpoint_folder_path, student, optimizer, lr_scheduler, ema, trainable_params, accelerator + ) + else: + accelerator.load_state(checkpoint_folder_path) + accelerator.print("accelerator.load_state() completed for FSDP / ZeRO stage 3.") + if ema is not None: + ema.shadow_params = [parameter.detach().clone() for parameter in trainable_params] + print(f"EMA is re-seeded from the loaded FSDP / ZeRO weights under {checkpoint_folder_path}.") + print(f"Resumed training from {checkpoint_folder_path} at step {global_step}.") + else: + initial_global_step = 0 + + if ema is not None: + ema.to(device) + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + disable=not accelerator.is_local_main_process, + ) + + train_loss = 0.0 + train_video_loss = 0.0 + train_audio_loss = 0.0 + step_started = time.time() + + while global_step < args.max_train_steps: + with accelerator.accumulate(transformer): + if trajectory.index is None or trajectory.index >= args.pdd_num_steps: + trajectory.reset() + start = trajectory.index + + # Sample the intra-block indices the loss is evaluated at, `k ~ U{n, ..., min(n + L_max, N) - 1}`, without + # replacement so that several targets always supervise several distinct heads. + reach = min(start + args.pdd_max_block_size, args.pdd_num_steps) + targets = sorted( + target_rng.choice( + np.arange(start, reach), size=min(args.pdd_num_targets, reach - start), replace=False + ).tolist() + ) + + # One student evaluation yields, per target, the displacement to `X_k` and the velocity `u_k` the loss + # regresses, plus the `L_min` advance of the carried state (the paper's layer fusion, §3.1). + set_parallel_plan( + student, + pdd_training_plan(grids[2], start, targets, args.pdd_block_size).float(), + pdd_training_plan(grids[3], start, targets, args.pdd_block_size).float(), + ) + video_output, audio_output = transformer( + hidden_states=trajectory.video[None], + audio_hidden_states=trajectory.audio[None], + **trajectory.forward_kwargs(video_grid[start], audio_grid[start]), + ) + # The heads run over every row and the modality rows are selected afterwards, so a ref2va output still + # carries the conditioning rows in front. Only the generated tail is rolled forward and supervised. + video_output = video_output[0].unflatten(-1, (-1, latent_channels * math.prod(patch_size))) + audio_output = audio_output[0].unflatten(-1, (-1, audio_channels)) + video_output, audio_output = trajectory.generated(video_output, audio_output) + state_video_tail, state_audio_tail = trajectory.generated(trajectory.video, trajectory.audio) + + video_loss = video_output.new_zeros(()) + audio_loss = audio_output.new_zeros(()) + for position, target in enumerate(targets): + # The teacher's mean velocity is estimated on the student's own intra-block state (on-policy), and + # the state is a constant of the loss (eq. 11's stop-gradient). Conditioning rows are put back in + # front of the generated tail on ref2va; fl2va has no conditioning rows. + state_video, state_audio = trajectory.with_generated( + state_video_tail + video_output[:, 2 * position].detach(), + state_audio_tail + audio_output[:, 2 * position].detach(), + ) + with pdd_teacher_mode(student), torch.no_grad(): + target_video, target_audio = pdd_teacher_mean_velocity( + transformer, trajectory.forward_kwargs, state_video, state_audio, target, grids, args.pdd_solver + ) + target_video, target_audio = trajectory.generated(target_video, target_audio) + video_loss = video_loss + F.mse_loss(video_output[:, 2 * position + 1].float(), target_video) + audio_loss = audio_loss + F.mse_loss(audio_output[:, 2 * position + 1].float(), target_audio) + video_loss = video_loss / len(targets) + audio_loss = audio_loss / len(targets) + loss = args.video_loss_weight * video_loss + args.audio_loss_weight * audio_loss + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.detach()[None]).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + train_video_loss += ( + accelerator.gather(video_loss.detach()[None]).mean().item() / args.gradient_accumulation_steps + ) + train_audio_loss += ( + accelerator.gather(audio_loss.detach()[None]).mean().item() / args.gradient_accumulation_steps + ) + + # Backpropagate + accelerator.backward(loss) + if accelerator.sync_gradients: + max_grad_norm = linear_decay( + args.max_grad_norm * args.initial_grad_norm_ratio, + args.max_grad_norm, + args.abnormal_norm_clip_start, + global_step, + ) + accelerator.clip_grad_norm_(trainable_params, max_grad_norm) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad(set_to_none=True) + + trajectory.video, trajectory.audio = trajectory.with_generated( + state_video_tail + video_output[:, -1].detach(), + state_audio_tail + audio_output[:, -1].detach(), + ) + trajectory.index = start + args.pdd_block_size + del video_output, audio_output + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + if ema is not None: + ema.step(trainable_params) + progress_bar.update(1) + global_step += 1 + accelerator.log( + { + "train_loss": train_loss, + "video_loss": train_video_loss, + "audio_loss": train_audio_loss, + "grid_index": start, + "lr_heads": lr_scheduler.get_last_lr()[0], + "lr_lora": lr_scheduler.get_last_lr()[-1], + "step_seconds": time.time() - step_started, + }, + step=global_step, + ) + train_loss = 0.0 + train_video_loss = 0.0 + train_audio_loss = 0.0 + step_started = time.time() + + if global_step % args.checkpointing_steps == 0: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + if args.use_deepspeed or args.use_fsdp or args.save_state: + accelerator.save_state(save_path) + else: + save_resume_state(save_path, student, optimizer, lr_scheduler, ema, accelerator) + dump_pdd_config(args, save_path) + + if ema is not None: + ema.store(trainable_params) + ema.copy_to(trainable_params) + checkpoint_dir = os.path.join(args.output_dir, f"checkpoint-{global_step}") + if args.use_deepspeed or args.use_fsdp: + state_dict = gather_full_state_dict(transformer, accelerator) + if accelerator.is_main_process and state_dict is not None: + save_pdd_weights( + os.path.join(checkpoint_dir, PDD_EMA_WEIGHTS_NAME), + pdd_state_dict(unwrap_model(transformer), state_dict), + ) + dump_pdd_config(args, checkpoint_dir) + elif accelerator.is_main_process: + save_pdd_weights( + os.path.join(checkpoint_dir, PDD_EMA_WEIGHTS_NAME), + pdd_state_dict(unwrap_model(transformer)), + ) + ema.restore(trainable_params) + if accelerator.is_main_process: + logger.info(f"Saved state to {os.path.join(args.output_dir, f'checkpoint-{global_step}')}") + accelerator.wait_for_everyone() + + if global_step % args.validation_steps == 0 and val_cache: + if ema is not None: + ema.store(trainable_params) + ema.copy_to(trainable_params) + accelerator.wait_for_everyone() + log_validation( + vae, audio_vae, transformer, scheduler, audio_scheduler, args, accelerator, + val_cache, grids, global_step, + ) + accelerator.wait_for_everyone() + if ema is not None: + ema.restore(trainable_params) + step_started = time.time() + + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + # Create the pipeline using the trained modules and save it. + accelerator.wait_for_everyone() + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + if args.use_deepspeed or args.use_fsdp or args.save_state: + accelerator.save_state(save_path) + else: + save_resume_state(save_path, student, optimizer, lr_scheduler, ema, accelerator) + dump_pdd_config(args, save_path) + if ema is not None: + ema.copy_to(trainable_params) + if args.use_deepspeed or args.use_fsdp: + state_dict = gather_full_state_dict(transformer, accelerator) + if accelerator.is_main_process and state_dict is not None: + save_pdd_weights( + os.path.join(save_path, PDD_EMA_WEIGHTS_NAME), + pdd_state_dict(unwrap_model(transformer), state_dict), + ) + dump_pdd_config(args, save_path) + elif accelerator.is_main_process: + save_pdd_weights(os.path.join(save_path, PDD_EMA_WEIGHTS_NAME), pdd_state_dict(unwrap_model(transformer))) + if accelerator.is_main_process: + logger.info(f"Saved state to {save_path}") + accelerator.end_training() + + +if __name__ == "__main__": + main() diff --git a/scripts/minimax_h3/train_pdd_lora.sh b/scripts/minimax_h3/train_pdd_lora.sh new file mode 100644 index 0000000..d5e2116 --- /dev/null +++ b/scripts/minimax_h3/train_pdd_lora.sh @@ -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 diff --git a/videox_fun/__init__.py b/videox_fun/__init__.py index 2eac33a..ded8763 100644 --- a/videox_fun/__init__.py +++ b/videox_fun/__init__.py @@ -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() diff --git a/videox_fun/models/minimax_h3_pdd.py b/videox_fun/models/minimax_h3_pdd.py new file mode 100644 index 0000000..4a1e3df --- /dev/null +++ b/videox_fun/models/minimax_h3_pdd.py @@ -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 diff --git a/videox_fun/models/wan_transformer3d_self_forcing.py b/videox_fun/models/wan_transformer3d_self_forcing.py index 38a2919..c91bd8c 100644 --- a/videox_fun/models/wan_transformer3d_self_forcing.py +++ b/videox_fun/models/wan_transformer3d_self_forcing.py @@ -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) diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index d2cd082..fcb8280 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -52,4 +52,10 @@ WanFunPipeline = WanPipeline WanI2VPipeline = WanFunInpaintPipeline Wan2_2FunPipeline = Wan2_2Pipeline -Wan2_2I2VPipeline = Wan2_2FunInpaintPipeline \ No newline at end of file +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() diff --git a/videox_fun/pipeline/pipeline_minimax_h3.py b/videox_fun/pipeline/pipeline_minimax_h3.py index e4864a5..9791087 100644 --- a/videox_fun/pipeline/pipeline_minimax_h3.py +++ b/videox_fun/pipeline/pipeline_minimax_h3.py @@ -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( diff --git a/videox_fun/pipeline/pipeline_wan_self_forcing.py b/videox_fun/pipeline/pipeline_wan_self_forcing.py index 9f8d941..2a7c9b2 100644 --- a/videox_fun/pipeline/pipeline_wan_self_forcing.py +++ b/videox_fun/pipeline/pipeline_wan_self_forcing.py @@ -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 diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py index 432938f..73921f2 100755 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -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) diff --git a/videox_fun/utils/lora_utils_pdd.py b/videox_fun/utils/lora_utils_pdd.py new file mode 100644 index 0000000..651d267 --- /dev/null +++ b/videox_fun/utils/lora_utils_pdd.py @@ -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}.") diff --git a/videox_fun/utils/perf_metrics.py b/videox_fun/utils/perf_metrics.py new file mode 100644 index 0000000..5ef9dd9 --- /dev/null +++ b/videox_fun/utils/perf_metrics.py @@ -0,0 +1,1863 @@ +"""Environment-gated inference and training metrics for `videox_fun`. + +This module is instrumentation only: it times, it never computes. Nothing here touches a tensor that feeds a +model, so a run with metrics on and a run with metrics off produce bit-identical outputs. + +It is off unless `VIDEOX_PERF` is set, and "off" is literal: [`install`] and [`install_training`] return on their +first line, nothing is wrapped and no hook is registered, so a default run pays nothing at all. `import videox_fun` +does not even load this file with the variable unset, that bootstrap being guarded by it; importing +`videox_fun.utils` or `videox_fun.pipeline` does load it either way, which costs one bytecode load and no new +dependencies -- everything above is the standard library plus the `torch` those packages already import. + +The inference half wires itself up in two layers: + +* [`install`] wraps the `__call__` of every pipeline class exported by `videox_fun.pipeline`, which is what marks + a *request* boundary -- where the counters are reset and where the one and only `cuda.synchronize` of the whole + scheme happens, at a point the caller was about to synchronize anyway to save its output. +* At the start of every request, [`_attach_hooks`] walks `pipe.components` and hooks whichever module components + are not hooked yet. Attaching *lazily* rather than at construction is what makes this work under FSDP and + sequence parallel: the entry scripts reassign `pipeline.transformer = shard_fn(pipeline.transformer)` after + building the pipeline, so a hook attached at construction would sit on a discarded object, while one attached per + request lands on whatever the pipeline actually runs. + +Per-step timings therefore come out of the natural granularity of the denoising loop -- one transformer forward +per step, two under CFG -- without the loop itself being touched. + +The training half is [`install_training`], and it hangs off `accelerate` rather than off this repo's scripts, +because `Accelerator` is the one thing all of them have in common; see its docstring for the phase model. It adds +no synchronize at all. + +Environment variables: + VIDEOX_PERF: `1` for a per-request (inference) or per-window (training) summary plus an exit summary, `2` to + also dump every step. Unset or `0` disables the module entirely. + VIDEOX_PERF_JSON: path to append one JSON object per request / per window to. Suffixed with `.rank{N}` under + multi-GPU so ranks never share a file. Each object carries a `kind` telling the two apart. + VIDEOX_PERF_WARMUP: number of leading requests / global steps to exclude from the exit summary (they are still + logged). + VIDEOX_PERF_RANKS: `0` (default) to log from rank 0 only, `all` to log from every rank. + VIDEOX_PERF_PEAK_TFLOPS: per-device hardware bf16 peak to compute MFU against, overriding the built-in device + table. Under multi-GPU the MFU is taken against this times the world size. + VIDEOX_PERF_DIT_PARAMS: exact transformer parameter count, overriding the FSDP-aware inference below. + VIDEOX_PERF_FLOPS_ATTN: `0` to price only the linear layers, leaving out the quadratic core-attention term the + FLOPs figures include by default. Worth reaching for on the causal models, whose masked attention costs about + half of what the term charges them. + VIDEOX_PERF_EVERY: training only, default 50. Global steps per aggregated log line; a line per step is + unreadable over the tens of thousands of steps a real run takes. + VIDEOX_PERF_TOTAL_STEPS: training only. Total planned steps, enabling a remaining-time estimate. Not guessed + when unset -- `max_train_steps` lives in the entry script's argparse and cannot be read from here. + VIDEOX_PERF_FLOPS_COEF: training only. Overrides the automatically chosen FLOPs multiplier; see + [`flops_coef`]. Setting it collapses the reported MFU and HFU onto each other. + +Note that when enabled this module calls `torch.cuda.reset_peak_memory_stats()` on the compute device once per +request (inference) or once per window (training), so any caller reading the peak memory counters itself sees them +scoped to the same interval. +""" + +import atexit +import collections +import contextlib +import functools +import json +import logging +import os +import statistics +import sys +import threading +import time +from typing import Any, Dict, List, Optional, Tuple + +import torch + +# Dense bf16 tensor-core peaks in TFLOPS, matched as substrings against `torch.cuda.get_device_name()`. Only +# devices with a published dense figure are listed; an unmatched device reports no MFU rather than a made-up one, +# and `VIDEOX_PERF_PEAK_TFLOPS` covers anything missing. +_PEAK_TFLOPS_BF16 = { + "A100": 312.0, + "A800": 312.0, + "H100": 989.0, + "H800": 989.0, + "H200": 989.0, + "H20": 148.0, + "L40S": 362.0, + "L20": 119.5, + "RTX 4090": 165.2, +} + +logger = logging.getLogger("videox_fun.perf") + +_INSTALLED = False +_STATE: Optional["_MetricsState"] = None +# The request currently being measured, held per thread. The component hooks are permanent once attached, so they +# key off this to know whether they are inside a measured request -- and a `None` here is what makes them a no-op. +# It is thread-local because a server can have two requests in flight at once, and a shared slot would have them +# writing into one record: the hooks always fire on the thread that called the forward, so per-thread state keeps +# concurrent requests from corrupting each other's counts. +_TLS = threading.local() + + +def _current() -> Optional["_Request"]: + return getattr(_TLS, "request", None) + + +def _env_int(name: str, default: int) -> int: + raw = os.environ.get(name) + if raw is None or raw.strip() == "": + return default + try: + return int(raw) + except ValueError: + return default + + +def _env_float(name: str) -> Optional[float]: + raw = os.environ.get(name) + if raw is None or raw.strip() == "": + return None + try: + return float(raw) + except ValueError: + return None + + +class _StageStat: + """Timings of one component across one request, as raw marker pairs resolved only at the end.""" + + __slots__ = ("pairs", "tokens", "batch") + + def __init__(self): + self.pairs: List[Tuple[Any, Any]] = [] + self.tokens: Optional[int] = None + self.batch: int = 1 + + def ready(self) -> bool: + """Whether every event pair here has completed, so [`elapsed_ms`] can read all of them. + + Only the closing event of each pair is tested: the opening one was recorded earlier on the same stream, and + a stream retires events in the order they were recorded, so a completed end implies a completed start. + """ + for _, (_, event_end) in self.pairs: + if event_end is not None and not event_end.query(): + return False + return True + + def elapsed_ms(self) -> List[float]: + out = [] + for (host_start, event_start), (host_end, event_end) in self.pairs: + host_ms = (host_end - host_start) * 1000.0 + # `elapsed_time` on an event the device has not reached yet raises, so the readiness of the pair is a + # precondition for reading it, not an optimisation. The inference path gets there via the synchronize + # at the request boundary; the training path never synchronizes and instead defers settling a step + # until its events have retired, falling back to the host clock for the rare pair that never does. + if ( + event_start is not None + and event_end is not None + and event_start.query() + and event_end.query() + ): + # The larger of the two is the one that saw the work; see [`_mark`]. + out.append(max(host_ms, event_start.elapsed_time(event_end))) + else: + out.append(host_ms) + return out + + +class _Request: + __slots__ = ("index", "warmup", "t0", "device", "stages", "open_marks") + + def __init__(self, index: int, warmup: bool, device: Optional[torch.device]): + self.index = index + self.warmup = warmup + # Resolved once per request and carried here so that every CUDA call of the request -- the event records in + # the hooks, the memory counters, the closing synchronize -- names the same device explicitly. + self.device = device + self.t0 = time.perf_counter() + self.stages: Dict[str, _StageStat] = {} + # Start markers of forwards that have not returned yet, keyed by stage; a list so that a re-entrant + # component (a VAE decoding chunk by chunk inside an outer call) nests instead of losing its pair. + self.open_marks: Dict[str, List[Any]] = {} + + def stage(self, name: str) -> _StageStat: + stat = self.stages.get(name) + if stat is None: + stat = _StageStat() + self.stages[name] = stat + return stat + + +class _MetricsState: + """Process-wide configuration and the accumulated history the exit summary is built from.""" + + def __init__(self, level: int): + self.level = level + self._json_base = os.environ.get("VIDEOX_PERF_JSON") or None + self.warmup = _env_int("VIDEOX_PERF_WARMUP", 0) + self.log_all_ranks = (os.environ.get("VIDEOX_PERF_RANKS", "0").strip().lower() == "all") + self.peak_tflops_override = _env_float("VIDEOX_PERF_PEAK_TFLOPS") + self.dit_params_override = _env_int("VIDEOX_PERF_DIT_PARAMS", 0) or None + # On by default: leaving core attention out understates a video step by two thirds and, worse, by a factor + # that moves with the sequence length. Off is for the causal models, where charging the full square is an + # overcount of nearly two, and for anyone who wants the old linear-only bound back. + self.attn_flops = _env_int("VIDEOX_PERF_FLOPS_ATTN", 1) != 0 + self.num_requests = 0 + self.history: List[Dict[str, Any]] = [] + self._rank = 0 + self._world_size = 1 + self._topology_final = False + self._reported = False + self._configure_logger() + atexit.register(self.report_summary) + + def _resolve_topology(self) -> None: + """Settle rank and world size, preferring the process group over the environment once it exists. + + [`install`] runs while `videox_fun.pipeline` is being imported, which in the entry scripts is *before* + `set_multi_gpus_devices` brings the process group up, so at that point the launcher's environment is all + there is to go on. Re-resolving until the group is initialized means the numbers the FSDP parameter + recovery and the per-rank json paths depend on are the real ones by the time a request runs. + """ + if self._topology_final: + return + if torch.distributed.is_available() and torch.distributed.is_initialized(): + self._rank = torch.distributed.get_rank() + self._world_size = torch.distributed.get_world_size() + self._topology_final = True + else: + self._rank = _env_int("RANK", 0) + self._world_size = _env_int("WORLD_SIZE", 1) + + @property + def rank(self) -> int: + self._resolve_topology() + return self._rank + + @property + def world_size(self) -> int: + self._resolve_topology() + return self._world_size + + @property + def json_path(self) -> Optional[str]: + if self._json_base is None: + return None + # Every rank keeps its own file: ranks share a filesystem, and appending from eight processes to one path + # interleaves partial lines. + return f"{self._json_base}.rank{self.rank}" if self.world_size > 1 else self._json_base + + def _configure_logger(self): + # Own the handler outright instead of relying on the entry script's `basicConfig`: the metrics have to show + # up whether the process was started by python, torchrun, accelerate or a server framework, and exactly + # once. + if not any(getattr(h, "_videox_perf", False) for h in logger.handlers): + handler = logging.StreamHandler(sys.stderr) + handler.setFormatter(logging.Formatter("%(message)s")) + handler._videox_perf = True + logger.addHandler(handler) + logger.setLevel(logging.INFO) + logger.propagate = False + + @property + def tag(self) -> str: + return f"[Perf][rank{self.rank}]" + + def should_log(self) -> bool: + return self.log_all_ranks or self.rank == 0 + + def peak_tflops(self, device: Optional[torch.device] = None) -> Optional[float]: + if self.peak_tflops_override is not None: + return self.peak_tflops_override + if device is None: + # `get_device_name` with no argument reads device 0, and on a host with no context up it creates one; + # see [`_perf_device`]. + return None + name = torch.cuda.get_device_name(device) + for key, value in _PEAK_TFLOPS_BF16.items(): + if key in name: + return value + return None + + def report_summary(self): + # Guarded the way [`_TrainState.finish`] is: reporting explicitly should not then be reported again by the + # `atexit` hook over the very same requests. + if self._reported: + return + measured = [record for record in self.history if not record["warmup"]] + if not measured or not self.should_log(): + return + self._reported = True + e2e = sorted(record["e2e_s"] for record in measured) + steps = [ms for record in measured for ms in record.get("_step_ms", [])] + parts = [ + f"{len(measured)} reqs", + f"e2e p50 {_percentile(e2e, 50):.1f}s p95 {_percentile(e2e, 95):.1f}s", + ] + skipped = len(self.history) - len(measured) + if skipped: + parts[0] = f"{len(measured)} reqs ({skipped} warmup excluded)" + if steps: + parts.append(f"transformer mean {statistics.fmean(steps):.0f}ms/step") + total = sum(record["e2e_s"] for record in measured) + if total > 0: + parts.append(f"{len(measured) / total * 3600.0:.1f} reqs/hour") + logger.info(f"{self.tag} === {' | '.join(parts)} ===") + + +def _percentile(values_sorted: List[float], q: float) -> float: + if not values_sorted: + return float("nan") + if len(values_sorted) == 1: + return values_sorted[0] + pos = (len(values_sorted) - 1) * q / 100.0 + low = int(pos) + high = min(low + 1, len(values_sorted) - 1) + return values_sorted[low] + (values_sorted[high] - values_sorted[low]) * (pos - low) + + +def _perf_device(pipe) -> Optional[torch.device]: + """The CUDA device a pipeline computes on, or `None` to keep this module off the CUDA APIs altogether. + + Read off the pipeline's own modules rather than from `torch.cuda.current_device()`, because under multi-GPU the + two are different devices here. `set_multi_gpus_devices` hands the entry script a `cuda:{local_rank}` to move + the weights to but never calls `torch.cuda.set_device`, and the `set_device` inside xfuser's + `init_distributed_environment` is guarded by `if not torch.distributed.is_initialized()`, which is already + false by the time it runs -- the process group was brought up on the line before. So every rank keeps + `current_device() == 0` while its model runs on `cuda:{local_rank}`, and a default-argument `synchronize`, + `Event.record`, `reset_peak_memory_stats` or `max_memory_allocated` would every one of them aim at device 0: + the synchronize would wait on an idle device instead of the busy one, the events would time an empty stream, + the memory counters would report a device the pipeline never wrote to -- and each would in passing bring a + context up on GPU 0 from all eight processes, taking memory from the rank that does have a model there. + + `None` covers a cpu-only run and a pipeline whose weights are offloaded to the host: no CUDA call is made at + all, so metrics never create a context that the same run without them would not have. + """ + if not (torch.cuda.is_available() and torch.cuda.is_initialized()): + return None + for name in ("transformer", "transformer_2", "unet", "vae"): + module = getattr(pipe, name, None) + if isinstance(module, torch.nn.Module): + for param in module.parameters(): + if param.device.type == "cuda": + return param.device + try: + # Under cpu offload the parameters rest on the host and only accelerate's hooks know where they run. + device = pipe._execution_device + except Exception: + return None + return device if getattr(device, "type", None) == "cuda" else None + + +def _mark(device: Optional[torch.device]) -> Tuple[float, Any]: + """A timing marker: a host timestamp, plus a recorded CUDA event when the request runs on a device. + + Both are taken because neither alone is right for every component. A device forward returns to the host as soon + as its kernels are *queued*, so the host clock on its own reports submission time; conversely a pipeline can + hold a module that genuinely runs on the host -- a text encoder left on cpu, a vae under offload -- and a CUDA + event pair around such a forward measures nothing, the two events executing back to back with no work between + them. Resolving the pair with a `max` picks whichever of the two actually saw the work, per forward, with no + need to guess where a module lives. + + The event is only *recorded* here, never waited on -- `cudaEventRecord` is a few microseconds against a step of + several hundred milliseconds, and the elapsed time is read at the end of the request. + """ + if device is not None: + event = torch.cuda.Event(enable_timing=True) + # Named stream rather than the default one: it pins the event to the pipeline's device (see + # [`_perf_device`]) and follows any `with torch.cuda.stream(...)` the pipeline put the forward on. + event.record(torch.cuda.current_stream(device)) + return time.perf_counter(), event + return time.perf_counter(), None + + +def _infer_tokens(module: torch.nn.Module, args, kwargs) -> Optional[Tuple[int, int]]: + """`(tokens per sample, batch)` of one transformer forward, or `None` when the inputs do not say. + + This only scales the analytic FLOPs, so it has to reflect what the blocks actually run over. Three input + conventions appear across the families in this repo: + + * The Wan models are handed an explicit `seq_len`, the padded length their blocks run on, which is + authoritative over anything derived from the latent's shape. + * A packed sequence arrives already tokenized as `(B, tokens, D)`, as in MiniMax-H3. + * A video latent arrives as `(B, C, F, H, W)` and only becomes tokens through the model's own patch size. + + The batch is returned separately rather than folded in, because it is a multiplier on the cost but not on the + sequence length the log reports -- classifier-free guidance doubles the former and leaves the latter alone. + Anything unrecognized reports nothing, so the FLOPs line is dropped rather than guessed. + """ + hidden = kwargs.get("hidden_states") + if hidden is None: + hidden = kwargs.get("x") # the Wan models name it `x` + if hidden is None and args: + hidden = args[0] + if not torch.is_tensor(hidden) or hidden.ndim < 3: + return None + batch = int(hidden.shape[0]) + + seq_len = kwargs.get("seq_len") + if isinstance(seq_len, int) and seq_len > 0: + return seq_len, batch + if hidden.ndim == 3: + return int(hidden.shape[1]), batch + if hidden.ndim == 5: + patch = getattr(getattr(module, "config", None), "patch_size", None) + if isinstance(patch, int): + patch = (1, patch, patch) + if not isinstance(patch, (tuple, list)) or len(patch) != 3: + return None + _, _, frames, height, width = hidden.shape + return int((frames // patch[0]) * (height // patch[1]) * (width // patch[2])), batch + return None + + +def _is_sharded(module: torch.nn.Module) -> bool: + for sub in module.modules(): + if hasattr(sub, "_fsdp_wrapped_module") or type(sub).__name__ == "FullyShardedDataParallel": + return True + for param in module.parameters(): + if type(param).__name__ in ("FlatParameter", "DTensor"): + return True + return False + + +def _count_params(module: torch.nn.Module, world_size: int) -> int: + """Total parameter count of a component, undoing FSDP sharding. + + Under FSDP each rank only holds its shard, so the local `numel` is the full count divided by the world size; + multiplying it back is right for the flat even sharding `shard_model` sets up. `VIDEOX_PERF_DIT_PARAMS` is the + escape hatch for any layout where it is not. + """ + total = sum(param.numel() for param in module.parameters()) + if world_size > 1 and _is_sharded(module): + total *= world_size + return total + + +# The naming conventions the attention projections go by across the families here: `q`/`k`/`v` in the Wan models, +# `to_q`/`to_k`/`to_v` in the diffusers ones, `q_proj`/`k_proj`/`v_proj` in the ones that came from a transformers +# tower. +_ATTN_PROJECTIONS = (("q", "to_q", "q_proj"), ("k", "to_k", "k_proj"), ("v", "to_v", "v_proj")) + +# The fused form, where one matrix produces all three: `nn.Linear(dim, dim * 3)` in the LongCat and HiDream blocks. +# The query is the first of three equal shares of its output, which is why an output width that is not a multiple of +# three is not read as one of these. +_ATTN_FUSED = ("qkv", "to_qkv", "qkv_proj") + +# Attention stacks that refine the *text* embedding before the blocks run. They are attention by structure, but they +# never see the latent sequence -- they run over a few hundred text rows -- so charging them the latent square is a +# pure overcount. Named in full rather than matched on "refiner": Z-Image has a `noise_refiner` beside its +# `context_refiner`, and that one does run over the latents, so a substring test would drop real work. +_ATTN_TEXT_TOWERS = ("token_refiner", "context_refiner") + +# `QK^T` and `AV`, each a multiply-accumulate. The two products are what "core attention" means here, as against the +# q, k, v and output projections, whose cost is linear in the sequence and already inside the parameter count. +_ATTN_CORE_FACTOR = 4.0 + + +def _query_width(module: torch.nn.Module) -> Optional[int]: + """The width of one module's query projection if it is an attention, or `None` if it is not. + + Duck-typed on `out_features` rather than `isinstance(nn.Linear)`, so that a wrapped projection -- a peft + `lora.Linear`, or anything else holding a base layer -- is still recognized. Requiring all three of q, k and v + keeps the looser test from matching a module that merely happens to own an attribute named `v`. + + The separate q, k and v are looked for first and the fused matrix only after, because a model may have both: the + HiDream tower fuses its own attention and holds an unfused `q_proj` elsewhere, and the unfused reading is the + one that needs no assumption about how the output is divided. + """ + found = [] + for candidates in _ATTN_PROJECTIONS: + for attr in candidates: + child = getattr(module, attr, None) + if isinstance(getattr(child, "out_features", None), int): + found.append(child) + break + if len(found) == len(_ATTN_PROJECTIONS): + return found[0].out_features + for attr in _ATTN_FUSED: + width = getattr(getattr(module, attr, None), "out_features", None) + if isinstance(width, int) and width % 3 == 0: + return width // 3 + return None + + +def _attn_widths(module: torch.nn.Module) -> Dict[str, int]: + """The query-projection width of a model's attention, split by what its keys and values run over. + + The analytic `2 * params * tokens` cost prices the linear layers and nothing else, and at video lengths the + attention it leaves out is the larger half of the work: core attention grows with the square of the sequence + while the linear part grows linearly, so it is a third of a step at 8k tokens and two thirds at 28k. That is + also why a bound that omits it cannot be used to compare two runs at different lengths, which is what it was + previously documented as being good for. This reads the widths needed to price it back, the way Megatron keeps + its `self_attn_core_term` separate from its per-token terms rather than folding attention into a parameter + count. + + Widths come from `in_features` / `out_features` rather than from any config, because those are set in + `nn.Linear.__init__` and survive what a config does not: the families here name their dimensions a dozen + different ways, and under FSDP the weights have been flattened away while these remain. They are also the + *logical* widths, so unlike a parameter count they need no unsharding -- sharding in this repo splits weights + across ranks (FSDP) or the sequence across ranks (ulysses, ring), and neither narrows a projection. + + Which modules count, and what their query width is, is [`_query_width`]; both the separate and the fused forms + of the projection are recognized. The query width is what both `QK^T` and `AV` are wide in, including under GQA: + the fewer key heads are repeated up to the query heads before the product, so the key width does not enter. + + The text-refiner stacks are skipped -- `token_refiner` in MiniMax-H3 and HunyuanVideo, `context_refiner` in + Z-Image. They are attention over a few hundred text rows, not over the latents, so pricing them by the latent + sequence overstates them by the ratio of the two lengths squared. They are matched by name in full because + Z-Image also has a `noise_refiner`, which does run over the latents and must keep counting. + + Keys over the latent sequence and keys over the text are counted apart because only the first is quadratic. A + module is taken to be cross-attention when its qualified name says so -- `cross_attn` in the Wan blocks, `attn2` + in the diffusers convention -- and self-attention otherwise. Guessing self-attention is the conservative + direction for the joint attention that MMDiT runs over text and latents concatenated: its true length is a + little above the latent count, so charging it the latent count alone understates rather than inflates. + + Nothing here can see whether the attention is masked, and that is the one direction in which this overcounts. + The models in this repo are bidirectional over the latent sequence and so pay the full square, but the causal + variants -- `wan_flex_causal_attn`, the self-forcing transformers -- compute about half of it, and Megatron + halves its own core term for exactly that reason. Reach for `VIDEOX_PERF_FLOPS_ATTN=0` on those runs to fall + back to the linear-only bound rather than read a figure that is too high by nearly a factor of two. + """ + widths = {"self": 0, "cross": 0, "modules": 0} + for name, sub in module.named_modules(): + lowered = name.lower() + if any(tower in lowered for tower in _ATTN_TEXT_TOWERS): + continue + query = _query_width(sub) + if query is None: + continue + cross = "cross" in lowered or lowered.rsplit(".", 1)[-1] == "attn2" + widths["cross" if cross else "self"] += query + widths["modules"] += 1 + return widths + + +def _attn_flops(widths: Optional[Dict[str, int]], tokens: int, text_tokens: int) -> float: + """Core attention FLOPs of one forward over `tokens` latent tokens, or zero when the widths were unreadable. + + Zero rather than a guess: a model whose attention this could not find falls back to the linear-only bound it + always reported, which is wrong in a known direction by a known mechanism, and the caller says which of the two + it used. + + The cross-attention term needs the text length, which is a property of the run and not of the module, so it is + dropped when unknown. It is worth much less than the self term -- a few percent of the core at video lengths, + being linear in the sequence where the other is quadratic -- so dropping it moves the total very little. + """ + if not widths or not widths["self"]: + return 0.0 + total = _ATTN_CORE_FACTOR * float(tokens) * float(tokens) * widths["self"] + if text_tokens: + total += _ATTN_CORE_FACTOR * float(tokens) * float(text_tokens) * widths["cross"] + return total + + +def _text_tokens(module: torch.nn.Module) -> int: + """The padded text length a model's cross-attention runs against, or 0 when it does not advertise one. + + Only the Wan family states it (`text_len`, 512), which is the family whose cross-attention is a separate module + and therefore the family where the distinction changes anything. + """ + value = getattr(module, "text_len", None) + return int(value) if isinstance(value, int) and value > 0 else 0 + + +def _attach_hooks(pipe) -> None: + """Hook the module components of a pipeline, picking up any that were swapped since the last request. + + Run at the start of every request rather than once, because a component can be replaced after the pipeline was + built: `pipeline.transformer = shard_fn(pipeline.transformer)` in the entry scripts is exactly that, and a + wrapper applied between two requests would otherwise never be hooked and would drop its stage from the log + without saying so. The guards below sit on the modules, so a repeat visit costs one `getattr` per component. + """ + state = _STATE + dit_params: Dict[str, int] = {} + dit_attn: Dict[str, Tuple[Optional[Dict[str, int]], int]] = {} + + try: + components = dict(pipe.components) + except Exception: # a pipeline may expose a component it cannot resolve; metrics must never break a run + components = {} + + for name, component in components.items(): + if not isinstance(component, torch.nn.Module): + continue # tokenizer, processor, scheduler + + # Counted once per module and cached on it: walking the parameters of a 14B transformer on every request + # would be wasted work, and the count cannot change once the sharding is in place. + if name.startswith("transformer"): + params = getattr(component, "_videox_perf_params", None) + if params is None: + params = state.dit_params_override or _count_params(component, state.world_size) + component._videox_perf_params = params + if params: + dit_params[name] = params + # Cached on the module beside the parameter count and for the same reason: reading the projection widths + # walks every submodule, and they cannot change once the model is built. + attn = getattr(component, "_videox_perf_attn", None) + if attn is None: + widths = _attn_widths(component) if state.attn_flops else None + widths = widths if widths and widths["self"] else None + attn = (widths, _text_tokens(component) if widths else 0) + component._videox_perf_attn = attn + dit_attn[name] = attn + + # The guard is on the module rather than on the pipeline, so that a module shared by two pipelines -- a + # base and a refiner over the same vae -- is hooked once. Hooking it twice would append two marker pairs + # per forward and double every count it appears in. + if getattr(component, "_videox_perf_hooked", False): + continue + component._videox_perf_hooked = True + + if "vae" in name: + # Pipelines call `vae.encode` / `vae.decode`; `vae.forward` never runs, so hooking it would report + # nothing. Wrapping the two bound methods on the instance also splits the two directions apart. The + # streaming pair shares those same two stages, being the same work done in chunks: `AutoencoderKLWan` is + # the one vae that has them, and `pipeline_wan_self_forcing` reaches it only through `decode_stream`, + # whose time would otherwise land in `other`. Neither streaming method calls the plain one, so sharing + # the stage cannot double-count. + for method_name, stage in ( + ("encode", f"{name}_enc"), + ("decode", f"{name}_dec"), + ("encode_stream", f"{name}_enc"), + ("decode_stream", f"{name}_dec"), + ): + method = getattr(component, method_name, None) + if callable(method): + setattr(component, method_name, _wrap_timed_method(method, stage)) + continue + + component.register_forward_pre_hook(_make_pre_hook(name), with_kwargs=True) + component.register_forward_hook(_make_post_hook(name), with_kwargs=True) + + pipe._videox_perf_dit_params = dit_params + pipe._videox_perf_dit_attn = dit_attn + + +def _make_pre_hook(name: str): + def pre_hook(module, args, kwargs): + request = _current() + if request is None: + return + stat = request.stage(name) + if name.startswith("transformer") and stat.tokens is None: + shape = _infer_tokens(module, args, kwargs) + if shape is not None: + stat.tokens, stat.batch = shape + request.open_marks.setdefault(name, []).append(_mark(request.device)) + + return pre_hook + + +def _make_post_hook(name: str): + def post_hook(module, args, kwargs, output): + request = _current() + if request is None: + return + marks = request.open_marks.get(name) + if marks: + request.stage(name).pairs.append((marks.pop(), _mark(request.device))) + + return post_hook + + +def _wrap_timed_method(method, stage: str): + @functools.wraps(method) + def wrapper(*args, **kwargs): + request = _current() + if request is None: + return method(*args, **kwargs) + start = _mark(request.device) + try: + return method(*args, **kwargs) + finally: + request.stage(stage).pairs.append((start, _mark(request.device))) + + wrapper._videox_perf = True + return wrapper + + +def _begin_request(pipe) -> Optional[_Request]: + if _current() is not None: + # A pipeline invoked from inside another (e.g. a latent upsampler): let the outer request own the timings + # instead of overwriting them. + return None + try: + _attach_hooks(pipe) + state = _STATE + device = _perf_device(pipe) + request = _Request(state.num_requests, warmup=state.num_requests < state.warmup, device=device) + state.num_requests += 1 + if device is not None: + torch.cuda.reset_peak_memory_stats(device) + except Exception as error: # never let measurement take down a generation run + logger.warning(f"[Perf] failed to start metrics: {error!r}") + return None + _TLS.request = request + return request + + +def _end_request(pipe, request: Optional[_Request]) -> None: + if request is None: + return + _TLS.request = None + try: + if request.device is not None: + # Before the wall clock is read, not after. A pipeline `__call__` returns once the last kernel is + # *queued*, so an unsynchronized reading would time how long the host took to submit the work, not how + # long the device took to do it -- and would come out below the device time it is meant to bound. + # This is the only synchronize of the scheme. On the usual path it is free, the caller being about to + # read the generated frames on the host anyway; with `output_type="latent"` it does add a wait the same + # run without metrics would not have, which is the one place this module is not entirely free. It is a + # local wait on one device, never a collective, so it cannot deadlock or skew a rank group. + torch.cuda.synchronize(request.device) + e2e = time.perf_counter() - request.t0 + _settle_request(pipe, request, e2e) + except Exception as error: # never let measurement take down a generation run + logger.warning(f"{_STATE.tag} failed to report metrics: {error!r}") + + +def _settle_request(pipe, request: _Request, e2e: float) -> None: + state = _STATE + stages: Dict[str, Dict[str, float]] = {} + step_ms: List[float] = [] + for name, stat in request.stages.items(): + values = stat.elapsed_ms() + if not values: + continue + stages[name] = { + "count": len(values), + "total_ms": sum(values), + "mean_ms": statistics.fmean(values), + "std_ms": statistics.stdev(values) if len(values) > 1 else 0.0, + "min_ms": min(values), + "max_ms": max(values), + } + if name.startswith("transformer"): + step_ms.extend(values) + + dit = _dit_throughput(pipe, request, stages) + record: Dict[str, Any] = { + "kind": "request", + "ts": time.time(), + "rank": state.rank, + "world_size": state.world_size, + "req": request.index, + "warmup": request.warmup, + "e2e_s": e2e, + "stages": stages, + "dit": dit, + "_step_ms": step_ms, + } + if request.device is not None: + record["device"] = str(request.device) + record["peak_alloc_bytes"] = torch.cuda.max_memory_allocated(request.device) + record["peak_reserved_bytes"] = torch.cuda.max_memory_reserved(request.device) + state.history.append(record) + + if state.should_log(): + logger.info(_format_record(state, record)) + if state.level >= 2: + for name in sorted(stages): + detail = ", ".join(f"{ms:.1f}" for ms in request.stages[name].elapsed_ms()) + logger.info(f"{state.tag} {name} steps(ms): {detail}") + if state.json_path: + _append_json(state, record) + + +def _dit_throughput(pipe, request: _Request, stages: Dict[str, Dict[str, float]]) -> Optional[Dict[str, Any]]: + """Achieved TFLOPS and MFU from an analytic cost of a transformer forward. + + The cost is `2 * params * tokens` for the linear layers plus a separate quadratic term for core attention, read + from the model's projection widths by [`_attn_widths`] and dropped when those cannot be read -- in which case the + figure is the linear-only lower bound this reported for every model before, understating a long video sequence by + roughly two thirds. There is no MFU / HFU split here: inference has no backward and nothing to recompute, so the + model and hardware costs coincide. + + The rate is a *job* rate, not a per-device one. Multi-GPU inference here splits the sequence across ranks with + ulysses / ring attention, so the length one rank is handed covers the whole job while that rank computes only + its shard of it -- which is also why the MFU below is taken against the aggregate peak of every device in the + group. Comparing a single-device peak against a rate that eight devices produced would report an impossible + figure well above 100%. + """ + params_by_stage = getattr(pipe, "_videox_perf_dit_params", {}) or {} + attn_by_stage = getattr(pipe, "_videox_perf_dit_attn", {}) or {} + total_flops = 0.0 + attn_flops = 0.0 + total_ms = 0.0 + total_count = 0 + tokens_seen: List[int] = [] + batches_seen: List[int] = [] + for name, params in params_by_stage.items(): + stat = request.stages.get(name) + summary = stages.get(name) + if stat is None or summary is None or not stat.tokens: + continue + widths, text_tokens = attn_by_stage.get(name, (None, 0)) + per_forward = 2.0 * params * stat.tokens + _attn_flops(widths, stat.tokens, text_tokens) + total_flops += per_forward * stat.batch * summary["count"] + attn_flops += _attn_flops(widths, stat.tokens, text_tokens) * stat.batch * summary["count"] + total_ms += summary["total_ms"] + total_count += summary["count"] + tokens_seen.append(stat.tokens) + batches_seen.append(stat.batch) + if not total_count or total_ms <= 0: + return None + achieved = total_flops / (total_ms / 1000.0) / 1e12 + devices = _STATE.world_size + per_device_peak = _STATE.peak_tflops(request.device) + peak = per_device_peak * devices if per_device_peak else None + return { + "params": sum(params_by_stage.values()), + "tokens": max(tokens_seen), + "batch": max(batches_seen), + "devices": devices, + "flops_per_fwd": total_flops / total_count, + "attn_share": (attn_flops / total_flops) if total_flops else 0.0, + "tflops": achieved, + "peak_tflops": peak, + "mfu": (achieved / peak) if peak else None, + } + + +def _format_record(state: _MetricsState, record: Dict[str, Any]) -> str: + parts = [f"req#{record['req']}{' warmup' if record['warmup'] else ''}", f"e2e {record['e2e_s']:.2f}s"] + for name in sorted(record["stages"]): + summary = record["stages"][name] + seconds = summary["total_ms"] / 1000.0 + if summary["count"] == 1: + parts.append(f"{name} {seconds:.2f}s(1x)") + else: + parts.append( + f"{name} {seconds:.2f}s({summary['count']}x, {summary['mean_ms']:.0f}+-{summary['std_ms']:.0f}ms, " + f"min {summary['min_ms']:.0f} max {summary['max_ms']:.0f})" + ) + if "peak_alloc_bytes" in record: + gib = 1024.0 ** 3 + parts.append( + f"peak_alloc {record['peak_alloc_bytes'] / gib:.1f}GiB " + f"peak_reserved {record['peak_reserved_bytes'] / gib:.1f}GiB" + ) + dit = record.get("dit") + if dit: + segment = ( + f"DiT {dit['flops_per_fwd']:.2e} FLOPs/fwd (attn {dit['attn_share'] * 100:.0f}%) " + f"-> {dit['tflops']:.1f} TFLOPS" + ) + if dit["mfu"] is not None: + over = f" over {dit['devices']} GPUs" if dit["devices"] > 1 else "" + segment += f" (MFU {dit['mfu'] * 100:.1f}% @{dit['peak_tflops']:.0f}{over})" + else: + segment += " (MFU n/a)" + parts.append(segment) + return f"{state.tag} " + " | ".join(parts) + + +def _append_json(state: _MetricsState, record: Dict[str, Any]) -> None: + payload = {key: value for key, value in record.items() if not key.startswith("_")} + directory = os.path.dirname(state.json_path) + if directory: + os.makedirs(directory, exist_ok=True) + with open(state.json_path, "a") as handle: + handle.write(json.dumps(payload) + "\n") + + +def _wrap_call(fn): + @functools.wraps(fn) + def wrapper(self, *args, **kwargs): + request = _begin_request(self) + try: + return fn(self, *args, **kwargs) + finally: + # `finally` so a failed request still releases the current-request slot; leaving it set would silently + # attribute the next request's forwards to a dead record. + _end_request(self, request) + + wrapper._videox_perf = True + return wrapper + + +def _instrument_class(cls) -> bool: + """Wrap the `__call__` a pipeline class defines itself. Returns whether anything was wrapped.""" + call = cls.__dict__.get("__call__") + if call is None or getattr(call, "_videox_perf", False): + # No own `__call__` means it inherits one, which its owner has already had wrapped -- wrapping here too + # would time the same request twice. + return False + cls.__call__ = _wrap_call(call) + return True + + +def install(level: Optional[int] = None) -> bool: + """Wrap every `videox_fun` pipeline class so that requests are measured. No-op unless `VIDEOX_PERF` is set. + + Called at the end of `videox_fun.pipeline.__init__`, by which point every pipeline class is present in that + module's namespace. + """ + global _INSTALLED, _STATE + if _INSTALLED: + return True + if level is None: + level = _env_int("VIDEOX_PERF", 0) + if level <= 0: + return False + + # Marked installed before the import below, which can re-enter this function: called directly rather than from + # the tail of `videox_fun.pipeline.__init__`, the import runs that module for the first time and its tail calls + # `install` again. Without the flag set, that inner call would build a second state with a second atexit + # summary and then have this one overwrite it, splitting the history across two objects. + _INSTALLED = True + # Adopted rather than replaced: in a training job both halves install, `videox_fun.__init__` running + # `install_training` first and the entry script's `import videox_fun.pipeline` reaching here second -- and all + # 110 training scripts do import it. Overwriting would leave `_TrainState` holding a state that is no longer the + # module's, with two atexit summaries registered over it. The level is raised to the more verbose of the two + # requests instead of being dropped, so that an explicit `install(level=2)` over a level-1 state still gets the + # per-step detail its own log line promises. + if _STATE is None: + _STATE = _MetricsState(level) + else: + _STATE.level = max(_STATE.level, level) + + from diffusers import DiffusionPipeline + + from .. import pipeline as pipeline_package + + seen = set() + wrapped = 0 + for obj in list(vars(pipeline_package).values()): + # Dedupe on the class object, not its name: the package ends with aliases (`WanFunPipeline = WanPipeline`) + # that would otherwise get the same class wrapped twice. + if not isinstance(obj, type) or id(obj) in seen: + continue + seen.add(id(obj)) + if not issubclass(obj, DiffusionPipeline): + continue # the namespace also holds schedulers, models and helpers + if _instrument_class(obj): + wrapped += 1 + + if _STATE.should_log(): + logger.info(f"{_STATE.tag} inference metrics enabled (level {level}) on {wrapped} pipeline classes") + return True + + +def instrument_pipeline(pipe, level: int = 1): + """Measure one pipeline instance regardless of `VIDEOX_PERF`, for programmatic use. + + Wraps the class that actually owns the `__call__` being run, then attaches the component hooks eagerly, so + this should be called after any FSDP / offload wrapping has been applied. + """ + global _STATE + if _STATE is None: + _STATE = _MetricsState(level) + for cls in type(pipe).__mro__: + if "__call__" in cls.__dict__: + _instrument_class(cls) + break + _attach_hooks(pipe) + return pipe + + +# --------------------------------------------------------------------------- +# Training +# --------------------------------------------------------------------------- + +# A global step is partitioned into these, in the order they occur. They are non-overlapping by construction and +# `other` is the residual of the step's wall clock against the sum of the rest, so the parts always add up to the +# whole: an unmeasured cost shows up as a fat `other` rather than disappearing. +_TRAIN_PHASES = ( + "data", + "prep", + "vae_enc", + "vae_dec", + "text_encoder", + "aux", + "fwd", + "bwd", + "clip", + "opt", + "other", +) + +# How many later steps may close before an unsettled one is read off the host clock instead of its CUDA events. +# Zero extra lag would suffice for the loops in this repo, which sync on `gather(loss).item()` every micro-step; +# the margin is for a loop that does not. +_SETTLE_LAG = 2 + +_TRAIN_INSTALLED = False +_TRAIN: Optional["_TrainState"] = None + + +class _StepAccum: + """Phase timings of one global step, accumulated over however many micro-steps it spans.""" + + __slots__ = ("index", "t0", "total_s", "stages", "micro_steps", "tokens", "samples", "fwd_calls") + + def __init__(self, index: int, t0: float): + self.index = index + # The close of the *previous* step, not the top of this one's first micro-step, so that the dataloader wait + # and the loop's own bookkeeping fall inside a step instead of between two of them. + self.t0 = t0 + self.total_s = 0.0 + self.stages: Dict[str, _StageStat] = {} + self.micro_steps = 0 + self.tokens: Optional[int] = None + self.samples = 0 + self.fwd_calls = 0 + + def add(self, name: str, start, end) -> None: + stat = self.stages.get(name) + if stat is None: + stat = _StageStat() + self.stages[name] = stat + stat.pairs.append((start, end)) + + def ready(self) -> bool: + return all(stat.ready() for stat in self.stages.values()) + + +class _TrainState: + """Everything the training path keeps between steps: phase accumulation, windowing, and the model's size.""" + + def __init__(self, state: _MetricsState): + self.state = state + self.every = max(1, _env_int("VIDEOX_PERF_EVERY", 50)) + self.total_steps = _env_int("VIDEOX_PERF_TOTAL_STEPS", 0) or None + self.coef_override = _env_float("VIDEOX_PERF_FLOPS_COEF") + + self.num_steps = 0 + self.step: Optional[_StepAccum] = None + self.pending: Any = collections.deque() + self.window: List[Dict[str, Any]] = [] + self.measured_step_s: List[float] = [] + self.wrapped_classes = 0 + + # Re-entrancy and attribution flags. Plain attributes rather than the thread-local state the inference + # path needs: a training loop runs on one thread, and the dataloader's parallelism is in worker + # *processes*, which hold their own copy of this module and never reach any of this. + self.depth = 0 + self.in_backward = False + self.step_owner: Optional[int] = None + + self.head_open = False + self.micro_enter = 0.0 + self.prev_exit: Optional[float] = None + + self.accelerator = None + self.dit_module: Optional[torch.nn.Module] = None + self.dit_params = 0 + self.dit_trainable = 0 + self.attn_widths: Optional[Dict[str, int]] = None + self.text_tokens = 0 + self._dit_resolved = False + self._ckpt = False + self._device: Optional[torch.device] = None + self._device_seen = False + self._finished = False + atexit.register(self.finish) + + # -- device ------------------------------------------------------------ + + @property + def device(self) -> Optional[torch.device]: + """The CUDA device this rank trains on, or `None` to stay off the CUDA APIs entirely. + + `accelerator.device` is both authoritative and, unlike the inference path (see [`_perf_device`]), backed by + a real `torch.cuda.set_device`: accelerate's `PartialState` points the process at its own local index when + it comes up. The device is still named explicitly on every call below rather than left to default, so that + a script which never built an `Accelerator` cannot quietly aim this at device 0. + """ + if self._device_seen: + return self._device + if not (torch.cuda.is_available() and torch.cuda.is_initialized()): + return None # not resolved, not cached: CUDA may still come up later in the run + device = getattr(self.accelerator, "device", None) + if getattr(device, "type", None) == "cuda": + self._device = device + self._device_seen = True + return self._device + + # -- step assembly ----------------------------------------------------- + + def current(self, t0: Optional[float] = None) -> _StepAccum: + if self.step is None: + self.step = _StepAccum(self.num_steps, t0 if t0 is not None else time.perf_counter()) + return self.step + + def close_head(self, now: float) -> None: + """End the `prep` window of the current micro-step, at the first measured work to start inside it. + + `prep` is thus literally what runs before any model does -- the noise and timestep sampling, the latent + bookkeeping. It is taken off the host clock: the tensors it creates are small, and an event pair here would + instead measure when the *previous* phase's queued kernels drained. + """ + if self.head_open: + self.head_open = False + self.current().add("prep", (self.micro_enter, None), (now, None)) + + def enter_micro(self) -> None: + now = time.perf_counter() + step = self.current(now) + if self.prev_exit is not None: + # From leaving the previous `accumulate` block to entering this one: the host blocked on the + # dataloader, plus the loop bookkeeping around it. Host clock by definition -- nothing is submitted to + # the device in that gap. It measures the main process's *wait*, not the work inside the workers, + # which multiprocessing puts out of reach from here; a near-zero `data` still proves the input + # pipeline is keeping up. + step.add("data", (self.prev_exit, None), (now, None)) + step.micro_steps += 1 + self.head_open = True + self.micro_enter = now + + def exit_micro(self) -> None: + self.head_open = False + self.prev_exit = time.perf_counter() + + def note_optimizer_step(self, optimizer, t_end: float) -> None: + """Close the global step, if this call was the real one. + + The entry scripts call `optimizer.step()` on every micro-step and leave the decision to accelerate: the + body of `AcceleratedOptimizer.step` is guarded by `sync_gradients`, so on an accumulation micro-step it is + a no-op. Reading that same flag is what separates a genuine boundary from a pass-through, and it is still + valid here -- `accumulate` sets it on entry and the loops themselves read it again after stepping. + + A script with several optimizers (the distillation and preference-tuning ones) steps more than one per + iteration. The first instance to close a step owns the boundary from then on, so the count follows one of + them instead of counting the same iteration several times over; the others' time still lands in `opt`. + """ + try: + sync = bool(optimizer.gradient_state.sync_gradients) + except AttributeError: + sync = True # not an accelerate-managed optimizer, so every call is a step + if not sync: + return + key = id(optimizer) + if self.step_owner is None: + self.step_owner = key + if self.step_owner == key: + self.close_step(t_end) + + def close_step(self, t_close: float) -> None: + step = self.step + if step is None: + return + step.total_s = t_close - step.t0 + self.num_steps += 1 + self.step = _StepAccum(self.num_steps, t_close) + self.head_open = False + self.pending.append(step) + self._drain() + + # -- settling ---------------------------------------------------------- + + def _drain(self, force: bool = False) -> None: + """Settle the steps whose CUDA events have retired, and only those. + + This is where the promise of adding no synchronize is kept. A step's closing events are recorded a few + microseconds before it closes and cannot be read yet; waiting on them would be exactly the stall this path + exists to avoid, and reading them unretired raises. So a finished step waits in a queue until a later step + proves the device has moved past it -- which the loops here do on their own, syncing on + `gather(loss).item()` once per micro-step. A step that is still unready after `_SETTLE_LAG` more have + closed is read off the host clock instead, which loses the device-side precision for that step but never + blocks and never drops it. + """ + while self.pending: + step = self.pending[0] + if not (force or step.ready() or self.num_steps - step.index > _SETTLE_LAG): + break + self.pending.popleft() + self._settle(step) + + def _settle(self, step: _StepAccum) -> None: + phases: Dict[str, float] = {} + for name, stat in step.stages.items(): + values = stat.elapsed_ms() + if values: + phases[name] = sum(values) / 1000.0 + residual = step.total_s - sum(phases.values()) + phases["other"] = max(0.0, residual) + record = { + "step": step.index, + "warmup": step.index < self.state.warmup, + "total_s": step.total_s, + "micro_steps": step.micro_steps, + "phases": phases, + # The measured phases can outrun the step when device work from one phase drains inside the next, since + # each pair is resolved to whichever of its host and device spans is longer. Reported rather than + # folded away, because a large one means the breakdown below should not be read too closely. + "overrun_s": max(0.0, -residual), + "tokens": step.tokens, + "samples": step.samples, + "fwd_calls": step.fwd_calls, + } + self.window.append(record) + if not record["warmup"]: + self.measured_step_s.append(step.total_s) + if len(self.window) >= self.every: + self.emit_window() + + # -- the model being trained ------------------------------------------- + + def note_models(self, objects) -> None: + """Size the trained model from what `prepare` handed back, not from the module the script built. + + Under FSDP the wrapper owns the flat sharded parameters and the inner module's own are emptied, so counting + through the class the forward wrapper sees would report zero. The return value of `prepare` is the wrapper + itself, which is the one place the shard and its `requires_grad` can both be read. + + The largest module wins, which is the DiT in every script here: the vae and the text encoder are not passed + to `prepare` at all, and where a second network is (a discriminator, a fake-score model) it is the smaller. + """ + for obj in objects: + if not isinstance(obj, torch.nn.Module): + continue + total = self.state.dit_params_override or _count_params(obj, self.state.world_size) + if total > self.dit_params: + self.dit_params = total + self.dit_module = obj + self._dit_resolved = False + + def _resolve_dit(self) -> None: + """Read the trainable fraction and the checkpointing flag, once, at the first window rather than at + `prepare`. + + Deferred because a script is free to freeze weights or call `enable_gradient_checkpointing` after preparing, + and by the first window -- fifty steps in by default -- whatever it was going to do it has done. + """ + module = self.dit_module + if self._dit_resolved or module is None: + return + self._dit_resolved = True + world = self.state.world_size + scale = world if world > 1 and _is_sharded(module) else 1 + self.dit_trainable = sum(p.numel() for p in module.parameters() if p.requires_grad) * scale + self._ckpt = any(getattr(sub, "gradient_checkpointing", False) for sub in module.modules()) + if self.state.attn_flops: + # Walked once, here, for the same reason the trainable fraction is: it is a walk of every submodule of a + # 14B transformer and nothing about it changes from one step to the next. + widths = _attn_widths(module) + self.attn_widths = widths if widths["self"] else None + self.text_tokens = _text_tokens(module) if self.attn_widths else 0 + + @property + def dit_ckpt(self) -> bool: + self._resolve_dit() + return self._ckpt + + def flops_coef(self) -> Tuple[float, float, str]: + """The two multipliers on a forward's cost that a whole step comes to, and why. + + A forward costs one pass over the model. A full-parameter backward costs two more -- one for the input + gradients, one for the weight gradients -- so a full step is 3x the forward, which is the flat 3 Megatron + applies as its `forward_backward_expansion_factor`. Freezing the base weights, as LoRA does, drops the + weight-gradient pass over them and brings the backward to about 1x, for 2x total. + + Gradient checkpointing adds one more forward, recomputed during the backward. It is returned as a *second* + coefficient rather than folded into the first, because the two answer different questions and the industry + gave them different names. MFU, as PaLM defined it, is the work the model needs against the hardware peak, + and it deliberately excludes recomputation: a run that recomputes has not become more useful for it. HFU is + the work the hardware actually issued, recomputation included. Reporting one figure under the name of the + other is what this did: every run here trains with checkpointing on, so every `mfu` it ever printed was an + HFU, a quarter high. Megatron sidesteps the distinction by never counting recomputation at all -- its + `num_floating_point_operations` has no term for it -- and by reporting TFLOP/s rather than a utilization. + + The trainable fraction decides between 3x and 2x, with 0.5 as the split: every LoRA configuration here trains + well under a percent of the weights and every full-parameter one trains all of them, so nothing real lands + near the threshold. `VIDEOX_PERF_FLOPS_COEF` overrides both coefficients at once, which collapses `mfu` and + `hfu` onto each other by construction; it is what to reach for when a run mixes the two -- or when FSDP has + flattened frozen and trainable weights into one parameter, where the fraction cannot be read apart. + + The fraction is `requires_grad`, not the set of weights the optimizer updates, and those come apart: the + multiviews scripts pass `--trainable_modules view` yet flip `requires_grad_(True)` over the whole stack under + FSDP, so that every wrapped unit gets a post-backward reshard. Those weight gradients really are computed + before being discarded, so `requires_grad` is the fraction the FLOPs follow -- reading `trainable 100%` next + to a narrow `--trainable_modules` is that, not a contradiction. + """ + self._resolve_dit() + ckpt = self._ckpt + full = not self.dit_params or self.dit_trainable >= 0.5 * self.dit_params + coef = 3.0 if full else 2.0 + hw_coef = coef + (1.0 if ckpt else 0.0) + if self.attn_widths: + attn = f"attn from dims ({self.attn_widths['modules']} modules)" + elif self.state.attn_flops: + attn = "attn omitted (widths unreadable)" + else: + attn = "attn omitted (disabled)" + reason = f"{'full' if full else 'lora'}, ckpt {'on' if ckpt else 'off'} -> {coef:g}x/{hw_coef:g}x, {attn}" + if self.coef_override is not None: + override = self.coef_override + return override, override, f"{reason}, overridden to {override:g}x" + return coef, hw_coef, reason + + # -- windowing --------------------------------------------------------- + + def emit_window(self) -> None: + window = self.window + self.window = [] + if not window: + return + record = self._aggregate(window) + if self.state.should_log(): + logger.info(_format_train_window(self.state, record)) + if self.state.level >= 2: + for entry in window: + detail = " ".join(f"{name} {value:.3f}" for name, value in entry["phases"].items()) + logger.info(f"{self.state.tag} step {entry['step']} {entry['total_s']:.3f}s: {detail}") + if self.state.json_path: + _append_json(self.state, record) + device = self.device + if device is not None: + # Scope the next window's peaks to the next window, the way the inference path scopes them to a request. + torch.cuda.reset_peak_memory_stats(device) + + def _aggregate(self, window: List[Dict[str, Any]]) -> Dict[str, Any]: + # Warmup steps are logged but kept out of the statistics. If a whole window is warmup there is nothing to + # fall back on but the window itself, which is better than printing nothing at all. + measured = [entry for entry in window if not entry["warmup"]] or window + totals = sorted(entry["total_s"] for entry in measured) + phases = {} + for name in _TRAIN_PHASES: + values = [entry["phases"].get(name, 0.0) for entry in measured] + if any(values): + phases[name] = statistics.fmean(values) + record: Dict[str, Any] = { + "kind": "train_window", + "ts": time.time(), + "rank": self.state.rank, + "world_size": self.state.world_size, + "step_first": window[0]["step"], + "step_last": window[-1]["step"], + "steps": len(measured), + "warmup_excluded": len(window) - len(measured), + "step_s_p50": _percentile(totals, 50), + "step_s_p95": _percentile(totals, 95), + "phases_s": phases, + "micro_steps": statistics.fmean([entry["micro_steps"] for entry in measured]), + "overrun_s": statistics.fmean([entry["overrun_s"] for entry in measured]), + } + device = self.device + if device is not None: + record["device"] = str(device) + record["peak_alloc_bytes"] = torch.cuda.max_memory_allocated(device) + record["peak_reserved_bytes"] = torch.cuda.max_memory_reserved(device) + record["throughput"] = self._throughput(measured) + if self.total_steps: + remaining = max(0, self.total_steps - window[-1]["step"] - 1) + record["eta_s"] = remaining * record["step_s_p50"] + return record + + def _throughput(self, measured: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + """Samples per second, achieved TFLOPS and MFU for the window. + + Every rate here is one sum over another -- the window's own work over the window's own wall clock -- and + never an aggregate of one quantity divided by an aggregate of the other. The two agree while the steps in a + window are alike and diverge badly when they are not, and under `--enable_bucket` they are not: the sampler + mixes image steps with 81-frame video steps, so one window holds token counts that differ several-fold and + step times with them. Taking the last step's tokens over the window's p50 step, as this did, reported 181.8 + TFLOPS for a window whose honest rate was a third of that -- the tokens having come from a video step and + the p50 from an image one. The p50 the window reports is a description of the step times, not a divisor. + + A step that captured no token count contributes neither work nor time, rather than contributing time alone, + which would pull the rate down by whatever share of the window it occupied. + + `flops_per_step` and `tokens` are window means, reported to say what the window contained. Dividing the + first by the logged p50 will not reproduce `tflops` and is not meant to: a mean of products is not the + product of the means, and the p50 spans steps that the FLOPs figure excludes. The two clocks the rates did + divide by are reported as `wall_s` and `priced_s`, so that either rate can be checked against the record it + came from. They differ only by the steps that carried no token count. + + The cost model is `2 * params * tokens` for the linear layers plus a separate quadratic term for core + attention, read from the model's projection widths by [`_attn_widths`]. The split is Megatron's: it multiplies + its per-token terms by the token count and its `self_attn_core_term` by the sum of the squared lengths, + because one grows linearly in the sequence and the other quadratically. Folding attention into a parameter + count, as this used to, understates a 28k-token video step by about two thirds and an 8k-token one by about a + third -- which is why the old figure could not be used to compare two runs at different lengths, whatever its + docstring claimed. When the widths cannot be read the attention term is dropped and `flops_coef_reason` says + so; the figure is then the old lower bound. + + The linear half is priced from parameters, and that is loose in the opposite direction: it charges every + weight against every latent token, while the cross-attention key and value projections run on the few hundred + text tokens and the timestep and text embeddings on fewer still. On Wan that is worth about a fifth of the + linear half, which the attention term now dwarfs. + + Two utilizations are reported and they are not interchangeable. `mfu` counts the forward and backward the + model needs; `hfu` also counts the forward that gradient checkpointing recomputes. `tflops` pairs with the + first and `hw_tflops` with the second, so that either rate divided by `peak_tflops` reproduces its own + utilization. See [`flops_coef`] for why the two are kept apart. + + Two widths matter and they are not the same. `samples` is what this rank alone consumed, and it is reported + as such: under sequence parallel one sample is spread over several ranks, so summing it across the world + would claim several times the samples that were actually trained on. The FLOPs rate is instead reported for + the whole job -- the local cost times the data-parallel width -- because that is the figure the aggregate + peak below is comparable against. Under plain FSDP or DDP, where every rank holds its own samples, the two + multiplications cancel and the MFU is exactly the per-device one. + """ + wall_s = sum(entry["total_s"] for entry in measured) + if wall_s <= 0: + return None + samples = statistics.fmean([entry["samples"] for entry in measured]) + dp = _dp_degree(self.state) + result: Dict[str, Any] = { + "samples_per_s": sum(entry["samples"] for entry in measured) / wall_s, + "samples": samples, + "wall_s": wall_s, + "dp": dp, + } + priced = [entry for entry in measured if entry["tokens"] and entry["samples"]] + priced_s = sum(entry["total_s"] for entry in priced) + if not (priced and priced_s > 0 and self.dit_params): + return result + coef, hw_coef, reason = self.flops_coef() + linear = 0.0 + attention = 0.0 + for entry in priced: + tokens = entry["tokens"] + linear += 2.0 * self.dit_params * tokens * entry["samples"] + attention += _attn_flops(self.attn_widths, tokens, self.text_tokens) * entry["samples"] + per_forward = linear + attention + flops_total = coef * per_forward + hw_flops_total = hw_coef * per_forward + achieved = flops_total * dp / priced_s / 1e12 + hw_achieved = hw_flops_total * dp / priced_s / 1e12 + per_device_peak = self.state.peak_tflops(self.device) + peak = per_device_peak * self.state.world_size if per_device_peak else None + result.update( + { + "params": self.dit_params, + "trainable": self.dit_trainable, + "tokens": statistics.fmean([entry["tokens"] for entry in priced]), + "priced_steps": len(priced), + "priced_s": priced_s, + "flops_coef": coef, + "flops_coef_hw": hw_coef, + "flops_coef_reason": reason, + "flops_per_step": flops_total / len(priced), + # What share of the figure is the quadratic term, so that a reader can see how much of it rests on + # the width introspection rather than on the parameter count. + "attn_share": (attention / per_forward) if per_forward else 0.0, + "devices": self.state.world_size, + "tflops": achieved, + "hw_tflops": hw_achieved, + "peak_tflops": peak, + "mfu": (achieved / peak) if peak else None, + "hfu": (hw_achieved / peak) if peak else None, + } + ) + return result + + def finish(self) -> None: + """Flush at exit: force the queued steps out, emit the partial window, then one summary line. + + Guarded so that flushing explicitly does not then get flushed again by the `atexit` hook, which would print + the summary twice over the same steps. + """ + if self._finished: + return + self._finished = True + try: + self._drain(force=True) + self.emit_window() + if self.measured_step_s and self.state.should_log(): + totals = sorted(self.measured_step_s) + logger.info( + f"{self.state.tag} === {len(totals)} steps | " + f"p50 {_percentile(totals, 50):.2f}s/step p95 {_percentile(totals, 95):.2f}s | " + f"{sum(totals) / 3600.0:.2f}h in-loop ===" + ) + except Exception as error: # an exit handler must not turn a finished run into a failed one + logger.warning(f"[Perf] failed to report the training summary: {error!r}") + + +def _dp_degree(state: _MetricsState) -> int: + """How many ranks hold *different* samples on a step. + + The world splits along two axes at once. Sequence parallel shares one sample across a group of ranks, each + computing a slice of its tokens; data parallel gives each group its own samples. Only the latter multiplies the + samples a step trains on, so it is the only one the throughput may be scaled by. + + xfuser is asked rather than assumed, and only if it is already imported: reaching into `sys.modules` avoids + importing it in a run that does not use it, and the query is guarded because the accessor raises before the + parallel state is initialized. + """ + sp = 1 + module = sys.modules.get("xfuser.core.distributed.parallel_state") + if module is not None: + try: + sp = max(1, int(module.get_sequence_parallel_world_size())) + except Exception: + sp = 1 + return max(1, state.world_size // sp) + + +def _model_stage(cls) -> str: + """Which phase a model class's forward belongs to. + + Classified from the class rather than from a variable name in a script, because the training scripts are what + this must not touch. The order matters: an autoencoder is one before it is anything else, and the vision towers + are pulled out ahead of the text-encoder test they would otherwise pass. + """ + from diffusers.models.modeling_utils import ModelMixin + + name = cls.__name__ + if "Autoencoder" in name or "AutoEncoder" in name or "VAE" in name: + return "vae" + if "Vision" in name: + return "aux" + if "T5Encoder" in name or "TextEncoder" in name or "CLIPTextModel" in name: + return "text_encoder" + if not issubclass(cls, ModelMixin): + # What is left having reached here is a `transformers.PreTrainedModel`, which across this repo means an LLM + # or T5 tower standing in for a text encoder. + return "text_encoder" + if "Transformer" in name or "UNet" in name or "Unet" in name or "LatentUpsampler" in name: + return "fwd" + # Audio encoders, vocoders, projection bridges and connectors: real cost, but none of the phases above. + return "aux" + + +def _wrap_model_method(fn, stage: str, capture_tokens: bool = False): + """Time one model method into `stage`, but only when it is the outermost such call of a training step. + + The nesting guard is what keeps the phases a true partition. Wrapped classes do contain each other -- an audio + encoder inside a DiT, a decoder inside an autoencoder -- and timing both would count the inner one twice, once + on its own and once inside its caller, so the phases would sum past the step they came from. Only the outermost + is measured and the inner ones fold into it. + + Timing is likewise suppressed inside `Accelerator.backward`. Gradient checkpointing recomputes forwards during + the backward pass; those are a real cost of the backward and belong in `bwd`, which is why `fwd` and `bwd` come + out near 1:2 with checkpointing on rather than the 1:2 of the arithmetic being a coincidence. + """ + + @functools.wraps(fn) + def wrapper(self, *args, **kwargs): + train = _TRAIN + if train is None or train.depth or train.in_backward: + return fn(self, *args, **kwargs) + start = _mark(train.device) + train.close_head(start[0]) + shape = _infer_tokens(self, args, kwargs) if capture_tokens else None + train.depth += 1 + try: + return fn(self, *args, **kwargs) + finally: + train.depth -= 1 + step = train.current() + step.add(stage, start, _mark(train.device)) + if shape is not None: + if step.tokens is None: + step.tokens = shape[0] + step.samples += shape[1] + step.fwd_calls += 1 + + wrapper._videox_perf = True + return wrapper + + +def _wrap_model_attr(cls, name: str, stage: str, capture_tokens: bool = False) -> bool: + """Wrap `cls.name` unless it is already wrapped, resolving it through the class's bases. + + Resolving through the bases is what dedupes an inheritance chain for free: a subclass that does not define its + own `forward` finds the parent's, which carries the marker if the parent has been done, and is skipped. Should + the subclass be reached first instead, the guard inside [`_wrap_model_method`] keeps the resulting double + wrapper from double-counting. + """ + fn = getattr(cls, name, None) + if not callable(fn) or getattr(fn, "_videox_perf", False): + return False + setattr(cls, name, _wrap_model_method(fn, stage, capture_tokens)) + return True + + +def _instrument_model_classes() -> Tuple[int, int]: + """Wrap the model classes exported by `videox_fun.models`. Returns `(wrapped, skipped)`. + + Done at the *class* level, and at import time, because training has no pipeline object to walk: the entry + scripts build their models themselves and hand them to `accelerate`, so there is no single place an instance + can be caught. Wrapping the class before any instance exists also means it survives everything applied + afterwards -- FSDP, DDP, peft -- since all of those end up calling the original class's forward. + + Membership is `ModelMixin` or `transformers.PreTrainedModel`, which is what separates a whole model from a + building block: it is what keeps `WanSelfAttention` and `WanRMSNorm` out, and wrapping either of those would + have the counters tick once per layer per step. The same test also leaves out four classes that are whole models + but subclass neither -- `AutoencoderKLWan_`, `AutoencoderKLWan2_2_`, `MOVAModel` and `Wav2Vec2ModelWrapper` -- + and none of the four is a gap, each being held by an instrumented class that does the calling: the two inner + vaes as the `self.model` of `AutoencoderKLWan` and `AutoencoderKLWan3_8`, and `Wav2Vec2ModelWrapper` as the + `self.audio_encoder` of `LongCatVideoAudioEncoder`, so their time already lands in the holder's stage. + `MOVAModel` defines no `forward` at all, only a `__call__` routing to four submodels that are themselves + wrapped. + """ + from diffusers.models.modeling_utils import ModelMixin + from transformers import PreTrainedModel + + from .. import models as models_package + + seen = set() + wrapped = 0 + skipped = 0 + for obj in list(vars(models_package).values()): + # Dedupe on the class object rather than the name: the package exports aliases of the same class, and + # wrapping one twice would append two marker pairs per call. + if not isinstance(obj, type) or id(obj) in seen: + continue + seen.add(id(obj)) + if not issubclass(obj, torch.nn.Module): + continue + if not (issubclass(obj, ModelMixin) or issubclass(obj, PreTrainedModel)): + skipped += 1 + continue + stage = _model_stage(obj) + if stage == "vae": + # `forward` never runs on these: callers use `encode` and `decode`, and wrapping the two separately is + # also what splits the two directions apart in the log. The streaming pair folds into those same two + # stages, being the same work done in chunks, and neither of them calls the plain method, so nothing is + # counted twice. + done = _wrap_model_attr(obj, "encode", "vae_enc") + done |= _wrap_model_attr(obj, "decode", "vae_dec") + done |= _wrap_model_attr(obj, "encode_stream", "vae_enc") + done |= _wrap_model_attr(obj, "decode_stream", "vae_dec") + else: + done = _wrap_model_attr(obj, "forward", stage, capture_tokens=(stage == "fwd")) + wrapped += bool(done) + return wrapped, skipped + + +def _patch_method(cls, name: str, factory) -> int: + """Replace a method a class defines itself, once. Returns whether it did.""" + fn = cls.__dict__.get(name) + if fn is None or getattr(fn, "_videox_perf", False): + return 0 + setattr(cls, name, factory(fn)) + return 1 + + +def _wrap_accumulate(fn): + @functools.wraps(fn) + @contextlib.contextmanager + def accumulate(self, *models): + train = _TRAIN + if train is None: + with fn(self, *models) as value: + yield value + return + train.accelerator = self + train.enter_micro() + try: + # Re-entered as a context manager, not called as a plain function. `Accelerator.accumulate` is a + # `@contextmanager` generator, so `fn(self, *models)` hands back a context manager that has not yet run + # a line of its body; driving it with `with` is what keeps `no_sync` wrapped around the micro-step, and + # yielding its value through keeps an exception raised inside the block propagating as it did before. + with fn(self, *models) as value: + yield value + finally: + train.exit_micro() + + accumulate._videox_perf = True + return accumulate + + +def _wrap_backward(fn): + @functools.wraps(fn) + def backward(self, *args, **kwargs): + train = _TRAIN + if train is None: + return fn(self, *args, **kwargs) + train.accelerator = self + device = train.device + start = _mark(device) + train.close_head(start[0]) + train.in_backward = True + try: + return fn(self, *args, **kwargs) + finally: + train.in_backward = False + train.current().add("bwd", start, _mark(device)) + + backward._videox_perf = True + return backward + + +def _wrap_clip(fn): + @functools.wraps(fn) + def clip_grad_norm_(self, *args, **kwargs): + train = _TRAIN + if train is None: + return fn(self, *args, **kwargs) + device = train.device + start = _mark(device) + try: + return fn(self, *args, **kwargs) + finally: + # Worth its own phase rather than being left in `other`: under FSDP the global norm needs an all-reduce + # across the shards, so this is where that collective becomes visible. + train.current().add("clip", start, _mark(device)) + + clip_grad_norm_._videox_perf = True + return clip_grad_norm_ + + +def _wrap_optimizer_step(fn): + @functools.wraps(fn) + def step(self, *args, **kwargs): + train = _TRAIN + if train is None: + return fn(self, *args, **kwargs) + device = train.device + start = _mark(device) + try: + return fn(self, *args, **kwargs) + finally: + end = _mark(device) + train.current().add("opt", start, end) + train.note_optimizer_step(self, end[0]) + + step._videox_perf = True + return step + + +def _wrap_prepare(fn): + @functools.wraps(fn) + def prepare(self, *args, **kwargs): + result = fn(self, *args, **kwargs) + train = _TRAIN + if train is not None: + train.accelerator = self + try: + train.note_models(result if isinstance(result, tuple) else (result,)) + except Exception as error: # never let measurement take down a training run + logger.warning(f"[Perf] failed to size the prepared model: {error!r}") + return result + + prepare._videox_perf = True + return prepare + + +def _patch_accelerate() -> int: + """Hang the step model off `accelerate`. Returns the number of methods patched. + + `accelerate` is the seam because it is the one thing all 113 training scripts in this repo share: every one of + them builds an `Accelerator` and calls `prepare`, and all but a handful use `accumulate`, `backward` and + `clip_grad_norm_`. + + The step boundary is `AcceleratedOptimizer.step` and deliberately not `torch.optim.AdamW.step`. Three optimizers + are in use across these scripts -- torch's `AdamW`, bitsandbytes' `AdamW8bit` and `CAME` -- and accelerate's + wrapper is the one point all three pass through. It also steps around a trap: `AdamW.step` overrides + `Optimizer.step`, so patching the base class would silently miss it. + """ + from accelerate import Accelerator + from accelerate.optimizer import AcceleratedOptimizer + + patched = _patch_method(Accelerator, "accumulate", _wrap_accumulate) + patched += _patch_method(Accelerator, "backward", _wrap_backward) + patched += _patch_method(Accelerator, "clip_grad_norm_", _wrap_clip) + patched += _patch_method(Accelerator, "prepare", _wrap_prepare) + patched += _patch_method(AcceleratedOptimizer, "step", _wrap_optimizer_step) + try: + from accelerate.utils import DeepSpeedOptimizerWrapper + except ImportError: # pragma: no cover - depends on the accelerate build + pass + else: + # It overrides `step` with a no-op, deepspeed having done the stepping inside `backward`, so the base class + # patch above never runs for it and it needs its own to mark the boundary. + patched += _patch_method(DeepSpeedOptimizerWrapper, "step", _wrap_optimizer_step) + return patched + + +def _format_train_window(state: _MetricsState, record: Dict[str, Any]) -> str: + span = f"step {record['step_first']}-{record['step_last']}" + if record["warmup_excluded"]: + span += f" ({record['warmup_excluded']} warmup excluded)" + parts = [ + span, + f"{record['step_s_p50']:.2f}s/step p95 {record['step_s_p95']:.2f}s", + " ".join(f"{name} {seconds:.2f}s" for name, seconds in record["phases_s"].items()), + ] + if record["micro_steps"] > 1.0: + parts.append(f"{record['micro_steps']:.0f} micro-steps") + if record["overrun_s"] > 0.02 * max(record["step_s_p50"], 1e-9): + # The phases came to more than the step. Said out loud rather than hidden, because past a couple of percent + # it means the phase boundaries are blurred by device work draining across them. + parts.append(f"overrun {record['overrun_s']:.2f}s") + + throughput = record.get("throughput") + if throughput: + parts.append(f"{throughput['samples_per_s']:.3f} samples/s (local, dp={throughput['dp']})") + if "peak_alloc_bytes" in record: + gib = 1024.0 ** 3 + parts.append( + f"peak_alloc {record['peak_alloc_bytes'] / gib:.1f}GiB " + f"peak_reserved {record['peak_reserved_bytes'] / gib:.1f}GiB" + ) + if throughput and "params" in throughput: + trainable = throughput["trainable"] / throughput["params"] * 100.0 if throughput["params"] else 0.0 + segment = ( + f"DiT {throughput['params'] / 1e9:.1f}B params " + f"(trainable {trainable:.3g}%, {throughput['flops_coef_reason']}, " + f"attn {throughput['attn_share'] * 100:.0f}%) " + f"{throughput['flops_per_step']:.2e} FLOPs/step -> {throughput['tflops']:.1f} TFLOPS" + ) + if throughput["mfu"] is not None: + over = f" over {throughput['devices']} GPUs" if throughput["devices"] > 1 else "" + # MFU first because it is the figure that compares across runs, HFU beside it because with gradient + # checkpointing on the hardware really did issue that much and the gap between the two is the + # recomputation. They coincide when checkpointing is off. + segment += ( + f" (MFU {throughput['mfu'] * 100:.1f}% / HFU {throughput['hfu'] * 100:.1f}%" + f" @{throughput['peak_tflops']:.0f}{over})" + ) + else: + segment += " (MFU n/a)" + if throughput["priced_steps"] < record["steps"]: + # Both the FLOPs and the rate cover only the steps that reported a token count. Said out loud when that is + # fewer than the window held, so the figure is not read as a rate over the whole window. + segment += f" [{throughput['priced_steps']}/{record['steps']} steps priced]" + parts.append(segment) + if "eta_s" in record: + parts.append(f"ETA {record['eta_s'] / 3600.0:.1f}h") + return f"{state.tag} " + " | ".join(parts) + + +def install_training(level: Optional[int] = None) -> bool: + """Measure the training loop by wrapping `accelerate` and the `videox_fun.models` classes. + + No-op unless `VIDEOX_PERF` is set. Called at the end of `videox_fun.__init__`, which is early enough to wrap the + model classes before any instance of one exists and covers every training script here -- the three that never + import `videox_fun.pipeline` still import `videox_fun.models`. + + A global step is timed from the close of the previous optimizer step to the close of this one, and split into + the phases in [`_TRAIN_PHASES`]. Nothing in this path calls `torch.cuda.synchronize`: the phase timings are CUDA + events, which are recorded and then left alone until a later step has shown the device to be past them (see + [`_TrainState._drain`]). Wall clock alone would not do here, because these loops synchronize on + `gather(loss).item()` *before* `backward`, so the host returns from the backward long before the device is done + with it and a host-only reading would push that work into the following step. + """ + global _TRAIN_INSTALLED, _TRAIN, _STATE + if _TRAIN_INSTALLED: + return True + if level is None: + level = _env_int("VIDEOX_PERF", 0) + if level <= 0: + return False + + _TRAIN_INSTALLED = True + if _STATE is None: + _STATE = _MetricsState(level) + train = _TrainState(_STATE) + + try: + patched = _patch_accelerate() + wrapped, skipped = _instrument_model_classes() + except Exception as error: + # A run that cannot be measured is still a run. Leave `_TRAIN` unset so the wrappers that did land, if any, + # stay inert rather than half-reporting. + logger.warning(f"[Perf] training metrics disabled: {error!r}") + return False + + # Published last: every wrapper above reads this global and does nothing while it is `None`, so nothing is + # measured until the whole set is in place and no half-installed state can produce a partial step. + _TRAIN = train + train.wrapped_classes = wrapped + if _STATE.should_log(): + logger.info( + f"{_STATE.tag} training metrics enabled (level {level}) on {patched} accelerate methods and " + f"{wrapped} model classes ({skipped} non-model classes skipped), " + f"reporting every {train.every} steps" + ) + return True +