"""MiniMax H3 learned latent upscaler inference. Adapted from LBH-123-AI/Comfyui_Minimax_h3_latent_Upscaler (Apache-2.0). """ from pathlib import Path import re import torch from safetensors.torch import load_file from torch import nn from torch.nn import functional as F LATENTS_MEAN = ( 0.858090341091156, -0.9606591463088989, 1.0661640167236328, -0.5090325474739075, -0.2727581858634949, -1.3675414323806763, -0.2553254961967468, -0.26907554268836975, -0.5376840829849243, -0.0464097298681736, 0.6657370328903198, 0.19690127670764923, -0.5460608005523682, -0.4035342037677765, -0.23683024942874908, 0.25928452610969543, -0.30133944749832153, 0.211341992020607, -1.1206848621368408, 0.3581933379173279, -0.04225143790245056, 0.2604829967021942, 0.22864092886447906, 0.7056031823158264, ) LATENTS_STD = ( 1.2223774194717407, 1.2767263650894165, 1.6831774711608887, 1.7549455165863037, 1.5636216402053833, 2.194143533706665, 0.9653137922286987, 1.0569885969161987, 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, 0.7996809482574463, 0.44988900423049927, 0.7197399735450745, 0.6936293244361877, 2.961095094680786, 2.7694199085235596, 3.0496184825897217, 2.1088054180145264, 3.276226282119751, 3.1627357006073, 2.2816812992095947, 2.6127843856811523, ) def _normalization(channels: int) -> nn.GroupNorm: return nn.GroupNorm(32, channels) class ResBlockEmb3D(nn.Module): def __init__(self, channels: int, emb_channels: int, dropout: float = 0.0): super().__init__() self.in_layers = nn.Sequential( _normalization(channels), nn.SiLU(), nn.Conv3d(channels, channels, 3, padding=1), ) self.emb_layers = nn.Sequential(nn.SiLU(), nn.Linear(emb_channels, 2 * channels)) self.out_norm = _normalization(channels) self.out_layers = nn.Sequential( nn.SiLU(), nn.Dropout(p=dropout), nn.Conv3d(channels, channels, 3, padding=1), ) nn.init.zeros_(self.out_layers[-1].weight) nn.init.zeros_(self.out_layers[-1].bias) def forward(self, x: torch.Tensor, emb: torch.Tensor) -> torch.Tensor: hidden = self.in_layers(x) scale, shift = self.emb_layers(emb).to(hidden.dtype).chunk(2, dim=1) hidden = self.out_norm(hidden) * (1 + scale[:, :, None, None, None]) + shift[:, :, None, None, None] return x + self.out_layers(hidden) class TemporalConv(nn.Module): def __init__(self, channels: int, kernel_size: int = 5): super().__init__() self.norm = _normalization(channels) self.dwconv = nn.Conv3d( channels, channels, kernel_size=(kernel_size, 1, 1), padding=(kernel_size // 2, 0, 0), groups=channels, ) self.pwconv = nn.Conv3d(channels, channels, 1) nn.init.zeros_(self.pwconv.weight) nn.init.zeros_(self.pwconv.bias) def forward(self, x: torch.Tensor) -> torch.Tensor: hidden = self.dwconv(F.silu(self.norm(x))) return x + self.pwconv(hidden) class H3LatentResizer3D(nn.Module): def __init__( self, in_channels: int = 24, in_blocks: int = 12, out_blocks: int = 12, channels: int = 512, dropout: float = 0.1, temporal_every: int = 2, temporal_kernel: int = 5, ): super().__init__() self.conv_in = nn.Conv3d(in_channels, channels, 3, padding=1) embed_dim = 64 self.embed = nn.Sequential(nn.Linear(1, embed_dim), nn.SiLU(), nn.Linear(embed_dim, embed_dim)) self.in_blocks = self._make_blocks(in_blocks, channels, embed_dim, dropout, temporal_every, temporal_kernel) self.out_blocks = self._make_blocks(out_blocks, channels, embed_dim, dropout, temporal_every, temporal_kernel) self.norm_out = _normalization(channels) self.conv_out = nn.Conv3d(channels, in_channels, 3, padding=1) @staticmethod def _make_blocks(count, channels, embed_dim, dropout, temporal_every, temporal_kernel): blocks = nn.ModuleList() for index in range(count): blocks.append(ResBlockEmb3D(channels, embed_dim, dropout)) if temporal_every > 0 and index % temporal_every == 0: blocks.append(TemporalConv(channels, temporal_kernel)) return blocks def forward(self, x: torch.Tensor, *, scale: float, target_size: tuple[int, int, int]) -> torch.Tensor: emb = self.embed(torch.tensor([[scale - 1]], dtype=x.dtype, device=x.device)).expand(x.shape[0], -1) hidden = self.conv_in(x) for block in self.in_blocks: hidden = block(hidden, emb) if isinstance(block, ResBlockEmb3D) else block(hidden) hidden = F.interpolate(hidden, size=target_size, mode="trilinear", align_corners=False) for block in self.out_blocks: hidden = block(hidden, emb) if isinstance(block, ResBlockEmb3D) else block(hidden) return self.conv_out(F.silu(self.norm_out(hidden))) def _checkpoint_config(state: dict[str, torch.Tensor]) -> dict: in_ids = {int(match.group(1)) for key in state if (match := re.match(r"in_blocks\.(\d+)\.in_layers\.", key))} out_ids = {int(match.group(1)) for key in state if (match := re.match(r"out_blocks\.(\d+)\.in_layers\.", key))} temporal_keys = [key for key in state if key.endswith("dwconv.weight")] conv_in = state["conv_in.weight"] return { "in_channels": conv_in.shape[1], "in_blocks": len(in_ids), "out_blocks": len(out_ids), "channels": conv_in.shape[0], "temporal_every": 2 if temporal_keys else 0, "temporal_kernel": state[temporal_keys[0]].shape[2] if temporal_keys else 5, } def load_h3_latent_upscaler( path: str | Path, *, device: str | torch.device = "cuda", dtype: torch.dtype = torch.float16, ) -> H3LatentResizer3D: state = load_file(str(path), device="cpu") if any(key.startswith("upscaler.") for key in state): state = {key.removeprefix("upscaler."): value for key, value in state.items() if key.startswith("upscaler.")} model = H3LatentResizer3D(**_checkpoint_config(state)) model.load_state_dict(state, strict=True) return model.to(device=device, dtype=dtype).eval().requires_grad_(False) @torch.inference_mode() def upscale_h3_latent(model: H3LatentResizer3D, latent: torch.Tensor, *, scale: float = 2.0) -> torch.Tensor: if latent.ndim != 5 or latent.shape[1] != 24: raise ValueError("H3 latent must have shape [B,24,T,H,W]") if not 1.0 <= scale <= 4.0: raise ValueError("scale must be between 1.0 and 4.0") dtype = next(model.parameters()).dtype device = next(model.parameters()).device source = latent.to(device=device, dtype=dtype) if scale == 1.0: return source mean = torch.tensor(LATENTS_MEAN, device=device, dtype=dtype).view(1, 24, 1, 1, 1) std = torch.tensor(LATENTS_STD, device=device, dtype=dtype).view(1, 24, 1, 1, 1) target_size = (source.shape[2], round(source.shape[3] * scale), round(source.shape[4] * scale)) result = model((source - mean) / std, scale=scale, target_size=target_size) return result.mul_(std).add_(mean)