from functools import lru_cache import glob import math import os import re import comfy.samplers import folder_paths import torch try: import torch.nn as nn import torch.nn.functional as F except Exception: # pragma: no cover - import-time fallback for the test shim nn = None F = None 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.1088054180265264, 3.276226282119751, 3.1627357006073, 2.2816812992095947, 2.6127843856811523, ] _LATENT_UPSCALE_FOLDER = "latent_upscale_models" MP_UNIT = 1024 * 1024 RES_MULTIPLE = 32 def _models_dir(): try: if _LATENT_UPSCALE_FOLDER not in folder_paths.folder_names_and_paths: folder_paths.add_model_folder_path( _LATENT_UPSCALE_FOLDER, os.path.join(folder_paths.models_dir, _LATENT_UPSCALE_FOLDER), ) return folder_paths.get_folder_paths(_LATENT_UPSCALE_FOLDER)[0] except Exception: return os.path.join(getattr(folder_paths, "models_dir", ""), _LATENT_UPSCALE_FOLDER) def _scan_models(): try: model_dir = _models_dir() files = [] for ext in ("*.pth", "*.safetensors"): files.extend(glob.glob(os.path.join(model_dir, ext))) names = sorted(os.path.basename(path) for path in files) return ["none"] + names if names else [f"(no upscale models found in: {model_dir})"] except Exception: return ["none"] def _make_norm_tensors(device, dtype): mean = torch.tensor(LATENTS_MEAN, dtype=dtype, device=device).view(1, -1, 1, 1, 1) std = torch.tensor(LATENTS_STD, dtype=dtype, device=device).view(1, -1, 1, 1, 1) return mean, std if nn is not None: def _normalization(channels): return nn.GroupNorm(32, channels) def _zero_module(module): for p in module.parameters(): p.detach().zero_() return module class _AttnBlock3D(nn.Module): def __init__(self, in_channels): super().__init__() self.norm = _normalization(in_channels) self.q = nn.Conv3d(in_channels, in_channels, 1) self.k = nn.Conv3d(in_channels, in_channels, 1) self.v = nn.Conv3d(in_channels, in_channels, 1) self.proj_out = nn.Conv3d(in_channels, in_channels, 1) def forward(self, x): h = self.norm(x) b, c, t, hh, w = h.shape q = self.q(h).flatten(2).transpose(1, 2) k = self.k(h).flatten(2).transpose(1, 2) v = self.v(h).flatten(2).transpose(1, 2) h = F.scaled_dot_product_attention(q, k, v) h = h.transpose(1, 2).view(b, c, t, hh, w) return x + self.proj_out(h) class _ResBlockEmb3D(nn.Module): def __init__(self, channels, emb_channels, dropout=0, out_channels=None): super().__init__() self.out_channels = out_channels or channels self.in_layers = nn.Sequential( _normalization(channels), nn.SiLU(), nn.Conv3d(channels, self.out_channels, 3, padding=1), ) self.emb_layers = nn.Sequential( nn.SiLU(), nn.Linear(emb_channels, 2 * self.out_channels), ) self.out_norm = _normalization(self.out_channels) self.out_layers = nn.Sequential( nn.SiLU(), nn.Dropout(p=dropout), _zero_module(nn.Conv3d(self.out_channels, self.out_channels, 3, padding=1)), ) self.skip = ( nn.Conv3d(channels, self.out_channels, 1) if self.out_channels != channels else nn.Identity() ) def forward(self, x, emb): h = self.in_layers(x) emb_out = self.emb_layers(emb).type(h.dtype) while len(emb_out.shape) < len(h.shape): emb_out = emb_out[..., None] scale, shift = torch.chunk(emb_out, 2, dim=1) h = self.out_norm(h) * (1 + scale) + shift h = self.out_layers(h) return self.skip(x) + h class _TemporalConv(nn.Module): def __init__(self, channels, kernel_size=5): super().__init__() padding = kernel_size // 2 self.norm = _normalization(channels) self.dwconv = nn.Conv3d( channels, channels, kernel_size=(kernel_size, 1, 1), padding=(padding, 0, 0), groups=channels, ) self.pwconv = nn.Conv3d(channels, channels, kernel_size=1) nn.init.zeros_(self.pwconv.weight) nn.init.zeros_(self.pwconv.bias) def forward(self, x): identity = x h = self.norm(x) h = F.silu(h) h = self.dwconv(h) h = self.pwconv(h) return identity + h class _LatentResizer3D(nn.Module): def __init__(self, in_channels=24, in_blocks=12, out_blocks=12, channels=512, dropout=0.1, attn=False, temporal_every=2, temporal_kernel=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 = nn.ModuleList() for b in range(in_blocks): if (b == 1 or b == in_blocks - 1) and attn: self.in_blocks.append(_AttnBlock3D(channels)) self.in_blocks.append(_ResBlockEmb3D(channels, embed_dim, dropout)) if temporal_every > 0 and b % temporal_every == 0: self.in_blocks.append(_TemporalConv(channels, temporal_kernel)) self.out_blocks = nn.ModuleList() for b in range(out_blocks): if (b == 1 or b == out_blocks - 1) and attn: self.out_blocks.append(_AttnBlock3D(channels)) self.out_blocks.append(_ResBlockEmb3D(channels, embed_dim, dropout)) if temporal_every > 0 and b % temporal_every == 0: self.out_blocks.append(_TemporalConv(channels, temporal_kernel)) self.norm_out = _normalization(channels) self.conv_out = nn.Conv3d(channels, in_channels, 3, padding=1) def forward(self, x, scale=None, target_size=None): if target_size is not None: size = target_size elif scale is not None: size = tuple(int(round(s * scale)) for s in x.shape[-3:]) else: return x if size == x.shape[-3:]: return x scale_emb = torch.tensor( [scale - 1 if scale is not None else 0.0], dtype=x.dtype, device=x.device).unsqueeze(0) emb = self.embed(scale_emb) x = self.conv_in(x) for block in self.in_blocks: if isinstance(block, _ResBlockEmb3D): x = block(x, emb.expand(x.shape[0], -1)) else: x = block(x) x = F.interpolate(x, size=size, mode="trilinear", align_corners=False) for block in self.out_blocks: if isinstance(block, _ResBlockEmb3D): x = block(x, emb.expand(x.shape[0], -1)) else: x = block(x) x = self.norm_out(x) x = F.silu(x) x = self.conv_out(x) return x else: # pragma: no cover - import-time fallback for the test shim _LatentResizer3D = None _MODEL_CACHE = {} def _load_raw_sd(path): if path.endswith(".safetensors"): from safetensors.torch import load_file sd = load_file(path, device="cpu") else: sd = torch.load(path, map_location="cpu", weights_only=False) if isinstance(sd, dict) and "model" in sd: sd = sd["model"] float8 = getattr(torch, "float8_e4m3fn", None) if float8 is not None: sd = {k: v.to(torch.float16) if getattr(v, "dtype", None) == float8 else v for k, v in sd.items()} return sd def _extract_upscaler_sd(sd): if any(k.startswith("upscaler.") for k in sd): return {k[len("upscaler."):]: v for k, v in sd.items() if k.startswith("upscaler.")} return sd def _detect_arch(sd): cfg = { "in_channels": 24, "in_blocks": 12, "out_blocks": 12, "channels": 512, "dropout": 0.1, "attn": False, "temporal_every": 2, "temporal_kernel": 5, } conv_key = "conv_in.weight" if conv_key in sd: cfg["in_channels"] = sd[conv_key].shape[1] cfg["channels"] = sd[conv_key].shape[0] in_ids, out_ids = set(), set() temporal_in_indices, temporal_out_indices = set(), set() for k in sd.keys(): m = re.match(r"in_blocks\.(\d+)\.in_layers\.", k) if m: in_ids.add(int(m.group(1))) m = re.match(r"out_blocks\.(\d+)\.in_layers\.", k) if m: out_ids.add(int(m.group(1))) m = re.match(r"in_blocks\.(\d+)\.dwconv\.weight", k) if m: temporal_in_indices.add(int(m.group(1))) m = re.match(r"out_blocks\.(\d+)\.dwconv\.weight", k) if m: temporal_out_indices.add(int(m.group(1))) if in_ids: cfg["in_blocks"] = len(in_ids) if out_ids: cfg["out_blocks"] = len(out_ids) if temporal_in_indices or temporal_out_indices: cfg["temporal_every"] = 2 for k in sd.keys(): if "dwconv.weight" in k and k.endswith("dwconv.weight"): cfg["temporal_kernel"] = sd[k].shape[2] break else: cfg["temporal_every"] = 0 cfg["attn"] = False return cfg def load_upscale_model(name, device, precision): if _LatentResizer3D is None: raise RuntimeError("latent upscaler requires torch.nn") cache_key = f"{name}::{device}::{precision}" if cache_key in _MODEL_CACHE: return _MODEL_CACHE[cache_key].to(device) path = os.path.join(_models_dir(), name) if not os.path.exists(path): raise FileNotFoundError(f"Model file not found: {path}") raw_sd = _load_raw_sd(path) up_sd = _extract_upscaler_sd(raw_sd) cfg = _detect_arch(up_sd) if cfg["in_channels"] != 24: raise ValueError( f"Checkpoint '{name}' is not an H3 latent upscaler (expected 24 input channels, got {cfg['in_channels']})." ) model = _LatentResizer3D( in_channels=cfg["in_channels"], in_blocks=cfg["in_blocks"], out_blocks=cfg["out_blocks"], channels=cfg["channels"], dropout=cfg["dropout"], attn=cfg["attn"], temporal_every=cfg["temporal_every"], temporal_kernel=cfg["temporal_kernel"], ) model.load_state_dict(up_sd, strict=True) dtype = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}.get(precision, torch.float32) model = model.to(device).eval().requires_grad_(False) if dtype != torch.float32: model = model.to(dtype) _MODEL_CACHE[cache_key] = model return model def unload_upscale_model(name, device, precision): cache_key = f"{name}::{device}::{precision}" model = _MODEL_CACHE.get(cache_key) if model is not None and str(next(model.parameters()).device) != "cpu": model.to("cpu") if str(device) == "cuda" and hasattr(torch, "cuda"): try: torch.cuda.empty_cache() except Exception: pass def _compute_upscale_target(width, height, h_in, w_in): ds = 16 w_px = float(width) h_px = float(height) eff = (w_px / (w_in * ds) + h_px / (h_in * ds)) / 2.0 w_px_f = round(w_px / ds) * ds h_px_f = round(h_px / ds) * ds w_out = max(1, int(w_px_f // ds)) h_out = max(1, int(h_px_f // ds)) return h_out, w_out, eff def _scale_to_megapixels(w, h, mp, multiple=RES_MULTIPLE): if not mp or mp <= 0 or w <= 0 or h <= 0: return int(h), int(w) multiple = max(1, int(multiple)) scale = math.sqrt((float(mp) * MP_UNIT) / float(w * h)) nw = max(multiple, int(round(w * scale / multiple)) * multiple) nh = max(multiple, int(round(h * scale / multiple)) * multiple) return nh, nw def _resolve_target_size(param, h_in, w_in): width = int(param.get("width", 0) or 0) height = int(param.get("height", 0) or 0) megapixels = float(param.get("megapixels", 0.0) or 0.0) if width > 0 and height > 0: return height, width if megapixels > 0: return _scale_to_megapixels(w_in, h_in, megapixels) return int(h_in), int(w_in) def upscale_video_model(video, param): model_name = param["model_name"] device = param.get("device", "cuda") precision = param.get("precision", "fp16") orig_dtype = video.dtype dev = torch.device(device if (device == "cpu" or torch.cuda.is_available()) else "cpu") compute_dtype = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}[precision] _, c, t, h_in, w_in = video.shape h_out, w_out = _resolve_target_size(param, h_in, w_in) eff = (w_out / float(w_in) + h_out / float(h_in)) / 2.0 if w_in and h_in else 1.0 if eff < 1.0 and (w_out < w_in or h_out < h_in): raise ValueError("This model only supports upscaling (effective scale >= 1.0).") if w_out == w_in and h_out == h_in: return video, h_in, w_in if str(model_name).startswith("("): raise ValueError("Please place H3 upscale model files into the latent_upscale_models directory") s = video.to(device=dev, dtype=compute_dtype, copy=True) model = load_upscale_model(model_name, dev, precision) norm_mean, norm_std = _make_norm_tensors(dev, compute_dtype) with torch.inference_mode(): s = s.sub(norm_mean).div(norm_std) out = model(s, scale=eff, target_size=(t, h_out, w_out)) del s out = out.mul(norm_std).add(norm_mean) out = out.to(device="cpu", dtype=orig_dtype) unload_upscale_model(model_name, dev, precision) return out, h_out, w_out def upscale_video_interp(video, param): method = str(param.get("method") or "bilinear") _, c, t, h_in, w_in = video.shape h_out, w_out = _resolve_target_size(param, h_in, w_in) if h_out == h_in and w_out == w_in: return video, h_in, w_in video_bt = video.permute(0, 2, 1, 3, 4).reshape(-1, c, h_in, w_in) up = F.interpolate(video_bt, size=(h_out, w_out), mode=method) up = up.reshape(video.shape[0], t, c, h_out, w_out).permute(0, 2, 1, 3, 4).contiguous() return up, h_out, w_out def upscale_latent_video(video, param): mode = str(param.get("mode") or "off") if mode == "off": return video, video.shape[-2], video.shape[-1] if mode == "model": return upscale_video_model(video, param) return upscale_video_interp(video, param) class H3LatentUpscaleParams: CATEGORY = "Dumas/MiniMax" FUNCTION = "build" RETURN_TYPES = ("DUMAS_H3_LATENT_UPSCALE_PARAM",) RETURN_NAMES = ("latent_upscale_param",) @classmethod def INPUT_TYPES(cls): return { "required": { "mode": (["off", "model", "interp"], {"default": "off", "tooltip": "Latent refinement mode. 'off' skips the stage, 'model' uses the H3 latent upscaler model, 'interp' uses model-free interpolation."}), "model_name": (_scan_models(), { "default": "none", "tooltip": "H3 latent upscale checkpoint from models/latent_upscale_models, used when mode = model."}), "method": (["nearest-exact", "bilinear", "area", "bicubic"], {"default": "bilinear", "tooltip": "Interpolation method used when mode = interp."}), "width": ("INT", {"default": 0, "min": 0, "max": 4096, "step": 32, "tooltip": "Explicit target width for the latent refinement stage. Leave at 0 to let megapixels choose the size instead."}), "height": ("INT", {"default": 0, "min": 0, "max": 4096, "step": 32, "tooltip": "Explicit target height for the latent refinement stage. Leave at 0 to let megapixels choose the size instead."}), "device": (["cuda", "cpu"], {"default": "cuda", "tooltip": "Device used by the H3 latent upscaler model when mode = model."}), "precision": (["fp16", "fp32", "bf16"], {"default": "fp16", "tooltip": "Computation precision used by the H3 latent upscaler model when mode = model."}), "sampler_name": (comfy.samplers.KSampler.SAMPLERS, {"default": "euler_ancestral", "tooltip": "Sampler used for the latent refinement pass. Default matches the current H3 preference."}), "scheduler": (comfy.samplers.KSampler.SCHEDULERS, {"default": "simple", "tooltip": "Scheduler used for the latent refinement pass."}), "steps": ("INT", {"default": 2, "min": 1, "max": 50, "step": 1, "tooltip": "Number of refinement steps applied after the latent upscaler stage."}), "denoise": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "How much the refinement pass may rewrite the upscaled latent. Lower = safer, higher = freer."}), "megapixels": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 4.0, "step": 0.01, "tooltip": "Primary target size for the latent refinement stage. If width and height are both set, they win; otherwise the node scales the current shot to this pixel budget while preserving aspect ratio. 0 keeps the incoming latent size."}), } } def build(self, mode, model_name, method, width, height, device, precision, sampler_name, scheduler, steps, denoise, megapixels): width = int(width) height = int(height) steps = int(steps) if mode == "off": return ({ "mode": "off", "width": width, "height": height, "sampler_name": sampler_name, "scheduler": scheduler, "steps": steps, "denoise": float(denoise), "refine_denoise": float(denoise), "megapixels": float(megapixels), },) if width > 0: width = int(round(width / 32.0)) * 32 if height > 0: height = int(round(height / 32.0)) * 32 return ({ "mode": mode, "model_name": model_name, "method": method, "width": width, "height": height, "device": device, "precision": precision, "sampler_name": sampler_name, "scheduler": scheduler, "steps": steps, "denoise": float(denoise), "refine_denoise": float(denoise), "megapixels": float(megapixels), },) NODE_CLASS_MAPPINGS = {"DumasH3LatentUpscaleParams": H3LatentUpscaleParams} NODE_DISPLAY_NAME_MAPPINGS = {"DumasH3LatentUpscaleParams": "Dumas H3 Latent Upscale Params"} __all__ = [ "NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "upscale_latent_video", ]