Download training_code/model_architecture.py from SleepMastger/pusht-fastwam: direct link, hf CLI and curl.
- Browser
- Download file 35.6 kB
-
https://huggingface.co/SleepMastger/pusht-fastwam/resolve/main/training_code/model_architecture.py
- Command line
-
hf download hf://SleepMastger/pusht-fastwam/training_code/model_architecture.py
-
curl -L -o model_architecture.py https://huggingface.co/SleepMastger/pusht-fastwam/resolve/main/training_code/model_architecture.py
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 | |