Instructions to use Cccccz/HY with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use Cccccz/HY with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("Cccccz/HY", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download models/predictor.py from Cccccz/HY: direct link, hf CLI and curl.
- Browser
- Download file 13.8 kB
-
https://huggingface.co/Cccccz/HY/resolve/main/models/predictor.py
- Command line
-
hf download hf://Cccccz/HY/models/predictor.py
-
curl -L -o predictor.py https://huggingface.co/Cccccz/HY/resolve/main/models/predictor.py
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 | |
| 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) | |
| 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 | |
| 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() | |
| } | |