Diffusers
Safetensors
HY / models /predictor.py
Cccccz's picture
Upload batch 46: 223 files (7.45 GiB)
ba798d3 verified
Raw History Blame Contribute Delete
13.8 kB
"""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()
}