pusht-fastwam / training_code /model_architecture.py
SleepMastger's picture
add model card, conditioning, and training-time processing
222aa51 verified
Raw History Blame Contribute Delete
35.6 kB
# Author: Rui Heng Yang
import logging
import os
import inspect
from pathlib import Path
import torch
from hydra.utils import instantiate
from omegaconf import DictConfig
from PIL import Image
import numpy as np
from einops import repeat
from omegaconf import OmegaConf
from .trainer import Wan22Trainer, create_accelerator_from_cfg
from .utils.logging_config import get_logger, setup_logging
from .utils.video_io import save_mp4
from .utils import misc
logger = get_logger(__name__)
def _normalize_mixed_precision(mixed_precision: str) -> str:
if not isinstance(mixed_precision, str):
raise ValueError(f"`mixed_precision` must be str, got {type(mixed_precision)}")
key = mixed_precision.strip().lower()
if key not in {"no", "fp16", "bf16"}:
raise ValueError(
f"Unsupported mixed_precision: {mixed_precision}. "
"Expected one of: ['no', 'fp16', 'bf16']."
)
return key
def _mixed_precision_to_model_dtype(mixed_precision: str) -> torch.dtype:
precision = _normalize_mixed_precision(mixed_precision)
if precision == "no":
return torch.float32
if precision == "fp16":
return torch.float16
return torch.bfloat16
def create_wan22_model(
model_id: str,
tokenizer_model_id: str,
dit_config,
tokenizer_max_len: int = 512,
train_shift: float = 5.0,
infer_shift: float = 5.0,
num_train_timesteps: int = 1000,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
):
from .models.wan22.wan22 import Wan22Core
if isinstance(dit_config, DictConfig):
dit_config = OmegaConf.to_container(dit_config, resolve=True)
if not isinstance(dit_config, dict):
raise ValueError(f"`dit_config` must resolve to a dict, got {type(dit_config)}")
return Wan22Core.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
redirect_common_files=bool(redirect_common_files),
dit_config=dit_config,
train_shift=float(train_shift),
infer_shift=float(infer_shift),
num_train_timesteps=int(num_train_timesteps),
)
def create_fastwam(
model_id: str,
tokenizer_model_id: str,
video_dit_config,
tokenizer_max_len: int = 512,
load_text_encoder: bool = True,
proprio_dim: int | None = None,
action_dit_config=None,
action_dit_pretrained_path: str | None = None,
skip_dit_load_from_pretrain: bool = False,
freeze_video_backbone: bool = False,
video_scheduler=None,
action_scheduler=None,
loss=None,
mot_checkpoint_mixed_attn: bool = True,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
):
from .models.wan22.fastwam import FastWAM
if isinstance(video_dit_config, DictConfig):
video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True)
if not isinstance(video_dit_config, dict):
raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}")
if isinstance(action_dit_config, DictConfig):
action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True)
if action_dit_config is None:
action_dit_config = {}
if not isinstance(action_dit_config, dict):
raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}")
if isinstance(video_scheduler, DictConfig):
video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True)
if video_scheduler is None:
video_scheduler = {}
if not isinstance(video_scheduler, dict):
raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}")
if isinstance(action_scheduler, DictConfig):
action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True)
if action_scheduler is None:
raise ValueError("`action_scheduler` is required for FastWAM.")
if not isinstance(action_scheduler, dict):
raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}")
required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"}
missing_keys = required_action_scheduler_keys - set(action_scheduler.keys())
if missing_keys:
raise ValueError(
f"`action_scheduler` missing required keys: {sorted(missing_keys)}. "
"Expected keys: train_shift, infer_shift, num_train_timesteps."
)
if isinstance(loss, DictConfig):
loss = OmegaConf.to_container(loss, resolve=True)
if loss is None:
loss = {}
if not isinstance(loss, dict):
raise ValueError(f"`loss` must be dict-like, got {type(loss)}")
return FastWAM.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
load_text_encoder=bool(load_text_encoder),
proprio_dim=(None if proprio_dim is None else int(proprio_dim)),
redirect_common_files=bool(redirect_common_files),
video_dit_config=video_dit_config,
action_dit_config=action_dit_config,
action_dit_pretrained_path=action_dit_pretrained_path,
skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain),
freeze_video_backbone=bool(freeze_video_backbone),
mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn),
video_train_shift=float(video_scheduler.get("train_shift", 5.0)),
video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)),
video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)),
action_train_shift=float(action_scheduler["train_shift"]),
action_infer_shift=float(action_scheduler["infer_shift"]),
action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]),
loss_lambda_video=float(loss.get("lambda_video", 1.0)),
loss_lambda_action=float(loss.get("lambda_action", 1.0)),
)
def create_fastwam_joint(
model_id: str,
tokenizer_model_id: str,
video_dit_config,
tokenizer_max_len: int = 512,
load_text_encoder: bool = True,
proprio_dim: int | None = None,
action_dit_config=None,
action_dit_pretrained_path: str | None = None,
skip_dit_load_from_pretrain: bool = False,
freeze_video_backbone: bool = False,
video_scheduler=None,
action_scheduler=None,
loss=None,
mot_checkpoint_mixed_attn: bool = True,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
):
from .models.wan22.fastwam_joint import FastWAMJoint
if isinstance(video_dit_config, DictConfig):
video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True)
if not isinstance(video_dit_config, dict):
raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}")
if isinstance(action_dit_config, DictConfig):
action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True)
if action_dit_config is None:
action_dit_config = {}
if not isinstance(action_dit_config, dict):
raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}")
if isinstance(video_scheduler, DictConfig):
video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True)
if video_scheduler is None:
video_scheduler = {}
if not isinstance(video_scheduler, dict):
raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}")
if isinstance(action_scheduler, DictConfig):
action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True)
if action_scheduler is None:
raise ValueError("`action_scheduler` is required for FastWAM.")
if not isinstance(action_scheduler, dict):
raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}")
required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"}
missing_keys = required_action_scheduler_keys - set(action_scheduler.keys())
if missing_keys:
raise ValueError(
f"`action_scheduler` missing required keys: {sorted(missing_keys)}. "
"Expected keys: train_shift, infer_shift, num_train_timesteps."
)
if isinstance(loss, DictConfig):
loss = OmegaConf.to_container(loss, resolve=True)
if loss is None:
loss = {}
if not isinstance(loss, dict):
raise ValueError(f"`loss` must be dict-like, got {type(loss)}")
return FastWAMJoint.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
load_text_encoder=bool(load_text_encoder),
proprio_dim=(None if proprio_dim is None else int(proprio_dim)),
redirect_common_files=bool(redirect_common_files),
video_dit_config=video_dit_config,
action_dit_config=action_dit_config,
action_dit_pretrained_path=action_dit_pretrained_path,
skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain),
freeze_video_backbone=bool(freeze_video_backbone),
mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn),
video_train_shift=float(video_scheduler.get("train_shift", 5.0)),
video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)),
video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)),
action_train_shift=float(action_scheduler["train_shift"]),
action_infer_shift=float(action_scheduler["infer_shift"]),
action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]),
loss_lambda_video=float(loss.get("lambda_video", 1.0)),
loss_lambda_action=float(loss.get("lambda_action", 1.0)),
)
def create_fastwam_idm(
model_id: str,
tokenizer_model_id: str,
video_dit_config,
tokenizer_max_len: int = 512,
load_text_encoder: bool = True,
proprio_dim: int | None = None,
action_dit_config=None,
action_dit_pretrained_path: str | None = None,
skip_dit_load_from_pretrain: bool = False,
freeze_video_backbone: bool = False,
video_scheduler=None,
action_scheduler=None,
loss=None,
mot_checkpoint_mixed_attn: bool = True,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
):
from .models.wan22.fastwam_idm import (
FastWAMIDM,
)
if isinstance(video_dit_config, DictConfig):
video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True)
if not isinstance(video_dit_config, dict):
raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}")
if isinstance(action_dit_config, DictConfig):
action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True)
if action_dit_config is None:
action_dit_config = {}
if not isinstance(action_dit_config, dict):
raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}")
if isinstance(video_scheduler, DictConfig):
video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True)
if video_scheduler is None:
video_scheduler = {}
if not isinstance(video_scheduler, dict):
raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}")
if isinstance(action_scheduler, DictConfig):
action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True)
if action_scheduler is None:
raise ValueError("`action_scheduler` is required for FastWAM.")
if not isinstance(action_scheduler, dict):
raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}")
required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"}
missing_keys = required_action_scheduler_keys - set(action_scheduler.keys())
if missing_keys:
raise ValueError(
f"`action_scheduler` missing required keys: {sorted(missing_keys)}. "
"Expected keys: train_shift, infer_shift, num_train_timesteps."
)
if isinstance(loss, DictConfig):
loss = OmegaConf.to_container(loss, resolve=True)
if loss is None:
loss = {}
if not isinstance(loss, dict):
raise ValueError(f"`loss` must be dict-like, got {type(loss)}")
return FastWAMIDM.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
load_text_encoder=bool(load_text_encoder),
proprio_dim=(None if proprio_dim is None else int(proprio_dim)),
redirect_common_files=bool(redirect_common_files),
video_dit_config=video_dit_config,
action_dit_config=action_dit_config,
action_dit_pretrained_path=action_dit_pretrained_path,
skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain),
freeze_video_backbone=bool(freeze_video_backbone),
mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn),
video_train_shift=float(video_scheduler.get("train_shift", 5.0)),
video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)),
video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)),
action_train_shift=float(action_scheduler["train_shift"]),
action_infer_shift=float(action_scheduler["infer_shift"]),
action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]),
loss_lambda_video=float(loss.get("lambda_video", 1.0)),
loss_lambda_action=float(loss.get("lambda_action", 1.0)),
)
def create_fastwam_decoupled(
model_id: str,
tokenizer_model_id: str,
video_dit_config,
tokenizer_max_len: int = 512,
load_text_encoder: bool = True,
proprio_dim: int | None = None,
action_dit_config=None,
action_dit_pretrained_path: str | None = None,
skip_dit_load_from_pretrain: bool = False,
freeze_video_backbone: bool = False,
decoupled: bool = True,
video_scheduler=None,
action_scheduler=None,
loss=None,
mot_checkpoint_mixed_attn: bool = True,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
kv_source_mode: str = "final_only",
fusion_hidden_dim: int = 64,
fusion_use_norm: bool = True,
):
"""Factory function for the Decoupled MoT variant (FastWAMDecoupled).
Author: Rui Heng Yang
Mirrors ``create_fastwam()`` but constructs ``FastWAMDecoupled`` with
``decoupled=True``, enabling asymmetric expert layer counts (e.g.,
video=30, action=5). The action expert loads from
``action_dit_pretrained_path`` when provided, optionally with
layer-selective initialization from ``kv_source_mapping``; otherwise it
starts from random init.
Args:
model_id: HuggingFace model ID for the Wan2.2-TI2V-5B backbone.
tokenizer_model_id: HuggingFace model ID for the tokenizer.
video_dit_config: Video DiT architecture config dict.
tokenizer_max_len: Maximum tokenizer sequence length.
load_text_encoder: Whether to load the T5 text encoder.
proprio_dim: Proprioceptive state dimension (None to disable).
action_dit_config: Action DiT architecture config dict.
action_dit_pretrained_path: Path to pretrained ActionDiT weights.
skip_dit_load_from_pretrain: If True, skip pretrained DiT loads (for full checkpoint override).
decoupled: Must be True for this factory (passed to from_wan22_pretrained).
video_scheduler: Video scheduler config dict.
action_scheduler: Action scheduler config dict.
loss: Loss config dict with lambda_video and lambda_action.
mot_checkpoint_mixed_attn: Whether to use gradient checkpointing in MoT.
redirect_common_files: Whether to redirect common HuggingFace files.
model_dtype: Model dtype for parameters.
device: Device to place the model on.
kv_source_mode: KV source routing mode for the action expert.
fusion_hidden_dim: MLP hidden dim for KVFusionModule (only used when kv_source_mode="fused_mlp").
fusion_use_norm: Whether KVFusionModule applies RMSNorm after each fusion MLP.
Returns:
FastWAMDecoupled instance with decoupled video/action experts.
"""
from .models.wan22.fastwam_decoupled import FastWAMDecoupled
if not decoupled:
# This factory only builds decoupled models. Previously the flag was
# accepted but ignored (decoupled=True was hardcoded downstream), so a
# config setting decoupled=false silently got a decoupled model anyway.
raise ValueError(
"create_fastwam_decoupled only builds decoupled models; got "
"decoupled=False. Use fastwam.runtime.create_fastwam (model=fastwam) "
"for the non-decoupled variant."
)
if isinstance(video_dit_config, DictConfig):
video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True)
if not isinstance(video_dit_config, dict):
raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}")
if isinstance(action_dit_config, DictConfig):
action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True)
if action_dit_config is None:
action_dit_config = {}
if not isinstance(action_dit_config, dict):
raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}")
if isinstance(video_scheduler, DictConfig):
video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True)
if video_scheduler is None:
video_scheduler = {}
if not isinstance(video_scheduler, dict):
raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}")
if isinstance(action_scheduler, DictConfig):
action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True)
if action_scheduler is None:
raise ValueError("`action_scheduler` is required for FastWAM.")
if not isinstance(action_scheduler, dict):
raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}")
required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"}
missing_keys = required_action_scheduler_keys - set(action_scheduler.keys())
if missing_keys:
raise ValueError(
f"`action_scheduler` missing required keys: {sorted(missing_keys)}. "
"Expected keys: train_shift, infer_shift, num_train_timesteps."
)
if isinstance(loss, DictConfig):
loss = OmegaConf.to_container(loss, resolve=True)
if loss is None:
loss = {}
if not isinstance(loss, dict):
raise ValueError(f"`loss` must be dict-like, got {type(loss)}")
from .models.wan22.mot_decoupled import compute_kv_source_mapping
from .models.wan22.action_width import resolve_action_width
kv_fusion = None # set by fused_mlp branch below
# Resolve the opt-in wide action expert before anything downstream reads the action
# dimensions. Under standard (or an absent key) this returns the very same dict with
# no writes, so the default path is unchanged; under wide it returns a copy, which is
# what makes this safe here, since this factory does not copy action_dit_config.
action_dit_config = resolve_action_width(action_dit_config)
layer_selected = action_dit_config.pop("layer_selected", None)
if layer_selected is not None:
if kv_source_mode == "fused_mlp":
raise ValueError("layer_selected and kv_source_mode='fused_mlp' are incompatible")
action_dit_config["num_layers"] = len(layer_selected)
kv_source_mapping = list(layer_selected)
logger.info(f"layer_selected={layer_selected}, overriding num_layers={len(layer_selected)}, kv_source_mapping={kv_source_mapping}")
elif kv_source_mode == "fused_mlp":
from .models.wan22.kv_fusion import KVFusionModule
N = int(video_dit_config.get("num_layers", 30))
M = int(action_dit_config.get("num_layers", 5))
num_heads = int(video_dit_config.get("num_heads", 24))
attn_head_dim = int(video_dit_config.get("attn_head_dim", 128))
init_mode = "uniform_end" if M <= N else "final_only"
kv_source_mapping = compute_kv_source_mapping(
mode=init_mode, video_num_layers=N, action_num_layers=M,
)
kv_fusion = KVFusionModule(
num_action_layers=M,
num_video_layers=N,
attn_hidden_dim=num_heads * attn_head_dim,
fusion_hidden_dim=fusion_hidden_dim,
fusion_use_norm=fusion_use_norm,
dtype=model_dtype,
)
logger.info(f"KV source mode: fused_mlp (all {N} video layers fused via MLP), "
f"fusion params: {sum(p.numel() for p in kv_fusion.parameters()) / 1e3:.1f}K")
else:
kv_source_mapping = compute_kv_source_mapping(
mode=kv_source_mode,
video_num_layers=int(video_dit_config.get("num_layers", 30)),
action_num_layers=int(action_dit_config.get("num_layers", 5)),
)
logger.info(f"KV source mode: {kv_source_mode}, mapping: {kv_source_mapping}")
return FastWAMDecoupled.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
load_text_encoder=bool(load_text_encoder),
proprio_dim=(None if proprio_dim is None else int(proprio_dim)),
redirect_common_files=bool(redirect_common_files),
video_dit_config=video_dit_config,
action_dit_config=action_dit_config,
action_dit_pretrained_path=action_dit_pretrained_path,
skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain),
freeze_video_backbone=bool(freeze_video_backbone),
mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn),
video_train_shift=float(video_scheduler.get("train_shift", 5.0)),
video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)),
video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)),
action_train_shift=float(action_scheduler["train_shift"]),
action_infer_shift=float(action_scheduler["infer_shift"]),
action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]),
loss_lambda_video=float(loss.get("lambda_video", 1.0)),
loss_lambda_action=float(loss.get("lambda_action", 1.0)),
decoupled=True,
kv_source_mapping=kv_source_mapping,
kv_source_mode=kv_source_mode,
kv_fusion=kv_fusion,
)
def create_fastwam_action_mlp(
model_id: str,
tokenizer_model_id: str,
video_dit_config,
tokenizer_max_len: int = 512,
load_text_encoder: bool = True,
proprio_dim: int | None = None,
action_mlp_config=None,
skip_dit_load_from_pretrain: bool = False,
video_scheduler=None,
action_scheduler=None,
loss=None,
redirect_common_files: bool = True,
model_dtype: torch.dtype = torch.bfloat16,
device: str = "cuda",
):
"""Factory for FastWAM with action DiT replaced by mean pooling + MLP."""
from .models.wan22.fastwam_action_mlp import FastWAMActionMLP
if isinstance(video_dit_config, DictConfig):
video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True)
if not isinstance(video_dit_config, dict):
raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}")
if isinstance(action_mlp_config, DictConfig):
action_mlp_config = OmegaConf.to_container(action_mlp_config, resolve=True)
if action_mlp_config is None:
action_mlp_config = {}
if not isinstance(action_mlp_config, dict):
raise ValueError(f"`action_mlp_config` must resolve to a dict, got {type(action_mlp_config)}")
if isinstance(video_scheduler, DictConfig):
video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True)
if video_scheduler is None:
video_scheduler = {}
if not isinstance(video_scheduler, dict):
raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}")
if isinstance(action_scheduler, DictConfig):
action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True)
if action_scheduler is None:
raise ValueError("`action_scheduler` is required for FastWAMActionMLP.")
if not isinstance(action_scheduler, dict):
raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}")
required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"}
missing_keys = required_action_scheduler_keys - set(action_scheduler.keys())
if missing_keys:
raise ValueError(
f"`action_scheduler` missing required keys: {sorted(missing_keys)}. "
"Expected keys: train_shift, infer_shift, num_train_timesteps."
)
if isinstance(loss, DictConfig):
loss = OmegaConf.to_container(loss, resolve=True)
if loss is None:
loss = {}
if not isinstance(loss, dict):
raise ValueError(f"`loss` must be dict-like, got {type(loss)}")
return FastWAMActionMLP.from_wan22_pretrained(
device=device,
torch_dtype=model_dtype,
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
tokenizer_max_len=int(tokenizer_max_len),
load_text_encoder=bool(load_text_encoder),
proprio_dim=(None if proprio_dim is None else int(proprio_dim)),
redirect_common_files=bool(redirect_common_files),
video_dit_config=video_dit_config,
action_mlp_config=action_mlp_config,
skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain),
video_train_shift=float(video_scheduler.get("train_shift", 5.0)),
video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)),
video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)),
action_train_shift=float(action_scheduler["train_shift"]),
action_infer_shift=float(action_scheduler["infer_shift"]),
action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]),
loss_lambda_video=float(loss.get("lambda_video", 1.0)),
loss_lambda_action=float(loss.get("lambda_action", 1.0)),
)
def _resolve_model_vae_identity(
model_cfg: DictConfig, model_dtype: torch.dtype
) -> dict[str, str]:
"""Resolve and hash the exact VAE artifact selected by the model factory."""
import hashlib
from .datasets.video_latent_cache import vae_implementation_sha256
from .models.wan22.helpers.loader import _resolve_configs
if model_dtype != torch.bfloat16:
raise ValueError(
"The configured video latent cache was encoded with a BF16 VAE; "
f"cached training requires mixed_precision=bf16, got {model_dtype}."
)
model_id = str(model_cfg.get("model_id", "Wan-AI/Wan2.2-TI2V-5B"))
tokenizer_model_id = str(
model_cfg.get("tokenizer_model_id", "Wan-AI/Wan2.1-T2V-1.3B")
)
_, _, vae_config, _ = _resolve_configs(
model_id=model_id,
tokenizer_model_id=tokenizer_model_id,
redirect_common_files=bool(model_cfg.get("redirect_common_files", True)),
)
vae_config.download_if_necessary()
vae_path = Path(str(vae_config.path)).resolve()
digest = hashlib.sha256()
with vae_path.open("rb") as file_handle:
for chunk in iter(lambda: file_handle.read(8 * 1024 * 1024), b""):
digest.update(chunk)
return {
"model_id": model_id,
"sha256": digest.hexdigest(),
"implementation_sha256": vae_implementation_sha256(),
"encode_dtype": "bfloat16",
}
def build_datasets(
data_cfg: DictConfig,
model_cfg: DictConfig | None = None,
model_dtype: torch.dtype = torch.bfloat16,
):
resolved_data_cfg = OmegaConf.create(OmegaConf.to_container(data_cfg, resolve=True))
if (
model_cfg is not None
and resolved_data_cfg.train.get("video_latent_cache_dir") is not None
):
resolved_data_cfg.train.video_latent_cache_vae_identity = _resolve_model_vae_identity(
model_cfg, model_dtype
)
train_ds = instantiate(resolved_data_cfg.train)
if resolved_data_cfg.get("val") is None:
if resolved_data_cfg.train.get("video_latent_cache_dir") is None:
val_ds = train_ds
else:
# Periodic evaluation needs original pixels for rollout and VAE
# reconstruction metrics. Keep training video-free while building
# one online-video dataset for the infrequent evaluation sample.
online_val_cfg = OmegaConf.create(
OmegaConf.to_container(resolved_data_cfg.train, resolve=True)
)
online_val_cfg.video_latent_cache_dir = None
online_val_cfg.video_latent_cache_vae_identity = None
online_val_cfg.strict_getitem = False
online_val_cfg.is_training_set = False
stats_path = resolved_data_cfg.train.get("pretrained_norm_stats")
online_val_cfg.pretrained_norm_stats = (
stats_path
if stats_path is not None
else os.path.join(misc.get_work_dir(), "dataset_stats.json")
)
logger.info(
"Building an online-video validation dataset while training uses cached latents."
)
val_ds = instantiate(online_val_cfg)
else:
train_stats_path = resolved_data_cfg.train.get("pretrained_norm_stats")
default_stats_path = os.path.join(misc.get_work_dir(), "dataset_stats.json")
val_stats_path = resolved_data_cfg.val.get("pretrained_norm_stats")
pretrained_norm_stats = val_stats_path or train_stats_path or default_stats_path
logger.info("Building val dataset with pretrained_norm_stats: %s", pretrained_norm_stats)
val_ds = instantiate(resolved_data_cfg.val, pretrained_norm_stats=pretrained_norm_stats)
return train_ds, val_ds
def _resolve_train_device() -> str:
if not torch.cuda.is_available():
return "cpu"
device_count = torch.cuda.device_count()
if device_count <= 1:
return "cuda:0"
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
if local_rank < 0 or local_rank >= device_count:
return "cuda:0"
return f"cuda:{local_rank}"
def _validate_pre_instantiation_training_config(cfg: DictConfig) -> None:
"""Reject IDM-only contradictions before the model factory can allocate."""
target = str(OmegaConf.select(cfg, "model._target_", default=""))
idm_targets = {
"fastwam.models.wan22.fasterwam_idm.create_fasterwam_idm",
(
"fastwam.models.wan22.fasterwam_idm_piha."
"create_fasterwam_idm_piha"
),
}
if target not in idm_targets:
return
if bool(OmegaConf.select(cfg, "freeze_video_backbone", default=False)):
raise ValueError(
"The IDM variant requires freeze_video_backbone=false because its "
"action loss must backpropagate through the conditional video tower"
)
def run_training(cfg: DictConfig):
setup_logging(
log_level=logging.INFO,
is_main_process=torch.distributed.get_rank() == 0 if torch.distributed.is_initialized() else True,
)
misc.register_work_dir(cfg.output_dir)
config_payload = OmegaConf.to_container(cfg, resolve=True)
with open(Path(cfg.output_dir) / "config.yaml", "w") as f:
OmegaConf.save(config_payload, f)
mixed_precision = _normalize_mixed_precision(cfg.mixed_precision)
model_dtype = _mixed_precision_to_model_dtype(mixed_precision)
accelerator = create_accelerator_from_cfg(cfg)
model_device = str(accelerator.device)
_validate_pre_instantiation_training_config(cfg)
model = instantiate(cfg.model, model_dtype=model_dtype, device=model_device)
train_ds, val_ds = build_datasets(cfg.data, cfg.model, model_dtype)
trainer = Wan22Trainer(
cfg=cfg,
model=model,
train_dataset=train_ds,
val_dataset=val_ds,
accelerator=accelerator,
)
trainer.train()
def run_inference(cfg: DictConfig):
setup_logging(log_level=logging.INFO)
inference_cfg = cfg.inference
mixed_precision = _normalize_mixed_precision(cfg.mixed_precision)
model_dtype = _mixed_precision_to_model_dtype(mixed_precision)
model = instantiate(cfg.model, model_dtype=model_dtype, device=str(inference_cfg.device))
checkpoint_path = inference_cfg.get("checkpoint_path")
if checkpoint_path:
ckpt = Path(checkpoint_path)
if ckpt.exists():
logger.info("Loading finetuned checkpoint: %s", checkpoint_path)
model.load_checkpoint(checkpoint_path)
else:
logger.warning("Checkpoint not found, skipping load: %s", checkpoint_path)
model.eval()
def center_crop_resize(img: Image, width: int, height: int) -> Image.Image:
src_w, src_h = img.size
scale = max(width / src_w, height / src_h)
resized = img.resize((round(src_w * scale), round(src_h * scale)), resample=Image.BILINEAR)
rw, rh = resized.size
left = max((rw - width) // 2, 0)
top = max((rh - height) // 2, 0)
return resized.crop((left, top, left + width, top + height))
input_image = Image.open(str(inference_cfg.input_image_path)).convert("RGB")
input_image = center_crop_resize(input_image, width=inference_cfg.width, height=inference_cfg.height)
arr = np.array(input_image, dtype=np.float32)
x = torch.from_numpy(arr)
x = x.to(device=model.device, dtype=model.torch_dtype)
x = x * (2.0 / 255.0) - 1.0
x = repeat(x, "H W C -> B C H W", B=1)
output_mp4 = str(inference_cfg.output_mp4)
infer_kwargs = {
"prompt": str(inference_cfg.prompt),
"negative_prompt": str(inference_cfg.negative_prompt),
"text_cfg_scale": float(inference_cfg.text_cfg_scale),
"action_cfg_scale": float(inference_cfg.action_cfg_scale),
"input_image": x,
"num_frames": int(inference_cfg.num_frames),
"num_inference_steps": int(inference_cfg.num_inference_steps),
"sigma_shift": None if inference_cfg.get("sigma_shift") is None else float(inference_cfg.sigma_shift),
"seed": int(inference_cfg.seed),
"rand_device": str(inference_cfg.rand_device),
"tiled": bool(inference_cfg.tiled),
}
infer_out = model.infer(**infer_kwargs)
video = infer_out["video"]
save_mp4(video, output_mp4, fps=15)
logger.info("Saved inference video to %s", output_mp4)
return output_mp4