"""DisCa-style dense image/text feature Predictor for HY-WorldPlay AR.""" from __future__ import annotations from dataclasses import asdict, dataclass from math import prod from pathlib import Path from typing import Any import torch from einops import repeat from safetensors import safe_open from torch import nn from torch.utils.checkpoint import checkpoint from hyvideo.models.transformers.modules.activation_layers import get_activation_layer from hyvideo.models.transformers.modules.embed_layers import PatchEmbed from hyvideo.models.transformers.modules.mlp_layers import FinalLayer from hyvideo.models.transformers.modules.posemb_layers import get_nd_rotary_pos_embed from hyvideo.models.transformers.worldplay_1_5_transformer import MMDoubleStreamBlock @dataclass(frozen=True) class PredictorConfig: hidden_size: int = 2048 heads_num: int = 16 mlp_width_ratio: float = 4.0 mlp_act_type: str = "gelu_tanh" qkv_bias: bool = True qk_norm: bool = True qk_norm_type: str = "rms" attn_mode: str = "flash" patch_size: tuple[int, int, int] = (1, 1, 1) in_channels: int = 32 out_channels: int = 32 concat_condition: bool = True rope_dim_list: tuple[int, int, int] = (16, 56, 56) rope_theta: float = 256.0 source_block_ids: tuple[int, int] = (1, 52) latent_height: int = 30 latent_width: int = 52 latent_frames: int = 4 class FeatureFusion(nn.Module): def __init__(self, hidden_size: int) -> None: super().__init__() self.current_norm = nn.LayerNorm(hidden_size, eps=1e-6) self.cached_norm = nn.LayerNorm(hidden_size, eps=1e-6) self.mlp = nn.Sequential( nn.Linear(2 * hidden_size, hidden_size), nn.SiLU(), nn.Linear(hidden_size, hidden_size), ) def forward(self, current: torch.Tensor, cached: torch.Tensor) -> torch.Tensor: if current.shape != cached.shape: raise ValueError(f"Fusion shape mismatch: {current.shape} != {cached.shape}") return self.mlp( torch.cat([self.current_norm(current), self.cached_norm(cached)], dim=-1) ) class HYWorldPlayPredictor(nn.Module): """Two full double-stream blocks over fused current/cached image and txt features.""" def __init__(self, config: PredictorConfig | None = None) -> None: super().__init__() self.predictor_config = config or PredictorConfig() cfg = self.predictor_config self.img_in = PatchEmbed( list(cfg.patch_size), cfg.in_channels, cfg.hidden_size, is_reshape_temporal_channels=False, concat_condition=cfg.concat_condition, ) self.img_fusion = FeatureFusion(cfg.hidden_size) self.txt_fusion = FeatureFusion(cfg.hidden_size) self.double_blocks = nn.ModuleList( [ MMDoubleStreamBlock( cfg.hidden_size, cfg.heads_num, mlp_width_ratio=cfg.mlp_width_ratio, mlp_act_type=cfg.mlp_act_type, attn_mode=cfg.attn_mode, qk_norm=cfg.qk_norm, qk_norm_type=cfg.qk_norm_type, qkv_bias=cfg.qkv_bias, ) for _ in cfg.source_block_ids ] ) # HY-WorldPlay action checkpoints add the ProPE output projection after # constructing the base Hunyuan blocks (see transformer.add_action_parameters). for block in self.double_blocks: block.img_attn_prope_proj = nn.Linear( cfg.hidden_size, cfg.hidden_size, bias=cfg.qkv_bias ) self.residual_out = nn.Linear(cfg.hidden_size, cfg.hidden_size) self.final_layer = FinalLayer( cfg.hidden_size, list(cfg.patch_size), cfg.out_channels, get_activation_layer("silu"), ) nn.init.zeros_(self.residual_out.weight) nn.init.zeros_(self.residual_out.bias) self.gradient_checkpointing = False self.attn_param: dict[str, Any] = { "thw": [cfg.latent_frames, cfg.latent_height, cfg.latent_width], "win_type": "fixed", "win_ratio": 0, } self.img_in.requires_grad_(False) self.final_layer.requires_grad_(False) @property def config_dict(self) -> dict[str, Any]: return asdict(self.predictor_config) def enable_gradient_checkpointing(self, enabled: bool = True) -> None: self.gradient_checkpointing = enabled def train(self, mode: bool = True): super().train(mode) self.img_in.eval() self.final_layer.eval() return self @staticmethod def _load_prefixed_module( module: nn.Module, handle, prefix: str, ) -> None: target = module.state_dict() loaded = {} for name in target: key = prefix + name if key not in handle.keys(): raise KeyError(f"Missing Teacher checkpoint key: {key}") loaded[name] = handle.get_tensor(key) result = module.load_state_dict(loaded, strict=True) if result.missing_keys or result.unexpected_keys: raise RuntimeError(f"Unexpected load result for {prefix}: {result}") def load_teacher_initialization(self, checkpoint_path: str | Path) -> None: path = str(Path(checkpoint_path).resolve()) with safe_open(path, framework="pt", device="cpu") as handle: self._load_prefixed_module(self.img_in, handle, "img_in.") self._load_prefixed_module(self.final_layer, handle, "final_layer.") for predictor_block, teacher_id in zip( self.double_blocks, self.predictor_config.source_block_ids ): self._load_prefixed_module( predictor_block, handle, f"double_blocks.{teacher_id}." ) self.img_in.requires_grad_(False).eval() self.final_layer.requires_grad_(False).eval() nn.init.zeros_(self.residual_out.weight) nn.init.zeros_(self.residual_out.bias) def _vision_rope( self, rope_temporal_size: int, start_rope_start_idx: int, *, device: torch.device, dtype: torch.dtype, ) -> tuple[torch.Tensor, torch.Tensor]: cfg = self.predictor_config cos, sin = get_nd_rotary_pos_embed( list(cfg.rope_dim_list), (rope_temporal_size, cfg.latent_height, cfg.latent_width), theta=cfg.rope_theta, use_real=True, theta_rescale_factor=1, ) tokens_per_frame = cfg.latent_height * cfg.latent_width start = start_rope_start_idx * tokens_per_frame end = (start_rope_start_idx + cfg.latent_frames) * tokens_per_frame cos = cos[start:end].to(device=device, dtype=dtype) sin = sin[start:end].to(device=device, dtype=dtype) expected = cfg.latent_frames * tokens_per_frame if cos.shape[0] != expected: raise ValueError(f"RoPE tokens {cos.shape[0]} != {expected}") return cos, sin def _run_block( self, block: MMDoubleStreamBlock, source_block_id: int, img: torch.Tensor, txt: torch.Tensor, vec: torch.Tensor, vec_txt: torch.Tensor, freqs_cis: tuple[torch.Tensor, torch.Tensor], viewmats: torch.Tensor, Ks: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: def block_forward(img_arg: torch.Tensor, txt_arg: torch.Tensor): return block( bi_inference=True, ar_txt_inference=False, ar_vision_inference=False, img=img_arg, txt=txt_arg, vec_txt=vec_txt, vec=vec, freqs_cis=freqs_cis, text_mask=None, attn_param=self.attn_param, is_flash=False, block_idx=source_block_id, viewmats=viewmats, Ks=Ks, ) if self.gradient_checkpointing and self.training: return checkpoint(block_forward, img, txt, use_reentrant=False) return block_forward(img, txt) def unpatchify(self, x: torch.Tensor) -> torch.Tensor: cfg = self.predictor_config batch = x.shape[0] expected = cfg.latent_frames * cfg.latent_height * cfg.latent_width if x.shape[1] != expected: raise ValueError(f"Output tokens {x.shape[1]} != {expected}") x = x.reshape( batch, cfg.latent_frames, cfg.latent_height, cfg.latent_width, cfg.out_channels, *cfg.patch_size, ) x = torch.einsum("nthwcopq->nctohpwq", x) return x.reshape( batch, cfg.out_channels, cfg.latent_frames * cfg.patch_size[0], cfg.latent_height * cfg.patch_size[1], cfg.latent_width * cfg.patch_size[2], ) def hidden_to_velocity( self, hidden: torch.Tensor, frame_condition: torch.Tensor, ) -> torch.Tensor: cfg = self.predictor_config spatial_tokens = cfg.latent_height * cfg.latent_width condition_tokens = frame_condition.repeat_interleave(spatial_tokens, dim=1) vec = condition_tokens.reshape(-1, cfg.hidden_size) return self.unpatchify(self.final_layer(hidden, vec)) def forward( self, *, target_model_input: torch.Tensor, anchor_hidden: torch.Tensor, current_txt: torch.Tensor, cached_txt: torch.Tensor, target_frame_condition: torch.Tensor, vec_txt: torch.Tensor, target_viewmats: torch.Tensor, target_Ks: torch.Tensor, rope_temporal_size: torch.Tensor | int, start_rope_start_idx: torch.Tensor | int, ) -> dict[str, torch.Tensor]: cfg = self.predictor_config batch = target_model_input.shape[0] if batch != 1: raise ValueError("Predictor v1 currently requires micro-batch 1") if target_model_input.shape[1:] != ( 65, cfg.latent_frames, cfg.latent_height, cfg.latent_width, ): raise ValueError(f"Unexpected target_model_input: {target_model_input.shape}") with torch.no_grad(): current_img = self.img_in(target_model_input) if current_img.shape != anchor_hidden.shape: raise ValueError(f"Current/anchor mismatch: {current_img.shape} != {anchor_hidden.shape}") if current_txt.shape[0] != batch: current_txt = current_txt.expand(batch, -1, -1) if cached_txt.shape[0] != batch: cached_txt = cached_txt.expand(batch, -1, -1) if vec_txt.shape[0] != batch: vec_txt = vec_txt.expand(batch, -1) img = self.img_fusion(current_img, anchor_hidden) txt = self.txt_fusion(current_txt, cached_txt) spatial_tokens = cfg.latent_height * cfg.latent_width if target_frame_condition.shape != (batch, cfg.latent_frames, cfg.hidden_size): raise ValueError(f"Unexpected frame condition: {target_frame_condition.shape}") condition_tokens = target_frame_condition.repeat_interleave(spatial_tokens, dim=1) vec = condition_tokens.reshape(-1, cfg.hidden_size) viewmats = repeat( target_viewmats, "B T M N -> B (T H W) M N", H=cfg.latent_height, W=cfg.latent_width, ) Ks = repeat( target_Ks, "B T M N -> B (T H W) M N", H=cfg.latent_height, W=cfg.latent_width, ) rope_size = int(rope_temporal_size.reshape(-1)[0].item()) if torch.is_tensor(rope_temporal_size) else int(rope_temporal_size) rope_start = int(start_rope_start_idx.reshape(-1)[0].item()) if torch.is_tensor(start_rope_start_idx) else int(start_rope_start_idx) freqs_cis = self._vision_rope( rope_size, rope_start, device=img.device, dtype=img.dtype, ) self.attn_param["thw"] = [cfg.latent_frames, cfg.latent_height, cfg.latent_width] for block, source_id in zip(self.double_blocks, cfg.source_block_ids): img, txt = self._run_block( block, source_id, img, txt, vec, vec_txt, freqs_cis, viewmats, Ks, ) delta_hidden = self.residual_out(img) pred_hidden = anchor_hidden + delta_hidden # Frozen parameters still allow the velocity loss to backpropagate to pred_hidden. pred_tokens = self.final_layer(pred_hidden, vec) pred_velocity = self.unpatchify(pred_tokens) return { "pred_hidden": pred_hidden, "pred_velocity": pred_velocity, "delta_hidden": delta_hidden, "pred_txt": txt, } def trainable_parameter_count(self) -> int: return sum(parameter.numel() for parameter in self.parameters() if parameter.requires_grad) def trainable_parameter_breakdown(self) -> dict[str, int]: groups = { "img_fusion": self.img_fusion, "txt_fusion": self.txt_fusion, "double_blocks": self.double_blocks, "residual_out": self.residual_out, } return { name: sum(parameter.numel() for parameter in module.parameters() if parameter.requires_grad) for name, module in groups.items() }