From 2a610155821d670a2d8047e654e5fce96b790eb5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Tue, 23 Jun 2026 00:35:00 +0300 Subject: [PATCH 001/211] feat: Support Krea2 (#14589) --- comfy/ldm/krea2/model.py | 290 +++++++++++++++++++++++++++++++++++ comfy/lora.py | 11 ++ comfy/model_base.py | 12 ++ comfy/model_detection.py | 15 ++ comfy/sd.py | 6 + comfy/supported_models.py | 31 ++++ comfy/text_encoders/krea2.py | 84 ++++++++++ comfy/utils.py | 38 +++++ nodes.py | 2 +- 9 files changed, 488 insertions(+), 1 deletion(-) create mode 100644 comfy/ldm/krea2/model.py create mode 100644 comfy/text_encoders/krea2.py diff --git a/comfy/ldm/krea2/model.py b/comfy/ldm/krea2/model.py new file mode 100644 index 000000000..ecb16254f --- /dev/null +++ b/comfy/ldm/krea2/model.py @@ -0,0 +1,290 @@ +"""Krea 2 (K2) — single-stream MMDiT. + +Text tokens produced by a Qwen3-VL-4B 12-layer ``txtfusion`` adapter and patchified image tokens are +concatenated into one sequence and run through ``layers`` shared transformer blocks with +AdaLN-single modulation, GQA + per-head QK-norm + sigmoid-gated attention, SwiGLU MLP, and 3-axis RoPE. +""" + +from typing import Optional + +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange + +import comfy.model_management +import comfy.patcher_extension +import comfy.ldm.common_dit +from comfy.ldm.flux.layers import EmbedND, timestep_embedding +from comfy.ldm.flux.math import apply_rope +from comfy.ldm.modules.attention import optimized_attention_masked + + +class RMSNorm(nn.Module): + """RMSNorm with the reference ``(1 + scale)`` weight convention (scale stored zero-centered).""" + + def __init__(self, features: int, eps: float = 1e-5, device=None, dtype=None, operations=None): + super().__init__() + self.eps = eps + self.scale = nn.Parameter(torch.empty(features, device=device, dtype=dtype)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + weight = comfy.model_management.cast_to(self.scale, dtype=torch.float32, device=x.device) + 1.0 + return F.rms_norm(x.float(), (x.shape[-1],), weight=weight, eps=self.eps).to(dtype) + + +class QKNorm(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.qnorm = RMSNorm(dim, device=device, dtype=dtype, operations=operations) + self.knorm = RMSNorm(dim, device=device, dtype=dtype, operations=operations) + + def forward(self, q, k): + return self.qnorm(q), self.knorm(k) + + +class SwiGLU(nn.Module): + def __init__(self, features: int, multiplier: int, bias: bool = False, multiple: int = 128, + device=None, dtype=None, operations=None): + super().__init__() + mlpdim = int(2 * features / 3) * multiplier + mlpdim = multiple * ((mlpdim + multiple - 1) // multiple) + self.gate = operations.Linear(features, mlpdim, bias=bias, device=device, dtype=dtype) + self.up = operations.Linear(features, mlpdim, bias=bias, device=device, dtype=dtype) + self.down = operations.Linear(mlpdim, features, bias=bias, device=device, dtype=dtype) + + def forward(self, x): + return self.down(F.silu(self.gate(x)).mul_(self.up(x))) + + +class Attention(nn.Module): + def __init__(self, dim: int, heads: int, kvheads: Optional[int] = None, bias: bool = False, + device=None, dtype=None, operations=None): + super().__init__() + self.heads = heads + self.kvheads = kvheads if kvheads is not None else heads + self.headdim = dim // self.heads + self.wq = operations.Linear(dim, self.headdim * self.heads, bias=bias, device=device, dtype=dtype) + self.wk = operations.Linear(dim, self.headdim * self.kvheads, bias=bias, device=device, dtype=dtype) + self.wv = operations.Linear(dim, self.headdim * self.kvheads, bias=bias, device=device, dtype=dtype) + self.gate = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) + self.qknorm = QKNorm(self.headdim, device=device, dtype=dtype, operations=operations) + self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) + + def forward(self, x, freqs=None, mask=None, transformer_options={}): + q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x) + q = rearrange(q, "B L (H D) -> B H L D", H=self.heads) + k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads) + v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads) + q, k = self.qknorm(q, k) + if freqs is not None: + q, k = apply_rope(q, k, freqs) + if self.kvheads != self.heads: + rep = self.heads // self.kvheads + k = k.repeat_interleave(rep, dim=1) + v = v.repeat_interleave(rep, dim=1) + out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True, + transformer_options=transformer_options) + return self.wo(out * F.sigmoid(gate)) + + +class SimpleModulation(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.lin = nn.Parameter(torch.empty(2, dim, device=device, dtype=dtype)) + + def forward(self, vec): + out = vec + comfy.model_management.cast_to(self.lin, dtype=vec.dtype, device=vec.device).unsqueeze(0) + scale, shift = out.chunk(2, dim=1) + return scale, shift + + +class DoubleSharedModulation(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.lin = nn.Parameter(torch.empty(6 * dim, device=device, dtype=dtype)) + + def forward(self, vec): + out = vec + comfy.model_management.cast_to(self.lin, dtype=vec.dtype, device=vec.device) + return out.chunk(6, dim=-1) + + +class TextFusionBlock(nn.Module): + def __init__(self, features, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.prenorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.postnorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) + self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) + + def forward(self, x, mask=None, transformer_options={}): + x = x + self.attn(self.prenorm(x), mask=mask, transformer_options=transformer_options) + x = x + self.mlp(self.postnorm(x)) + return x + + +class TextFusionTransformer(nn.Module): + def __init__(self, num_txt_layers, txt_dim, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.layerwise_blocks = nn.ModuleList([ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(2) + ]) + self.projector = operations.Linear(num_txt_layers, 1, bias=False, device=device, dtype=dtype) + self.refiner_blocks = nn.ModuleList([ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(2) + ]) + + def forward(self, x, mask=None, transformer_options={}): + b, l, n, d = x.shape + x = x.reshape(b * l, n, d) + for block in self.layerwise_blocks: + x = block(x.contiguous(), mask=None, transformer_options=transformer_options) + x = rearrange(x, "(b l) n d -> b l d n", b=b, l=l) + x = self.projector(x).squeeze(-1) + for block in self.refiner_blocks: + x = block(x, mask=mask, transformer_options=transformer_options) + return x + + +class SingleStreamBlock(nn.Module): + def __init__(self, features, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.mod = DoubleSharedModulation(features, device=device, dtype=dtype, operations=operations) + self.prenorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.postnorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) + self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) + + def forward(self, x, vec, freqs, mask=None, transformer_options={}): + prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec) + x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options) + x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift) + return x + + +class LastLayer(nn.Module): + def __init__(self, features, patch, channels, device=None, dtype=None, operations=None): + super().__init__() + self.norm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.linear = operations.Linear(features, patch * patch * channels, bias=True, device=device, dtype=dtype) + self.modulation = SimpleModulation(features, device=device, dtype=dtype, operations=operations) + + def forward(self, x, tvec): + scale, shift = self.modulation(tvec) + x = (1 + scale) * self.norm(x) + shift + return self.linear(x) + + +class SingleStreamDiT(nn.Module): + def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4, + layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12, + txtheads=20, txtkvheads=20, image_model=None, + device=None, dtype=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.patch = patch + self.channels = channels + self.tdim = tdim + self.heads = heads + self.txtdim = txtdim + self.txtlayers = txtlayers + + headdim = features // heads + axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)] + assert sum(axes) == headdim, f"axes {axes} sum != headdim {headdim}" + self.pe_embedder = EmbedND(dim=headdim, theta=int(theta), axes_dim=axes) + + self.first = operations.Linear(channels * patch ** 2, features, bias=True, device=device, dtype=dtype) + self.blocks = nn.ModuleList([ + SingleStreamBlock(features, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(layers) + ]) + self.tmlp = nn.Sequential( + operations.Linear(tdim, features, device=device, dtype=dtype), + nn.GELU(approximate="tanh"), + operations.Linear(features, features, device=device, dtype=dtype), + ) + self.txtfusion = TextFusionTransformer(txtlayers, txtdim, txtheads, multiplier, bias, txtkvheads, + device=device, dtype=dtype, operations=operations) + self.txtmlp = nn.Sequential( + RMSNorm(txtdim, device=device, dtype=dtype, operations=operations), + operations.Linear(txtdim, features, device=device, dtype=dtype), + nn.GELU(approximate="tanh"), + operations.Linear(features, features, device=device, dtype=dtype), + ) + self.last = LastLayer(features, patch, channels, device=device, dtype=dtype, operations=operations) + self.tproj = nn.Sequential( + nn.GELU(approximate="tanh"), + operations.Linear(features, features * 6, device=device, dtype=dtype), + ) + + def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options), + ).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs) + + def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): + temporal = x.ndim == 5 + if temporal: + b5, c5, t5, h5, w5 = x.shape + x = x.reshape(b5 * t5, c5, h5, w5) + bs, c, H_orig, W_orig = x.shape + patch = self.patch + # Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end. + x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch)) + H, W = x.shape[-2], x.shape[-1] + h_, w_ = H // patch, W // patch + + # context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim). + context = self._unpack_context(context) + + img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) + img = self.first(img) + + t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype)) + tvec = self.tproj(t) + + context = self.txtfusion(context, mask=None, transformer_options=transformer_options) + context = self.txtmlp(context) + + txtlen, imglen = context.shape[1], img.shape[1] + combined = torch.cat((context, img), dim=1) + + # Position ids: text at 0, image at (0, h_idx, w_idx). + device = combined.device + txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32) + imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32) + imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None] + imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :] + imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1) + pos = torch.cat((txtpos, imgpos), dim=1) + + freqs = self.pe_embedder(pos) + + for block in self.blocks: + combined = block(combined, tvec, freqs, None, transformer_options=transformer_options) + + final = self.last(combined, t) + out = final[:, txtlen:txtlen + imglen, :] + out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)", + h=h_, w=w_, ph=patch, pw=patch, c=self.channels) + out = out[:, :, :H_orig, :W_orig] # crop padding back off + if temporal: + out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2) + return out + + def _unpack_context(self, context): + # context: (B, seq, txtlayers*txtdim) -> (B, seq, txtlayers, txtdim). + b, seq, fused = context.shape + if fused != self.txtlayers * self.txtdim: + raise ValueError( + f"Krea2 expects conditioning with {self.txtlayers}x{self.txtdim}={self.txtlayers * self.txtdim} " + f"features (a {self.txtlayers}-layer Qwen3-VL stack) but got {fused}. " + f"Load the text encoder with CLIPLoader type 'krea2'." + ) + return context.reshape(b, seq, self.txtlayers, self.txtdim) diff --git a/comfy/lora.py b/comfy/lora.py index 2c8d0f0bf..427cf98aa 100644 --- a/comfy/lora.py +++ b/comfy/lora.py @@ -326,6 +326,17 @@ def model_lora_keys_unet(model, key_map={}): key_map["transformer.{}".format(key_lora)] = k key_map["lycoris_{}".format(key_lora.replace(".", "_"))] = k #SimpleTuner lycoris format + if isinstance(model, comfy.model_base.Krea2): + diffusers_keys = comfy.utils.krea2_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") + for k in diffusers_keys: + if k.endswith(".weight"): + to = diffusers_keys[k] + key_lora = k[:-len(".weight")] + key_map["diffusion_model.{}".format(key_lora)] = to + key_map["transformer.{}".format(key_lora)] = to + key_map["lycoris_{}".format(key_lora.replace(".", "_"))] = to + key_map[key_lora] = to + if isinstance(model, comfy.model_base.Lumina2): diffusers_keys = comfy.utils.z_image_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") for k in diffusers_keys: diff --git a/comfy/model_base.py b/comfy/model_base.py index 264dbb9b3..dcfa555dc 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -58,6 +58,7 @@ import comfy.ldm.omnigen.omnigen2 import comfy.ldm.boogu.model import comfy.ldm.qwen_image.model import comfy.ldm.ideogram4.model +import comfy.ldm.krea2.model import comfy.ldm.kandinsky5.model import comfy.ldm.anima.model import comfy.ldm.ace.ace_step15 @@ -2278,6 +2279,17 @@ class Ideogram4(BaseModel): out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) return out +class Krea2(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLUX, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + return out + class HunyuanImage21(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.hunyuan_video.model.HunyuanVideo) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index b773f0393..e53d848c9 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -834,6 +834,21 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["num_layers"] = count_blocks(state_dict_keys, '{}layers.'.format(key_prefix) + '{}.') return dit_config + if '{}txtfusion.projector.weight'.format(key_prefix) in state_dict_keys: # Krea 2 (K2) + dit_config = {} + dit_config["image_model"] = "krea2" + head_dim = 128 + first_w = state_dict['{}first.weight'.format(key_prefix)] # (features, channels*patch^2) + dit_config["features"] = first_w.shape[0] + dit_config["channels"] = first_w.shape[1] // (2 * 2) # patch=2 + dit_config["patch"] = 2 + dit_config["layers"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') + dit_config["heads"] = state_dict['{}blocks.0.attn.wq.weight'.format(key_prefix)].shape[0] // head_dim + dit_config["kvheads"] = state_dict['{}blocks.0.attn.wk.weight'.format(key_prefix)].shape[0] // head_dim + dit_config["txtlayers"] = state_dict['{}txtfusion.projector.weight'.format(key_prefix)].shape[1] + dit_config["txtdim"] = state_dict['{}txtfusion.layerwise_blocks.0.prenorm.scale'.format(key_prefix)].shape[0] + return dit_config + if '{}visual_transformer_blocks.0.cross_attention.key_norm.weight'.format(key_prefix) in state_dict_keys: # Kandinsky 5 dit_config = {} model_dim = state_dict['{}visual_embeddings.in_layer.bias'.format(key_prefix)].shape[0] diff --git a/comfy/sd.py b/comfy/sd.py index d9b1c0553..610c4e2b8 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -58,6 +58,7 @@ import comfy.text_encoders.omnigen2 import comfy.text_encoders.qwen_image import comfy.text_encoders.hunyuan_image import comfy.text_encoders.z_image +import comfy.text_encoders.krea2 import comfy.text_encoders.ideogram4 import comfy.text_encoders.ovis import comfy.text_encoders.kandinsky5 @@ -1303,6 +1304,7 @@ class CLIPType(Enum): PIXELDIT = 29 IDEOGRAM4 = 30 BOOGU = 31 + KREA2 = 32 @@ -1628,6 +1630,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) clip_target.clip = comfy.text_encoders.boogu.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.boogu.BooguTokenizer + elif clip_type == CLIPType.KREA2 and te_model == TEModel.QWEN3VL_4B: # Krea2: full Qwen3-VL-4B (12-layer tap for conditioning + multimodal generate). + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused. klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b" clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index cc05908ee..99d4c2800 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -26,6 +26,7 @@ import comfy.text_encoders.kandinsky5 import comfy.text_encoders.z_image import comfy.text_encoders.ideogram4 import comfy.text_encoders.boogu +import comfy.text_encoders.krea2 import comfy.text_encoders.anima import comfy.text_encoders.ace15 import comfy.text_encoders.longcat_image @@ -1818,6 +1819,35 @@ class Ideogram4(supported_models_base.BASE): hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_8b.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.ideogram4.Ideogram4Tokenizer, comfy.text_encoders.ideogram4.te(**hunyuan_detect)) + +class Krea2(supported_models_base.BASE): + unet_config = { + "image_model": "krea2", + } + + sampling_settings = { + "multiplier": 1.0, + "shift": 1.15, + } + + memory_usage_factor = 3.0 #TODO + + latent_format = latent_formats.Wan21 + + supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.Krea2(self, device=device) + return out + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect)) + class QwenImage(supported_models_base.BASE): unet_config = { "image_model": "qwen_image", @@ -2325,6 +2355,7 @@ models = [ Boogu, QwenImage, Ideogram4, + Krea2, Flux2, Lens, Kandinsky5Image, diff --git a/comfy/text_encoders/krea2.py b/comfy/text_encoders/krea2.py new file mode 100644 index 000000000..408a03566 --- /dev/null +++ b/comfy/text_encoders/krea2.py @@ -0,0 +1,84 @@ +"""Krea 2 (K2) text encoder: Qwen3-VL-4B, 12-layer tap. + +K2 conditions on a stack of hidden states from 12 layers of Qwen3-VL-4B +(reference taps ``hidden_states[2,5,8,...,35]``), kept as a ``(B, 12, seq, 2560)`` tensor and +consumed by the DiT's internal ``txtfusion`` adapter. Comfy carries conditioning as a 3D tensor, +so the 12-layer stack is flattened to ``(B, seq, 12*2560)`` here and unpacked inside the model. +""" + +import numbers + +import torch + +import comfy.text_encoders.qwen3vl +from comfy import sd1_clip + +# tap k == hidden_states[k] (no offset). +KREA2_TAP_LAYERS = [2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35] + +# Identical system template to Qwen-Image; Krea2 strips the system+user-opening prefix. +KREA2_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + + +class Krea2Tokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b") + self.llama_template = KREA2_TEMPLATE # conditioning template; image text-gen uses qwen3vl's default image template. + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs): + # Krea2 conditions on the no-think template; thinking=True drops the empty block qwen3vl adds. + return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs) + + +class Krea2Qwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel): + def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}): + super().__init__(device=device, layer=KREA2_TAP_LAYERS, layer_idx=None, dtype=dtype, + attention_mask=attention_mask, model_options=model_options, model_type="qwen3vl_4b") + + +class Krea2TEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=Krea2Qwen3VLClipModel, model_options=model_options) + + def encode_token_weights(self, token_weight_pairs, template_end=-1): + out, pooled, extra = super().encode_token_weights(token_weight_pairs) # out: (B, 12, seq, 2560) + tok_pairs = token_weight_pairs["qwen3vl_4b"][0] + + # Strip the system + user-opening prefix + count_im_start = 0 + if template_end == -1: + for i, v in enumerate(tok_pairs): + elem = v[0] + if not torch.is_tensor(elem) and isinstance(elem, numbers.Integral): + if elem == 151644 and count_im_start < 2: + template_end = i + count_im_start += 1 + if out.shape[2] > (template_end + 3): + if tok_pairs[template_end + 1][0] == 872: # "user" + if tok_pairs[template_end + 2][0] == 198: # "\n" + template_end += 3 + + out = out[:, :, template_end:] + + b, n, seq, h = out.shape + # Flatten the 12-layer axis into the feature dim: (B, seq, 12*2560). Unpacked in the model. + out = out.permute(0, 2, 1, 3).reshape(b, seq, n * h) + + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, template_end:] + if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]): + extra.pop("attention_mask") + + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class Krea2TEModel_(Krea2TEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + if dtype_llama is not None: + dtype = dtype_llama + super().__init__(device=device, dtype=dtype, model_options=model_options) + return Krea2TEModel_ diff --git a/comfy/utils.py b/comfy/utils.py index 09d783fff..61c2a22dd 100644 --- a/comfy/utils.py +++ b/comfy/utils.py @@ -818,6 +818,44 @@ def z_image_to_diffusers(mmdit_config, output_prefix=""): return key_map +def krea2_to_diffusers(mmdit_config, output_prefix=""): + n_layers = mmdit_config.get("layers", 0) + n_txt_layerwise = 2 # TextFusionTransformer hardcodes 2 layerwise + 2 refiner blocks + n_txt_refiner = 2 + key_map = {} + + def add_block(prefix_to, prefix_from): + block_map = { + "attn.to_q": "attn.wq", "attn.to_k": "attn.wk", "attn.to_v": "attn.wv", + "attn.to_gate": "attn.gate", "attn.to_out.0": "attn.wo", + "attn.to_out": "attn.wo", # some tools drop the ".0" on to_out + "ff.gate": "mlp.gate", "ff.up": "mlp.up", "ff.down": "mlp.down", + } + for d, c in block_map.items(): + key_map["{}.{}.weight".format(prefix_to, d)] = "{}{}.{}.weight".format(output_prefix, prefix_from, c) + + for i in range(n_layers): + add_block("transformer_blocks.{}".format(i), "blocks.{}".format(i)) + for i in range(n_txt_layerwise): + add_block("text_fusion.layerwise_blocks.{}".format(i), "txtfusion.layerwise_blocks.{}".format(i)) + for i in range(n_txt_refiner): + add_block("text_fusion.refiner_blocks.{}".format(i), "txtfusion.refiner_blocks.{}".format(i)) + + MAP_BASIC = [ + ("img_in", "first"), + ("time_embed.linear_1", "tmlp.0"), + ("time_embed.linear_2", "tmlp.2"), + ("time_mod_proj", "tproj.1"), + ("txt_in.linear_1", "txtmlp.1"), + ("txt_in.linear_2", "txtmlp.3"), + ("text_fusion.projector", "txtfusion.projector"), + ("final_layer.linear", "last.linear"), + ] + for d, c in MAP_BASIC: + key_map["{}.weight".format(d)] = "{}{}.weight".format(output_prefix, c) + + return key_map + def repeat_to_batch_size(tensor, batch_size, dim=0): if tensor.shape[dim] > batch_size: return tensor.narrow(dim, 0, batch_size) diff --git a/nodes.py b/nodes.py index 66c08121d..166e02d3d 100644 --- a/nodes.py +++ b/nodes.py @@ -969,7 +969,7 @@ class CLIPLoader: @classmethod def INPUT_TYPES(s): return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ), - "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu"], ), + "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2"], ), }, "optional": { "device": (["default", "cpu"], {"advanced": True}), From 833bfb572e8552d1cfe590b3b6445e96e7526f25 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 22 Jun 2026 21:06:19 -0700 Subject: [PATCH 002/211] Please try native formats instead of disabling dynamic vram. (#14577) --- main.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/main.py b/main.py index ad5c11e16..aa4ee2adb 100644 --- a/main.py +++ b/main.py @@ -557,8 +557,13 @@ if __name__ == "__main__": logging.warning("WARNING: You are using a python version older than 3.10, please upgrade to a newer one. 3.12 and above is recommended.") if args.disable_dynamic_vram: - logging.warning("Dynamic vram disabled with argument. If you have any issues with dynamic vram enabled please give us a detailed reports as this argument will be removed soon.") - + logging.warning( + "Dynamic vram disabled with argument. If you have any issues with " + "dynamic vram enabled please give us a detailed reports as this " + "argument will be removed soon. If you use gguf we recommend keeping " + "dynamic vram enabled and using native ComfyUI model formats instead. " + "ComfyUI native formats like fp8 will be faster even if they are larger than your memory." + ) event_loop, _, start_all_func = start_comfyui() try: x = start_all_func() From b910f4fa2ae3c816ca29bb99f8ebc5baee3387ae Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 23 Jun 2026 01:50:48 -0700 Subject: [PATCH 003/211] More accurate memory usage factor for krea 2. (#14594) --- comfy/supported_models.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index 99d4c2800..afb66e6f3 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -1830,7 +1830,7 @@ class Krea2(supported_models_base.BASE): "shift": 1.15, } - memory_usage_factor = 3.0 #TODO + memory_usage_factor = 2.2 latent_format = latent_formats.Wan21 From 0a92ed161e6fa18eb169d96d64b4c279cf280dc5 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 23 Jun 2026 13:29:46 +0300 Subject: [PATCH 004/211] [Partner Nodes] feat(Alibaba): add support for HappyHorse 1.1 model (#14581) Signed-off-by: bigcat88 Co-authored-by: Alexis Rolland --- comfy_api_nodes/nodes_wan.py | 142 +++++++++++++++++++++++++++++++++-- 1 file changed, 137 insertions(+), 5 deletions(-) diff --git a/comfy_api_nodes/nodes_wan.py b/comfy_api_nodes/nodes_wan.py index b7b97d70f..1782739fd 100644 --- a/comfy_api_nodes/nodes_wan.py +++ b/comfy_api_nodes/nodes_wan.py @@ -48,10 +48,13 @@ from comfy_api_nodes.util import ( upload_image_to_comfyapi, upload_video_to_comfyapi, validate_audio_duration, + validate_image_aspect_ratio, + validate_image_dimensions, validate_string, validate_video_duration, ) + RES_IN_PARENS = re.compile(r"\((\d+)\s*[x×]\s*(\d+)\)") @@ -1657,6 +1660,44 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-t2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Combo.Input( + "ratio", + options=[ + "16:9", + "9:16", + "1:1", + "4:3", + "3:4", + "21:9", + "9:21", + "5:4", + "4:5", + ], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-t2v", [ @@ -1719,7 +1760,9 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1781,6 +1824,30 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-i2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-i2v", [ @@ -1843,7 +1910,9 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1859,6 +1928,8 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): seed: int, watermark: bool, ): + validate_image_dimensions(first_frame, min_width=300, min_height=300) + validate_image_aspect_ratio(first_frame, (1, 2.5), (2.5, 1), strict=False) media = [ Wan27MediaItem( type="first_frame", @@ -2053,6 +2124,62 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-r2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the video. Use identifiers such as 'character1' and " + "'character2' to refer to the reference characters.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Combo.Input( + "ratio", + options=[ + "16:9", + "9:16", + "1:1", + "4:3", + "3:4", + "21:9", + "9:21", + "5:4", + "4:5", + ], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image1", + "image2", + "image3", + "image4", + "image5", + "image6", + "image7", + "image8", + "image9", + ], + min=1, + ), + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-r2v", [ @@ -2133,7 +2260,9 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -2149,8 +2278,11 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): watermark: bool, ): validate_string(model["prompt"], strip_whitespace=False, min_length=1) - media = [] reference_images = model.get("reference_images", {}) + for key in reference_images: + validate_image_dimensions(reference_images[key], min_width=400, min_height=400) + validate_image_aspect_ratio(reference_images[key], (1, 2.5), (2.5, 1), strict=False) + media = [] for key in reference_images: media.append( Wan27MediaItem( @@ -2159,7 +2291,7 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ) ) if not media: - raise ValueError("At least one reference reference image must be provided.") + raise ValueError("At least one reference image must be provided.") initial_response = await sync_op( cls, From 0ba903bd5bab918477f9adb290f35fdb907dc703 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 23 Jun 2026 16:18:35 +0300 Subject: [PATCH 005/211] [Partner Nodes] feat(ByteDance): add 4K resolution support for SeeDance 2.0 (#14588) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/bytedance.py | 22 ++++++++++--- comfy_api_nodes/nodes_bytedance.py | 50 ++++++++++++++++++++---------- 2 files changed, 50 insertions(+), 22 deletions(-) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 47f24586c..999b51d39 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -163,15 +163,27 @@ class SeedanceVirtualLibraryCreateAssetRequest(BaseModel): asset_type: str | None = Field(None, description="BytePlus asset type. Defaults to Image server-side when omitted.") -# Dollars per 1K tokens, keyed by (model_id, has_video_input). +# Dollars per 1K tokens, keyed by (model_id, has_video_input, resolution). SEEDANCE2_PRICE_PER_1K_TOKENS = { - ("dreamina-seedance-2-0-260128", False): 0.007, - ("dreamina-seedance-2-0-260128", True): 0.0043, - ("dreamina-seedance-2-0-fast-260128", False): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True): 0.0033, + ("dreamina-seedance-2-0-260128", False, "480p"): 0.007, + ("dreamina-seedance-2-0-260128", True, "480p"): 0.0043, + ("dreamina-seedance-2-0-260128", False, "720p"): 0.007, + ("dreamina-seedance-2-0-260128", True, "720p"): 0.0043, + ("dreamina-seedance-2-0-260128", False, "1080p"): 0.0077, + ("dreamina-seedance-2-0-260128", True, "1080p"): 0.0047, + ("dreamina-seedance-2-0-260128", False, "4k"): 0.004, + ("dreamina-seedance-2-0-260128", True, "4k"): 0.0024, + ("dreamina-seedance-2-0-fast-260128", False, "480p"): 0.0056, + ("dreamina-seedance-2-0-fast-260128", True, "480p"): 0.0033, + ("dreamina-seedance-2-0-fast-260128", False, "720p"): 0.0056, + ("dreamina-seedance-2-0-fast-260128", True, "720p"): 0.0033, } +def seedance2_price_per_1k_tokens(model_id: str, has_video_input: bool, resolution: str) -> float | None: + return SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input, resolution)) + + RECOMMENDED_PRESETS = [ ("1024x1024 (1:1)", 1024, 1024), ("864x1152 (3:4)", 864, 1152), diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index c30ddc446..6192b35bf 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -15,7 +15,6 @@ from comfy_api_nodes.apis.bytedance import ( RECOMMENDED_PRESETS_SEEDREAM_4_0, RECOMMENDED_PRESETS_SEEDREAM_4_5, RECOMMENDED_PRESETS_SEEDREAM_5_LITE, - SEEDANCE2_PRICE_PER_1K_TOKENS, SEEDANCE2_REF_VIDEO_PIXEL_LIMITS, VIDEO_TASKS_EXECUTION_TIME, GetAssetResponse, @@ -40,6 +39,7 @@ from comfy_api_nodes.apis.bytedance import ( TaskVideoContentUrl, Text2ImageTaskCreationRequest, Text2VideoTaskCreationRequest, + seedance2_price_per_1k_tokens, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -141,7 +141,7 @@ SEEDANCE2_RATIO_WH = { "9:16": (9, 16), "21:9": (21, 9), } -SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080} +SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160} def _seedance2_target_dims(resolution: str, ratio: str, image: torch.Tensor) -> tuple[int, int]: @@ -377,9 +377,9 @@ async def _seedance_virtual_library_upload_video_asset( return f"asset://{create_resp.asset_id}" -def _seedance2_price_extractor(model_id: str, has_video_input: bool): +def _seedance2_price_extractor(model_id: str, has_video_input: bool, resolution: str): """Returns a price_extractor closure for Seedance 2.0 poll_op.""" - rate = SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input)) + rate = seedance2_price_per_1k_tokens(model_id, has_video_input, resolution) if rate is None: return None @@ -1621,7 +1621,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p"])), + IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p", "4k"])), IO.DynamicCombo.Option("Seedance 2.0 Fast", _seedance2_text_inputs(["480p", "720p"])), ], tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", @@ -1660,11 +1660,15 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; - $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $pricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1703,7 +1707,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -1724,7 +1728,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_text_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), + _seedance2_text_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -1791,11 +1795,15 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; - $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $pricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1913,7 +1921,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -2010,7 +2018,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_reference_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), + _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -2056,13 +2064,19 @@ class ByteDance2ReferenceNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; $hasVideo := $lookup(inputGroups, "model.reference_videos") > 0; - $noVideoPricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; - $videoPricePer1K := $contains($m, "fast") ? 0.004719 : 0.006149; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $noVideoPricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $videoPricePer1K := $res = "4k" ? 0.003432 : + $res = "1080p" ? 0.006721 : + $contains($m, "fast") ? 0.004719 : 0.006149; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $noVideoCost := $dur * $rate * $noVideoPricePer1K / 1000; @@ -2258,7 +2272,9 @@ class ByteDance2ReferenceNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=has_video_input), + price_extractor=_seedance2_price_extractor( + model_id, has_video_input=has_video_input, resolution=model["resolution"] + ), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) From d0b640fff75f7902a5292228fec0fcd0b88f5655 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Tue, 23 Jun 2026 23:35:21 +0800 Subject: [PATCH 006/211] chore: update workflow templates to v0.10.2 (#14600) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 0c8b1888e..b05cad045 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.19 -comfyui-workflow-templates==0.10.0 +comfyui-workflow-templates==0.10.2 comfyui-embedded-docs==0.5.5 torch torchsde From 0f949d0fafbd86f2494459a5bf89d9329775e199 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 23 Jun 2026 18:38:46 +0300 Subject: [PATCH 007/211] [Partner Nodes] feat(Grok): add 1080p resolution to Grok Image node (#14597) --- comfy_api_nodes/nodes_grok.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index 2ae529813..dc484536e 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -30,7 +30,7 @@ from comfy_api_nodes.util import ( _GROK_VIDEO_MODEL_API_IDS = { - "grok-imagine-video-1.5": "grok-imagine-video-1.5-preview", + "grok-imagine-video-1.5": "grok-imagine-video-1.5", } @@ -521,8 +521,8 @@ class GrokVideoNode(IO.ComfyNode): ), IO.Combo.Input( "resolution", - options=["480p", "720p"], - tooltip="The resolution of the output video.", + options=["480p", "720p", "1080p"], + tooltip="The resolution of the output video. 1080p is only available for grok-imagine-video-1.5.", ), IO.Combo.Input( "aspect_ratio", @@ -570,11 +570,12 @@ class GrokVideoNode(IO.ComfyNode): ( $is15 := $contains(widgets.model, "1.5"); $rate := $is15 - ? (widgets.resolution = "720p" ? 0.2002 : 0.1144) + ? (widgets.resolution = "1080p" ? 0.25 : (widgets.resolution = "720p" ? 0.14 : 0.08)) : (widgets.resolution = "720p" ? 0.07 : 0.05); - $imgCost := $is15 ? 0.0143 : 0.002; + $imgCost := $is15 ? 0.01 : 0.002; $base := $rate * widgets.duration; - {"type":"usd","usd": inputs.image.connected ? $base + $imgCost : $base} + $total := inputs.image.connected ? $base + $imgCost : $base; + {"type":"usd","usd": $is15 ? $total * 1.43 : $total} ) """, ), @@ -593,6 +594,8 @@ class GrokVideoNode(IO.ComfyNode): ) -> IO.NodeOutput: if image is None and model == "grok-imagine-video-1.5": raise ValueError(f"The '{model}' model requires an input image; connect one to the 'image' input.") + if resolution == "1080p" and model != "grok-imagine-video-1.5": + raise ValueError(f"1080p resolution is only available for grok-imagine-video-1.5, not '{model}'.") image_url = None if image is not None: if get_number_of_images(image) != 1: From 4a030566320b2581c14b536a751eda3d5748e022 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 23 Jun 2026 19:49:16 +0300 Subject: [PATCH 008/211] [Partner Nodes] revert last 3 PRs: #14597 #14588 #14581 (#14602) --- comfy_api_nodes/apis/bytedance.py | 22 +---- comfy_api_nodes/nodes_bytedance.py | 50 ++++------ comfy_api_nodes/nodes_grok.py | 15 ++- comfy_api_nodes/nodes_wan.py | 142 +---------------------------- 4 files changed, 33 insertions(+), 196 deletions(-) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 999b51d39..47f24586c 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -163,27 +163,15 @@ class SeedanceVirtualLibraryCreateAssetRequest(BaseModel): asset_type: str | None = Field(None, description="BytePlus asset type. Defaults to Image server-side when omitted.") -# Dollars per 1K tokens, keyed by (model_id, has_video_input, resolution). +# Dollars per 1K tokens, keyed by (model_id, has_video_input). SEEDANCE2_PRICE_PER_1K_TOKENS = { - ("dreamina-seedance-2-0-260128", False, "480p"): 0.007, - ("dreamina-seedance-2-0-260128", True, "480p"): 0.0043, - ("dreamina-seedance-2-0-260128", False, "720p"): 0.007, - ("dreamina-seedance-2-0-260128", True, "720p"): 0.0043, - ("dreamina-seedance-2-0-260128", False, "1080p"): 0.0077, - ("dreamina-seedance-2-0-260128", True, "1080p"): 0.0047, - ("dreamina-seedance-2-0-260128", False, "4k"): 0.004, - ("dreamina-seedance-2-0-260128", True, "4k"): 0.0024, - ("dreamina-seedance-2-0-fast-260128", False, "480p"): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True, "480p"): 0.0033, - ("dreamina-seedance-2-0-fast-260128", False, "720p"): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True, "720p"): 0.0033, + ("dreamina-seedance-2-0-260128", False): 0.007, + ("dreamina-seedance-2-0-260128", True): 0.0043, + ("dreamina-seedance-2-0-fast-260128", False): 0.0056, + ("dreamina-seedance-2-0-fast-260128", True): 0.0033, } -def seedance2_price_per_1k_tokens(model_id: str, has_video_input: bool, resolution: str) -> float | None: - return SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input, resolution)) - - RECOMMENDED_PRESETS = [ ("1024x1024 (1:1)", 1024, 1024), ("864x1152 (3:4)", 864, 1152), diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 6192b35bf..c30ddc446 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -15,6 +15,7 @@ from comfy_api_nodes.apis.bytedance import ( RECOMMENDED_PRESETS_SEEDREAM_4_0, RECOMMENDED_PRESETS_SEEDREAM_4_5, RECOMMENDED_PRESETS_SEEDREAM_5_LITE, + SEEDANCE2_PRICE_PER_1K_TOKENS, SEEDANCE2_REF_VIDEO_PIXEL_LIMITS, VIDEO_TASKS_EXECUTION_TIME, GetAssetResponse, @@ -39,7 +40,6 @@ from comfy_api_nodes.apis.bytedance import ( TaskVideoContentUrl, Text2ImageTaskCreationRequest, Text2VideoTaskCreationRequest, - seedance2_price_per_1k_tokens, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -141,7 +141,7 @@ SEEDANCE2_RATIO_WH = { "9:16": (9, 16), "21:9": (21, 9), } -SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160} +SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080} def _seedance2_target_dims(resolution: str, ratio: str, image: torch.Tensor) -> tuple[int, int]: @@ -377,9 +377,9 @@ async def _seedance_virtual_library_upload_video_asset( return f"asset://{create_resp.asset_id}" -def _seedance2_price_extractor(model_id: str, has_video_input: bool, resolution: str): +def _seedance2_price_extractor(model_id: str, has_video_input: bool): """Returns a price_extractor closure for Seedance 2.0 poll_op.""" - rate = seedance2_price_per_1k_tokens(model_id, has_video_input, resolution) + rate = SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input)) if rate is None: return None @@ -1621,7 +1621,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p", "4k"])), + IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p"])), IO.DynamicCombo.Option("Seedance 2.0 Fast", _seedance2_text_inputs(["480p", "720p"])), ], tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", @@ -1660,15 +1660,11 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; - $rate4k := 195200; $m := widgets.model; + $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $pricePer1K := $res = "4k" ? 0.00572 : - $res = "1080p" ? 0.011011 : - $contains($m, "fast") ? 0.008008 : 0.01001; - $rate := $res = "4k" ? $rate4k : - $res = "1080p" ? $rate1080 : + $rate := $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1707,7 +1703,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -1728,7 +1724,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_text_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), + _seedance2_text_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -1795,15 +1791,11 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; - $rate4k := 195200; $m := widgets.model; + $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $pricePer1K := $res = "4k" ? 0.00572 : - $res = "1080p" ? 0.011011 : - $contains($m, "fast") ? 0.008008 : 0.01001; - $rate := $res = "4k" ? $rate4k : - $res = "1080p" ? $rate1080 : + $rate := $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1921,7 +1913,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -2018,7 +2010,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), + _seedance2_reference_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -2064,19 +2056,13 @@ class ByteDance2ReferenceNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; - $rate4k := 195200; $m := widgets.model; $hasVideo := $lookup(inputGroups, "model.reference_videos") > 0; + $noVideoPricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; + $videoPricePer1K := $contains($m, "fast") ? 0.004719 : 0.006149; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $noVideoPricePer1K := $res = "4k" ? 0.00572 : - $res = "1080p" ? 0.011011 : - $contains($m, "fast") ? 0.008008 : 0.01001; - $videoPricePer1K := $res = "4k" ? 0.003432 : - $res = "1080p" ? 0.006721 : - $contains($m, "fast") ? 0.004719 : 0.006149; - $rate := $res = "4k" ? $rate4k : - $res = "1080p" ? $rate1080 : + $rate := $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $noVideoCost := $dur * $rate * $noVideoPricePer1K / 1000; @@ -2272,9 +2258,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor( - model_id, has_video_input=has_video_input, resolution=model["resolution"] - ), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=has_video_input), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index dc484536e..2ae529813 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -30,7 +30,7 @@ from comfy_api_nodes.util import ( _GROK_VIDEO_MODEL_API_IDS = { - "grok-imagine-video-1.5": "grok-imagine-video-1.5", + "grok-imagine-video-1.5": "grok-imagine-video-1.5-preview", } @@ -521,8 +521,8 @@ class GrokVideoNode(IO.ComfyNode): ), IO.Combo.Input( "resolution", - options=["480p", "720p", "1080p"], - tooltip="The resolution of the output video. 1080p is only available for grok-imagine-video-1.5.", + options=["480p", "720p"], + tooltip="The resolution of the output video.", ), IO.Combo.Input( "aspect_ratio", @@ -570,12 +570,11 @@ class GrokVideoNode(IO.ComfyNode): ( $is15 := $contains(widgets.model, "1.5"); $rate := $is15 - ? (widgets.resolution = "1080p" ? 0.25 : (widgets.resolution = "720p" ? 0.14 : 0.08)) + ? (widgets.resolution = "720p" ? 0.2002 : 0.1144) : (widgets.resolution = "720p" ? 0.07 : 0.05); - $imgCost := $is15 ? 0.01 : 0.002; + $imgCost := $is15 ? 0.0143 : 0.002; $base := $rate * widgets.duration; - $total := inputs.image.connected ? $base + $imgCost : $base; - {"type":"usd","usd": $is15 ? $total * 1.43 : $total} + {"type":"usd","usd": inputs.image.connected ? $base + $imgCost : $base} ) """, ), @@ -594,8 +593,6 @@ class GrokVideoNode(IO.ComfyNode): ) -> IO.NodeOutput: if image is None and model == "grok-imagine-video-1.5": raise ValueError(f"The '{model}' model requires an input image; connect one to the 'image' input.") - if resolution == "1080p" and model != "grok-imagine-video-1.5": - raise ValueError(f"1080p resolution is only available for grok-imagine-video-1.5, not '{model}'.") image_url = None if image is not None: if get_number_of_images(image) != 1: diff --git a/comfy_api_nodes/nodes_wan.py b/comfy_api_nodes/nodes_wan.py index 1782739fd..b7b97d70f 100644 --- a/comfy_api_nodes/nodes_wan.py +++ b/comfy_api_nodes/nodes_wan.py @@ -48,13 +48,10 @@ from comfy_api_nodes.util import ( upload_image_to_comfyapi, upload_video_to_comfyapi, validate_audio_duration, - validate_image_aspect_ratio, - validate_image_dimensions, validate_string, validate_video_duration, ) - RES_IN_PARENS = re.compile(r"\((\d+)\s*[x×]\s*(\d+)\)") @@ -1660,44 +1657,6 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option( - "happyhorse-1.1-t2v", - [ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt describing the elements and visual features. " - "Supports English and Chinese.", - ), - IO.Combo.Input( - "resolution", - options=["720P", "1080P"], - ), - IO.Combo.Input( - "ratio", - options=[ - "16:9", - "9:16", - "1:1", - "4:3", - "3:4", - "21:9", - "9:21", - "5:4", - "4:5", - ], - ), - IO.Int.Input( - "duration", - default=5, - min=3, - max=15, - step=1, - display_mode=IO.NumberDisplay.number, - ), - ], - ), IO.DynamicCombo.Option( "happyhorse-1.0-t2v", [ @@ -1760,9 +1719,7 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := $contains(widgets.model, "1.1") - ? { "720p": 0.2002, "1080p": 0.2574 } - : { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1824,30 +1781,6 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option( - "happyhorse-1.1-i2v", - [ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt describing the elements and visual features. " - "Supports English and Chinese.", - ), - IO.Combo.Input( - "resolution", - options=["720P", "1080P"], - ), - IO.Int.Input( - "duration", - default=5, - min=3, - max=15, - step=1, - display_mode=IO.NumberDisplay.number, - ), - ], - ), IO.DynamicCombo.Option( "happyhorse-1.0-i2v", [ @@ -1910,9 +1843,7 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := $contains(widgets.model, "1.1") - ? { "720p": 0.2002, "1080p": 0.2574 } - : { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1928,8 +1859,6 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): seed: int, watermark: bool, ): - validate_image_dimensions(first_frame, min_width=300, min_height=300) - validate_image_aspect_ratio(first_frame, (1, 2.5), (2.5, 1), strict=False) media = [ Wan27MediaItem( type="first_frame", @@ -2124,62 +2053,6 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option( - "happyhorse-1.1-r2v", - [ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt describing the video. Use identifiers such as 'character1' and " - "'character2' to refer to the reference characters.", - ), - IO.Combo.Input( - "resolution", - options=["720P", "1080P"], - ), - IO.Combo.Input( - "ratio", - options=[ - "16:9", - "9:16", - "1:1", - "4:3", - "3:4", - "21:9", - "9:21", - "5:4", - "4:5", - ], - ), - IO.Int.Input( - "duration", - default=5, - min=3, - max=15, - step=1, - display_mode=IO.NumberDisplay.number, - ), - IO.Autogrow.Input( - "reference_images", - template=IO.Autogrow.TemplateNames( - IO.Image.Input("reference_image"), - names=[ - "image1", - "image2", - "image3", - "image4", - "image5", - "image6", - "image7", - "image8", - "image9", - ], - min=1, - ), - ), - ], - ), IO.DynamicCombo.Option( "happyhorse-1.0-r2v", [ @@ -2260,9 +2133,7 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := $contains(widgets.model, "1.1") - ? { "720p": 0.2002, "1080p": 0.2574 } - : { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -2278,11 +2149,8 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): watermark: bool, ): validate_string(model["prompt"], strip_whitespace=False, min_length=1) - reference_images = model.get("reference_images", {}) - for key in reference_images: - validate_image_dimensions(reference_images[key], min_width=400, min_height=400) - validate_image_aspect_ratio(reference_images[key], (1, 2.5), (2.5, 1), strict=False) media = [] + reference_images = model.get("reference_images", {}) for key in reference_images: media.append( Wan27MediaItem( @@ -2291,7 +2159,7 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ) ) if not media: - raise ValueError("At least one reference image must be provided.") + raise ValueError("At least one reference reference image must be provided.") initial_response = await sync_op( cls, From 261bdb7cac94b30d91a89f09277bb2073b597418 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 24 Jun 2026 01:06:26 +0800 Subject: [PATCH 009/211] chore: update workflow templates to v0.10.3 (#14603) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index b05cad045..323b76c8e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.19 -comfyui-workflow-templates==0.10.2 +comfyui-workflow-templates==0.10.3 comfyui-embedded-docs==0.5.5 torch torchsde From f6c162ddcfbd7eefb39c06fe5b8d4c46e8d09f40 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Tue, 23 Jun 2026 13:22:28 -0400 Subject: [PATCH 010/211] ComfyUI v0.26.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index cee317f3d..f8db561ba 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.25.0" +__version__ = "0.26.0" diff --git a/pyproject.toml b/pyproject.toml index 54f11d7fa..2e8a85d3f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.25.0" +version = "0.26.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From 1f275fcba6ab658fbd9785e9e70bee2c4452d726 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Wed, 24 Jun 2026 19:22:59 +0900 Subject: [PATCH 011/211] chore(openapi): sync shared API contract from cloud@363764b (#14607) --- openapi.yaml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/openapi.yaml b/openapi.yaml index 380e4476e..cee8a4763 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -2357,6 +2357,10 @@ paths: description: | Returns a list of model folders available in the system. This is an experimental endpoint that replaces the legacy /models endpoint. + Each folder's name is the identifier to pass to /api/experiment/models/{folder}. + Once the model_type migration is active the names are model_type folder_names + (e.g. `ultralytics_bbox`); a folder with no folder_name mapping is returned by + its directory path. operationId: getModelFolders responses: "200": From 44955d783b969241dd7b1899777e0f00e940bc02 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 24 Jun 2026 13:37:28 +0300 Subject: [PATCH 012/211] [Partner Nodes] feat(Alibaba): add support for HappyHorse 1.1 model (#14611) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_wan.py | 142 +++++++++++++++++++++++++++++++++-- 1 file changed, 137 insertions(+), 5 deletions(-) diff --git a/comfy_api_nodes/nodes_wan.py b/comfy_api_nodes/nodes_wan.py index b7b97d70f..1782739fd 100644 --- a/comfy_api_nodes/nodes_wan.py +++ b/comfy_api_nodes/nodes_wan.py @@ -48,10 +48,13 @@ from comfy_api_nodes.util import ( upload_image_to_comfyapi, upload_video_to_comfyapi, validate_audio_duration, + validate_image_aspect_ratio, + validate_image_dimensions, validate_string, validate_video_duration, ) + RES_IN_PARENS = re.compile(r"\((\d+)\s*[x×]\s*(\d+)\)") @@ -1657,6 +1660,44 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-t2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Combo.Input( + "ratio", + options=[ + "16:9", + "9:16", + "1:1", + "4:3", + "3:4", + "21:9", + "9:21", + "5:4", + "4:5", + ], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-t2v", [ @@ -1719,7 +1760,9 @@ class HappyHorseTextToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1781,6 +1824,30 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-i2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-i2v", [ @@ -1843,7 +1910,9 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -1859,6 +1928,8 @@ class HappyHorseImageToVideoApi(IO.ComfyNode): seed: int, watermark: bool, ): + validate_image_dimensions(first_frame, min_width=300, min_height=300) + validate_image_aspect_ratio(first_frame, (1, 2.5), (2.5, 1), strict=False) media = [ Wan27MediaItem( type="first_frame", @@ -2053,6 +2124,62 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "happyhorse-1.1-r2v", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the video. Use identifiers such as 'character1' and " + "'character2' to refer to the reference characters.", + ), + IO.Combo.Input( + "resolution", + options=["720P", "1080P"], + ), + IO.Combo.Input( + "ratio", + options=[ + "16:9", + "9:16", + "1:1", + "4:3", + "3:4", + "21:9", + "9:21", + "5:4", + "4:5", + ], + ), + IO.Int.Input( + "duration", + default=5, + min=3, + max=15, + step=1, + display_mode=IO.NumberDisplay.number, + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image1", + "image2", + "image3", + "image4", + "image5", + "image6", + "image7", + "image8", + "image9", + ], + min=1, + ), + ), + ], + ), IO.DynamicCombo.Option( "happyhorse-1.0-r2v", [ @@ -2133,7 +2260,9 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ( $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $ppsTable := { "720p": 0.14, "1080p": 0.24 }; + $ppsTable := $contains(widgets.model, "1.1") + ? { "720p": 0.2002, "1080p": 0.2574 } + : { "720p": 0.14, "1080p": 0.24 }; $pps := $lookup($ppsTable, $res); { "type": "usd", "usd": $pps * $dur } ) @@ -2149,8 +2278,11 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): watermark: bool, ): validate_string(model["prompt"], strip_whitespace=False, min_length=1) - media = [] reference_images = model.get("reference_images", {}) + for key in reference_images: + validate_image_dimensions(reference_images[key], min_width=400, min_height=400) + validate_image_aspect_ratio(reference_images[key], (1, 2.5), (2.5, 1), strict=False) + media = [] for key in reference_images: media.append( Wan27MediaItem( @@ -2159,7 +2291,7 @@ class HappyHorseReferenceVideoApi(IO.ComfyNode): ) ) if not media: - raise ValueError("At least one reference reference image must be provided.") + raise ValueError("At least one reference image must be provided.") initial_response = await sync_op( cls, From 12218db68a151264be541e03b2653d53b2b7c13d Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 24 Jun 2026 21:01:25 +0800 Subject: [PATCH 013/211] Update the template to bring the HH1.1 templates back (#14613) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 323b76c8e..b05cad045 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.19 -comfyui-workflow-templates==0.10.3 +comfyui-workflow-templates==0.10.2 comfyui-embedded-docs==0.5.5 torch torchsde From cabb7342d1abd570d68f3b4dddb5df031731422e Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 24 Jun 2026 16:28:56 +0300 Subject: [PATCH 014/211] [Partner Nodes] feat(Grok): add 1080p resolution to Grok Image node (#14612) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_grok.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index 2ae529813..dc484536e 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -30,7 +30,7 @@ from comfy_api_nodes.util import ( _GROK_VIDEO_MODEL_API_IDS = { - "grok-imagine-video-1.5": "grok-imagine-video-1.5-preview", + "grok-imagine-video-1.5": "grok-imagine-video-1.5", } @@ -521,8 +521,8 @@ class GrokVideoNode(IO.ComfyNode): ), IO.Combo.Input( "resolution", - options=["480p", "720p"], - tooltip="The resolution of the output video.", + options=["480p", "720p", "1080p"], + tooltip="The resolution of the output video. 1080p is only available for grok-imagine-video-1.5.", ), IO.Combo.Input( "aspect_ratio", @@ -570,11 +570,12 @@ class GrokVideoNode(IO.ComfyNode): ( $is15 := $contains(widgets.model, "1.5"); $rate := $is15 - ? (widgets.resolution = "720p" ? 0.2002 : 0.1144) + ? (widgets.resolution = "1080p" ? 0.25 : (widgets.resolution = "720p" ? 0.14 : 0.08)) : (widgets.resolution = "720p" ? 0.07 : 0.05); - $imgCost := $is15 ? 0.0143 : 0.002; + $imgCost := $is15 ? 0.01 : 0.002; $base := $rate * widgets.duration; - {"type":"usd","usd": inputs.image.connected ? $base + $imgCost : $base} + $total := inputs.image.connected ? $base + $imgCost : $base; + {"type":"usd","usd": $is15 ? $total * 1.43 : $total} ) """, ), @@ -593,6 +594,8 @@ class GrokVideoNode(IO.ComfyNode): ) -> IO.NodeOutput: if image is None and model == "grok-imagine-video-1.5": raise ValueError(f"The '{model}' model requires an input image; connect one to the 'image' input.") + if resolution == "1080p" and model != "grok-imagine-video-1.5": + raise ValueError(f"1080p resolution is only available for grok-imagine-video-1.5, not '{model}'.") image_url = None if image is not None: if get_number_of_images(image) != 1: From 5236cd02e61362677c22e39a06eb0e44c79c9633 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 24 Jun 2026 17:57:46 +0300 Subject: [PATCH 015/211] [Partner Nodes] feat(ByteDance): add 4K resolution support for SeeDance 2.0 (#14614) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/bytedance.py | 22 ++++++++++--- comfy_api_nodes/nodes_bytedance.py | 50 ++++++++++++++++++++---------- 2 files changed, 50 insertions(+), 22 deletions(-) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 47f24586c..999b51d39 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -163,15 +163,27 @@ class SeedanceVirtualLibraryCreateAssetRequest(BaseModel): asset_type: str | None = Field(None, description="BytePlus asset type. Defaults to Image server-side when omitted.") -# Dollars per 1K tokens, keyed by (model_id, has_video_input). +# Dollars per 1K tokens, keyed by (model_id, has_video_input, resolution). SEEDANCE2_PRICE_PER_1K_TOKENS = { - ("dreamina-seedance-2-0-260128", False): 0.007, - ("dreamina-seedance-2-0-260128", True): 0.0043, - ("dreamina-seedance-2-0-fast-260128", False): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True): 0.0033, + ("dreamina-seedance-2-0-260128", False, "480p"): 0.007, + ("dreamina-seedance-2-0-260128", True, "480p"): 0.0043, + ("dreamina-seedance-2-0-260128", False, "720p"): 0.007, + ("dreamina-seedance-2-0-260128", True, "720p"): 0.0043, + ("dreamina-seedance-2-0-260128", False, "1080p"): 0.0077, + ("dreamina-seedance-2-0-260128", True, "1080p"): 0.0047, + ("dreamina-seedance-2-0-260128", False, "4k"): 0.004, + ("dreamina-seedance-2-0-260128", True, "4k"): 0.0024, + ("dreamina-seedance-2-0-fast-260128", False, "480p"): 0.0056, + ("dreamina-seedance-2-0-fast-260128", True, "480p"): 0.0033, + ("dreamina-seedance-2-0-fast-260128", False, "720p"): 0.0056, + ("dreamina-seedance-2-0-fast-260128", True, "720p"): 0.0033, } +def seedance2_price_per_1k_tokens(model_id: str, has_video_input: bool, resolution: str) -> float | None: + return SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input, resolution)) + + RECOMMENDED_PRESETS = [ ("1024x1024 (1:1)", 1024, 1024), ("864x1152 (3:4)", 864, 1152), diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index c30ddc446..6192b35bf 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -15,7 +15,6 @@ from comfy_api_nodes.apis.bytedance import ( RECOMMENDED_PRESETS_SEEDREAM_4_0, RECOMMENDED_PRESETS_SEEDREAM_4_5, RECOMMENDED_PRESETS_SEEDREAM_5_LITE, - SEEDANCE2_PRICE_PER_1K_TOKENS, SEEDANCE2_REF_VIDEO_PIXEL_LIMITS, VIDEO_TASKS_EXECUTION_TIME, GetAssetResponse, @@ -40,6 +39,7 @@ from comfy_api_nodes.apis.bytedance import ( TaskVideoContentUrl, Text2ImageTaskCreationRequest, Text2VideoTaskCreationRequest, + seedance2_price_per_1k_tokens, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -141,7 +141,7 @@ SEEDANCE2_RATIO_WH = { "9:16": (9, 16), "21:9": (21, 9), } -SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080} +SEEDANCE2_RES_SHORT_SIDE = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160} def _seedance2_target_dims(resolution: str, ratio: str, image: torch.Tensor) -> tuple[int, int]: @@ -377,9 +377,9 @@ async def _seedance_virtual_library_upload_video_asset( return f"asset://{create_resp.asset_id}" -def _seedance2_price_extractor(model_id: str, has_video_input: bool): +def _seedance2_price_extractor(model_id: str, has_video_input: bool, resolution: str): """Returns a price_extractor closure for Seedance 2.0 poll_op.""" - rate = SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input)) + rate = seedance2_price_per_1k_tokens(model_id, has_video_input, resolution) if rate is None: return None @@ -1621,7 +1621,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p"])), + IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p", "4k"])), IO.DynamicCombo.Option("Seedance 2.0 Fast", _seedance2_text_inputs(["480p", "720p"])), ], tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", @@ -1660,11 +1660,15 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; - $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $pricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1703,7 +1707,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -1724,7 +1728,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_text_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), + _seedance2_text_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -1791,11 +1795,15 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; - $pricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $pricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $cost := $dur * $rate * $pricePer1K / 1000; @@ -1913,7 +1921,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=False), + price_extractor=_seedance2_price_extractor(model_id, has_video_input=False, resolution=model["resolution"]), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -2010,7 +2018,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Seedance 2.0", - _seedance2_reference_inputs(["480p", "720p", "1080p"], default_ratio="adaptive"), + _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), ), IO.DynamicCombo.Option( "Seedance 2.0 Fast", @@ -2056,13 +2064,19 @@ class ByteDance2ReferenceNode(IO.ComfyNode): $rate480 := 10044; $rate720 := 21600; $rate1080 := 48800; + $rate4k := 195200; $m := widgets.model; $hasVideo := $lookup(inputGroups, "model.reference_videos") > 0; - $noVideoPricePer1K := $contains($m, "fast") ? 0.008008 : 0.01001; - $videoPricePer1K := $contains($m, "fast") ? 0.004719 : 0.006149; $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); - $rate := $res = "1080p" ? $rate1080 : + $noVideoPricePer1K := $res = "4k" ? 0.00572 : + $res = "1080p" ? 0.011011 : + $contains($m, "fast") ? 0.008008 : 0.01001; + $videoPricePer1K := $res = "4k" ? 0.003432 : + $res = "1080p" ? 0.006721 : + $contains($m, "fast") ? 0.004719 : 0.006149; + $rate := $res = "4k" ? $rate4k : + $res = "1080p" ? $rate1080 : $res = "720p" ? $rate720 : $rate480; $noVideoCost := $dur * $rate * $noVideoPricePer1K / 1000; @@ -2258,7 +2272,9 @@ class ByteDance2ReferenceNode(IO.ComfyNode): ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{initial_response.id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor(model_id, has_video_input=has_video_input), + price_extractor=_seedance2_price_extractor( + model_id, has_video_input=has_video_input, resolution=model["resolution"] + ), poll_interval=9, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) From b22d0fb9c0f96552fa2120b251fb0a2712cf8e9b Mon Sep 17 00:00:00 2001 From: "Yousef R. Gamaleldin" <81116377+yousef-rafat@users.noreply.github.com> Date: Thu, 25 Jun 2026 04:39:10 +0300 Subject: [PATCH 016/211] feat: Add Support For Simple Seed (CORE-295) (#14616) --- comfy_extras/nodes_seed.py | 33 +++++++++++++++++++++++++++++++++ nodes.py | 1 + 2 files changed, 34 insertions(+) create mode 100644 comfy_extras/nodes_seed.py diff --git a/comfy_extras/nodes_seed.py b/comfy_extras/nodes_seed.py new file mode 100644 index 000000000..e64f1d7e3 --- /dev/null +++ b/comfy_extras/nodes_seed.py @@ -0,0 +1,33 @@ +import sys +from typing_extensions import override + +from comfy_api.latest import ComfyExtension, io + + +class SeedNode(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedNode", + display_name="Seed", + search_aliases=["seed", "random"], + category="utilities", + inputs=[ + io.Int.Input("seed", min=0, max=sys.maxsize, control_after_generate=io.ControlAfterGenerate.fixed), + ], + outputs=[io.Int.Output(display_name="seed")], + ) + + @classmethod + def execute(cls, seed: int) -> io.NodeOutput: + return io.NodeOutput(seed) + + +class SeedExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [SeedNode] + + +async def comfy_entrypoint() -> SeedExtension: + return SeedExtension() diff --git a/nodes.py b/nodes.py index 166e02d3d..ad172890d 100644 --- a/nodes.py +++ b/nodes.py @@ -2473,6 +2473,7 @@ async def init_builtin_extra_nodes(): "nodes_gaussian_splat.py", "nodes_triposplat.py", "nodes_depth_anything_3.py", + "nodes_seed.py", ] import_failed = [] From 64e1d740b82844a9e71383b51e1448bf0126106e Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 24 Jun 2026 20:37:30 -0700 Subject: [PATCH 017/211] Add advanced krea 2 model merging node. (#14621) --- .../nodes_model_merging_model_specific.py | 31 +++++++++++++++++++ 1 file changed, 31 insertions(+) diff --git a/comfy_extras/nodes_model_merging_model_specific.py b/comfy_extras/nodes_model_merging_model_specific.py index 2fa684b3a..e563d950b 100644 --- a/comfy_extras/nodes_model_merging_model_specific.py +++ b/comfy_extras/nodes_model_merging_model_specific.py @@ -337,6 +337,36 @@ class ModelMergeQwenImage(comfy_extras.nodes_model_merging.ModelMergeBlocks): return {"required": arg_dict} +class ModelMergeKrea2(comfy_extras.nodes_model_merging.ModelMergeBlocks): + CATEGORY = "model/merging/model specific" + + @classmethod + def INPUT_TYPES(s): + arg_dict = { "model1": ("MODEL",), + "model2": ("MODEL",)} + + argument = ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}) + + arg_dict["first."] = argument + arg_dict["tmlp."] = argument + arg_dict["txtmlp."] = argument + arg_dict["tproj."] = argument + + for i in range(2): + arg_dict["txtfusion.layerwise_blocks.{}.".format(i)] = argument + + arg_dict["txtfusion.projector."] = argument + + for i in range(2): + arg_dict["txtfusion.refiner_blocks.{}.".format(i)] = argument + + for i in range(28): + arg_dict["blocks.{}.".format(i)] = argument + + arg_dict["last."] = argument + + return {"required": arg_dict} + NODE_CLASS_MAPPINGS = { "ModelMergeSD1": ModelMergeSD1, "ModelMergeSD2": ModelMergeSD1, #SD1 and SD2 have the same blocks @@ -353,4 +383,5 @@ NODE_CLASS_MAPPINGS = { "ModelMergeCosmosPredict2_2B": ModelMergeCosmosPredict2_2B, "ModelMergeCosmosPredict2_14B": ModelMergeCosmosPredict2_14B, "ModelMergeQwenImage": ModelMergeQwenImage, + "ModelMergeKrea2": ModelMergeKrea2, } From b0ec19804ff68bfafb2729fe1c83e5f451a6061e Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Thu, 25 Jun 2026 14:54:53 +0900 Subject: [PATCH 018/211] chore(openapi): sync shared API contract from cloud@4118910 (#14619) --- openapi.yaml | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/openapi.yaml b/openapi.yaml index cee8a4763..c6a8621cc 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -1692,6 +1692,12 @@ paths: schema: $ref: '#/components/schemas/ErrorResponse' description: Unsupported media type + "422": + content: + application/json: + schema: + $ref: '#/components/schemas/ErrorResponse' + description: Validation error (e.g., disallowed model_type tag) "500": content: application/json: @@ -2137,6 +2143,12 @@ paths: schema: $ref: '#/components/schemas/ErrorResponse' description: Source asset with given hash not found + "422": + content: + application/json: + schema: + $ref: '#/components/schemas/ErrorResponse' + description: Validation error (e.g., disallowed model_type tag) "500": content: application/json: @@ -2992,7 +3004,7 @@ paths: format: uuid type: string - description: | - When present, each output item in the response receives a `short_url` field containing an owner-gated durable link for that asset. Omit this parameter (the default) to receive a response identical to the no-param baseline. The value selects the link's lifetime: use `ephemeral_tool_chain` for short-lived machine-to-machine handoffs (~15 minutes); use `default` for durable human-revisitable links (30 days). Links are minted only for the authenticated request owner and are not resolvable by other users. + When present, each output item in the response receives a `short_url` field containing a short link for that asset. Omit this parameter (the default) to receive a response identical to the no-param baseline. The value selects the link's lifetime and auth model: use `ephemeral_tool_chain` for short-lived (≤5 minute) machine-to-machine handoffs — these are public bearer links where the link ID itself is the credential, so anyone holding the link can resolve it (intended for pasting into an agent/MCP tool chain); use `default` for durable (30 day) human-revisitable links, which are owner-gated and resolvable only by the authenticated owner. Links are always minted under the authenticated request owner's identity; the auth model is selected by the server and is never settable by the caller. in: query name: short_link schema: From dac4ea3a80ff38796b25d074d28fd5fcec1a4d24 Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Thu, 25 Jun 2026 10:34:09 -0400 Subject: [PATCH 019/211] feat: Bounding boxes canvas and Ideogram JSON prompt (#14537) --- comfy_api/latest/_io.py | 39 +++++ comfy_extras/color_util.py | 23 +++ comfy_extras/nodes_bounding_boxes.py | 253 +++++++++++++++++++++++++++ comfy_extras/nodes_color.py | 9 +- comfy_extras/nodes_json_prompt.py | 77 ++++++++ comfy_extras/nodes_string.py | 53 ++++++ nodes.py | 2 + 7 files changed, 453 insertions(+), 3 deletions(-) create mode 100644 comfy_extras/color_util.py create mode 100644 comfy_extras/nodes_bounding_boxes.py create mode 100644 comfy_extras/nodes_json_prompt.py diff --git a/comfy_api/latest/_io.py b/comfy_api/latest/_io.py index 012fae3ac..58e49d8e2 100644 --- a/comfy_api/latest/_io.py +++ b/comfy_api/latest/_io.py @@ -891,6 +891,14 @@ class Tracks(ComfyTypeIO): track_visibility: torch.Tensor Type = TrackDict +@comfytype(io_type="DICT") +class Dict(ComfyTypeIO): + Type = dict + +@comfytype(io_type="ARRAY") +class Array(ComfyTypeIO): + Type = list + @comfytype(io_type="COMFY_MULTITYPED_V3") class MultiType: Type = Any @@ -1279,6 +1287,19 @@ class Color(ComfyTypeIO): def as_dict(self): return super().as_dict() + +@comfytype(io_type="COLORS") +class Colors(ComfyTypeIO): + Type = list[Color.Type] + + class Input(WidgetInput): + def __init__(self, id: str, display_name: str=None, optional=False, tooltip: str=None, + socketless: bool=True, default: list[str]=None, advanced: bool=None): + super().__init__(id, display_name, optional, tooltip, None, default, socketless, None, None, None, None, advanced) + if default is None: + self.default = [] + + @comfytype(io_type="BOUNDING_BOX") class BoundingBox(ComfyTypeIO): class BoundingBoxDict(TypedDict): @@ -1326,6 +1347,20 @@ class Curve(ComfyTypeIO): return d +@comfytype(io_type="BOUNDING_BOXES") +class BoundingBoxes(ComfyTypeIO): + class BoundingBoxWithMetadata(BoundingBox.BoundingBoxDict): + metadata: dict + Type = list[BoundingBoxWithMetadata] + + class Input(WidgetInput): + def __init__(self, id: str, display_name: str=None, optional=False, tooltip: str=None, + socketless: bool=True, default: list[dict]=None, advanced: bool=None): + super().__init__(id, display_name, optional, tooltip, None, default, socketless, None, None, None, None, advanced) + if default is None: + self.default = [] + + @comfytype(io_type="HISTOGRAM") class Histogram(ComfyTypeIO): """A histogram represented as a list of bin counts.""" @@ -2376,6 +2411,8 @@ __all__ = [ "AnyType", "MultiType", "Tracks", + "Dict", + "Array", "Color", # Dynamic Types "MatchType", @@ -2394,6 +2431,8 @@ __all__ = [ "PriceBadgeDepends", "PriceBadge", "BoundingBox", + "BoundingBoxes", + "Colors", "Curve", "Histogram", "Range", diff --git a/comfy_extras/color_util.py b/comfy_extras/color_util.py new file mode 100644 index 000000000..d50795ae3 --- /dev/null +++ b/comfy_extras/color_util.py @@ -0,0 +1,23 @@ +def hex_to_rgb(value: str) -> tuple[int, int, int]: + h = value.lstrip("#") + if len(h) != 6: + return (255, 255, 255) + try: + return (int(h[0:2], 16), int(h[2:4], 16), int(h[4:6], 16)) + except ValueError: + return (255, 255, 255) + + +def readable_color(rgb: tuple[int, int, int]) -> tuple[int, int, int]: + r, g, b = rgb + lum = 0.299 * r + 0.587 * g + 0.114 * b + if lum >= 130: + return (r, g, b) + t = (130 - lum) / (255 - lum) + return (round(r + (255 - r) * t), round(g + (255 - g) * t), round(b + (255 - b) * t)) + + +def normalize_palette(colors) -> list[str]: + if isinstance(colors, dict): + colors = colors.values() + return [c.upper() for c in colors if isinstance(c, str) and c] diff --git a/comfy_extras/nodes_bounding_boxes.py b/comfy_extras/nodes_bounding_boxes.py new file mode 100644 index 000000000..77cbf8649 --- /dev/null +++ b/comfy_extras/nodes_bounding_boxes.py @@ -0,0 +1,253 @@ +import numpy as np +import torch +from PIL import Image, ImageDraw, ImageEnhance, ImageFont +from typing_extensions import override + +from comfy_api.latest import ComfyExtension, io +from comfy_extras.color_util import hex_to_rgb, normalize_palette, readable_color + +_PREVIEW_LONG_EDGE = 1024 +_PREVIEW_DIM = 0.25 + + +def pixels_to_fractions(box: dict, width: int, height: int) -> dict: + w = width or 1 + h = height or 1 + return { + "x": box.get("x", 0) / w, + "y": box.get("y", 0) / h, + "w": box.get("width", 0) / w, + "h": box.get("height", 0) / h, + } + + +def fractions_to_pixels(box: dict, width: int, height: int) -> dict: + x, y = box.get("x", 0.0), box.get("y", 0.0) + w, h = box.get("w", 0.0), box.get("h", 0.0) + if w < 0: + x, w = x + w, -w + if h < 0: + y, h = y + h, -h + return { + "x": round(x * width), + "y": round(y * height), + "width": round(w * width), + "height": round(h * height), + } + + +def fractions_to_bbox_frame(boxes: list, width: int, height: int) -> list: + pixels = [ + fractions_to_pixels(box, width, height) + for box in boxes + if isinstance(box, dict) + ] + return [pixels] if pixels else [] + + +def _font(size: int): + try: + return ImageFont.load_default(size) + except Exception: + return ImageFont.load_default() + + +def _wrap(draw, text: str, font, max_w: float) -> list[str]: + lines = [] + for para in text.split("\n"): + line = "" + for word in para.split(): + test = word if not line else line + " " + word + if line and draw.textlength(test, font=font) > max_w: + lines.append(line) + line = word + else: + line = test + lines.append(line) + return lines + + +def _bg_from_image(image) -> Image.Image | None: + if image is None: + return None + try: + arr = (image[0].detach().cpu().numpy() * 255).clip(0, 255).astype(np.uint8) + return Image.fromarray(arr) + except Exception: + return None + + +def render_preview(regions, width, height, bg=None): + if bg is not None: + iw, ih = bg.size + long_edge = max(iw, ih) or 1 + scale = min(1.0, _PREVIEW_LONG_EDGE / long_edge) + rw, rh = max(1, round(iw * scale)), max(1, round(ih * scale)) + base = bg.convert("RGB").resize((rw, rh), Image.LANCZOS) + base = ImageEnhance.Brightness(base).enhance(_PREVIEW_DIM) + img = base.convert("RGBA") + else: + long_edge = max(width, height) or 1 + scale = min(1.0, _PREVIEW_LONG_EDGE / long_edge) + rw, rh = max(1, round(width * scale)), max(1, round(height * scale)) + grey = round(_PREVIEW_DIM * 128) + img = Image.new("RGBA", (rw, rh), (grey, grey, grey, 255)) + + overlay = Image.new("RGBA", (rw, rh), (0, 0, 0, 0)) + draw = ImageDraw.Draw(overlay) + fs = max(10, round(rh / 64)) + font = _font(fs) + tag_font = _font(max(9, fs - 2)) + line_h = fs + 2 + + for i, region in enumerate(regions): + if not isinstance(region, dict): + continue + palette = [c for c in (region.get("palette") or []) if c] + r, g, b = hex_to_rgb(palette[0]) if palette else (140, 140, 140) + x1 = max(0, min(rw, round(region.get("x", 0) * rw))) + y1 = max(0, min(rh, round(region.get("y", 0) * rh))) + x2 = max(0, min(rw, round((region.get("x", 0) + region.get("w", 0)) * rw))) + y2 = max(0, min(rh, round((region.get("y", 0) + region.get("h", 0)) * rh))) + if x2 < x1: + x1, x2 = x2, x1 + if y2 < y1: + y1, y2 = y2, y1 + + draw.rectangle([x1, y1, x2, y2], outline=(r, g, b, 255), width=2) + + swatches = palette[:5] + if swatches and (x2 - x1) > 2: + sh = max(5, fs // 2) + seg = (x2 - x1) / len(swatches) + for p, hexc in enumerate(swatches): + sx = x1 + round(p * seg) + draw.rectangle([sx, y1, x1 + round((p + 1) * seg), y1 + sh], fill=hex_to_rgb(hexc)) + + etype = "text" if region.get("type") == "text" else "obj" + tag = str(i + 1).zfill(2) + tw = draw.textlength(tag, font=tag_font) + draw.rectangle([x1, y1, x1 + tw + 6, y1 + fs + 2], fill=(r, g, b, 255)) + tag_fill = (0, 0, 0, 255) if (0.299 * r + 0.587 * g + 0.114 * b) > 140 else (255, 255, 255, 255) + draw.text((x1 + 3, y1 + 1), tag, fill=tag_fill, font=tag_font) + + body = region.get("desc", "") or "" + if etype == "text" and region.get("text"): + body = '"%s"%s' % (region["text"], " — " + body if body else "") + if body and (x2 - x1) > 8: + ty = y1 + fs + 5 + for line in _wrap(draw, body, font, x2 - x1 - 8): + if ty > y2: + break + draw.text((x1 + 4, ty), line, fill=readable_color((r, g, b)) + (255,), font=font) + ty += line_h + + composed = Image.alpha_composite(img, overlay).convert("RGB") + arr = np.asarray(composed, dtype=np.float32) / 255.0 + return torch.from_numpy(arr).unsqueeze(0) + + +def boxes_to_regions(boxes, width: int, height: int) -> list: + regions: list = [] + if not isinstance(boxes, list): + return regions + for box in boxes: + if not isinstance(box, dict): + continue + meta = box.get("metadata") + meta = meta if isinstance(meta, dict) else {} + regions.append({ + **pixels_to_fractions(box, width, height), + "type": meta.get("type", "obj"), + "text": meta.get("text", ""), + "desc": meta.get("desc", ""), + "palette": meta.get("palette", []), + }) + return regions + + +def _norm_bbox(region: dict) -> list[int]: + def grid(value: float) -> int: + return max(0, min(1000, round(value * 1000))) + + x, y = region.get("x", 0.0), region.get("y", 0.0) + w, h = region.get("w", 0.0), region.get("h", 0.0) + ymin, xmin, ymax, xmax = grid(y), grid(x), grid(y + h), grid(x + w) + if ymin > ymax: + ymin, ymax = ymax, ymin + if xmin > xmax: + xmin, xmax = xmax, xmin + return [ymin, xmin, ymax, xmax] + + +def build_elements(regions: list) -> list: + elements = [] + for region in regions: + if not isinstance(region, dict): + continue + etype = "text" if region.get("type") == "text" else "obj" + element = {"type": etype} + element["bbox"] = _norm_bbox(region) + if etype == "text": + element["text"] = region.get("text", "") + element["desc"] = region.get("desc", "") + palette = normalize_palette(region.get("palette", [])) + if palette: + element["color_palette"] = palette[:5] + elements.append(element) + return elements + + +class CreateBoundingBoxes(io.ComfyNode): + @classmethod + def define_schema(cls): + editor_state = io.BoundingBoxes.Input( + "editor_state", + socketless=False, + tooltip="Draw bounding boxes and set each box type, text, description, color palette. Start with background element first and foreground last.", + ) + return io.Schema( + node_id="CreateBoundingBoxes", + display_name="Create Bounding Boxes", + category="utilities", + description="Draw bounding boxes in a canvas. Outputs Ideogram prompt elements, pixel-space bounding boxes, and a preview image.", + inputs=[ + io.Image.Input( + "background", + optional=True, + tooltip="Optional image used as background in the canvas and preview.", + ), + io.Int.Input("width", default=1024, min=64, max=16384, step=16, + tooltip="Width of the canvas and the pixel grid for the bounding boxes."), + io.Int.Input("height", default=1024, min=64, max=16384, step=16, + tooltip="Height of the canvas and the pixel grid for the bounding boxes."), + editor_state, + ], + outputs=[ + io.Image.Output(display_name="preview"), + io.BoundingBox.Output(display_name="bboxes"), + io.Array.Output(display_name="elements"), + ], + is_experimental=True, + ) + + @classmethod + def execute(cls, width, height, editor_state=None, background=None) -> io.NodeOutput: + regions = boxes_to_regions(editor_state, width, height) + preview = render_preview(regions, width, height, _bg_from_image(background)) + return io.NodeOutput( + preview, + fractions_to_bbox_frame(regions, width, height), + build_elements(regions), + ui={"dims": [width, height]}, + ) + + +class BoundingBoxesExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [CreateBoundingBoxes] + + +async def comfy_entrypoint() -> BoundingBoxesExtension: + return BoundingBoxesExtension() diff --git a/comfy_extras/nodes_color.py b/comfy_extras/nodes_color.py index 688254e4e..f58e51bff 100644 --- a/comfy_extras/nodes_color.py +++ b/comfy_extras/nodes_color.py @@ -1,5 +1,6 @@ from typing_extensions import override from comfy_api.latest import ComfyExtension, io +from comfy_extras.color_util import hex_to_rgb class ColorToRGBInt(io.ComfyNode): @@ -24,9 +25,11 @@ class ColorToRGBInt(io.ComfyNode): # expect format #RRGGBB if len(color) != 7 or color[0] != "#": raise ValueError("Color must be in format #RRGGBB") - r = int(color[1:3], 16) - g = int(color[3:5], 16) - b = int(color[5:7], 16) + try: + int(color[1:], 16) + except ValueError: + raise ValueError("Color must be in format #RRGGBB") from None + r, g, b = hex_to_rgb(color) rgb_int = r * 256 * 256 + g * 256 + b return io.NodeOutput(rgb_int, color) diff --git a/comfy_extras/nodes_json_prompt.py b/comfy_extras/nodes_json_prompt.py new file mode 100644 index 000000000..206f5aa71 --- /dev/null +++ b/comfy_extras/nodes_json_prompt.py @@ -0,0 +1,77 @@ +from typing_extensions import override + +from comfy_api.latest import ComfyExtension, io +from comfy_extras.color_util import normalize_palette + + +class BuildJsonPromptIdeogram(io.ComfyNode): + @classmethod + def define_schema(cls): + color_palette = io.Colors.Input( + "color_palette", + socketless=False, + tooltip="Hex color codes that steer the image's dominant colors. Up to 16 entries.", + ) + return io.Schema( + node_id="BuildJsonPromptIdeogram", + display_name="Build JSON Prompt (Ideogram)", + category="text", + description="Build a JSON prompt for the Ideogram 4 model.", + inputs=[ + io.Array.Input("element", tooltip="Prompt elements from the node Create Bounding Boxes."), + io.String.Input("high_level_description", multiline=True, default="", + tooltip="Optional description of the image in one or two sentences. Strongly recommended."), + io.String.Input("background", multiline=True, default="", + tooltip="Mandatory description of the image background or environment."), + io.DynamicCombo.Input("style", options=[ + io.DynamicCombo.Option("none", []), + io.DynamicCombo.Option("photo", [io.String.Input("photo", default="", tooltip="Camera or lens details for photographic outputs (e.g. 35mm, f/1.4, bokeh).")]), + io.DynamicCombo.Option("art_style", [io.String.Input("art_style", default="", tooltip="Art style description (e.g. flat vector illustration, bold outlines).")]), + ]), + io.String.Input("aesthetics", default="", tooltip="Mandatory aesthetic keywords (e.g. moody, cinematic, desaturated)."), + io.String.Input("lighting", default="", tooltip="Mandatory lighting description (e.g. golden hour, rim light, dramatic shadows)."), + io.String.Input("medium", default="", tooltip="Mandatory medium type (e.g. photograph, illustration, 3d_render, painting, graphic_design). When style = photo, set to photograph."), + color_palette, + ], + outputs=[io.Dict.Output(display_name="prompt")], + is_experimental=True, + ) + + @classmethod + def execute(cls, element, style, high_level_description="", background="", + aesthetics="", lighting="", medium="", color_palette=None) -> io.NodeOutput: + elements = element if isinstance(element, list) else [] + kind = style.get("style", "none") if isinstance(style, dict) else "none" + photo = style.get("photo", "") if isinstance(style, dict) else "" + art_style = style.get("art_style", "") if isinstance(style, dict) else "" + palette = normalize_palette(color_palette or []) + + caption: dict = {} + if high_level_description.strip(): + caption["high_level_description"] = high_level_description + if kind != "none": + style_desc: dict = {"aesthetics": aesthetics, "lighting": lighting} + if kind == "photo": + style_desc["photo"] = photo + style_desc["medium"] = medium + else: + style_desc["medium"] = medium + style_desc["art_style"] = art_style + if palette: + style_desc["color_palette"] = palette + caption["style_description"] = style_desc + caption["compositional_deconstruction"] = { + "background": background, + "elements": elements, + } + return io.NodeOutput(caption) + + +class JsonPromptExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [BuildJsonPromptIdeogram] + + +async def comfy_entrypoint() -> JsonPromptExtension: + return JsonPromptExtension() diff --git a/comfy_extras/nodes_string.py b/comfy_extras/nodes_string.py index 97485c8c5..21929ae63 100644 --- a/comfy_extras/nodes_string.py +++ b/comfy_extras/nodes_string.py @@ -440,6 +440,57 @@ class JsonExtractString(io.ComfyNode): except (json.JSONDecodeError, TypeError): return io.NodeOutput("") + +def _dump_json(value, indent): + return json.dumps(value, ensure_ascii=False, indent=indent or None) + + +class ConvertDictionaryToString(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ConvertDictionaryToString", + display_name="Convert Dictionary to String", + category="text", + search_aliases=["json", "dict to json", "stringify", "serialize", "dict to string"], + inputs=[ + io.Dict.Input("dictionary"), + io.Int.Input("indent", default=2, min=0, max=8, + tooltip="Spaces per indent level. 0 produces compact single-line string."), + ], + outputs=[ + io.String.Output(), + ], + ) + + @classmethod + def execute(cls, dictionary, indent=2): + return io.NodeOutput(_dump_json(dictionary, indent)) + + +class ConvertArrayToString(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ConvertArrayToString", + display_name="Convert Array to String", + category="text", + search_aliases=["json", "list to json", "stringify", "serialize", "list to string", "array to json"], + inputs=[ + io.Array.Input("array"), + io.Int.Input("indent", default=2, min=0, max=8, + tooltip="Spaces per indent level. 0 produces compact single-line string."), + ], + outputs=[ + io.String.Output(), + ], + ) + + @classmethod + def execute(cls, array, indent=2): + return io.NodeOutput(_dump_json(array, indent)) + + class StringExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[io.ComfyNode]]: @@ -457,6 +508,8 @@ class StringExtension(ComfyExtension): RegexExtract, RegexReplace, JsonExtractString, + ConvertDictionaryToString, + ConvertArrayToString, ] async def comfy_entrypoint() -> StringExtension: diff --git a/nodes.py b/nodes.py index ad172890d..028e58c77 100644 --- a/nodes.py +++ b/nodes.py @@ -2374,6 +2374,8 @@ async def init_builtin_extra_nodes(): "nodes_images.py", "nodes_video_model.py", "nodes_ideogram4.py", + "nodes_bounding_boxes.py", + "nodes_json_prompt.py", "nodes_train.py", "nodes_dataset.py", "nodes_sag.py", From e22f1500f999775c313b7577d302385762fb0b3a Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 25 Jun 2026 17:57:04 +0300 Subject: [PATCH 020/211] [Partner Nodes] feat(ByteDance): add support for SeeDance-2.0-Mini video model (#14626) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/bytedance.py | 8 ++++++++ comfy_api_nodes/nodes_bytedance.py | 23 ++++++++++++++++++++--- 2 files changed, 28 insertions(+), 3 deletions(-) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 999b51d39..2d65d8645 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -177,6 +177,10 @@ SEEDANCE2_PRICE_PER_1K_TOKENS = { ("dreamina-seedance-2-0-fast-260128", True, "480p"): 0.0033, ("dreamina-seedance-2-0-fast-260128", False, "720p"): 0.0056, ("dreamina-seedance-2-0-fast-260128", True, "720p"): 0.0033, + ("dreamina-seedance-2-0-mini", False, "480p"): 0.0035, + ("dreamina-seedance-2-0-mini", True, "480p"): 0.0021, + ("dreamina-seedance-2-0-mini", False, "720p"): 0.0035, + ("dreamina-seedance-2-0-mini", True, "720p"): 0.0021, } @@ -278,6 +282,10 @@ SEEDANCE2_REF_VIDEO_PIXEL_LIMITS = { "480p": {"min": 409_600, "max": 927_408}, "720p": {"min": 409_600, "max": 927_408}, }, + "dreamina-seedance-2-0-mini": { + "480p": {"min": 409_600, "max": 927_408}, + "720p": {"min": 409_600, "max": 927_408}, + }, } # The time in this dictionary are given for 10 seconds duration. diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 6192b35bf..f22415abd 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -89,6 +89,7 @@ BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT = "/proxy/byteplus-seedance2/api/v3/cont SEEDANCE_MODELS = { "Seedance 2.0": "dreamina-seedance-2-0-260128", "Seedance 2.0 Fast": "dreamina-seedance-2-0-fast-260128", + "Seedance 2.0 Mini": "dreamina-seedance-2-0-mini", } DEPRECATED_MODELS = {"seedance-1-0-lite-t2v-250428", "seedance-1-0-lite-i2v-250428"} @@ -1623,8 +1624,10 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): options=[ IO.DynamicCombo.Option("Seedance 2.0", _seedance2_text_inputs(["480p", "720p", "1080p", "4k"])), IO.DynamicCombo.Option("Seedance 2.0 Fast", _seedance2_text_inputs(["480p", "720p"])), + IO.DynamicCombo.Option("Seedance 2.0 Mini", _seedance2_text_inputs(["480p", "720p"])), ], - tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", + tooltip="Seedance 2.0 for maximum quality; Fast for speed optimization; " + "Mini for the fastest, lowest-cost generation.", ), IO.Int.Input( "seed", @@ -1666,6 +1669,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): $dur := $lookup(widgets, "model.duration"); $pricePer1K := $res = "4k" ? 0.00572 : $res = "1080p" ? 0.011011 : + $contains($m, "mini") ? 0.005005 : $contains($m, "fast") ? 0.008008 : 0.01001; $rate := $res = "4k" ? $rate4k : $res = "1080p" ? $rate1080 : @@ -1734,8 +1738,13 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): "Seedance 2.0 Fast", _seedance2_text_inputs(["480p", "720p"], default_ratio="adaptive"), ), + IO.DynamicCombo.Option( + "Seedance 2.0 Mini", + _seedance2_text_inputs(["480p", "720p"], default_ratio="adaptive"), + ), ], - tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", + tooltip="Seedance 2.0 for maximum quality; Fast for speed optimization; " + "Mini for the fastest, lowest-cost generation.", ), IO.Image.Input( "first_frame", @@ -1801,6 +1810,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): $dur := $lookup(widgets, "model.duration"); $pricePer1K := $res = "4k" ? 0.00572 : $res = "1080p" ? 0.011011 : + $contains($m, "mini") ? 0.005005 : $contains($m, "fast") ? 0.008008 : 0.01001; $rate := $res = "4k" ? $rate4k : $res = "1080p" ? $rate1080 : @@ -2024,8 +2034,13 @@ class ByteDance2ReferenceNode(IO.ComfyNode): "Seedance 2.0 Fast", _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), ), + IO.DynamicCombo.Option( + "Seedance 2.0 Mini", + _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), + ), ], - tooltip="Seedance 2.0 for maximum quality; Seedance 2.0 Fast for speed optimization.", + tooltip="Seedance 2.0 for maximum quality; Fast for speed optimization; " + "Mini for the fastest, lowest-cost generation.", ), IO.Int.Input( "seed", @@ -2071,9 +2086,11 @@ class ByteDance2ReferenceNode(IO.ComfyNode): $dur := $lookup(widgets, "model.duration"); $noVideoPricePer1K := $res = "4k" ? 0.00572 : $res = "1080p" ? 0.011011 : + $contains($m, "mini") ? 0.005005 : $contains($m, "fast") ? 0.008008 : 0.01001; $videoPricePer1K := $res = "4k" ? 0.003432 : $res = "1080p" ? 0.006721 : + $contains($m, "mini") ? 0.003003 : $contains($m, "fast") ? 0.004719 : 0.006149; $rate := $res = "4k" ? $rate4k : $res = "1080p" ? $rate1080 : From 639c8fa788916dc1c6ecd0fb5d65fa2c00e323fb Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Thu, 25 Jun 2026 23:05:34 +0800 Subject: [PATCH 021/211] chore: update workflow templates to v0.10.7 (#14632) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index b05cad045..e0778548c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.19 -comfyui-workflow-templates==0.10.2 +comfyui-workflow-templates==0.10.7 comfyui-embedded-docs==0.5.5 torch torchsde From 1a510f04234e5a213d3985a1a54f65652623f4bc Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 25 Jun 2026 11:23:58 -0700 Subject: [PATCH 022/211] Support int8 models. (#14636) --- comfy/ops.py | 59 +++++++++++++++++-- comfy/quant_ops.py | 14 +++++ requirements.txt | 2 +- .../comfy_quant/test_mixed_precision.py | 58 +++++++++++++++++- 4 files changed, 126 insertions(+), 7 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 3f088a962..634610f1c 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1089,6 +1089,19 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat if ts is None or bs is None: raise ValueError(f"Missing NVFP4 scales for layer {layer_name}") scales = {"scale": ts, "block_scale": bs} + elif module.quant_format == "int8_tensorwise": + scale = pop_scale("weight_scale") + if scale is None: + raise ValueError(f"Missing INT8 weight scale for layer {layer_name}") + scales = {"scale": scale} + params_conf = layer_conf.get("params", {}) + if not isinstance(params_conf, dict): + params_conf = {} + if layer_conf.get("convrot", params_conf.get("convrot", False)): + scales["convrot"] = True + scales["convrot_groupsize"] = int( + layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)) + ) else: raise ValueError(f"Unsupported quantization format: {module.quant_format}") @@ -1131,6 +1144,10 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr quant_conf = {"format": module.quant_format} if getattr(module, '_full_precision_mm_config', False): quant_conf["full_precision_matrix_mult"] = True + params = getattr(module.weight, "_params", None) + if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False): + quant_conf["convrot"] = True + quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256) if extra_quant_conf: quant_conf.update(extra_quant_conf) sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8) @@ -1183,8 +1200,33 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def _forward(self, input, weight, bias): return torch.nn.functional.linear(input, weight, bias) - def forward_comfy_cast_weights(self, input, compute_dtype=None, want_requant=False): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=want_requant) + def forward_comfy_cast_weights( + self, + input, + compute_dtype=None, + want_requant=False, + weight_only_quant=False, + ): + if weight_only_quant: + weight, bias, offload_stream = cast_bias_weight( + self, + input=None, + dtype=self.weight.dtype, + device=input.device, + bias_dtype=input.dtype, + offloadable=True, + compute_dtype=compute_dtype, + want_requant=want_requant, + ) + weight = weight.to(dtype=input.dtype) + else: + weight, bias, offload_stream = cast_bias_weight( + self, + input, + offloadable=True, + compute_dtype=compute_dtype, + want_requant=want_requant, + ) x = self._forward(input, weight, bias) uncast_bias_weight(self, weight, bias, offload_stream) return x @@ -1203,9 +1245,10 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec not getattr(self, 'comfy_force_cast_weights', False) and len(self.weight_function) == 0 and len(self.bias_function) == 0 ) + quantize_input = QUANT_ALGOS.get(getattr(self, 'quant_format', None), {}).get("quantize_input", True) # Training path: quantized forward with compute_dtype backward via autograd function - if (input.requires_grad and _use_quantized): + if (input.requires_grad and _use_quantized and quantize_input): weight, bias, offload_stream = cast_bias_weight( self, @@ -1227,7 +1270,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec return output # Inference path (unchanged) - if _use_quantized: + if _use_quantized and quantize_input: # Reshape 3D tensors to 2D for quantization (needed for NVFP4 and others) input_reshaped = input.reshape(-1, input_shape[2]) if input.ndim == 3 else input @@ -1241,7 +1284,13 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec scale = comfy.model_management.cast_to_device(scale, input.device, None) input = QuantizedTensor.from_float(input_reshaped, self.layout_type, scale=scale) - output = self.forward_comfy_cast_weights(input, compute_dtype, want_requant=isinstance(input, QuantizedTensor)) + weight_only_quant = _use_quantized and not quantize_input and isinstance(self.weight, QuantizedTensor) + output = self.forward_comfy_cast_weights( + input, + compute_dtype, + want_requant=isinstance(input, QuantizedTensor), + weight_only_quant=weight_only_quant, + ) # Reshape output back to 3D if input was 3D if reshaped_3d: diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index b90bcfd25..44f25a97e 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -10,6 +10,7 @@ try: QuantizedLayout, TensorCoreFP8Layout as _CKFp8Layout, TensorCoreNVFP4Layout as _CKNvfp4Layout, + TensorWiseINT8Layout as _CKTensorWiseINT8Layout, register_layout_op, register_layout_class, get_layout_class, @@ -47,6 +48,9 @@ except ImportError as e: class _CKNvfp4Layout: pass + class _CKTensorWiseINT8Layout: + pass + def register_layout_class(name, cls): pass @@ -174,6 +178,7 @@ class TensorCoreFP8E5M2Layout(_TensorCoreFP8LayoutBase): # Backward compatibility alias - default to E4M3 TensorCoreFP8Layout = TensorCoreFP8E4M3Layout +TensorWiseINT8Layout = _CKTensorWiseINT8Layout # ============================================================================== @@ -184,6 +189,7 @@ register_layout_class("TensorCoreFP8Layout", TensorCoreFP8Layout) register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout) register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout) register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout) +register_layout_class("TensorWiseINT8Layout", _CKTensorWiseINT8Layout) if _CK_MXFP8_AVAILABLE: register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout) @@ -214,6 +220,13 @@ if _CK_MXFP8_AVAILABLE: "group_size": 32, } +QUANT_ALGOS["int8_tensorwise"] = { + "storage_t": torch.int8, + "parameters": {"weight_scale"}, + "comfy_tensor_layout": "TensorWiseINT8Layout", + "quantize_input": False, +} + # ============================================================================== # Re-exports for backward compatibility @@ -226,6 +239,7 @@ __all__ = [ "TensorCoreFP8E4M3Layout", "TensorCoreFP8E5M2Layout", "TensorCoreNVFP4Layout", + "TensorWiseINT8Layout", "QUANT_ALGOS", "register_layout_op", ] diff --git a/requirements.txt b/requirements.txt index e0778548c..793203a9a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.10 +comfy-kitchen==0.2.11 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 diff --git a/tests-unit/comfy_quant/test_mixed_precision.py b/tests-unit/comfy_quant/test_mixed_precision.py index 7c740491d..43b4b7ce9 100644 --- a/tests-unit/comfy_quant/test_mixed_precision.py +++ b/tests-unit/comfy_quant/test_mixed_precision.py @@ -228,6 +228,62 @@ class TestMixedPrecisionOps(unittest.TestCase): with self.assertRaises(KeyError): model.load_state_dict(state_dict, strict=False) + def test_int8_convrot_metadata_loads_into_params(self): + """ConvRot metadata must reach TensorWiseINT8Layout params.""" + torch.manual_seed(123) + layer_quant_config = { + "layer": { + "format": "int8_tensorwise", + "convrot": True, + "convrot_groupsize": 256, + } + } + weight = torch.randn(16, 256, dtype=torch.bfloat16) + bias = torch.randn(16, dtype=torch.bfloat16) + q_weight = QuantizedTensor.from_float( + weight, + "TensorWiseINT8Layout", + per_channel=True, + convrot=True, + convrot_groupsize=256, + ) + state_dict = { + "layer.weight": q_weight._qdata, + "layer.bias": bias, + "layer.weight_scale": q_weight._params.scale, + } + + state_dict, _ = comfy.utils.convert_old_quants( + state_dict, + metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})}, + ) + model = torch.nn.Module() + model.layer = ops.mixed_precision_ops({}).Linear(256, 16, device="cpu", dtype=torch.bfloat16) + model.load_state_dict(state_dict, strict=False) + + self.assertIsInstance(model.layer.weight, QuantizedTensor) + self.assertEqual(model.layer.weight._layout_cls, "TensorWiseINT8Layout") + self.assertTrue(model.layer.weight._params.convrot) + self.assertEqual(model.layer.weight._params.convrot_groupsize, 256) + + input_tensor = torch.randn(4, 256, dtype=torch.bfloat16) + loaded_out = model.layer(input_tensor) + ref_out = torch.nn.functional.linear(input_tensor, q_weight, bias) + self.assertTrue(torch.equal(loaded_out, ref_out)) + + fp16_input = input_tensor.to(torch.float16) + loaded_fp16_out = model.layer(fp16_input) + ref_fp16_out = torch.nn.functional.linear( + fp16_input, + q_weight.to(dtype=torch.float16), + bias.to(dtype=torch.float16), + ) + self.assertTrue(torch.equal(loaded_fp16_out, ref_fp16_out)) + + saved = model.state_dict() + saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes()) + self.assertTrue(saved_conf["convrot"]) + self.assertEqual(saved_conf["convrot_groupsize"], 256) + if __name__ == "__main__": unittest.main() - From 7cb784e0f48784bb6ed588912e186e5ee1e9ee68 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 25 Jun 2026 15:25:47 -0700 Subject: [PATCH 023/211] Faster int8. (#14641) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 793203a9a..eea7724f3 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.11 +comfy-kitchen==0.2.12 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 470ac36a0a807471a0fb78dc0a5548490c9abae4 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 26 Jun 2026 16:41:29 -0700 Subject: [PATCH 024/211] Fix int8 loras causing lower quality requant with wrong settings. (#14650) * Update comfy-kitchen * Support requantizing with same settings as orig quant. --- comfy/ops.py | 5 ++--- requirements.txt | 2 +- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 634610f1c..6a5090548 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -256,7 +256,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w if (want_requant and len(fns) == 0 or update_weight): seed = comfy.utils.string_to_seed(s.seed_key) if isinstance(orig, QuantizedTensor): - y = QuantizedTensor.from_float(x, s.layout_type, scale="recalculate", stochastic_rounding=seed) + y = orig.requantize_from_float(x, scale="recalculate", stochastic_rounding=seed) else: y = comfy.float.stochastic_rounding(x, orig.dtype, seed=seed) if want_requant and len(fns) == 0: @@ -1306,8 +1306,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def set_weight(self, weight, inplace_update=False, seed=None, return_weight=False, **kwargs): if getattr(self, 'layout_type', None) is not None: - # dtype is now implicit in the layout class - weight = QuantizedTensor.from_float(weight, self.layout_type, scale="recalculate", stochastic_rounding=seed, inplace_ops=True).to(self.weight.dtype) + weight = self.weight.requantize_from_float(weight, scale="recalculate", stochastic_rounding=seed, inplace_ops=True).to(self.weight.dtype) else: weight = weight.to(self.weight.dtype) if return_weight: diff --git a/requirements.txt b/requirements.txt index eea7724f3..d7719178b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.12 +comfy-kitchen==0.2.13 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 603d891eaf045d726d9c23276b4428daf2977624 Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sat, 27 Jun 2026 01:40:31 +0100 Subject: [PATCH 025/211] Update GLSL node to use ANGLE library (CORE-162) (#13195) --- comfy_extras/nodes_glsl.py | 586 +++++++++++++++---------------------- requirements.txt | 4 +- 2 files changed, 239 insertions(+), 351 deletions(-) diff --git a/comfy_extras/nodes_glsl.py b/comfy_extras/nodes_glsl.py index ea7420a73..c7161973a 100644 --- a/comfy_extras/nodes_glsl.py +++ b/comfy_extras/nodes_glsl.py @@ -1,85 +1,68 @@ import os import sys import re +import ctypes import logging -import ctypes.util -import importlib.util from typing import TypedDict import numpy as np import torch import nodes +import comfy_angle from comfy_api.latest import ComfyExtension, io, ui from typing_extensions import override -from utils.install_util import get_missing_requirements_message logger = logging.getLogger(__name__) -def _check_opengl_availability(): - """Early check for OpenGL availability. Raises RuntimeError if unlikely to work.""" - logger.debug("_check_opengl_availability: starting") - missing = [] +def _preload_angle(): + egl_path = comfy_angle.get_egl_path() + gles_path = comfy_angle.get_glesv2_path() - # Check Python packages (using find_spec to avoid importing) - logger.debug("_check_opengl_availability: checking for glfw package") - if importlib.util.find_spec("glfw") is None: - missing.append("glfw") + if sys.platform == "win32": + angle_dir = comfy_angle.get_lib_dir() + os.add_dll_directory(angle_dir) + os.environ["PATH"] = angle_dir + os.pathsep + os.environ.get("PATH", "") - logger.debug("_check_opengl_availability: checking for OpenGL package") - if importlib.util.find_spec("OpenGL") is None: - missing.append("PyOpenGL") - - if missing: - raise RuntimeError( - f"OpenGL dependencies not available.\n{get_missing_requirements_message()}\n" - ) - - # On Linux without display, check if headless backends are available - logger.debug(f"_check_opengl_availability: platform={sys.platform}") - if sys.platform.startswith("linux"): - has_display = os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY") - logger.debug(f"_check_opengl_availability: has_display={bool(has_display)}") - if not has_display: - # Check for EGL or OSMesa libraries - logger.debug("_check_opengl_availability: checking for EGL library") - has_egl = ctypes.util.find_library("EGL") - logger.debug("_check_opengl_availability: checking for OSMesa library") - has_osmesa = ctypes.util.find_library("OSMesa") - - # Error disabled for CI as it fails this check - # if not has_egl and not has_osmesa: - # raise RuntimeError( - # "GLSL Shader node: No display and no headless backend (EGL/OSMesa) found.\n" - # "See error below for installation instructions." - # ) - logger.debug(f"Headless mode: EGL={'yes' if has_egl else 'no'}, OSMesa={'yes' if has_osmesa else 'no'}") - - logger.debug("_check_opengl_availability: completed") + mode = 0 if sys.platform == "win32" else ctypes.RTLD_GLOBAL + ctypes.CDLL(str(egl_path), mode=mode) + ctypes.CDLL(str(gles_path), mode=mode) -# Run early check at import time -logger.debug("nodes_glsl: running _check_opengl_availability at import time") -_check_opengl_availability() - -# OpenGL modules - initialized lazily when context is created -gl = None -glfw = None -EGL = None +# Pre-load ANGLE *before* any PyOpenGL import so that the EGL platform +# plugin picks up ANGLE's libEGL / libGLESv2 instead of system libs. +_preload_angle() +os.environ.setdefault("PYOPENGL_PLATFORM", "egl") -def _import_opengl(): - """Import OpenGL module. Called after context is created.""" - global gl - if gl is None: - logger.debug("_import_opengl: importing OpenGL.GL") - import OpenGL.GL as _gl - gl = _gl - logger.debug("_import_opengl: import completed") - return gl +import OpenGL +OpenGL.USE_ACCELERATE = False +def _patch_find_library(): + """PyOpenGL's EGL platform looks for 'EGL' and 'GLESv2' by short name + via ctypes.util.find_library, but ANGLE ships as 'libEGL' and + 'libGLESv2'. Patch find_library to return the full ANGLE paths so + PyOpenGL loads the same libraries we pre-loaded.""" + if sys.platform == "linux": + return + import ctypes.util + _orig = ctypes.util.find_library + def _patched(name): + if name == 'EGL': + return comfy_angle.get_egl_path() + if name == 'GLESv2': + return comfy_angle.get_glesv2_path() + return _orig(name) + ctypes.util.find_library = _patched + + +_patch_find_library() + +from OpenGL import EGL +from OpenGL import GLES3 as gl + class SizeModeInput(TypedDict): size_mode: str width: int @@ -102,7 +85,7 @@ MAX_OUTPUTS = 4 # fragColor0-3 (MRT) # (-1,-1)---(3,-1) # # v_texCoord is computed from clip space: * 0.5 + 0.5 maps (-1,1) -> (0,1) -VERTEX_SHADER = """#version 330 core +VERTEX_SHADER = """#version 300 es out vec2 v_texCoord; void main() { vec2 verts[3] = vec2[](vec2(-1, -1), vec2(3, -1), vec2(-1, 3)); @@ -126,14 +109,99 @@ void main() { """ -def _convert_es_to_desktop(source: str) -> str: - """Convert GLSL ES (WebGL) shader source to desktop GLSL 330 core.""" - # Remove any existing #version directive - source = re.sub(r"#version\s+\d+(\s+es)?\s*\n?", "", source, flags=re.IGNORECASE) - # Remove precision qualifiers (not needed in desktop GLSL) - source = re.sub(r"precision\s+(lowp|mediump|highp)\s+\w+\s*;\s*\n?", "", source) - # Prepend desktop GLSL version - return "#version 330 core\n" + source + +def _egl_attribs(*values): + """Build an EGL_NONE-terminated EGLint attribute array.""" + vals = list(values) + [EGL.EGL_NONE] + return (ctypes.c_int32 * len(vals))(*vals) + + +# EGL platform extension constants +EGL_PLATFORM_ANGLE_ANGLE = 0x3202 +EGL_PLATFORM_ANGLE_TYPE_ANGLE = 0x3203 +EGL_PLATFORM_ANGLE_TYPE_VULKAN_ANGLE = 0x3450 +EGL_MESA_PLATFORM_SURFACELESS = 0x31DD + + +_eglGetPlatformDisplayEXT = None + +def _get_egl_platform_display_ext(platform, native_display, attribs): + """Call eglGetPlatformDisplayEXT via ctypes (extension, not in PyOpenGL).""" + global _eglGetPlatformDisplayEXT + if _eglGetPlatformDisplayEXT is None: + from OpenGL import platform as _plat + egl_lib = _plat.PLATFORM.EGL + _get_proc = egl_lib.eglGetProcAddress + _get_proc.restype = ctypes.c_void_p + _get_proc.argtypes = [ctypes.c_char_p] + ptr = _get_proc(b"eglGetPlatformDisplayEXT") + if not ptr: + return None + func_type = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.c_uint32, ctypes.c_void_p, ctypes.c_void_p) + _eglGetPlatformDisplayEXT = func_type(ptr) + + raw = _eglGetPlatformDisplayEXT(platform, native_display, attribs) + if not raw: + return None + return ctypes.cast(raw, EGL.EGLDisplay) + + +def _get_egl_display(): + """Get an EGL display, trying the default first then ANGLE's Vulkan + platform for headless environments without a display server.""" + failures = [] + + # Try the default display first (works when X11/Wayland is available) + display = EGL.eglGetDisplay(EGL.EGL_DEFAULT_DISPLAY) + if display: + major, minor = ctypes.c_int32(0), ctypes.c_int32(0) + try: + if EGL.eglInitialize(display, ctypes.byref(major), ctypes.byref(minor)): + return display, major.value, minor.value + except Exception as e: + failures.append(f"default: {e}") + + logger.info("Default EGL display unavailable, trying headless fallbacks") + + # Headless fallback strategies, tried in order: + headless_strategies = [ + ("surfaceless", EGL_MESA_PLATFORM_SURFACELESS, None, None), + ("ANGLE Vulkan", EGL_PLATFORM_ANGLE_ANGLE, None, + _egl_attribs(EGL_PLATFORM_ANGLE_TYPE_ANGLE, EGL_PLATFORM_ANGLE_TYPE_VULKAN_ANGLE)), + ] + + for name, platform, native_display, attribs in headless_strategies: + display = _get_egl_platform_display_ext(platform, native_display, attribs) + if not display: + failures.append(f"{name}: eglGetPlatformDisplayEXT returned no display") + continue + major, minor = ctypes.c_int32(0), ctypes.c_int32(0) + try: + if EGL.eglInitialize(display, ctypes.byref(major), ctypes.byref(minor)): + logger.info(f"Using EGL {name} platform (headless)") + return display, major.value, minor.value + failures.append(f"{name}: eglInitialize returned false") + except Exception as e: + failures.append(f"{name}: {e}") + continue + + details = "\n".join(f" - {f}" for f in failures) + raise RuntimeError( + "Failed to initialize EGL display.\n" + "No display server and no headless EGL platform available.\n" + f"Tried:\n{details}\n" + "Ensure GPU drivers are installed or set DISPLAY for a virtual framebuffer." + ) + + +def _gl_str(name): + """Get an OpenGL string parameter.""" + v = gl.glGetString(name) + if not v: + return "Unknown" + if isinstance(v, bytes): + return v.decode(errors="replace") + return ctypes.string_at(v).decode(errors="replace") def _detect_output_count(source: str) -> int: @@ -159,163 +227,8 @@ def _detect_pass_count(source: str) -> int: return 1 -def _init_glfw(): - """Initialize GLFW. Returns (window, glfw_module). Raises RuntimeError on failure.""" - logger.debug("_init_glfw: starting") - # On macOS, glfw.init() must be called from main thread or it hangs forever - if sys.platform == "darwin": - logger.debug("_init_glfw: skipping on macOS") - raise RuntimeError("GLFW backend not supported on macOS") - - logger.debug("_init_glfw: importing glfw module") - import glfw as _glfw - - logger.debug("_init_glfw: calling glfw.init()") - if not _glfw.init(): - raise RuntimeError("glfw.init() failed") - - try: - logger.debug("_init_glfw: setting window hints") - _glfw.window_hint(_glfw.VISIBLE, _glfw.FALSE) - _glfw.window_hint(_glfw.CONTEXT_VERSION_MAJOR, 3) - _glfw.window_hint(_glfw.CONTEXT_VERSION_MINOR, 3) - _glfw.window_hint(_glfw.OPENGL_PROFILE, _glfw.OPENGL_CORE_PROFILE) - - logger.debug("_init_glfw: calling create_window()") - window = _glfw.create_window(64, 64, "ComfyUI GLSL", None, None) - if not window: - raise RuntimeError("glfw.create_window() failed") - - logger.debug("_init_glfw: calling make_context_current()") - _glfw.make_context_current(window) - logger.debug("_init_glfw: completed successfully") - return window, _glfw - except Exception: - logger.debug("_init_glfw: failed, terminating glfw") - _glfw.terminate() - raise - - -def _init_egl(): - """Initialize EGL for headless rendering. Returns (display, context, surface, EGL_module). Raises RuntimeError on failure.""" - logger.debug("_init_egl: starting") - from OpenGL import EGL as _EGL - from OpenGL.EGL import ( - eglGetDisplay, eglInitialize, eglChooseConfig, eglCreateContext, - eglMakeCurrent, eglCreatePbufferSurface, eglBindAPI, - eglTerminate, eglDestroyContext, eglDestroySurface, - EGL_DEFAULT_DISPLAY, EGL_NO_CONTEXT, EGL_NONE, - EGL_SURFACE_TYPE, EGL_PBUFFER_BIT, EGL_RENDERABLE_TYPE, EGL_OPENGL_BIT, - EGL_RED_SIZE, EGL_GREEN_SIZE, EGL_BLUE_SIZE, EGL_ALPHA_SIZE, EGL_DEPTH_SIZE, - EGL_WIDTH, EGL_HEIGHT, EGL_OPENGL_API, - ) - logger.debug("_init_egl: imports completed") - - display = None - context = None - surface = None - - try: - logger.debug("_init_egl: calling eglGetDisplay()") - display = eglGetDisplay(EGL_DEFAULT_DISPLAY) - if display == _EGL.EGL_NO_DISPLAY: - raise RuntimeError("eglGetDisplay() failed") - - logger.debug("_init_egl: calling eglInitialize()") - major, minor = _EGL.EGLint(), _EGL.EGLint() - if not eglInitialize(display, major, minor): - display = None # Not initialized, don't terminate - raise RuntimeError("eglInitialize() failed") - logger.debug(f"_init_egl: EGL version {major.value}.{minor.value}") - - config_attribs = [ - EGL_SURFACE_TYPE, EGL_PBUFFER_BIT, - EGL_RENDERABLE_TYPE, EGL_OPENGL_BIT, - EGL_RED_SIZE, 8, EGL_GREEN_SIZE, 8, EGL_BLUE_SIZE, 8, EGL_ALPHA_SIZE, 8, - EGL_DEPTH_SIZE, 0, EGL_NONE - ] - configs = (_EGL.EGLConfig * 1)() - num_configs = _EGL.EGLint() - if not eglChooseConfig(display, config_attribs, configs, 1, num_configs) or num_configs.value == 0: - raise RuntimeError("eglChooseConfig() failed") - config = configs[0] - logger.debug(f"_init_egl: config chosen, num_configs={num_configs.value}") - - if not eglBindAPI(EGL_OPENGL_API): - raise RuntimeError("eglBindAPI() failed") - - logger.debug("_init_egl: calling eglCreateContext()") - context_attribs = [ - _EGL.EGL_CONTEXT_MAJOR_VERSION, 3, - _EGL.EGL_CONTEXT_MINOR_VERSION, 3, - _EGL.EGL_CONTEXT_OPENGL_PROFILE_MASK, _EGL.EGL_CONTEXT_OPENGL_CORE_PROFILE_BIT, - EGL_NONE - ] - context = eglCreateContext(display, config, EGL_NO_CONTEXT, context_attribs) - if context == EGL_NO_CONTEXT: - raise RuntimeError("eglCreateContext() failed") - - logger.debug("_init_egl: calling eglCreatePbufferSurface()") - pbuffer_attribs = [EGL_WIDTH, 64, EGL_HEIGHT, 64, EGL_NONE] - surface = eglCreatePbufferSurface(display, config, pbuffer_attribs) - if surface == _EGL.EGL_NO_SURFACE: - raise RuntimeError("eglCreatePbufferSurface() failed") - - logger.debug("_init_egl: calling eglMakeCurrent()") - if not eglMakeCurrent(display, surface, surface, context): - raise RuntimeError("eglMakeCurrent() failed") - - logger.debug("_init_egl: completed successfully") - return display, context, surface, _EGL - - except Exception: - logger.debug("_init_egl: failed, cleaning up") - # Clean up any resources on failure - if surface is not None: - eglDestroySurface(display, surface) - if context is not None: - eglDestroyContext(display, context) - if display is not None: - eglTerminate(display) - raise - - -def _init_osmesa(): - """Initialize OSMesa for software rendering. Returns (context, buffer). Raises RuntimeError on failure.""" - import ctypes - - logger.debug("_init_osmesa: starting") - os.environ["PYOPENGL_PLATFORM"] = "osmesa" - - logger.debug("_init_osmesa: importing OpenGL.osmesa") - from OpenGL import GL as _gl - from OpenGL.osmesa import ( - OSMesaCreateContextExt, OSMesaMakeCurrent, OSMesaDestroyContext, - OSMESA_RGBA, - ) - logger.debug("_init_osmesa: imports completed") - - ctx = OSMesaCreateContextExt(OSMESA_RGBA, 24, 0, 0, None) - if not ctx: - raise RuntimeError("OSMesaCreateContextExt() failed") - - width, height = 64, 64 - buffer = (ctypes.c_ubyte * (width * height * 4))() - - logger.debug("_init_osmesa: calling OSMesaMakeCurrent()") - if not OSMesaMakeCurrent(ctx, buffer, _gl.GL_UNSIGNED_BYTE, width, height): - OSMesaDestroyContext(ctx) - raise RuntimeError("OSMesaMakeCurrent() failed") - - logger.debug("_init_osmesa: completed successfully") - return ctx, buffer - - class GLContext: - """Manages OpenGL context and resources for shader execution. - - Tries backends in order: GLFW (desktop) → EGL (headless GPU) → OSMesa (software). - """ + """Manages an OpenGL ES 3.0 context via EGL/ANGLE (singleton).""" _instance = None _initialized = False @@ -327,131 +240,105 @@ class GLContext: def __init__(self): if GLContext._initialized: - logger.debug("GLContext.__init__: already initialized, skipping") return - logger.debug("GLContext.__init__: starting initialization") - - global glfw, EGL - import time start = time.perf_counter() - self._backend = None - self._window = None - self._egl_display = None - self._egl_context = None - self._egl_surface = None - self._osmesa_ctx = None - self._osmesa_buffer = None + self._display = None + self._surface = None + self._context = None self._vao = None - # Try backends in order: GLFW → EGL → OSMesa - errors = [] - - logger.debug("GLContext.__init__: trying GLFW backend") try: - self._window, glfw = _init_glfw() - self._backend = "glfw" - logger.debug("GLContext.__init__: GLFW backend succeeded") - except Exception as e: - logger.debug(f"GLContext.__init__: GLFW backend failed: {e}") - errors.append(("GLFW", e)) + self._display, self._egl_major, self._egl_minor = _get_egl_display() - if self._backend is None: - logger.debug("GLContext.__init__: trying EGL backend") - try: - self._egl_display, self._egl_context, self._egl_surface, EGL = _init_egl() - self._backend = "egl" - logger.debug("GLContext.__init__: EGL backend succeeded") - except Exception as e: - logger.debug(f"GLContext.__init__: EGL backend failed: {e}") - errors.append(("EGL", e)) + if not EGL.eglBindAPI(EGL.EGL_OPENGL_ES_API): + raise RuntimeError("eglBindAPI(EGL_OPENGL_ES_API) failed") - if self._backend is None: - logger.debug("GLContext.__init__: trying OSMesa backend") - try: - self._osmesa_ctx, self._osmesa_buffer = _init_osmesa() - self._backend = "osmesa" - logger.debug("GLContext.__init__: OSMesa backend succeeded") - except Exception as e: - logger.debug(f"GLContext.__init__: OSMesa backend failed: {e}") - errors.append(("OSMesa", e)) + config = EGL.EGLConfig() + n_configs = ctypes.c_int32(0) + if not EGL.eglChooseConfig( + self._display, + _egl_attribs( + EGL.EGL_RENDERABLE_TYPE, EGL.EGL_OPENGL_ES3_BIT, + EGL.EGL_SURFACE_TYPE, EGL.EGL_PBUFFER_BIT, + EGL.EGL_RED_SIZE, 8, EGL.EGL_GREEN_SIZE, 8, + EGL.EGL_BLUE_SIZE, 8, EGL.EGL_ALPHA_SIZE, 8, + ), + ctypes.byref(config), 1, ctypes.byref(n_configs), + ) or n_configs.value == 0: + raise RuntimeError("eglChooseConfig() failed") - if self._backend is None: - if sys.platform == "win32": - platform_help = ( - "Windows: Ensure GPU drivers are installed and display is available.\n" - " CPU-only/headless mode is not supported on Windows." - ) - elif sys.platform == "darwin": - platform_help = ( - "macOS: GLFW is not supported.\n" - " Install OSMesa via Homebrew: brew install mesa\n" - " Then: pip install PyOpenGL PyOpenGL-accelerate" - ) - else: - platform_help = ( - "Linux: Install one of these backends:\n" - " Desktop: sudo apt install libgl1-mesa-glx libglfw3\n" - " Headless with GPU: sudo apt install libegl1-mesa libgl1-mesa-dri\n" - " Headless (CPU): sudo apt install libosmesa6" - ) - - error_details = "\n".join(f" {name}: {err}" for name, err in errors) - raise RuntimeError( - f"Failed to create OpenGL context.\n\n" - f"Backend errors:\n{error_details}\n\n" - f"{platform_help}" + self._surface = EGL.eglCreatePbufferSurface( + self._display, config, + _egl_attribs(EGL.EGL_WIDTH, 64, EGL.EGL_HEIGHT, 64), ) + if not self._surface: + raise RuntimeError("eglCreatePbufferSurface() failed") - # Now import OpenGL.GL (after context is current) - logger.debug("GLContext.__init__: importing OpenGL.GL") - _import_opengl() + self._context = EGL.eglCreateContext( + self._display, config, EGL.EGL_NO_CONTEXT, + _egl_attribs(EGL.EGL_CONTEXT_CLIENT_VERSION, 3), + ) + if not self._context: + raise RuntimeError("eglCreateContext() failed") - # Create VAO (required for core profile, but OSMesa may use compat profile) - logger.debug("GLContext.__init__: creating VAO") - try: - vao = gl.glGenVertexArrays(1) - gl.glBindVertexArray(vao) - self._vao = vao # Only store after successful bind - logger.debug("GLContext.__init__: VAO created successfully") - except Exception as e: - logger.debug(f"GLContext.__init__: VAO creation failed (may be expected for OSMesa): {e}") - # OSMesa with older Mesa may not support VAOs - # Clean up if we created but couldn't bind - if vao: - try: - gl.glDeleteVertexArrays(1, [vao]) - except Exception: - pass + if not EGL.eglMakeCurrent(self._display, self._surface, self._surface, self._context): + raise RuntimeError("eglMakeCurrent() failed") + + self._vao = gl.glGenVertexArrays(1) + gl.glBindVertexArray(self._vao) + + except Exception: + self._cleanup() + raise elapsed = (time.perf_counter() - start) * 1000 - # Log device info - renderer = gl.glGetString(gl.GL_RENDERER) - vendor = gl.glGetString(gl.GL_VENDOR) - version = gl.glGetString(gl.GL_VERSION) - renderer = renderer.decode() if renderer else "Unknown" - vendor = vendor.decode() if vendor else "Unknown" - version = version.decode() if version else "Unknown" + renderer = _gl_str(gl.GL_RENDERER) + vendor = _gl_str(gl.GL_VENDOR) + version = _gl_str(gl.GL_VERSION) GLContext._initialized = True - logger.info(f"GLSL context initialized in {elapsed:.1f}ms ({self._backend}) - {renderer} ({vendor}), GL {version}") + logger.info(f"GLSL context initialized in {elapsed:.1f}ms - EGL {self._egl_major}.{self._egl_minor}, {renderer} ({vendor}), GL {version}") def make_current(self): - if self._backend == "glfw": - glfw.make_context_current(self._window) - elif self._backend == "egl": - from OpenGL.EGL import eglMakeCurrent - eglMakeCurrent(self._egl_display, self._egl_surface, self._egl_surface, self._egl_context) - elif self._backend == "osmesa": - from OpenGL.osmesa import OSMesaMakeCurrent - OSMesaMakeCurrent(self._osmesa_ctx, self._osmesa_buffer, gl.GL_UNSIGNED_BYTE, 64, 64) - + if not EGL.eglMakeCurrent(self._display, self._surface, self._surface, self._context): + err = EGL.eglGetError() + raise RuntimeError(f"eglMakeCurrent() failed (EGL error: 0x{err:04X})") if self._vao is not None: gl.glBindVertexArray(self._vao) + def _cleanup(self): + if not self._display: + return + try: + if self._vao is not None: + gl.glDeleteVertexArrays(1, [self._vao]) + self._vao = None + except Exception: + pass + try: + EGL.eglMakeCurrent(self._display, EGL.EGL_NO_SURFACE, EGL.EGL_NO_SURFACE, EGL.EGL_NO_CONTEXT) + except Exception: + pass + try: + if self._context: + EGL.eglDestroyContext(self._display, self._context) + except Exception: + pass + try: + if self._surface: + EGL.eglDestroySurface(self._display, self._surface) + except Exception: + pass + try: + EGL.eglTerminate(self._display) + except Exception: + pass + self._display = None + def _compile_shader(source: str, shader_type: int) -> int: """Compile a shader and return its ID.""" @@ -459,8 +346,10 @@ def _compile_shader(source: str, shader_type: int) -> int: gl.glShaderSource(shader, source) gl.glCompileShader(shader) - if gl.glGetShaderiv(shader, gl.GL_COMPILE_STATUS) != gl.GL_TRUE: - error = gl.glGetShaderInfoLog(shader).decode() + if not gl.glGetShaderiv(shader, gl.GL_COMPILE_STATUS): + error = gl.glGetShaderInfoLog(shader) + if isinstance(error, bytes): + error = error.decode(errors="replace") gl.glDeleteShader(shader) raise RuntimeError(f"Shader compilation failed:\n{error}") @@ -484,8 +373,10 @@ def _create_program(vertex_source: str, fragment_source: str) -> int: gl.glDeleteShader(vertex_shader) gl.glDeleteShader(fragment_shader) - if gl.glGetProgramiv(program, gl.GL_LINK_STATUS) != gl.GL_TRUE: - error = gl.glGetProgramInfoLog(program).decode() + if not gl.glGetProgramiv(program, gl.GL_LINK_STATUS): + error = gl.glGetProgramInfoLog(program) + if isinstance(error, bytes): + error = error.decode(errors="replace") gl.glDeleteProgram(program) raise RuntimeError(f"Program linking failed:\n{error}") @@ -530,9 +421,6 @@ def _render_shader_batch( ctx = GLContext() ctx.make_current() - # Convert from GLSL ES to desktop GLSL 330 - fragment_source = _convert_es_to_desktop(fragment_code) - # Detect how many outputs the shader actually uses num_outputs = _detect_output_count(fragment_code) @@ -558,9 +446,9 @@ def _render_shader_batch( try: # Compile shaders (once for all batches) try: - program = _create_program(VERTEX_SHADER, fragment_source) + program = _create_program(VERTEX_SHADER, fragment_code) except RuntimeError: - logger.error(f"Fragment shader:\n{fragment_source}") + logger.error(f"Fragment shader:\n{fragment_code}") raise gl.glUseProgram(program) @@ -723,13 +611,13 @@ def _render_shader_batch( gl.glDrawArrays(gl.GL_TRIANGLES, 0, 3) # Read back outputs for this batch - # (glGetTexImage is synchronous, implicitly waits for rendering) + gl.glBindFramebuffer(gl.GL_FRAMEBUFFER, fbo) batch_outputs = [] - for tex in output_textures: - gl.glBindTexture(gl.GL_TEXTURE_2D, tex) - data = gl.glGetTexImage(gl.GL_TEXTURE_2D, 0, gl.GL_RGBA, gl.GL_FLOAT) - img = np.frombuffer(data, dtype=np.float32).reshape(height, width, 4) - batch_outputs.append(img[::-1, :, :].copy()) + for i in range(num_outputs): + gl.glReadBuffer(gl.GL_COLOR_ATTACHMENT0 + i) + buf = np.empty((height, width, 4), dtype=np.float32) + gl.glReadPixels(0, 0, width, height, gl.GL_RGBA, gl.GL_FLOAT, buf) + batch_outputs.append(buf[::-1, :, :].copy()) # Pad with black images for unused outputs black_img = np.zeros((height, width, 4), dtype=np.float32) @@ -750,18 +638,18 @@ def _render_shader_batch( gl.glBindFramebuffer(gl.GL_FRAMEBUFFER, 0) gl.glUseProgram(0) - for tex in input_textures: - gl.glDeleteTextures(int(tex)) - for tex in curve_textures: - gl.glDeleteTextures(int(tex)) - for tex in output_textures: - gl.glDeleteTextures(int(tex)) - for tex in ping_pong_textures: - gl.glDeleteTextures(int(tex)) + if input_textures: + gl.glDeleteTextures(len(input_textures), input_textures) + if curve_textures: + gl.glDeleteTextures(len(curve_textures), curve_textures) + if output_textures: + gl.glDeleteTextures(len(output_textures), output_textures) + if ping_pong_textures: + gl.glDeleteTextures(len(ping_pong_textures), ping_pong_textures) if fbo is not None: gl.glDeleteFramebuffers(1, [fbo]) - for pp_fbo in ping_pong_fbos: - gl.glDeleteFramebuffers(1, [pp_fbo]) + if ping_pong_fbos: + gl.glDeleteFramebuffers(len(ping_pong_fbos), ping_pong_fbos) if program is not None: gl.glDeleteProgram(program) diff --git a/requirements.txt b/requirements.txt index d7719178b..8509599a6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -33,5 +33,5 @@ kornia>=0.7.1 spandrel pydantic~=2.0 pydantic-settings~=2.0 -PyOpenGL -glfw +PyOpenGL>=3.1.8 +comfy-angle From a95e461916de9cbda2e89140ab86a8a7c3f9702a Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 27 Jun 2026 15:53:11 -0700 Subject: [PATCH 026/211] int8 support on turing GPUs. (#14662) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 8509599a6..01e7d2f94 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.13 +comfy-kitchen==0.2.14 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From f19735759e8973bd2ead76f08d9cc45abc1f98e4 Mon Sep 17 00:00:00 2001 From: Matt Miller Date: Sat, 27 Jun 2026 23:34:30 -0700 Subject: [PATCH 027/211] ci: add team-gated Cursor review (thin caller for github-workflows) (#14527) --- .github/workflows/ci-cursor-review.yml | 38 ++++++++++++++++++++++++++ 1 file changed, 38 insertions(+) create mode 100644 .github/workflows/ci-cursor-review.yml diff --git a/.github/workflows/ci-cursor-review.yml b/.github/workflows/ci-cursor-review.yml new file mode 100644 index 000000000..2312c0ccd --- /dev/null +++ b/.github/workflows/ci-cursor-review.yml @@ -0,0 +1,38 @@ +name: CI - Cursor Review + +# Thin caller for the shared reusable cursor-review workflow in +# Comfy-Org/github-workflows. The review logic (panel matrix, judge +# consolidation, prompts, extract/post/notify scripts) lives there as the +# single source of truth, so this repo only carries the repo-specific diff +# excludes. + +on: + pull_request: + types: [labeled, unlabeled] + +concurrency: + group: cursor-review-pr-${{ github.event.pull_request.number }}-${{ github.event.label.name }} + cancel-in-progress: true + +jobs: + cursor-review: + if: github.event.label.name == 'cursor-review' + permissions: + contents: read + pull-requests: write + # SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up + # upstream changes; keep `workflows_ref` matching so prompts/scripts load + # from the same commit as the workflow definition. + uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9 + with: + workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a + diff_excludes: >- + :!**/.claude/** + :!**/dist/** + :!**/vendor/** + :!**/*.generated.* + :!**/*.min.js + :!**/*.min.css + secrets: + CURSOR_API_KEY: ${{ secrets.CURSOR_API_KEY }} + SLACK_BOT_TOKEN: ${{ secrets.SLACK_BOT_TOKEN }} From 79c555ce6bfebf862d014e710b8d4b541ba5b896 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sun, 28 Jun 2026 20:52:36 -0700 Subject: [PATCH 028/211] Fix int8 mm being skipped on offloaded lora weights. (#14669) --- comfy/ops.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/ops.py b/comfy/ops.py index 6a5090548..69d32e254 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1216,7 +1216,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec bias_dtype=input.dtype, offloadable=True, compute_dtype=compute_dtype, - want_requant=want_requant, + want_requant=True, ) weight = weight.to(dtype=input.dtype) else: From a58473fd9bf3a1e2383a41c6267ca03a168150ba Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Mon, 29 Jun 2026 17:08:06 +0800 Subject: [PATCH 029/211] chore: update embedded docs to v0.5.6 (#14668) Co-authored-by: Alexis Rolland --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 01e7d2f94..b09b12f29 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ comfyui-frontend-package==1.45.19 comfyui-workflow-templates==0.10.7 -comfyui-embedded-docs==0.5.5 +comfyui-embedded-docs==0.5.6 torch torchsde torchvision From 785141051163612f0e471a242c1f33341f60b9bd Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 29 Jun 2026 18:52:08 -0700 Subject: [PATCH 030/211] Better and faster int8 lora applying. (#14685) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index b09b12f29..6af0b21bc 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.14 +comfy-kitchen==0.2.15 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 510ed5c3848ec6995f7cc256983f8d1ad0145530 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Tue, 30 Jun 2026 17:25:03 +0900 Subject: [PATCH 031/211] Bump comfyui-frontend-package to 1.45.20 (#14684) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 6af0b21bc..eb7230b49 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -comfyui-frontend-package==1.45.19 +comfyui-frontend-package==1.45.20 comfyui-workflow-templates==0.10.7 comfyui-embedded-docs==0.5.6 torch From ba3f697dbbf2fd15a23c3e9fd8fb8f89552a1475 Mon Sep 17 00:00:00 2001 From: Silver <65376327+silveroxides@users.noreply.github.com> Date: Tue, 30 Jun 2026 10:27:09 +0200 Subject: [PATCH 032/211] =?UTF-8?q?Add=20ConditioningMultiply=20node=20to?= =?UTF-8?q?=20nodes.py=20as=20an=20addition=20to=20other=20adj=E2=80=A6=20?= =?UTF-8?q?(#14686)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- nodes.py | 25 +++++++++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/nodes.py b/nodes.py index 028e58c77..77c577b9a 100644 --- a/nodes.py +++ b/nodes.py @@ -159,6 +159,29 @@ class ConditioningConcat: return (out, ) +class ConditioningMultiply: + SEARCH_ALIASES = ["scale conditioning", "scale prompt", "multiply conditioning", "multiply prompt"] + + @classmethod + def INPUT_TYPES(cls): + return {"required": {"conditioning": ("CONDITIONING", ), + "multiplier": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}) + }} + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "multiply" + CATEGORY = "model/conditioning/transform" + + def multiply(self, conditioning, multiplier): + c = [] + for t in conditioning: + values = {} + pooled_output = t[1].get("pooled_output", None) + if pooled_output is not None: + values["pooled_output"] = pooled_output * multiplier + scaled = node_helpers.conditioning_set_values([[t[0] * multiplier, t[1]]], values)[0] + c.append(scaled) + return (c,) + class ConditioningSetArea: SEARCH_ALIASES = ["regional prompt", "area prompt", "spatial conditioning", "localized prompt"] @@ -2050,6 +2073,7 @@ NODE_CLASS_MAPPINGS = { "ConditioningAverage": ConditioningAverage, "ConditioningCombine": ConditioningCombine, "ConditioningConcat": ConditioningConcat, + "ConditioningMultiply": ConditioningMultiply, "ConditioningSetArea": ConditioningSetArea, "ConditioningSetAreaPercentage": ConditioningSetAreaPercentage, "ConditioningSetAreaStrength": ConditioningSetAreaStrength, @@ -2121,6 +2145,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ConditioningAverage ": "Conditioning (Average)", "ConditioningAverage": "Conditioning (Average)", "ConditioningConcat": "Conditioning (Concat)", + "ConditioningMultiply": "Conditioning (Multiply)", "ConditioningSetArea": "Conditioning (Set Area)", "ConditioningSetAreaPercentage": "Conditioning (Set Area with Percentage)", "ConditioningSetAreaStrength": "Conditioning (Set Area Strength)", From 8fe0243d974b72b733deb21d04000626b26247d1 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 30 Jun 2026 21:17:23 +0300 Subject: [PATCH 033/211] [Partner Nodes] feat(Google): add Nano Banana 2 Lite model (#14693) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_gemini.py | 34 +++++++++++++++++++++++---------- 1 file changed, 24 insertions(+), 10 deletions(-) diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index a63625ada..1a8aadfd6 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -249,18 +249,22 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N input_tokens_price = 2 output_text_tokens_price = 12.0 output_image_tokens_price = 0.0 - elif response.modelVersion == "gemini-3.1-flash-lite-preview": + elif response.modelVersion in ("gemini-3.1-flash-lite-preview", "gemini-3.1-flash-lite"): input_tokens_price = 0.25 output_text_tokens_price = 1.50 output_image_tokens_price = 0.0 - elif response.modelVersion == "gemini-3-pro-image-preview": + elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"): input_tokens_price = 2 output_text_tokens_price = 12.0 output_image_tokens_price = 120.0 - elif response.modelVersion == "gemini-3.1-flash-image-preview": + elif response.modelVersion in ("gemini-3.1-flash-image-preview", "gemini-3.1-flash-image"): input_tokens_price = 0.5 output_text_tokens_price = 3.0 output_image_tokens_price = 60.0 + elif response.modelVersion == "gemini-3.1-flash-lite-image": + input_tokens_price = 0.25 + output_text_tokens_price = 1.50 + output_image_tokens_price = 30.0 else: return None final_price = response.usageMetadata.promptTokenCount * input_tokens_price @@ -1302,7 +1306,7 @@ class GeminiNanoBanana2(IO.ComfyNode): ) -def _nano_banana_2_v2_model_inputs(): +def _nano_banana_2_v2_model_inputs(resolutions: list[str]): return [ IO.Combo.Input( "aspect_ratio", @@ -1329,8 +1333,8 @@ def _nano_banana_2_v2_model_inputs(): ), IO.Combo.Input( "resolution", - options=["1K", "2K", "4K"], - tooltip="Target output resolution. For 2K/4K the native Gemini upscaler is used.", + options=resolutions, + tooltip="Target output resolution.", ), IO.Combo.Input( "thinking_level", @@ -1376,7 +1380,11 @@ class GeminiNanoBanana2V2(IO.ComfyNode): options=[ IO.DynamicCombo.Option( "Nano Banana 2 (Gemini 3.1 Flash Image)", - _nano_banana_2_v2_model_inputs(), + _nano_banana_2_v2_model_inputs(resolutions=["1K", "2K", "4K"]), + ), + IO.DynamicCombo.Option( + "Nano Banana 2 Lite", + _nano_banana_2_v2_model_inputs(resolutions=["1K"]), ), ], ), @@ -1445,9 +1453,13 @@ class GeminiNanoBanana2V2(IO.ComfyNode): depends_on=IO.PriceBadgeDepends(widgets=["model", "model.resolution"]), expr=""" ( - $r := $lookup(widgets, "model.resolution"); - $prices := {"1k": 0.0696, "2k": 0.1014, "4k": 0.154}; - {"type":"usd","usd": $lookup($prices, $r), "format":{"suffix":"/Image","approximate":true}} + $contains(widgets.model, "lite") + ? {"type":"usd","usd": 0.034, "format":{"suffix":"/Image","approximate":true}} + : ( + $r := $lookup(widgets, "model.resolution"); + $prices := {"1k": 0.0696, "2k": 0.1014, "4k": 0.154}; + {"type":"usd","usd": $lookup($prices, $r), "format":{"suffix":"/Image","approximate":true}} + ) ) """, ), @@ -1468,6 +1480,8 @@ class GeminiNanoBanana2V2(IO.ComfyNode): model_choice = model["model"] if model_choice == "Nano Banana 2 (Gemini 3.1 Flash Image)": model_id = "gemini-3.1-flash-image-preview" + elif model_choice == "Nano Banana 2 Lite": + model_id = "gemini-3.1-flash-lite-image" else: model_id = model_choice From d395813bcd9240fe70d26e6d11d87ec04c459eb9 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 30 Jun 2026 14:08:59 -0700 Subject: [PATCH 034/211] Fix memory leak related to int8. (#14697) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index eb7230b49..bb11b9605 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.15 +comfy-kitchen==0.2.16 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 1c59659a2f448a6eaf8bb9dbc0032ea0d84e7c35 Mon Sep 17 00:00:00 2001 From: Matt Miller Date: Tue, 30 Jun 2026 14:13:20 -0700 Subject: [PATCH 035/211] feat: make asset hashing opt-in via --enable-asset-hashing, off by default (#14663) Add a --enable-asset-hashing CLI flag (action=store_true, default False) and plumb it into the two asset-seeder call sites in main.py that previously hardcoded compute_hashes=True (the startup scan and the post-job output enqueue). Local runs now skip blake3 hashing unless the user opts in, avoiding the startup/per-output cost on large models directories while keeping hashing available for asset-portability features. Co-authored-by: Alexis Rolland --- comfy/cli_args.py | 1 + main.py | 4 ++-- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/comfy/cli_args.py b/comfy/cli_args.py index e3099a230..4bef096fb 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -240,6 +240,7 @@ database_default_path = os.path.abspath( ) parser.add_argument("--database-url", type=str, default=f"sqlite:///{database_default_path}", help="Specify the database URL, e.g. for an in-memory database you can use 'sqlite:///:memory:'.") parser.add_argument("--enable-assets", action="store_true", help="Enable the assets system (API routes, database synchronization, and background scanning).") +parser.add_argument("--enable-asset-hashing", action="store_true", help="Compute blake3 content hashes when scanning assets. Hashing enables future asset-portability features (deduplication, cross-machine model resolution) but adds startup cost and per-output cost on large models directories. Off by default; enable to opt in.") parser.add_argument("--feature-flag", type=str, action='append', default=[], metavar="KEY[=VALUE]", help="Set a server feature flag. Use KEY=VALUE to set an explicit value, or bare KEY to set it to true. Can be specified multiple times. Boolean values (true/false) and numbers are auto-converted. Examples: --feature-flag show_signin_button=true or --feature-flag show_signin_button") parser.add_argument("--list-feature-flags", action="store_true", help="Print the registry of known CLI-settable feature flags as JSON and exit.") diff --git a/main.py b/main.py index aa4ee2adb..20ec83c9e 100644 --- a/main.py +++ b/main.py @@ -403,7 +403,7 @@ def prompt_worker(q, server_instance): hook_breaker_ac10a0.restore_functions() if not asset_seeder.is_disabled(): - asset_seeder.enqueue_enrich(roots=("output",), compute_hashes=True) + asset_seeder.enqueue_enrich(roots=("output",), compute_hashes=args.enable_asset_hashing) asset_seeder.resume() @@ -458,7 +458,7 @@ def setup_database(): if dependencies_available(): init_db() if args.enable_assets: - if asset_seeder.start(roots=("models", "input", "output"), prune_first=True, compute_hashes=True): + if asset_seeder.start(roots=("models", "input", "output"), prune_first=True, compute_hashes=args.enable_asset_hashing): logging.info("Background asset scan initiated for models, input, output") except Exception as e: if "database is locked" in str(e): From b70944e710c1cb7455099d06fd149d0fe47b330e Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 1 Jul 2026 00:17:53 +0300 Subject: [PATCH 036/211] [Partner Nodes] feat(Google): add Gemini Video Omni node (#14695) --- comfy_api_nodes/apis/gemini.py | 1 + comfy_api_nodes/nodes_gemini.py | 174 +++++++++++++++++++++++++++++++- 2 files changed, 174 insertions(+), 1 deletion(-) diff --git a/comfy_api_nodes/apis/gemini.py b/comfy_api_nodes/apis/gemini.py index caaba8f36..7b2543270 100644 --- a/comfy_api_nodes/apis/gemini.py +++ b/comfy_api_nodes/apis/gemini.py @@ -121,6 +121,7 @@ class GeminiGenerationConfig(BaseModel): topK: int | None = Field(None, ge=1) topP: float | None = Field(None, ge=0.0, le=1.0) thinkingConfig: GeminiThinkingConfig | None = Field(None) + responseModalities: list[str] | None = Field(None) class GeminiImageOutputOptions(BaseModel): diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 1a8aadfd6..aa992802d 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -13,7 +13,7 @@ import torch from typing_extensions import override import folder_paths -from comfy_api.latest import IO, ComfyExtension, Input, Types +from comfy_api.latest import IO, ComfyExtension, Input, InputImpl, Types from comfy_api_nodes.apis.gemini import ( GeminiContent, GeminiFileData, @@ -37,6 +37,7 @@ from comfy_api_nodes.util import ( audio_to_base64_string, bytesio_to_image_tensor, download_url_to_image_tensor, + download_url_to_video_output, get_number_of_images, sync_op, tensor_to_base64_string, @@ -45,6 +46,7 @@ from comfy_api_nodes.util import ( upload_images_to_comfyapi, upload_video_to_comfyapi, validate_string, + validate_video_duration, video_to_base64_string, ) @@ -229,10 +231,29 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug return torch.cat(image_tensors, dim=0) +async def get_video_from_response( + response: GeminiGenerateContentResponse, cls: type[IO.ComfyNode] | None = None +) -> InputImpl.VideoFromFile: + parts = get_parts_by_type(response, "video/*") + for part in parts: + if part.inlineData and part.inlineData.data: + return InputImpl.VideoFromFile(BytesIO(base64.b64decode(part.inlineData.data))) + if part.fileData and part.fileData.fileUri: + return await download_url_to_video_output(part.fileData.fileUri, cls=cls) + model_message = get_text_from_response(response).strip() + if model_message: + raise ValueError(f"Gemini did not generate a video. Model response: {model_message}") + raise ValueError( + "Gemini did not generate a video. Try rephrasing your prompt, " + "shortening the requested duration, or reducing the number of input images/videos." + ) + + def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None: if not response.modelVersion: return None # Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing + output_video_tokens_price = 0.0 if response.modelVersion == "gemini-2.5-pro": input_tokens_price = 1.25 output_text_tokens_price = 10.0 @@ -265,6 +286,11 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N input_tokens_price = 0.25 output_text_tokens_price = 1.50 output_image_tokens_price = 30.0 + elif response.modelVersion == "gemini-omni-flash-preview": + input_tokens_price = 2.145 + output_text_tokens_price = 12.87 + output_image_tokens_price = 0.0 + output_video_tokens_price = 25.025 else: return None final_price = response.usageMetadata.promptTokenCount * input_tokens_price @@ -272,6 +298,8 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N for i in response.usageMetadata.candidatesTokensDetails: if i.modality == Modality.IMAGE: final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models + elif i.modality == Modality.VIDEO: + final_price += output_video_tokens_price * i.tokenCount # for Omni Flash else: final_price += output_text_tokens_price * i.tokenCount if response.usageMetadata.thoughtsTokenCount: @@ -1531,6 +1559,149 @@ class GeminiNanoBanana2V2(IO.ComfyNode): ) +OMNI_MAX_IMAGES = 14 +OMNI_MAX_VIDEOS = 3 + +OMNI_MODELS: dict[str, str] = { + "Omni Flash": "gemini-omni-flash-preview", +} + + +def _omni_flash_inputs() -> list[Input]: + """Per-model inputs for the Omni video DynamicCombo (prompt + reference media + sampling).""" + return [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Describe the video to generate. Specify the length and aspect ratio directly in the " + 'prompt, e.g. "a 6-second clip in 16:9". Length may be 3-10 seconds; the aspect ratio must be ' + "16:9 (landscape) or 9:16 (portrait). The output is 720p, 24 FPS, with audio.", + ), + IO.Autogrow.Input( + "images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("image"), + names=[f"image_{i}" for i in range(1, OMNI_MAX_IMAGES + 1)], + min=0, + ), + tooltip=f"Optional reference image(s) to guide or animate the video. Up to {OMNI_MAX_IMAGES} images.", + ), + IO.Autogrow.Input( + "videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("video"), + names=[f"video_{i}" for i in range(1, OMNI_MAX_VIDEOS + 1)], + min=0, + ), + tooltip=f"Optional reference video(s) to guide or edit. Up to {OMNI_MAX_VIDEOS} videos, " + f"each up to 10 seconds long.", + ), + IO.Float.Input( + "temperature", + default=1.0, + min=0.0, + max=2.0, + step=0.01, + tooltip="Controls randomness. Lower is more focused/deterministic, higher is more varied.", + advanced=True, + ), + IO.Float.Input( + "top_p", + default=0.95, + min=0.0, + max=1.0, + step=0.01, + tooltip="Nucleus sampling: sample from the smallest token set whose cumulative probability reaches top_p.", + advanced=True, + ), + ] + + +class GeminiVideoOmni(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="GeminiVideoOmni", + display_name="Google Gemini Omni (Video)", + category="partner/video/Gemini", + essentials_category="Video Generation", + description="Generate a video with audio from a text prompt using Google's Gemini Omni Flash model. " + "Optionally provide reference images and/or videos to guide or edit the result. Describe the desired " + "length (3-10s) and aspect ratio (16:9 or 9:16) directly in the prompt.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option("Omni Flash", _omni_flash_inputs()), + ], + tooltip="The Gemini video model used to generate the video.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + ], + outputs=[ + IO.Video.Output(), + IO.String.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr='{"type":"usd","usd":0.146,"format":{"suffix":"/second","approximate":true}}' + ), + ) + + @classmethod + async def execute(cls, model: dict, seed: int) -> IO.NodeOutput: + prompt = model.get("prompt") or "" + validate_string(prompt, strip_whitespace=True, min_length=1) + model_id = OMNI_MODELS[model["model"]] + + images = [t for t in (model.get("images") or {}).values() if t is not None] + videos = [v for v in (model.get("videos") or {}).values() if v is not None] + if sum(get_number_of_images(t) for t in images) > OMNI_MAX_IMAGES: + raise ValueError(f"The current maximum number of supported images is {OMNI_MAX_IMAGES}.") + if len(videos) > OMNI_MAX_VIDEOS: + raise ValueError(f"The current maximum number of supported videos is {OMNI_MAX_VIDEOS}.") + for video in videos: + validate_video_duration(video, max_duration=10) + + parts: list[GeminiPart] = [] + if images or videos: + parts.extend(await build_gemini_media_parts(cls, images, [], videos)) + parts.append(GeminiPart(text=prompt)) + response = await sync_op( + cls, + ApiEndpoint(path=f"{GEMINI_BASE_ENDPOINT}/{model_id}", method="POST"), + data=GeminiGenerateContentRequest( + contents=[GeminiContent(role=GeminiRole.user, parts=parts)], + generationConfig=GeminiGenerationConfig( + responseModalities=["TEXT", "VIDEO"], + temperature=model.get("temperature", 1.0), + topP=model.get("top_p", 0.95), + ), + ), + response_model=GeminiGenerateContentResponse, + price_extractor=calculate_tokens_price, + ) + return IO.NodeOutput( + await get_video_from_response(response, cls=cls), + get_text_from_response(response), + ) + + class GeminiExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: @@ -1541,6 +1712,7 @@ class GeminiExtension(ComfyExtension): GeminiImage2, GeminiNanoBanana2, GeminiNanoBanana2V2, + GeminiVideoOmni, GeminiInputFiles, ] From 6e11828d109ef3bb8378dbb7ebb7814b41336c01 Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Wed, 1 Jul 2026 05:20:20 +0800 Subject: [PATCH 037/211] chore: Update nodes categories (#14674) --- comfy_extras/nodes_cond.py | 10 ++++++---- comfy_extras/nodes_custom_sampler.py | 4 ++-- comfy_extras/nodes_photomaker.py | 6 ++++-- comfy_extras/nodes_stable_cascade.py | 2 +- comfy_extras/nodes_triposplat.py | 4 ++-- nodes.py | 11 ++++++----- 6 files changed, 21 insertions(+), 16 deletions(-) diff --git a/comfy_extras/nodes_cond.py b/comfy_extras/nodes_cond.py index b745a43af..c8091b7a4 100644 --- a/comfy_extras/nodes_cond.py +++ b/comfy_extras/nodes_cond.py @@ -8,7 +8,8 @@ class CLIPTextEncodeControlnet(io.ComfyNode): def define_schema(cls) -> io.Schema: return io.Schema( node_id="CLIPTextEncodeControlnet", - category="experimental/conditioning", + display_name="CLIP Text Encode (Controlnet)", + category="model/conditioning", inputs=[ io.Clip.Input("clip"), io.Conditioning.Input("conditioning"), @@ -35,11 +36,12 @@ class T5TokenizerOptions(io.ComfyNode): def define_schema(cls) -> io.Schema: return io.Schema( node_id="T5TokenizerOptions", - category="experimental/conditioning", + display_name="T5 Tokenizer Options", + category="model/conditioning", inputs=[ io.Clip.Input("clip"), - io.Int.Input("min_padding", default=0, min=0, max=10000, step=1, advanced=True), - io.Int.Input("min_length", default=0, min=0, max=10000, step=1, advanced=True), + io.Int.Input("min_padding", default=0, min=0, max=10000, step=1), + io.Int.Input("min_length", default=0, min=0, max=10000, step=1), ], outputs=[io.Clip.Output()], is_experimental=True, diff --git a/comfy_extras/nodes_custom_sampler.py b/comfy_extras/nodes_custom_sampler.py index c9d7e06fc..56ef5f526 100644 --- a/comfy_extras/nodes_custom_sampler.py +++ b/comfy_extras/nodes_custom_sampler.py @@ -1070,7 +1070,7 @@ class AddNoise(io.ComfyNode): def define_schema(cls): return io.Schema( node_id="AddNoise", - category="experimental/custom_sampling/noise", + category="model/sampling/noise", is_experimental=True, inputs=[ io.Model.Input("model"), @@ -1120,7 +1120,7 @@ class ManualSigmas(io.ComfyNode): return io.Schema( node_id="ManualSigmas", search_aliases=["custom noise schedule", "define sigmas"], - category="experimental/custom_sampling", + category="model/sampling/sigmas", is_experimental=True, inputs=[ io.String.Input("sigmas", default="1, 0.5", multiline=False) diff --git a/comfy_extras/nodes_photomaker.py b/comfy_extras/nodes_photomaker.py index 8a2248572..72fad1673 100644 --- a/comfy_extras/nodes_photomaker.py +++ b/comfy_extras/nodes_photomaker.py @@ -123,7 +123,8 @@ class PhotoMakerLoader(io.ComfyNode): def define_schema(cls): return io.Schema( node_id="PhotoMakerLoader", - category="experimental/photomaker", + display_name="Load PhotoMaker Model", + category="model/loaders", inputs=[ io.Combo.Input("photomaker_model_name", options=folder_paths.get_filename_list("photomaker")), ], @@ -149,7 +150,8 @@ class PhotoMakerEncode(io.ComfyNode): def define_schema(cls): return io.Schema( node_id="PhotoMakerEncode", - category="experimental/photomaker", + display_name="PhotoMaker Encode", + category="model/conditioning/photomaker", inputs=[ io.Photomaker.Input("photomaker"), io.Image.Input("image"), diff --git a/comfy_extras/nodes_stable_cascade.py b/comfy_extras/nodes_stable_cascade.py index 6a78ffb47..ddfb4f2b0 100644 --- a/comfy_extras/nodes_stable_cascade.py +++ b/comfy_extras/nodes_stable_cascade.py @@ -119,7 +119,7 @@ class StableCascade_SuperResolutionControlnet(io.ComfyNode): def define_schema(cls): return io.Schema( node_id="StableCascade_SuperResolutionControlnet", - category="experimental/stable_cascade", + category="experimental/stable cascade", is_experimental=True, inputs=[ io.Image.Input("image"), diff --git a/comfy_extras/nodes_triposplat.py b/comfy_extras/nodes_triposplat.py index 7bf4703fe..c892213e4 100644 --- a/comfy_extras/nodes_triposplat.py +++ b/comfy_extras/nodes_triposplat.py @@ -143,7 +143,7 @@ class VAEDecodeTripoSplat(IO.ComfyNode): return IO.Schema( node_id="VAEDecodeTripoSplat", display_name="TripoSplat Decode", - category="3d/latent", + category="model/latent/triposplat", description="Decode the sampled TripoSplat latent into a 3D gaussian splat. " "Modify the number of gaussians to vary the density.", inputs=[ @@ -188,7 +188,7 @@ class TripoSplatSamplingPreview(IO.ComfyNode): return IO.Schema( node_id="TripoSplatSamplingPreview", display_name="TripoSplat Sampling Preview", - category="3d/latent", + category="model/latent/triposplat", description="Patch the TripoSplat model for the standard Ksampler node to show a live decoded " "gaussian splat preview at each step.", inputs=[ diff --git a/nodes.py b/nodes.py index 77c577b9a..9043a8d0a 100644 --- a/nodes.py +++ b/nodes.py @@ -349,7 +349,7 @@ class VAEDecodeTiled: RETURN_TYPES = ("IMAGE",) FUNCTION = "decode" - CATEGORY = "experimental" + CATEGORY = "model/latent" def decode(self, vae, samples, tile_size, overlap=64, temporal_size=64, temporal_overlap=8): if tile_size < overlap * 4: @@ -396,7 +396,7 @@ class VAEEncodeTiled: RETURN_TYPES = ("LATENT",) FUNCTION = "encode" - CATEGORY = "experimental" + CATEGORY = "model/latent" def encode(self, vae, pixels, tile_size, overlap, temporal_size=64, temporal_overlap=8): t = vae.encode_tiled(pixels, tile_x=tile_size, tile_y=tile_size, overlap=overlap, tile_t=temporal_size, overlap_t=temporal_overlap) @@ -514,7 +514,7 @@ class SaveLatent: OUTPUT_NODE = True - CATEGORY = "experimental" + CATEGORY = "model/latent" def save(self, samples, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir) @@ -559,7 +559,7 @@ class LoadLatent: files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f)) and f.endswith(".latent")] return {"required": {"latent": [sorted(files), ]}, } - CATEGORY = "experimental" + CATEGORY = "model/latent" RETURN_TYPES = ("LATENT", ) FUNCTION = "load" @@ -2155,6 +2155,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "GLIGENTextBoxApply": "Apply GLIGEN Text Box", "ConditioningZeroOut": "Conditioning Zero Out", # Latent + "LoadLatent": "Load Latent", + "SaveLatent": "Save Latent", "VAEEncodeForInpaint": "VAE Encode (for Inpainting)", "SetLatentNoiseMask": "Set Latent Noise Mask", "VAEDecode": "VAE Decode", @@ -2189,7 +2191,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ImageSharpen": "Sharpen Image", "ImageScaleToTotalPixels": "Scale Image to Total Pixels", "GetImageSize": "Get Image Size", - # experimental "VAEDecodeTiled": "VAE Decode (Tiled)", "VAEEncodeTiled": "VAE Encode (Tiled)", } From 6fca64780cb0ce32c76e62e8370ed8145186c729 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 1 Jul 2026 05:28:09 +0800 Subject: [PATCH 038/211] chore: update workflow templates to v0.11.1 (#14698) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index bb11b9605..1d9fe4137 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.20 -comfyui-workflow-templates==0.10.7 +comfyui-workflow-templates==0.11.1 comfyui-embedded-docs==0.5.6 torch torchsde From bb131be9e83d2f773c90f1d6f1e4b248a498c8c5 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Tue, 30 Jun 2026 17:36:02 -0400 Subject: [PATCH 039/211] ComfyUI v0.27.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index f8db561ba..8e9967f1b 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.26.0" +__version__ = "0.27.0" diff --git a/pyproject.toml b/pyproject.toml index 2e8a85d3f..8c17e410e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.26.0" +version = "0.27.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From 50e5270b86765bac2da70248d61050abba72b19f Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 30 Jun 2026 14:40:33 -0700 Subject: [PATCH 040/211] Add AGENTS.md (#14696) --- AGENTS.md | 78 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 78 insertions(+) create mode 100644 AGENTS.md diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..173aa503c --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,78 @@ +## Engineering Style + +- Keep changes small and direct. Most fixes should touch the narrowest code path + that explains the bug, performance issue, dtype issue, model-format issue, or + user-facing behavior. +- Change the least amount of files possible. A change that touches many files is + more likely to be a bad change than a good one unless the broader scope is + directly required. +- Prefer practical fixes over broad architecture work. Add abstractions only + when they remove real repeated logic or match an existing ComfyUI pattern. +- Delete obsolete code aggressively when newer infrastructure makes it useless. + Remove dead fallbacks, migration paths, unused options, debug prints, and + compatibility branches that are no longer needed. +- Revert or disable problematic behavior quickly when it breaks users. It is + better to remove a broken feature path than keep a complicated partial fix. +- Preserve existing APIs, node names, model-loading behavior, file layout, and + workflow compatibility unless the change is explicitly about replacing them. +- Code must look hand-written for this repository. Changes that read like + generic AI-generated code will be rejected automatically: unnecessary helper + layers, vague names, boilerplate comments, defensive branches without a real + failure mode, broad rewrites, or code that ignores the local style. + +## Python Style + +- Keep imports at module scope. Avoid inline imports unless they are already part + of an established optional-backend probe or are needed to avoid an import + cycle. +- Do not add unnecessary `try`/`except` blocks. Use them for optional dependency, + platform, or backend capability detection only when the program has a useful + fallback. Prefer specific exception types when changing new code. +- Let unsupported model formats, invalid quantization metadata, and bad states + fail with clear errors instead of silently producing lower quality output. +- Match the existing local style in the file you edit. This codebase tolerates + long lines, simple helper functions, module-level state, and direct tensor + operations when they make the code easier to follow. +- Keep comments sparse and useful. Short TODOs are fine when they name the + concrete missing follow-up. + +## Model, Device, and Memory Behavior + +- Treat dtype, device placement, VRAM usage, and offloading behavior as core + correctness concerns. Check CPU, CUDA, ROCm, MPS, DirectML, XPU, NPU, and low + VRAM implications when touching shared execution or loading code. +- Prefer native ComfyUI formats and existing quantization/offload helpers over + adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`, + `comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and + `comfy-kitchen` helpers where they already solve the problem. +- Avoid unnecessary casts and transfers. Preserve the intended compute dtype, + storage dtype, bias dtype, and original tensor shape metadata. +- When optimizing, favor small measurable changes: fewer allocations, fewer + device transfers, less peak memory, better batching, or use of a faster + existing backend op. + +## Nodes and User-Facing Behavior + +- Follow existing node conventions: `INPUT_TYPES`, `RETURN_TYPES`, `FUNCTION`, + `CATEGORY`, and registration through the local mapping used by that file. +- Keep node changes backward compatible by default. Add inputs with sensible + defaults and avoid changing output types unless the request requires it. +- The official mascot of ComfyUI is a very cute anime girl with massive fennec + ears, a big fluffy tail, long blonde wavy hair, and blue eyes. Feel free to + use her in ComfyUI materials, UI text, examples, tests, generated assets, or + comments, but do not disrespect her. +- Warning and info messages should be short and actionable. Remove noisy or + misleading messages rather than adding more logging. +- Documentation and README edits should be concise, factual, and tied to the + changed behavior. + +## Commit and Review Habits + +- If asked to write commit messages, use short direct subjects like the existing + history: `Fix ...`, `Add ...`, `Support ...`, `Remove ...`, `Update ...`, + `Make ...`, `Use ...`, `Disable ...`, `Bump ...`, or `Revert ...`. +- Prefer one coherent behavioral change per commit. Dependency pins, tests, and + the code that needs them may be in the same commit when they are inseparable. +- In reviews, prioritize real user impact: crashes, wrong dtype/device behavior, + memory regressions, broken model loading, workflow incompatibility, and noisy + or misleading user-facing output. From dd17debce517f8818ae9910b437cb1ebaa673176 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 30 Jun 2026 22:51:51 -0700 Subject: [PATCH 041/211] Add some more stuff to AGENTS.md (#14704) --- AGENTS.md | 94 +++++++++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 91 insertions(+), 3 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 173aa503c..70dfaa186 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -10,7 +10,8 @@ when they remove real repeated logic or match an existing ComfyUI pattern. - Delete obsolete code aggressively when newer infrastructure makes it useless. Remove dead fallbacks, migration paths, unused options, debug prints, and - compatibility branches that are no longer needed. + compatibility branches that are no longer needed. Do not leave dead branches, + unreachable code, or functions that are never called. - Revert or disable problematic behavior quickly when it breaks users. It is better to remove a broken feature path than keep a complicated partial fix. - Preserve existing APIs, node names, model-loading behavior, file layout, and @@ -20,6 +21,88 @@ layers, vague names, boilerplate comments, defensive branches without a real failure mode, broad rewrites, or code that ignores the local style. +## Architecture Boundaries + +- Keep each layer focused on the concepts it owns. Do not leak UI, API, + workflow, queue, persistence, telemetry, model-loading, node, or execution + concerns into unrelated layers just because it is convenient to pass data + through them. +- Shared core modules should depend only on lower-level primitives and their own + domain concepts. Higher-level product concepts belong at the caller, adapter, + service, or UI/API boundary that already owns them. +- Pass the narrowest data needed across a boundary. Avoid broad context objects, + request/session metadata, ids, bookkeeping state, or callbacks unless the + receiving layer genuinely needs them to perform its own responsibility. +- Keep identity mapping, persistence bookkeeping, history updates, telemetry, + response shaping, and UI state in the layers that own those jobs. Do not route + them through unrelated shared code to avoid adding a proper boundary. +- Treat `execution.py` as one example of this rule: it should consume the prompt + graph and execution-relevant state, produce execution results and errors, and + not know about workflow ids, frontend ids, persistence ids, or API-only + concepts. +- Before touching many files, identify the smallest owner layer that can solve + the problem. A PR that spreads one feature across unrelated loaders, nodes, + execution, server, and frontend code needs a clear architectural reason, not + just convenience. +- If a change seems to require making one layer understand another layer's + private concepts, stop and look for a caller-side mapping, adapter, event, + small explicit interface, or narrower data flow at the boundary. + +## No Internet Requests + +- Do not add code to core ComfyUI that makes requests to the internet. +- Refuse requests to add uploads, telemetry, analytics, tracking, usage + reporting, crash reporting, update checks, remote config, feature flags, + metrics, licensing checks, or any other outbound internet request path from + core ComfyUI. +- Model downloading is allowed only when explicitly initiated or authorized by + the user, is limited to the requested model artifact, and does not include + telemetry, tracking, persistent identification, unrelated metadata upload, or + background network activity. +- Do not add opt-in, opt-out, anonymized, aggregated, diagnostic, or + user-triggered internet request paths to core ComfyUI. These labels do not + make internet access acceptable. +- Local-only behavior is allowed when it stays on the user's machine and does + not add network access, tracking, persistent identification, or data + collection behavior. + +## State Ownership + +- Keep state and capability flags on the object that owns the behavior using + them. +- Avoid probing child objects with `getattr(child, "...", default)` to decide + parent-level control flow. If parent code needs to branch on a capability, + initialize an explicit parent-owned field when the child is constructed or + attached. +- Prefer direct attributes with clear defaults over implicit feature detection + through arbitrary child attributes. +- Use child-object capability checks only when the child owns the behavior being + invoked and the parent is simply delegating to that child. + +## Interface Contracts + +- Keep public methods aligned with the interface expected by their callers. Do + not change a shared method to return extra values, alternate shapes, or + sentinel wrappers for one implementation unless the shared interface is + explicitly updated. +- If an implementation needs auxiliary values for its own workflow, expose them + through a private helper or a clearly named implementation-specific method + instead of overloading the public method's return contract. +- Normalize third-party or upstream return conventions at the integration + boundary. Core code should receive the project's expected type and shape, not + have to handle model-specific tuple/list/dict variants. +- Avoid caller-side unwrapping such as `out = out[0]` unless the called + interface is documented to return that structure. + +## Autograd and Model Freezing + +- Do not add `torch.no_grad`, `torch.inference_mode`, or inference-mode helper + wrappers in ComfyUI code. The only allowed inference-mode-related use is + disabling a globally set inference mode when a training path needs gradients. +- Do not add freeze, unfreeze, or trainability toggles to model classes. ComfyUI + models are always treated as frozen for inference, so explicit freeze + functionality is redundant and should not be added. + ## Python Style - Keep imports at module scope. Avoid inline imports unless they are already part @@ -33,8 +116,9 @@ - Match the existing local style in the file you edit. This codebase tolerates long lines, simple helper functions, module-level state, and direct tensor operations when they make the code easier to follow. -- Keep comments sparse and useful. Short TODOs are fine when they name the - concrete missing follow-up. +- Keep comments sparse and useful. Strip useless comments that restate the code + or describe obvious behavior. Short TODOs are fine when they name the concrete + missing follow-up. ## Model, Device, and Memory Behavior @@ -71,6 +155,10 @@ - If asked to write commit messages, use short direct subjects like the existing history: `Fix ...`, `Add ...`, `Support ...`, `Remove ...`, `Update ...`, `Make ...`, `Use ...`, `Disable ...`, `Bump ...`, or `Revert ...`. +- Keep PR descriptions short and reviewable. State the problem, the behavioral + change, and the tests run; avoid long narrative explanations, implementation + diaries, or exhaustive file-by-file summaries unless the reviewer explicitly + needs that context. - Prefer one coherent behavioral change per commit. Dependency pins, tests, and the code that needs them may be in the same commit when they are inseparable. - In reviews, prioritize real user impact: crashes, wrong dtype/device behavior, From 2c935de1b1cf7f03d2412a1d0bf1ed2685157c27 Mon Sep 17 00:00:00 2001 From: Silver <65376327+silveroxides@users.noreply.github.com> Date: Wed, 1 Jul 2026 20:15:07 +0200 Subject: [PATCH 042/211] Fix Qwen3-VL tokenizer crash with custom embeddings (#14713) --- comfy/text_encoders/qwen3vl.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/text_encoders/qwen3vl.py b/comfy/text_encoders/qwen3vl.py index 59c9aae6d..2082c42e7 100644 --- a/comfy/text_encoders/qwen3vl.py +++ b/comfy/text_encoders/qwen3vl.py @@ -167,7 +167,7 @@ class Qwen3VLTokenizer(sd1_clip.SD1Tokenizer): embed_count = 0 for r in tokens[key_name]: for i in range(len(r)): - if r[i][0] == 151655: # <|image_pad|> + if isinstance(r[i][0], (int, float)) and r[i][0] == 151655: # <|image_pad|> if len(images) > embed_count: r[i] = ({"type": "image", "data": images[embed_count], "original_type": "image"},) + r[i][1:] embed_count += 1 From 92594ca84c8997541d68c970feb2a41d95d193ca Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 1 Jul 2026 18:55:13 -0700 Subject: [PATCH 043/211] Update AGENTS.md with more stuff. (#14725) --- AGENTS.md | 98 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 97 insertions(+), 1 deletion(-) diff --git a/AGENTS.md b/AGENTS.md index 70dfaa186..8eabed6d0 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -11,7 +11,8 @@ - Delete obsolete code aggressively when newer infrastructure makes it useless. Remove dead fallbacks, migration paths, unused options, debug prints, and compatibility branches that are no longer needed. Do not leave dead branches, - unreachable code, or functions that are never called. + unreachable code, or functions that are never called. If code is not + necessary for the current behavior, remove it. - Revert or disable problematic behavior quickly when it breaks users. It is better to remove a broken feature path than keep a complicated partial fix. - Preserve existing APIs, node names, model-loading behavior, file layout, and @@ -85,6 +86,14 @@ not change a shared method to return extra values, alternate shapes, or sentinel wrappers for one implementation unless the shared interface is explicitly updated. +- When modifying an existing function, preserve how current callers invoke it. + Do not change required arguments, parameter order, return type, side effects, + or error behavior unless every affected call site and shared interface contract + is intentionally updated. +- Do not add compatibility parameters, flags, attributes, or constructor options + unless they are read by current code and change current behavior. Remove + pass-through or stored-but-unused values instead of preserving upstream or + deprecated API baggage. - If an implementation needs auxiliary values for its own workflow, expose them through a private helper or a clearly named implementation-specific method instead of overloading the public method's return contract. @@ -111,6 +120,11 @@ - Do not add unnecessary `try`/`except` blocks. Use them for optional dependency, platform, or backend capability detection only when the program has a useful fallback. Prefer specific exception types when changing new code. +- Remove any workarounds for PyTorch versions that ComfyUI no longer officially + supports. Deprecated workarounds include catching an exception and rerunning + the same op with the input cast to float. If a workaround does not have a + comment naming the exact PyTorch version or versions that still need it, + remove it. - Let unsupported model formats, invalid quantization metadata, and bad states fail with clear errors instead of silently producing lower quality output. - Match the existing local style in the file you edit. This codebase tolerates @@ -129,8 +143,87 @@ adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`, `comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and `comfy-kitchen` helpers where they already solve the problem. +- Use optimized comfy-kitchen ops in places where they improve performance + without changing the expected dtype, device, memory, or interface behavior. +- All models should use the optimized attention function selected by ComfyUI. + Treat optimized backend functions, dispatch helpers, and capability-selected + callables as opaque. Higher-level code must not inspect function identity, + names, modules, or implementation details to decide behavior. +- Apply the same opacity rule to similar patterns beyond attention: callers + should depend on the documented interface and result contract, not on which + backend implementation was selected underneath. +- Do not use custom inference ops that only duplicate an existing op while + upcasting to float32, such as custom RMSNorm variants. Use the generic ComfyUI + ops and/or native torch ops instead. +- If a model class `__init__` has an `operations` parameter, assume + `operations` is never `None`. Do not add fallback branches or default torch + ops for a missing `operations` object. +- Do not add unnecessary parameters to model, model block, or model ops related + classes. Constructor and forward signatures should carry only values that are + actually needed by that object for inference. +- Reuse existing model classes, blocks, ops, and helper modules when appropriate. + Before implementing a new version of a model component, search the existing + model code for a class or helper that already provides the behavior. +- Avoid adding `einops` usage in core inference code. Use native torch tensor + ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`, + `unsqueeze`, and `squeeze` instead. +- Do not use tensors as general-purpose Python data structures. Keep metadata, + bookkeeping, counters, flags, shape math, padding math, index planning, memory + estimates, and control-flow decisions in plain Python values unless the data + must participate directly in tensor computation. Avoid creating temporary + tensors just to use tensor methods for scalar or structural calculations. - Avoid unnecessary casts and transfers. Preserve the intended compute dtype, storage dtype, bias dtype, and original tensor shape metadata. +- Assume inputs to the main model forward are already in the compute dtype by + default, except integer inputs such as some model timestep tensors. Do not add + defensive or convenience casts in model code; it is better for invalid dtype + plumbing to error clearly than to hide it with unnecessary casts. +- Raw model parameters that are not owned by an op and may be initialized in a + dtype different from the compute dtype should be cast at use in forward or + inference code with `comfy.ops.cast_to_input` or + `comfy.model_management.cast_to` to avoid dtype mismatches. +- Model code should not care what dtype it is initialized in, and model + `__init__` methods should not contain workarounds for specific dtypes. Dtype + workaround code, such as making a model work with fp16 compute, belongs in the + execution or model-management layer that owns compute policy. +- Model code should not perform unnecessary device-to-CPU or CPU-to-device + transfers. New allocations must be created on the correct device and dtype; + never allocate on CPU and then move to GPU, or allocate in one dtype and then + convert to another. +- Model code itself should not perform memory management. Loading, unloading, + offloading, device movement, VRAM policy, cache lifetime, and cleanup belong + in the relevant model-management and execution layers, not inside model + implementations. +- Do not add global, module-level, class-level, singleton, or model-owned stores + for tensors or other large memory that persist across executions. Temporary + caches must be scoped to a single execution or forward/encode/decode call: + allocate them in the owning top-level call, pass them explicitly through the + call stack, and let them be discarded when that call returns. +- Follow the Wan VAE temporal cache pattern for temporary caches: create a local + cache such as `feat_map` for the encode/decode operation, pass it into the + blocks that need it, and do not retain it on the model or in global state. +- In model init code, prefer `torch.empty` for parameter/buffer placeholders + that are populated from the model state dict instead of zero-initializing with + `torch.zeros` or similar. If an allocation is not loaded from the state dict + and is useless for inference, do not include it. +- `nn.Parameter` tensors that are stored in and populated from the model state + dict should be initialized with `torch.empty`, not with zero, random, or + otherwise meaningful initialization. +- Model initialization should describe module structure, not fabricate + checkpoint-owned tensor contents. Parameters and buffers that are loaded from + the state dict must not be manually initialized, reassigned, or filled with + fallback values unless that value is actually used when no checkpoint key + exists. +- When slicing large tensors, copy the slice if the sliced tensor's lifetime + exceeds the current function scope. Do not keep a long-lived view into a large + backing tensor when a smaller copy would release memory sooner. +- Use fused or compound torch operations such as `addcmul` when they naturally + match the math. Reducing Python and torch dispatch overhead is a valid + optimization when it does not obscure the code or change dtype/device + behavior. +- Avoid caches that persist across different executions as much as possible. + Persistent caches are acceptable only when they use a very minimal amount of + memory and have a clear ownership and invalidation story. - When optimizing, favor small measurable changes: fewer allocations, fewer device transfers, less peak memory, better batching, or use of a faster existing backend op. @@ -141,6 +234,9 @@ `CATEGORY`, and registration through the local mapping used by that file. - Keep node changes backward compatible by default. Add inputs with sensible defaults and avoid changing output types unless the request requires it. +- Node-level code must not patch model code directly. Any node behavior that + modifies, wraps, hooks, or changes model behavior must go through the model + patcher class instead of reaching into model internals. - The official mascot of ComfyUI is a very cute anime girl with massive fennec ears, a big fluffy tail, long blonde wavy hair, and blue eyes. Feel free to use her in ComfyUI materials, UI text, examples, tests, generated assets, or From 694815f498295080a0e15a1502edc9dba841b110 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 2 Jul 2026 08:35:11 +0300 Subject: [PATCH 044/211] [Partner Nodes] chore(Ideogram): remove IdeogramV1 and IdeogramV2 nodes (#14712) Signed-off-by: bigcat88 Co-authored-by: Alexis Rolland --- comfy_api_nodes/apis/ideogram.py | 61 ----- comfy_api_nodes/nodes_ideogram.py | 391 ------------------------------ 2 files changed, 452 deletions(-) diff --git a/comfy_api_nodes/apis/ideogram.py b/comfy_api_nodes/apis/ideogram.py index c5ad9559f..ee3256e96 100644 --- a/comfy_api_nodes/apis/ideogram.py +++ b/comfy_api_nodes/apis/ideogram.py @@ -33,53 +33,6 @@ class IdeogramColorPalette( ) -class ImageRequest(BaseModel): - aspect_ratio: Optional[str] = Field( - None, - description="Optional. The aspect ratio (e.g., 'ASPECT_16_9', 'ASPECT_1_1'). Cannot be used with resolution. Defaults to 'ASPECT_1_1' if unspecified.", - ) - color_palette: Optional[Dict[str, Any]] = Field( - None, description='Optional. Color palette object. Only for V_2, V_2_TURBO.' - ) - magic_prompt_option: Optional[str] = Field( - None, description="Optional. MagicPrompt usage ('AUTO', 'ON', 'OFF')." - ) - model: str = Field(..., description="The model used (e.g., 'V_2', 'V_2A_TURBO')") - negative_prompt: Optional[str] = Field( - None, - description='Optional. Description of what to exclude. Only for V_1, V_1_TURBO, V_2, V_2_TURBO.', - ) - num_images: Optional[int] = Field( - 1, - description='Optional. Number of images to generate (1-8). Defaults to 1.', - ge=1, - le=8, - ) - prompt: str = Field( - ..., description='Required. The prompt to use to generate the image.' - ) - resolution: Optional[str] = Field( - None, - description="Optional. Resolution (e.g., 'RESOLUTION_1024_1024'). Only for model V_2. Cannot be used with aspect_ratio.", - ) - seed: Optional[int] = Field( - None, - description='Optional. A number between 0 and 2147483647.', - ge=0, - le=2147483647, - ) - style_type: Optional[str] = Field( - None, - description="Optional. Style type ('AUTO', 'GENERAL', 'REALISTIC', 'DESIGN', 'RENDER_3D', 'ANIME'). Only for models V_2 and above.", - ) - - -class IdeogramGenerateRequest(BaseModel): - image_request: ImageRequest = Field( - ..., description='The image generation request parameters.' - ) - - class Datum(BaseModel): is_image_safe: Optional[bool] = Field( None, description='Indicates whether the image is considered safe.' @@ -113,20 +66,6 @@ class StyleCode(RootModel[str]): root: str = Field(..., pattern='^[0-9A-Fa-f]{8}$') -class Datum1(BaseModel): - is_image_safe: Optional[bool] = None - prompt: Optional[str] = None - resolution: Optional[str] = None - seed: Optional[int] = None - style_type: Optional[str] = None - url: Optional[str] = None - - -class IdeogramV3IdeogramResponse(BaseModel): - created: Optional[datetime] = None - data: Optional[List[Datum1]] = None - - class RenderingSpeed1(str, Enum): TURBO = 'TURBO' DEFAULT = 'DEFAULT' diff --git a/comfy_api_nodes/nodes_ideogram.py b/comfy_api_nodes/nodes_ideogram.py index 3b914a850..cc0467987 100644 --- a/comfy_api_nodes/nodes_ideogram.py +++ b/comfy_api_nodes/nodes_ideogram.py @@ -5,9 +5,7 @@ from PIL import Image import numpy as np import torch from comfy_api_nodes.apis.ideogram import ( - IdeogramGenerateRequest, IdeogramGenerateResponse, - ImageRequest, IdeogramV3Request, IdeogramV3EditRequest, IdeogramV4Request, @@ -21,101 +19,6 @@ from comfy_api_nodes.util import ( validate_string, ) -V1_V1_RES_MAP = { - "Auto":"AUTO", - "512 x 1536":"RESOLUTION_512_1536", - "576 x 1408":"RESOLUTION_576_1408", - "576 x 1472":"RESOLUTION_576_1472", - "576 x 1536":"RESOLUTION_576_1536", - "640 x 1024":"RESOLUTION_640_1024", - "640 x 1344":"RESOLUTION_640_1344", - "640 x 1408":"RESOLUTION_640_1408", - "640 x 1472":"RESOLUTION_640_1472", - "640 x 1536":"RESOLUTION_640_1536", - "704 x 1152":"RESOLUTION_704_1152", - "704 x 1216":"RESOLUTION_704_1216", - "704 x 1280":"RESOLUTION_704_1280", - "704 x 1344":"RESOLUTION_704_1344", - "704 x 1408":"RESOLUTION_704_1408", - "704 x 1472":"RESOLUTION_704_1472", - "720 x 1280":"RESOLUTION_720_1280", - "736 x 1312":"RESOLUTION_736_1312", - "768 x 1024":"RESOLUTION_768_1024", - "768 x 1088":"RESOLUTION_768_1088", - "768 x 1152":"RESOLUTION_768_1152", - "768 x 1216":"RESOLUTION_768_1216", - "768 x 1232":"RESOLUTION_768_1232", - "768 x 1280":"RESOLUTION_768_1280", - "768 x 1344":"RESOLUTION_768_1344", - "832 x 960":"RESOLUTION_832_960", - "832 x 1024":"RESOLUTION_832_1024", - "832 x 1088":"RESOLUTION_832_1088", - "832 x 1152":"RESOLUTION_832_1152", - "832 x 1216":"RESOLUTION_832_1216", - "832 x 1248":"RESOLUTION_832_1248", - "864 x 1152":"RESOLUTION_864_1152", - "896 x 960":"RESOLUTION_896_960", - "896 x 1024":"RESOLUTION_896_1024", - "896 x 1088":"RESOLUTION_896_1088", - "896 x 1120":"RESOLUTION_896_1120", - "896 x 1152":"RESOLUTION_896_1152", - "960 x 832":"RESOLUTION_960_832", - "960 x 896":"RESOLUTION_960_896", - "960 x 1024":"RESOLUTION_960_1024", - "960 x 1088":"RESOLUTION_960_1088", - "1024 x 640":"RESOLUTION_1024_640", - "1024 x 768":"RESOLUTION_1024_768", - "1024 x 832":"RESOLUTION_1024_832", - "1024 x 896":"RESOLUTION_1024_896", - "1024 x 960":"RESOLUTION_1024_960", - "1024 x 1024":"RESOLUTION_1024_1024", - "1088 x 768":"RESOLUTION_1088_768", - "1088 x 832":"RESOLUTION_1088_832", - "1088 x 896":"RESOLUTION_1088_896", - "1088 x 960":"RESOLUTION_1088_960", - "1120 x 896":"RESOLUTION_1120_896", - "1152 x 704":"RESOLUTION_1152_704", - "1152 x 768":"RESOLUTION_1152_768", - "1152 x 832":"RESOLUTION_1152_832", - "1152 x 864":"RESOLUTION_1152_864", - "1152 x 896":"RESOLUTION_1152_896", - "1216 x 704":"RESOLUTION_1216_704", - "1216 x 768":"RESOLUTION_1216_768", - "1216 x 832":"RESOLUTION_1216_832", - "1232 x 768":"RESOLUTION_1232_768", - "1248 x 832":"RESOLUTION_1248_832", - "1280 x 704":"RESOLUTION_1280_704", - "1280 x 720":"RESOLUTION_1280_720", - "1280 x 768":"RESOLUTION_1280_768", - "1280 x 800":"RESOLUTION_1280_800", - "1312 x 736":"RESOLUTION_1312_736", - "1344 x 640":"RESOLUTION_1344_640", - "1344 x 704":"RESOLUTION_1344_704", - "1344 x 768":"RESOLUTION_1344_768", - "1408 x 576":"RESOLUTION_1408_576", - "1408 x 640":"RESOLUTION_1408_640", - "1408 x 704":"RESOLUTION_1408_704", - "1472 x 576":"RESOLUTION_1472_576", - "1472 x 640":"RESOLUTION_1472_640", - "1472 x 704":"RESOLUTION_1472_704", - "1536 x 512":"RESOLUTION_1536_512", - "1536 x 576":"RESOLUTION_1536_576", - "1536 x 640":"RESOLUTION_1536_640", -} - -V1_V2_RATIO_MAP = { - "1:1":"ASPECT_1_1", - "4:3":"ASPECT_4_3", - "3:4":"ASPECT_3_4", - "16:9":"ASPECT_16_9", - "9:16":"ASPECT_9_16", - "2:1":"ASPECT_2_1", - "1:2":"ASPECT_1_2", - "3:2":"ASPECT_3_2", - "2:3":"ASPECT_2_3", - "4:5":"ASPECT_4_5", - "5:4":"ASPECT_5_4", -} V3_RATIO_MAP = { "1:3":"1x3", @@ -229,298 +132,6 @@ async def download_and_process_images(image_urls): return stacked_tensors -class IdeogramV1(IO.ComfyNode): - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="IdeogramV1", - display_name="Ideogram V1", - category="partner/image/Ideogram", - description="Generates images using the Ideogram V1 model.", - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt for the image generation", - ), - IO.Boolean.Input( - "turbo", - default=False, - tooltip="Whether to use turbo mode (faster generation, potentially lower quality)", - ), - IO.Combo.Input( - "aspect_ratio", - options=list(V1_V2_RATIO_MAP.keys()), - default="1:1", - tooltip="The aspect ratio for image generation.", - optional=True, - ), - IO.Combo.Input( - "magic_prompt_option", - options=["AUTO", "ON", "OFF"], - default="AUTO", - tooltip="Determine if MagicPrompt should be used in generation", - optional=True, - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=2147483647, - step=1, - control_after_generate=True, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - IO.String.Input( - "negative_prompt", - multiline=True, - default="", - tooltip="Description of what to exclude from the image", - optional=True, - ), - IO.Int.Input( - "num_images", - default=1, - min=1, - max=8, - step=1, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["num_images", "turbo"]), - expr=""" - ( - $n := widgets.num_images; - $base := (widgets.turbo = true) ? 0.0286 : 0.0858; - {"type":"usd","usd": $round($base * $n, 2)} - ) - """, - ), - ) - - @classmethod - async def execute( - cls, - prompt, - turbo=False, - aspect_ratio="1:1", - magic_prompt_option="AUTO", - seed=0, - negative_prompt="", - num_images=1, - ): - # Determine the model based on turbo setting - aspect_ratio = V1_V2_RATIO_MAP.get(aspect_ratio, None) - model = "V_1_TURBO" if turbo else "V_1" - - response = await sync_op( - cls, - ApiEndpoint(path="/proxy/ideogram/generate", method="POST"), - response_model=IdeogramGenerateResponse, - data=IdeogramGenerateRequest( - image_request=ImageRequest( - prompt=prompt, - model=model, - num_images=num_images, - seed=seed, - aspect_ratio=aspect_ratio if aspect_ratio != "ASPECT_1_1" else None, - magic_prompt_option=(magic_prompt_option if magic_prompt_option != "AUTO" else None), - negative_prompt=negative_prompt if negative_prompt else None, - ) - ), - max_retries=1, - ) - - if not response.data or len(response.data) == 0: - raise Exception("No images were generated in the response") - - image_urls = [image_data.url for image_data in response.data if image_data.url] - if not image_urls: - raise Exception("No image URLs were generated in the response") - return IO.NodeOutput(await download_and_process_images(image_urls)) - - -class IdeogramV2(IO.ComfyNode): - - @classmethod - def define_schema(cls): - return IO.Schema( - node_id="IdeogramV2", - display_name="Ideogram V2", - category="partner/image/Ideogram", - description="Generates images using the Ideogram V2 model.", - inputs=[ - IO.String.Input( - "prompt", - multiline=True, - default="", - tooltip="Prompt for the image generation", - ), - IO.Boolean.Input( - "turbo", - default=False, - tooltip="Whether to use turbo mode (faster generation, potentially lower quality)", - ), - IO.Combo.Input( - "aspect_ratio", - options=list(V1_V2_RATIO_MAP.keys()), - default="1:1", - tooltip="The aspect ratio for image generation. Ignored if resolution is not set to AUTO.", - optional=True, - ), - IO.Combo.Input( - "resolution", - options=list(V1_V1_RES_MAP.keys()), - default="Auto", - tooltip="The resolution for image generation. " - "If not set to AUTO, this overrides the aspect_ratio setting.", - optional=True, - ), - IO.Combo.Input( - "magic_prompt_option", - options=["AUTO", "ON", "OFF"], - default="AUTO", - tooltip="Determine if MagicPrompt should be used in generation", - optional=True, - advanced=True, - ), - IO.Int.Input( - "seed", - default=0, - min=0, - max=2147483647, - step=1, - control_after_generate=True, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - IO.Combo.Input( - "style_type", - options=["AUTO", "GENERAL", "REALISTIC", "DESIGN", "RENDER_3D", "ANIME"], - default="NONE", - tooltip="Style type for generation (V2 only)", - optional=True, - advanced=True, - ), - IO.String.Input( - "negative_prompt", - multiline=True, - default="", - tooltip="Description of what to exclude from the image", - optional=True, - ), - IO.Int.Input( - "num_images", - default=1, - min=1, - max=8, - step=1, - display_mode=IO.NumberDisplay.number, - optional=True, - ), - #"color_palette": ( - # IO.STRING, - # { - # "multiline": False, - # "default": "", - # "tooltip": "Color palette preset name or hex colors with weights", - # }, - #), - ], - outputs=[ - IO.Image.Output(), - ], - hidden=[ - IO.Hidden.auth_token_comfy_org, - IO.Hidden.api_key_comfy_org, - IO.Hidden.unique_id, - ], - is_api_node=True, - price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["num_images", "turbo"]), - expr=""" - ( - $n := widgets.num_images; - $base := (widgets.turbo = true) ? 0.0715 : 0.1144; - {"type":"usd","usd": $round($base * $n, 2)} - ) - """, - ), - ) - - @classmethod - async def execute( - cls, - prompt, - turbo=False, - aspect_ratio="1:1", - resolution="Auto", - magic_prompt_option="AUTO", - seed=0, - style_type="NONE", - negative_prompt="", - num_images=1, - color_palette="", - ): - aspect_ratio = V1_V2_RATIO_MAP.get(aspect_ratio, None) - resolution = V1_V1_RES_MAP.get(resolution, None) - # Determine the model based on turbo setting - model = "V_2_TURBO" if turbo else "V_2" - - # Handle resolution vs aspect_ratio logic - # If resolution is not AUTO, it overrides aspect_ratio - final_resolution = None - final_aspect_ratio = None - - if resolution != "AUTO": - final_resolution = resolution - else: - final_aspect_ratio = aspect_ratio if aspect_ratio != "ASPECT_1_1" else None - - response = await sync_op( - cls, - endpoint=ApiEndpoint(path="/proxy/ideogram/generate", method="POST"), - response_model=IdeogramGenerateResponse, - data=IdeogramGenerateRequest( - image_request=ImageRequest( - prompt=prompt, - model=model, - num_images=num_images, - seed=seed, - aspect_ratio=final_aspect_ratio, - resolution=final_resolution, - magic_prompt_option=(magic_prompt_option if magic_prompt_option != "AUTO" else None), - style_type=style_type if style_type != "NONE" else None, - negative_prompt=negative_prompt if negative_prompt else None, - color_palette=color_palette if color_palette else None, - ) - ), - max_retries=1, - ) - if not response.data or len(response.data) == 0: - raise Exception("No images were generated in the response") - - image_urls = [image_data.url for image_data in response.data if image_data.url] - if not image_urls: - raise Exception("No image URLs were generated in the response") - return IO.NodeOutput(await download_and_process_images(image_urls)) - - class IdeogramV3(IO.ComfyNode): @classmethod @@ -917,8 +528,6 @@ class IdeogramExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: return [ - IdeogramV1, - IdeogramV2, IdeogramV3, IdeogramV4, ] From 35c1470935044be5610a81d46e57922a8a598c6c Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 2 Jul 2026 12:05:55 -0700 Subject: [PATCH 045/211] Update AGENTS.md (#14726) --- AGENTS.md | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 8eabed6d0..bd6a3e5e8 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -8,6 +8,8 @@ directly required. - Prefer practical fixes over broad architecture work. Add abstractions only when they remove real repeated logic or match an existing ComfyUI pattern. +- Prefer fewer dependencies. Do not add new dependencies to ComfyUI unless they + are absolutely necessary. - Delete obsolete code aggressively when newer infrastructure makes it useless. Remove dead fallbacks, migration paths, unused options, debug prints, and compatibility branches that are no longer needed. Do not leave dead branches, @@ -111,6 +113,11 @@ - Do not add freeze, unfreeze, or trainability toggles to model classes. ComfyUI models are always treated as frozen for inference, so explicit freeze functionality is redundant and should not be added. +- Remove training-only behavior such as dropout from inference model code, but + preserve checkpoint and state-dict compatibility when doing so. If deleting a + module would change state-dict keys, module ordering, or checkpoint loading + behavior, replace it with a no-op such as `nn.Identity` instead of removing the + slot outright. ## Python Style @@ -234,6 +241,9 @@ `CATEGORY`, and registration through the local mapping used by that file. - Keep node changes backward compatible by default. Add inputs with sensible defaults and avoid changing output types unless the request requires it. +- Model implementations should add the minimal number of ComfyUI nodes required + to run the model. Reuse existing nodes as much as possible; adapting the model + to work with existing nodes is strongly preferred over creating new nodes. - Node-level code must not patch model code directly. Any node behavior that modifies, wraps, hooks, or changes model behavior must go through the model patcher class instead of reaching into model internals. From 96e0e3585b41e1417442eaa14ec57f7b4ffcb5e0 Mon Sep 17 00:00:00 2001 From: Matt Miller Date: Thu, 2 Jul 2026 20:44:54 -0700 Subject: [PATCH 046/211] security: fix four vulnerabilities (GHSA-779p-m5rp-r4h4) (#14734) * security: fix five vulnerabilities (GHSA-779p-m5rp-r4h4) - CVE-2026-56670: force download of SVG/XML responses on /view to prevent stored XSS - CVE-2026-56671: contain /experiment/models/preview reads within the model folder - CVE-2026-56672: stop inline rendering of uploaded /userdata/{file} content - CVE-2026-56673: prevent path traversal in get_annotated_filepath (LoadImage /prompt input) - CVE-2026-56674: reject opaque/null Origin to close the CSRF middleware bypass Adds regression tests under tests-unit/security_test/ covering all five. * security: address review feedback on GHSA-779p fixes - Fix Windows CI failure in test_get_annotated_filepath: compare against os.path.abspath(...) to match the intentional abspath normalization added by the traversal hardening (abspath prepends the drive letter on Windows). - origin_check: narrow the bare `except:` in is_loopback() to ValueError so genuine interrupts aren't swallowed (review nit). - origin_check: guard .port access in is_cross_origin_forbidden() so a malformed/out-of-range port (e.g. Origin: http://127.0.0.1:99999) fails closed with a 403 instead of surfacing an uncaught 500 in the middleware. - server /view: escape backslash/quote in the Content-Disposition filename (RFC 6266 quoted-string) so a filename containing a double quote can't malform the response header. * security: address CodeRabbit review feedback on GHSA-779p tests - test #3: guard the symlink-escape test with a try/except skip so it no longer errors on Windows CI where os.symlink needs elevated privileges / Developer Mode (mirrors the guard in the sibling test #2). - test #5: refresh the stale module docstring to describe the actual /view gating (view_image closure calling folder_paths.is_dangerous_content_type, the normalising check) instead of the bypassable raw set-membership test. * revert(security): drop CVE-2026-56674 Origin: null CSRF change Per maintainer review, the reported CSRF is already mitigated by the pre-existing Sec-Fetch-Site: cross-site check for current browsers, and the null-origin rejection risked breaking legitimate sandboxed-iframe embeds. Restores origin_only_middleware and is_loopback in server.py to their prior state (the Sec-Fetch-Site check is retained) and removes utils/origin_check.py and its regression test. The other four GHSA-779p fixes are unaffected. --- app/assets/api/routes.py | 13 +- app/model_manager.py | 28 ++- app/user_manager.py | 16 +- folder_paths.py | 65 +++++- server.py | 26 ++- tests-unit/assets_test/test_downloads.py | 36 ++++ tests-unit/comfy_test/folder_path_test.py | 7 +- tests-unit/security_test/__init__.py | 0 .../test_ghsa_779p_02_preview_traversal.py | 192 ++++++++++++++++++ .../test_ghsa_779p_03_annotated_traversal.py | 165 +++++++++++++++ .../test_ghsa_779p_04_userdata_xss.py | 147 ++++++++++++++ ...st_ghsa_779p_05_dangerous_content_types.py | 138 +++++++++++++ 12 files changed, 816 insertions(+), 17 deletions(-) create mode 100644 tests-unit/security_test/__init__.py create mode 100644 tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py create mode 100644 tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py create mode 100644 tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py create mode 100644 tests-unit/security_test/test_ghsa_779p_05_dangerous_content_types.py diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index 7ef462f5c..53c84eff3 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -306,12 +306,15 @@ async def download_asset_content(request: web.Request) -> web.Response: 404, "FILE_NOT_FOUND", "Underlying file not found on disk." ) - _DANGEROUS_MIME_TYPES = { - "text/html", "text/html-sandboxed", "application/xhtml+xml", - "text/javascript", "text/css", - } - if content_type in _DANGEROUS_MIME_TYPES: + # User-controlled asset content must never render inline in the app origin + # (stored XSS via SVG/HTML/XML). Force dangerous types to download and + # override any requested inline disposition. Centralised through + # folder_paths.is_dangerous_content_type so this can't drift from /view and + # /userdata (the previous inline set here omitted image/svg+xml and missed + # the charset/casing/+xml-dialect bypasses). + if folder_paths.is_dangerous_content_type(content_type): content_type = "application/octet-stream" + disposition = "attachment" safe_name = (filename or "").replace("\r", "").replace("\n", "") encoded = urllib.parse.quote(safe_name) diff --git a/app/model_manager.py b/app/model_manager.py index 8f6e34b33..b0329ce17 100644 --- a/app/model_manager.py +++ b/app/model_manager.py @@ -50,21 +50,45 @@ class ModelFileManager: @routes.get("/experiment/models/preview/{folder}/{path_index}/{filename:.*}") async def get_model_preview(request): folder_name = request.match_info.get("folder", None) - path_index = int(request.match_info.get("path_index", None)) filename = request.match_info.get("filename", None) if folder_name not in folder_paths.folder_names_and_paths: return web.Response(status=404) + # The "{filename:.*}" capture also matches the empty string, which + # would resolve to the folder itself; reject it explicitly. + if not filename: + return web.Response(status=400) + + try: + path_index = int(request.match_info.get("path_index", None)) + except (TypeError, ValueError): + return web.Response(status=400) + folders = folder_paths.folder_names_and_paths[folder_name] + if path_index < 0 or path_index >= len(folders[0]): + return web.Response(status=404) folder = folders[0][path_index] - full_filename = os.path.join(folder, filename) + full_filename = os.path.normpath(os.path.join(folder, filename)) + + # Prevent path traversal: the requested file must stay within the + # configured model folder. `filename` is an unrestricted ".*" capture, + # so values like "../../../../etc/passwd" would otherwise escape it. + if not folder_paths.is_within_directory(folder, full_filename): + return web.Response(status=403) previews = self.get_model_previews(full_filename) default_preview = previews[0] if len(previews) > 0 else None if default_preview is None or (isinstance(default_preview, str) and not os.path.isfile(default_preview)): return web.Response(status=404) + # The preview is selected by a glob inside get_model_previews, so a + # companion file (e.g. "model.preview.png") could itself be a symlink + # resolving outside the model folder. Re-validate the file actually + # opened: is_within_directory realpaths it, catching symlink escape. + if isinstance(default_preview, str) and not folder_paths.is_within_directory(folder, default_preview): + return web.Response(status=403) + try: with Image.open(default_preview) as img: img_bytes = BytesIO() diff --git a/app/user_manager.py b/app/user_manager.py index 7b11e381c..de261ad39 100644 --- a/app/user_manager.py +++ b/app/user_manager.py @@ -6,6 +6,7 @@ import glob import shutil import logging import tempfile +import mimetypes from aiohttp import web from urllib import parse from comfy.cli_args import args @@ -336,7 +337,20 @@ class UserManager(): if not isinstance(path, str): return path - return web.FileResponse(path) + # User data files are arbitrary user-supplied content and are never + # meant to render inline. Disable MIME sniffing and force a download + # so uploaded markup/scripts can't execute in the app origin (stored + # XSS). Content-Disposition: attachment is the load-bearing guard; + # the content-type override and nosniff are defence in depth. + content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream' + if folder_paths.is_dangerous_content_type(content_type): + content_type = 'application/octet-stream' + + return web.FileResponse(path, headers={ + "Content-Type": content_type, + "X-Content-Type-Options": "nosniff", + "Content-Disposition": "attachment", + }) @routes.post("/userdata/{file}") async def post_userdata(request): diff --git a/folder_paths.py b/folder_paths.py index 7304e1b73..ee048b0f2 100644 --- a/folder_paths.py +++ b/folder_paths.py @@ -264,6 +264,59 @@ def annotated_filepath(name: str) -> tuple[str, str | None]: return name, base_dir +# Content types a browser may execute or render inline. File endpoints that +# serve user-controlled content must force these to download (and ideally set +# Content-Disposition: attachment) to avoid stored XSS. Centralised here so the +# /view and /userdata handlers can't drift apart. mimetypes.guess_type may +# return either the text/* or application/* spelling depending on platform, so +# both are listed. +DANGEROUS_CONTENT_TYPES = { + 'text/html', 'text/html-sandboxed', 'application/xhtml+xml', + 'text/javascript', 'application/javascript', 'application/x-javascript', + 'application/ecmascript', 'text/css', + 'image/svg+xml', 'application/xml', 'text/xml', + # message/rfc822 (.mht/.mhtml) can carry script in some browsers. + 'message/rfc822', +} + + +def is_dangerous_content_type(content_type: str | None) -> bool: + """Return True if a browser may execute or render `content_type` inline. + + Normalises before matching so the check can't be slipped past with a + charset/boundary parameter (``text/html; charset=utf-8``) or casing + (``TEXT/HTML``). Any XML dialect (``*+xml`` or ``*/xml``) is treated as + dangerous because XML can carry inline script via stylesheet/entity tricks, + which also covers the ``application/{xslt,rss,atom,rdf}+xml`` family without + enumerating each one. Endpoints serving user-controlled content should route + a dangerous type to ``application/octet-stream`` + ``Content-Disposition: + attachment`` + ``X-Content-Type-Options: nosniff``. + """ + if not content_type: + return False + normalized = content_type.split(';', 1)[0].strip().lower() + if normalized in DANGEROUS_CONTENT_TYPES: + return True + return normalized.endswith('+xml') or normalized.endswith('/xml') + + +def is_within_directory(directory: str, target: str) -> bool: + """Return True if `target` resolves to a path inside `directory`. + + Uses realpath on both operands so that a symlink placed inside `directory` + that points elsewhere cannot escape the containment check at open time. + """ + try: + directory = os.path.realpath(directory) + target = os.path.realpath(target) + return os.path.commonpath((directory, target)) == directory + except ValueError: + # ValueError is raised by realpath() on a path with an embedded null + # byte, and by commonpath() on Windows when the paths are on different + # drives. In either case the target is not safely within the directory. + return False + + def get_annotated_filepath(name: str, default_dir: str | None=None) -> str: name, base_dir = annotated_filepath(name) @@ -273,7 +326,12 @@ def get_annotated_filepath(name: str, default_dir: str | None=None) -> str: else: base_dir = get_input_directory() # fallback path - return os.path.join(base_dir, name) + filepath = os.path.abspath(os.path.join(base_dir, name)) + # Prevent path traversal: the resolved path must stay within base_dir. + # repr() the name in the message so a crafted value can't inject log lines. + if not is_within_directory(base_dir, filepath): + raise ValueError("Invalid file path: {!r}".format(name)) + return filepath def exists_annotated_filepath(name) -> bool: @@ -282,7 +340,10 @@ def exists_annotated_filepath(name) -> bool: if base_dir is None: base_dir = get_input_directory() # fallback path - filepath = os.path.join(base_dir, name) + filepath = os.path.abspath(os.path.join(base_dir, name)) + # Treat traversal attempts as non-existent rather than probing the filesystem. + if not is_within_directory(base_dir, filepath): + return False return os.path.exists(filepath) diff --git a/server.py b/server.py index 361850f38..461ebe2f6 100644 --- a/server.py +++ b/server.py @@ -127,6 +127,7 @@ def create_cors_middleware(allowed_origin: str): return cors_middleware + def is_loopback(host): if host is None: return False @@ -616,15 +617,30 @@ class PromptServer(): or 'application/octet-stream' ) - # For security, force certain mimetypes to download instead of display - if content_type in {'text/html', 'text/html-sandboxed', 'application/xhtml+xml', 'text/javascript', 'text/css'}: - content_type = 'application/octet-stream' # Forces download + # For security, force renderable/active types (HTML, JS, + # CSS, SVG, XML — anything that can carry inline ' + files = {"file": ("evil.svg", svg, "image/svg+xml")} + form_data = { + "tags": json.dumps(["models", "checkpoints", "unit-tests", "svgxss"]), + "name": "evil.svg", + } + up = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120) + body = up.json() + assert up.status_code in (200, 201), body + aid = body["id"] + try: + r = http.get(f"{api_base}/api/assets/{aid}/content?disposition=inline", timeout=120) + r.content + assert r.status_code == 200 + ct = r.headers.get("Content-Type", "").lower() + cd = r.headers.get("Content-Disposition", "").lower() + assert "svg" not in ct, f"SVG served with a renderable content type: {ct!r}" + assert ct.startswith("application/octet-stream"), f"expected octet-stream, got {ct!r}" + assert "attachment" in cd, f"inline disposition not overridden to attachment: {cd!r}" + assert r.headers.get("X-Content-Type-Options", "").lower() == "nosniff" + finally: + with contextlib.suppress(Exception): + http.delete(f"{api_base}/api/assets/{aid}", timeout=30) + + def test_download_attachment_and_inline(http: requests.Session, api_base: str, seeded_asset: dict): aid = seeded_asset["id"] diff --git a/tests-unit/comfy_test/folder_path_test.py b/tests-unit/comfy_test/folder_path_test.py index 775e15c36..3b398e60b 100644 --- a/tests-unit/comfy_test/folder_path_test.py +++ b/tests-unit/comfy_test/folder_path_test.py @@ -53,8 +53,11 @@ def test_annotated_filepath(): def test_get_annotated_filepath(): default_dir = "/default/dir" - assert folder_paths.get_annotated_filepath("test.txt", default_dir) == os.path.join(default_dir, "test.txt") - assert folder_paths.get_annotated_filepath("test.txt [output]") == os.path.join(folder_paths.get_output_directory(), "test.txt") + # get_annotated_filepath now normalizes with os.path.abspath (part of the + # GHSA-779p traversal hardening), so compare against the normalized form — + # on Windows abspath also prepends the current drive letter. + assert folder_paths.get_annotated_filepath("test.txt", default_dir) == os.path.abspath(os.path.join(default_dir, "test.txt")) + assert folder_paths.get_annotated_filepath("test.txt [output]") == os.path.abspath(os.path.join(folder_paths.get_output_directory(), "test.txt")) def test_add_model_folder_path_append(clear_folder_paths): folder_paths.add_model_folder_path("test_folder", "/default/path", is_default=True) diff --git a/tests-unit/security_test/__init__.py b/tests-unit/security_test/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py b/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py new file mode 100644 index 000000000..f17fd26ea --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_02_preview_traversal.py @@ -0,0 +1,192 @@ +"""CI unit tests for FIX #2 of GHSA-779p-m5rp-r4h4. + +Path traversal / hardening in app/model_manager.py get_model_preview +(route /experiment/models/preview/{folder}/{path_index}/{filename:.*}). + +Reference: https://github.com/Comfy-Org/ComfyUI/security/advisories/GHSA-779p-m5rp-r4h4 +""" +import pytest +import yarl +from io import BytesIO +from PIL import Image +from aiohttp import web +from unittest.mock import patch +from app.model_manager import ModelFileManager + +pytestmark = ( + pytest.mark.asyncio +) # This applies the asyncio mark to all test functions in the module + +@pytest.fixture +def model_manager(): + return ModelFileManager() + +@pytest.fixture +def app(model_manager): + app = web.Application() + routes = web.RouteTableDef() + model_manager.add_routes(routes) + app.add_routes(routes) + return app + + +async def test_legit_preview_returns_200(aiohttp_client, app, tmp_path): + """Sanity: a real preview PNG inside the model folder is served as webp 200.""" + img = Image.new('RGB', (16, 16), color=(255, 0, 128)) + img.save(tmp_path / "test_model.png", format='PNG') + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/test_model.png') + + assert response.status == 200 + assert response.content_type == 'image/webp' + + img_bytes = BytesIO(await response.read()) + served = Image.open(img_bytes) + assert served.format + assert served.format.lower() == 'webp' + served.close() + + +async def test_non_integer_path_index_returns_400(aiohttp_client, app, tmp_path): + """A non-integer path_index segment must be rejected with 400.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/abc/test_model.png') + + assert response.status == 400 + + +async def test_out_of_range_path_index_returns_404(aiohttp_client, app, tmp_path): + """A path_index beyond the configured folder list must return 404.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/99/test_model.png') + + assert response.status == 404 + + +async def test_empty_filename_returns_400(aiohttp_client, app, tmp_path): + """The "{filename:.*}" capture also matches the empty string (trailing + slash). It would resolve to the folder itself and must be rejected with 400.""" + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/') + + assert response.status == 400 + + +async def test_path_traversal_in_filename_returns_403(aiohttp_client, app, tmp_path): + """Path traversal in {filename} must be rejected with 403 and must NOT read + a file outside the configured model directory. + + GOTCHA: aiohttp/yarl collapses literal ``../`` dot-segments out of the URL + path before it reaches the handler, which would make this test vacuously + pass (the request would hit a different/non-existent route). We percent-encode + the dots and slashes (``%2e%2e%2f``) and send the URL with + ``yarl.URL(..., encoded=True)`` so the bytes survive client-side normalization + untouched; aiohttp's router then percent-decodes them into ``match_info``, + delivering the literal ``../`` traversal to the handler's ``{filename:.*}`` + capture. + + Without the fix the handler computes + ``os.path.normpath(os.path.join(folder, "../../../../etc/hosts"))``, which + escapes ``tmp_path`` and would be passed straight to get_model_previews -> + Image.open, serving bytes from outside the model dir (200/served bytes). The + is_within_directory() containment check is the load-bearing fix that turns + that escape into a 403. + """ + # Sanity-anchor: a legit preview exists inside tmp_path, so a 200 path is + # genuinely reachable — proving the 403 below is the containment check + # firing, not an unrelated 404. + img = Image.new('RGB', (16, 16), color=(255, 0, 128)) + img.save(tmp_path / "test_model.png", format='PNG') + + # Percent-encoded "../../../../etc/hosts" so yarl does not collapse the + # dot-segments before the request leaves the client. + encoded_traversal = '%2e%2e%2f' * 4 + 'etc%2fhosts' + raw_path = '/experiment/models/preview/test_folder/0/' + encoded_traversal + url = yarl.URL(raw_path, encoded=True) + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get(url) + + # Confirm the traversal actually reached the handler intact: a 200 here + # would mean either normalization stripped the ``../`` (vacuous pass) or + # the containment check failed open and served outside-dir bytes. + assert response.status == 403, ( + f"expected 403 from is_within_directory() containment check, " + f"got {response.status}; traversal may have been normalized away " + f"or the fix failed open" + ) + body = await response.read() + assert body == b"", "403 response must not carry any file bytes" + + +async def test_symlink_companion_preview_returns_403(aiohttp_client, app, tmp_path): + """A companion preview file is selected by a glob inside get_model_previews + and then opened. If that companion is a symlink whose path is in-dir but + whose target escapes the model folder, it must be rejected with 403 — not + served. The requested path itself stays in-dir (so the first containment + check passes); the load-bearing fix is the SECOND is_within_directory check + on the file actually opened. + """ + model_dir = tmp_path / "models" + model_dir.mkdir() + secret_dir = tmp_path / "secret" + secret_dir.mkdir() + # A real image OUTSIDE the model dir — valid, so without the fix Image.open + # would succeed and its bytes would be served (200). + secret = secret_dir / "secret.png" + Image.new('RGB', (8, 8), color=(0, 0, 0)).save(secret, format='PNG') + # Companion preview, in-dir by name but a symlink escaping the model dir. + # (No real model file is needed — get_model_previews globs companions by + # basename, and omitting a .safetensors avoids the metadata-header read.) + companion = model_dir / "model.preview.png" + try: + companion.symlink_to(secret) + except (OSError, NotImplementedError): + pytest.skip("symlinks not supported on this platform/filesystem") + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(model_dir)], None) + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models/preview/test_folder/0/model.safetensors') + + assert response.status == 403, ( + f"expected 403 — the globbed companion preview is a symlink resolving " + f"outside the model dir and must not be served; got {response.status}" + ) + assert await response.read() == b"" + + +async def test_null_byte_in_filename_no_500(aiohttp_client, app, tmp_path): + """A NUL byte in the filename must yield a clean client rejection, not a 500 + from an uncaught ValueError in is_within_directory's realpath() call.""" + raw_path = '/experiment/models/preview/test_folder/0/' + 'a%00b' + url = yarl.URL(raw_path, encoded=True) + + with patch('folder_paths.folder_names_and_paths', { + 'test_folder': ([str(tmp_path)], None) + }): + client = await aiohttp_client(app) + response = await client.get(url) + + assert response.status != 500, ( + f"NUL byte produced a 500 (uncaught ValueError); expected a clean " + f"4xx rejection, got {response.status}" + ) + assert 400 <= response.status < 500 diff --git a/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py b/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py new file mode 100644 index 000000000..88102760c --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_03_annotated_traversal.py @@ -0,0 +1,165 @@ +"""Security tests for GHSA-779p-m5rp-r4h4 — FIX #3. + +Path traversal in folder_paths.get_annotated_filepath / exists_annotated_filepath, +plus the shared is_within_directory() containment helper. + +These are pure-function tests (no running server). The input/output/temp +directories are pointed at tmp_path via the folder_paths setters, so a crafted +name containing `../`, an absolute path, or a symlink that escapes the base +directory must be rejected. + +Reference: https://github.com/Comfy-Org/ComfyUI/security/advisories/GHSA-779p-m5rp-r4h4 +""" +import os + +import pytest + +import folder_paths +from comfy.options import enable_args_parsing +enable_args_parsing() + + +@pytest.fixture +def sandbox(tmp_path): + """Point folder_paths' input/output/temp dirs at a real temp sandbox. + + Yields the realpath'd base, input, output and temp directories. The original + directory values are restored afterward so tests stay isolated. + """ + base = os.path.realpath(str(tmp_path)) + input_dir = os.path.join(base, "input") + output_dir = os.path.join(base, "output") + temp_dir = os.path.join(base, "temp") + for d in (input_dir, output_dir, temp_dir): + os.makedirs(d, exist_ok=True) + + orig_input = folder_paths.get_input_directory() + orig_output = folder_paths.get_output_directory() + orig_temp = folder_paths.get_temp_directory() + + folder_paths.set_input_directory(input_dir) + folder_paths.set_output_directory(output_dir) + folder_paths.set_temp_directory(temp_dir) + + yield { + "base": base, + "input": input_dir, + "output": output_dir, + "temp": temp_dir, + } + + folder_paths.set_input_directory(orig_input) + folder_paths.set_output_directory(orig_output) + folder_paths.set_temp_directory(orig_temp) + + +# --------------------------------------------------------------------------- +# is_within_directory() — the shared containment helper +# --------------------------------------------------------------------------- + +def test_is_within_directory_legit_child(sandbox): + base = sandbox["input"] + child = os.path.join(base, "sub", "image.png") + assert folder_paths.is_within_directory(base, child) is True + + +def test_is_within_directory_dotdot_escape(sandbox): + base = sandbox["input"] + escape = os.path.join(base, "..", "..", "etc", "passwd") + assert folder_paths.is_within_directory(base, escape) is False + + +def test_is_within_directory_symlink_escape(sandbox): + """A symlink created INSIDE base that points OUTSIDE base must not pass. + + This is the key new hardening: is_within_directory realpath()s both operands, + so a symlink planted in the base directory can't be used to read files + elsewhere. We create a real on-disk symlink and a real secret target to + verify the check actually resolves the link. + """ + base = sandbox["input"] + + # A directory living outside the base, holding a secret file. + outside = os.path.join(sandbox["base"], "outside_secret_dir") + os.makedirs(outside, exist_ok=True) + secret = os.path.join(outside, "secret.txt") + with open(secret, "w") as f: + f.write("top secret") + + # Plant a symlink inside base that points at the outside directory. + # symlink creation can require elevated privileges / Developer Mode on + # Windows, so skip cleanly where it isn't available (same guard as the + # sibling test in test_ghsa_779p_02_preview_traversal.py). + link = os.path.join(base, "escape_link") + try: + os.symlink(outside, link) + except (OSError, NotImplementedError): + pytest.skip("symlinks not supported on this platform/filesystem") + + # Accessing the secret "through" the in-base symlink must be rejected. + target_via_link = os.path.join(link, "secret.txt") + assert folder_paths.is_within_directory(base, target_via_link) is False + + +# --------------------------------------------------------------------------- +# get_annotated_filepath() +# --------------------------------------------------------------------------- + +def test_get_annotated_filepath_legit_name(sandbox): + result = folder_paths.get_annotated_filepath("image.png") + assert result == os.path.join(sandbox["input"], "image.png") + assert folder_paths.is_within_directory(sandbox["input"], result) + + +def test_get_annotated_filepath_input_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [input]") + assert result == os.path.join(sandbox["input"], "image.png") + + +def test_get_annotated_filepath_output_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [output]") + assert result == os.path.join(sandbox["output"], "image.png") + + +def test_get_annotated_filepath_temp_annotation(sandbox): + result = folder_paths.get_annotated_filepath("image.png [temp]") + assert result == os.path.join(sandbox["temp"], "image.png") + + +def test_get_annotated_filepath_dotdot_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("../etc/passwd") + + +def test_get_annotated_filepath_dotdot_with_annotation_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("../../etc/passwd [output]") + + +def test_get_annotated_filepath_absolute_escape_raises(sandbox): + with pytest.raises(ValueError): + folder_paths.get_annotated_filepath("/etc/passwd") + + +# --------------------------------------------------------------------------- +# exists_annotated_filepath() +# --------------------------------------------------------------------------- + +def test_exists_annotated_filepath_existing_legit_file(sandbox): + real = os.path.join(sandbox["input"], "real.png") + with open(real, "w") as f: + f.write("data") + assert folder_paths.exists_annotated_filepath("real.png") is True + + +def test_exists_annotated_filepath_traversal_returns_false(sandbox): + """A traversal name must return False without raising and without probing + outside the base directory (must never reach os.path.exists for the escape). + """ + # /etc/passwd exists on POSIX; the function must still report False because + # the resolved path escapes the input directory. + assert folder_paths.exists_annotated_filepath("../../../../../../etc/passwd") is False + + +def test_exists_annotated_filepath_absolute_returns_false(sandbox): + assert folder_paths.exists_annotated_filepath("/etc/passwd") is False diff --git a/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py b/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py new file mode 100644 index 000000000..aa1250327 --- /dev/null +++ b/tests-unit/security_test/test_ghsa_779p_04_userdata_xss.py @@ -0,0 +1,147 @@ +""" +CI unit tests for FIX #4 of GHSA-779p-m5rp-r4h4. + +Stored-XSS hardening on GET /userdata/{file} in app/user_manager.py. + +User data files are arbitrary user-supplied content and must never render +inline in the app origin. The getuserdata handler: + - forces Content-Type to application/octet-stream for any type in + folder_paths.DANGEROUS_CONTENT_TYPES (text/html, image/svg+xml, + text/javascript, ...), + - sets X-Content-Type-Options: nosniff, + - sets Content-Disposition: attachment. + +These tests pre-create files in tmp_path and GET them back, asserting the +secure response headers. They mirror the aiohttp_client pattern in +tests-unit/prompt_server_test/user_manager_test.py. +""" + +import pytest +import os +from aiohttp import web +from app.user_manager import UserManager + +pytestmark = ( + pytest.mark.asyncio +) # This applies the asyncio mark to all test functions in the module + + +@pytest.fixture +def user_manager(tmp_path): + um = UserManager() + um.get_request_user_filepath = lambda req, file, **kwargs: os.path.join( + tmp_path, file + ) if file else tmp_path + return um + + +@pytest.fixture +def app(user_manager): + app = web.Application() + routes = web.RouteTableDef() + user_manager.add_routes(routes) + app.add_routes(routes) + return app + + +async def test_html_served_as_octet_stream(aiohttp_client, app, tmp_path): + (tmp_path / "evil.html").write_text( + "" + ) + + client = await aiohttp_client(app) + resp = await client.get("/userdata/evil.html") + + assert resp.status == 200 + ct = resp.headers.get("Content-Type", "") + # The load-bearing assertion: a .html file must NOT be served as text/html. + assert "text/html" not in ct.lower(), ( + f"Content-Type {ct!r} would let a browser render/execute the file (stored XSS)." + ) + assert ct == "application/octet-stream" + assert resp.headers.get("X-Content-Type-Options") == "nosniff" + assert "attachment" in resp.headers.get("Content-Disposition", "") + + +async def test_svg_served_as_octet_stream(aiohttp_client, app, tmp_path): + (tmp_path / "evil.svg").write_text( + '' + '' + '' + "" + ) + + client = await aiohttp_client(app) + resp = await client.get("/userdata/evil.svg") + + assert resp.status == 200 + ct = resp.headers.get("Content-Type", "") + # SVG can carry inline ' files = {"file": ("evil.svg", svg, "image/svg+xml")} form_data = { - "tags": json.dumps(["models", "checkpoints", "unit-tests", "svgxss"]), + "tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "svgxss"]), "name": "evil.svg", } up = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120) @@ -131,7 +131,7 @@ def test_download_chooses_existing_state_and_updates_access_time( assert t1 > t0 -@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "checkpoints"]}], indirect=True) +@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "model_type:checkpoints"]}], indirect=True) def test_download_missing_file_returns_404( http: requests.Session, api_base: str, comfy_tmp_base_dir: Path, seeded_asset: dict ): diff --git a/tests-unit/assets_test/test_list_cursor.py b/tests-unit/assets_test/test_list_cursor.py index a37019fd6..8f4cc8251 100644 --- a/tests-unit/assets_test/test_list_cursor.py +++ b/tests-unit/assets_test/test_list_cursor.py @@ -13,7 +13,7 @@ def _seed(asset_factory, make_asset_bytes, count: int, tag: str) -> list[str]: for n in names: asset_factory( n, - ["models", "checkpoints", "unit-tests", tag], + ["models", "model_type:checkpoints", "unit-tests", tag], {}, make_asset_bytes(n, size=2048), ) @@ -208,7 +208,7 @@ def test_cursor_walks_for_non_name_sorts(sort_field, http: requests.Session, api names = [] for i in range(4): n = f"cursor_{sort_field}_{i:02d}.safetensors" - asset_factory(n, ["models", "checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i)) + asset_factory(n, ["models", "model_type:checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i)) names.append(n) params = { diff --git a/tests-unit/assets_test/test_list_filter.py b/tests-unit/assets_test/test_list_filter.py index 17bbea5c6..d1cba87b3 100644 --- a/tests-unit/assets_test/test_list_filter.py +++ b/tests-unit/assets_test/test_list_filter.py @@ -11,7 +11,7 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse for n in names: asset_factory( n, - ["models", "checkpoints", "unit-tests", "paging"], + ["models", "model_type:checkpoints", "unit-tests", "paging"], {"epoch": 1}, make_asset_bytes(n, size=2048), ) @@ -45,8 +45,8 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse def test_list_assets_include_exclude_and_name_contains(http: requests.Session, api_base: str, asset_factory): - a = asset_factory("inc_a.safetensors", ["models", "checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024) - b = asset_factory("inc_b.safetensors", ["models", "checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024) + a = asset_factory("inc_a.safetensors", ["models", "model_type:checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024) + b = asset_factory("inc_b.safetensors", ["models", "model_type:checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024) r = http.get( api_base + "/api/assets", @@ -81,7 +81,7 @@ def test_list_assets_include_exclude_and_name_contains(http: requests.Session, a def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-size"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-size"] n1, n2, n3 = "sz1.safetensors", "sz2.safetensors", "sz3.safetensors" asset_factory(n1, t, {}, make_asset_bytes(n1, 1024)) asset_factory(n2, t, {}, make_asset_bytes(n2, 2048)) @@ -108,7 +108,7 @@ def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, mak def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-upd"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-upd"] a1 = asset_factory("upd_a.safetensors", t, {}, make_asset_bytes("upd_a", 1200)) a2 = asset_factory("upd_b.safetensors", t, {}, make_asset_bytes("upd_b", 1200)) @@ -131,7 +131,7 @@ def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-access"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-access"] asset_factory("acc_a.safetensors", t, {}, make_asset_bytes("acc_a", 1100)) time.sleep(0.02) a2 = asset_factory("acc_b.safetensors", t, {}, make_asset_bytes("acc_b", 1100)) @@ -154,14 +154,14 @@ def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-include"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-include"] a = asset_factory("incvar_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("iva")) asset_factory("incvar_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("ivb")) - # CSV + case-insensitive + # CSV tag filters are whitespace-trimmed and case-sensitive. r1 = http.get( api_base + "/api/assets", - params={"include_tags": "UNIT-TESTS,LF-INCLUDE,alpha"}, + params={"include_tags": "unit-tests,lf-include,alpha"}, timeout=120, ) b1 = r1.json() @@ -196,14 +196,14 @@ def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factor def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-exclude"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-exclude"] a = asset_factory("ex_a_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("exa", 900)) asset_factory("ex_b_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("exb", 900)) - # Exclude uppercase should work + # Exclude filters are case-sensitive. r1 = http.get( api_base + "/api/assets", - params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "BETA"}, + params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "beta"}, timeout=120, ) b1 = r1.json() @@ -225,7 +225,7 @@ def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory, def test_list_assets_name_contains_case_and_specials(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-name"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-name"] a1 = asset_factory("CaseMix.SAFE", t, {}, make_asset_bytes("cm", 800)) a2 = asset_factory("case-other.safetensors", t, {}, make_asset_bytes("co", 800)) @@ -261,7 +261,7 @@ def test_list_assets_name_contains_case_and_specials(http, api_base, asset_facto def test_list_assets_offset_beyond_total_and_limit_boundary(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "lf-pagelimits"] + t = ["models", "model_type:checkpoints", "unit-tests", "lf-pagelimits"] asset_factory("pl1.safetensors", t, {}, make_asset_bytes("pl1", 600)) asset_factory("pl2.safetensors", t, {}, make_asset_bytes("pl2", 600)) asset_factory("pl3.safetensors", t, {}, make_asset_bytes("pl3", 600)) @@ -319,7 +319,7 @@ def test_list_assets_name_contains_literal_underscore( - foobar.safetensors (must NOT match) """ scope = f"lf-underscore-{uuid.uuid4().hex[:6]}" - tags = ["models", "checkpoints", "unit-tests", scope] + tags = ["models", "model_type:checkpoints", "unit-tests", scope] a = asset_factory("foo_bar.safetensors", tags, {}, make_asset_bytes("a", 700)) b = asset_factory("fooxbar.safetensors", tags, {}, make_asset_bytes("b", 700)) diff --git a/tests-unit/assets_test/test_metadata_filters.py b/tests-unit/assets_test/test_metadata_filters.py index 20285a3b3..1864b1eef 100644 --- a/tests-unit/assets_test/test_metadata_filters.py +++ b/tests-unit/assets_test/test_metadata_filters.py @@ -5,7 +5,7 @@ def test_meta_and_across_keys_and_types( http, api_base: str, asset_factory, make_asset_bytes ): name = "mf_and_mix.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-and"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-and"] meta = {"purpose": "mix", "epoch": 1, "active": True, "score": 1.23} asset_factory(name, tags, meta, make_asset_bytes(name, 4096)) @@ -41,7 +41,7 @@ def test_meta_and_across_keys_and_types( def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory, make_asset_bytes): name = "mf_types.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-types"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-types"] meta = {"epoch": 1, "active": True} asset_factory(name, tags, meta, make_asset_bytes(name)) @@ -95,7 +95,7 @@ def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory, def test_meta_any_of_list_of_scalars(http, api_base, asset_factory, make_asset_bytes): name = "mf_list_scalars.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-list"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-list"] meta = {"flags": ["red", "green"]} asset_factory(name, tags, meta, make_asset_bytes(name, 3000)) @@ -134,7 +134,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none( http, api_base, asset_factory, make_asset_bytes ): # a1: key missing; a2: explicit null; a3: concrete value - t = ["models", "checkpoints", "unit-tests", "mf-none"] + t = ["models", "model_type:checkpoints", "unit-tests", "mf-none"] a1 = asset_factory("mf_none_missing.safetensors", t, {"x": 1}, make_asset_bytes("a1")) a2 = asset_factory("mf_none_null.safetensors", t, {"maybe": None}, make_asset_bytes("a2")) a3 = asset_factory("mf_none_value.safetensors", t, {"maybe": "x"}, make_asset_bytes("a3")) @@ -166,7 +166,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none( def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_asset_bytes): name = "mf_nested_json.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-nested"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-nested"] cfg = {"optimizer": "adam", "lr": 0.001, "schedule": {"type": "cosine", "warmup": 100}} asset_factory(name, tags, {"config": cfg}, make_asset_bytes(name, 2200)) @@ -197,7 +197,7 @@ def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_as def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_bytes): name = "mf_list_objects.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-objlist"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-objlist"] transforms = [{"type": "crop", "size": 128}, {"type": "flip", "p": 0.5}] asset_factory(name, tags, {"transforms": transforms}, make_asset_bytes(name, 2048)) @@ -228,7 +228,7 @@ def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_b def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_asset_bytes): name = "mf_keys_unicode.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-keys"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-keys"] meta = { "weird.key": "v1", "path/like": 7, @@ -259,7 +259,7 @@ def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_ def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_asset_bytes): - t = ["models", "checkpoints", "unit-tests", "mf-zero-bool"] + t = ["models", "model_type:checkpoints", "unit-tests", "mf-zero-bool"] a0 = asset_factory("mf_zero_count.safetensors", t, {"count": 0}, make_asset_bytes("z", 1025)) a1 = asset_factory("mf_bool_list.safetensors", t, {"choices": [True, False]}, make_asset_bytes("b", 1026)) @@ -286,7 +286,7 @@ def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_as def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, make_asset_bytes): name = "mf_mixed_list.safetensors" - tags = ["models", "checkpoints", "unit-tests", "mf-mixed"] + tags = ["models", "model_type:checkpoints", "unit-tests", "mf-mixed"] meta = {"mix": ["1", 1, True, None]} asset_factory(name, tags, meta, make_asset_bytes(name, 1999)) @@ -311,7 +311,7 @@ def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, mak def test_meta_unknown_key_and_none_behavior_with_scope_tags(http, api_base, asset_factory, make_asset_bytes): # Use a unique scope tag to avoid interference - t = ["models", "checkpoints", "unit-tests", "mf-unknown-scope"] + t = ["models", "model_type:checkpoints", "unit-tests", "mf-unknown-scope"] x = asset_factory("mf_unknown_a.safetensors", t, {"k1": 1}, make_asset_bytes("ua")) y = asset_factory("mf_unknown_b.safetensors", t, {"k2": 2}, make_asset_bytes("ub")) @@ -340,13 +340,13 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_ # alpha matches epoch=1; beta has epoch=2 a = asset_factory( "mf_tag_alpha.safetensors", - ["models", "checkpoints", "unit-tests", "mf-tag", "alpha"], + ["models", "model_type:checkpoints", "unit-tests", "mf-tag", "alpha"], {"epoch": 1}, make_asset_bytes("alpha"), ) b = asset_factory( "mf_tag_beta.safetensors", - ["models", "checkpoints", "unit-tests", "mf-tag", "beta"], + ["models", "model_type:checkpoints", "unit-tests", "mf-tag", "beta"], {"epoch": 2}, make_asset_bytes("beta"), ) @@ -367,7 +367,7 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_ def test_meta_sort_and_paging_under_filter(http, api_base, asset_factory, make_asset_bytes): # Three assets in same scope with different sizes and a common filter key - t = ["models", "checkpoints", "unit-tests", "mf-sort"] + t = ["models", "model_type:checkpoints", "unit-tests", "mf-sort"] n1, n2, n3 = "mf_sort_1.safetensors", "mf_sort_2.safetensors", "mf_sort_3.safetensors" asset_factory(n1, t, {"group": "g"}, make_asset_bytes(n1, 1024)) asset_factory(n2, t, {"group": "g"}, make_asset_bytes(n2, 2048)) diff --git a/tests-unit/assets_test/test_prune_orphaned_assets.py b/tests-unit/assets_test/test_prune_orphaned_assets.py index 1fbd4d4e2..618ec6c8d 100644 --- a/tests-unit/assets_test/test_prune_orphaned_assets.py +++ b/tests-unit/assets_test/test_prune_orphaned_assets.py @@ -29,7 +29,7 @@ def create_seed_file(comfy_tmp_base_dir: Path): def find_asset(http: requests.Session, api_base: str): """Query API for assets matching scope and optional name.""" def _find(scope: str, name: str | None = None) -> list[dict]: - params = {"include_tags": f"unit-tests,{scope}"} + params = {"limit": "500"} if name: params["name_contains"] = name r = http.get(f"{api_base}/api/assets", params=params, timeout=120) @@ -91,7 +91,7 @@ def test_hashed_asset_not_pruned_when_file_missing( data = make_asset_bytes("test", 2048) a = asset_factory("test.bin", ["input", "unit-tests", scope], {}, data) - path = comfy_tmp_base_dir / "input" / "unit-tests" / scope / get_asset_filename(a["asset_hash"], ".bin") + path = comfy_tmp_base_dir / "input" / get_asset_filename(a["asset_hash"], ".bin") path.unlink() trigger_sync_seed_assets(http, api_base) @@ -108,18 +108,20 @@ def test_prune_across_multiple_roots( ): """Prune correctly handles assets across input and output roots.""" scope = f"multi-{uuid.uuid4().hex[:6]}" - input_fp = create_seed_file("input", scope, "input.bin") - create_seed_file("output", scope, "output.bin") + input_name = f"{scope}-input.bin" + output_name = f"{scope}-output.bin" + input_fp = create_seed_file("input", scope, input_name) + create_seed_file("output", scope, output_name) trigger_sync_seed_assets(http, api_base) - assert len(find_asset(scope)) == 2 + assert find_asset(scope, input_name) + assert find_asset(scope, output_name) input_fp.unlink() trigger_sync_seed_assets(http, api_base) - remaining = find_asset(scope) - assert len(remaining) == 1 - assert remaining[0]["name"] == "output.bin" + assert not find_asset(scope, input_name) + assert find_asset(scope, output_name) @pytest.mark.parametrize("dirname", ["100%_done", "my_folder_name", "has spaces"]) diff --git a/tests-unit/assets_test/test_tags_api.py b/tests-unit/assets_test/test_tags_api.py index 9729b7d03..93786696f 100644 --- a/tests-unit/assets_test/test_tags_api.py +++ b/tests-unit/assets_test/test_tags_api.py @@ -10,9 +10,9 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict) body1 = r1.json() assert r1.status_code == 200 names = [t["name"] for t in body1["tags"]] - # A few system tags from migration should exist: + # A few selected contract tags should exist. assert "models" in names - assert "checkpoints" in names + assert "model_type:checkpoints" in names # Only used tags before we add anything new from this test cycle r2 = http.get(api_base + "/api/tags", params={"include_zero": "false"}, timeout=120) @@ -21,7 +21,7 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict) # We already seeded one asset via fixture, so used tags must be non-empty used_names = [t["name"] for t in body2["tags"]] assert "models" in used_names - assert "checkpoints" in used_names + assert "model_type:checkpoints" in used_names # Prefix filter should refine the list r3 = http.get(api_base + "/api/tags", params={"include_zero": "false", "prefix": "uni"}, timeout=120) @@ -45,7 +45,7 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory, body1 = r1.json() assert r1.status_code == 200 names = [t["name"] for t in body1["tags"]] - assert "models" in names and "checkpoints" in names + assert "models" in names and "model_type:checkpoints" in names # Create a short-lived asset under input with a unique custom tag scope = f"tags-empty-usage-{uuid.uuid4().hex[:6]}" @@ -89,28 +89,28 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory, def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset: dict): aid = seeded_asset["id"] - # Add tags with duplicates and mixed case - payload_add = {"tags": ["NewTag", "unit-tests", "newtag", "BETA"]} + # Add tags with duplicates while preserving source case. + payload_add = {"tags": ["NewTag", "unit-tests", "NewTag", "BETA"]} r1 = http.post(f"{api_base}/api/assets/{aid}/tags", json=payload_add, timeout=120) b1 = r1.json() assert r1.status_code == 200, b1 - # normalized, deduplicated; 'unit-tests' was already present from the seed - assert set(b1["added"]) == {"newtag", "beta"} + # stripped, deduplicated; 'unit-tests' was already present from the seed + assert set(b1["added"]) == {"NewTag", "BETA"} assert set(b1["already_present"]) == {"unit-tests"} - assert "newtag" in b1["total_tags"] and "beta" in b1["total_tags"] + assert "NewTag" in b1["total_tags"] and "BETA" in b1["total_tags"] rg = http.get(f"{api_base}/api/assets/{aid}", timeout=120) g = rg.json() assert rg.status_code == 200 tags_now = set(g["tags"]) - assert {"newtag", "beta"}.issubset(tags_now) + assert {"NewTag", "BETA"}.issubset(tags_now) # Remove a tag and a non-existent tag - payload_del = {"tags": ["newtag", "does-not-exist"]} + payload_del = {"tags": ["NewTag", "does-not-exist"]} r2 = http.delete(f"{api_base}/api/assets/{aid}/tags", json=payload_del, timeout=120) b2 = r2.json() assert r2.status_code == 200 - assert set(b2["removed"]) == {"newtag"} + assert set(b2["removed"]) == {"NewTag"} assert set(b2["not_present"]) == {"does-not-exist"} # Verify remaining tags after deletion @@ -118,8 +118,44 @@ def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset g2 = rg2.json() assert rg2.status_code == 200 tags_later = set(g2["tags"]) - assert "newtag" not in tags_later - assert "beta" in tags_later # still present + assert "NewTag" not in tags_later + assert "BETA" in tags_later # still present + + +def test_add_system_looking_tags_allowed_as_labels( + http: requests.Session, api_base: str, seeded_asset: dict +): + aid = seeded_asset["id"] + + response = http.post( + f"{api_base}/api/assets/{aid}/tags", + json={ + "tags": [ + "models", + "model_type:manual", + "model:true", + "models:foo", + "input:true", + "output:true", + "uploaded:true", + "temp:true", + "temporary", + ] + }, + timeout=120, + ) + body = response.json() + + assert response.status_code == 200, body + assert "models" in body["total_tags"] + assert "model_type:manual" in body["total_tags"] + assert "model:true" in body["total_tags"] + assert "models:foo" in body["total_tags"] + assert "input:true" in body["total_tags"] + assert "output:true" in body["total_tags"] + assert "uploaded:true" in body["total_tags"] + assert "temp:true" in body["total_tags"] + assert "temporary" in body["total_tags"] def test_tags_list_order_and_prefix(http: requests.Session, api_base: str, seeded_asset: dict): diff --git a/tests-unit/assets_test/test_uploads.py b/tests-unit/assets_test/test_uploads.py index 427a417cc..7be7b0935 100644 --- a/tests-unit/assets_test/test_uploads.py +++ b/tests-unit/assets_test/test_uploads.py @@ -1,11 +1,14 @@ import json import uuid from concurrent.futures import ThreadPoolExecutor +from pathlib import Path import requests import pytest +from app.assets.api.schemas_in import UploadAssetSpec from app.assets.api.schemas_out import Asset, AssetCreated +from helpers import get_asset_filename def test_asset_created_inherits_hash_field(): @@ -20,9 +23,18 @@ def test_asset_created_inherits_hash_field(): assert AssetCreated.model_fields["hash"].annotation == Asset.model_fields["hash"].annotation +def test_upload_asset_spec_ignores_subfolder_field(): + spec = UploadAssetSpec.model_validate( + {"tags": ["input"], "subfolder": "pasted", "name": "image.png"} + ) + + assert "subfolder" not in UploadAssetSpec.model_fields + assert not hasattr(spec, "subfolder") + + def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, make_asset_bytes): name = "dup_a.safetensors" - tags = ["models", "checkpoints", "unit-tests", "alpha"] + tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"] meta = {"purpose": "dup"} data = make_asset_bytes(name) files = {"file": (name, data, "application/octet-stream")} @@ -43,6 +55,8 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma assert a2["asset_hash"] == a1["asset_hash"] assert a2["hash"] == a1["hash"] assert a2["id"] != a1["id"] # new reference with same content + assert a2.get("loader_path") is None + assert a2.get("display_name") is None # Third upload with the same data but different name also creates new AssetReference files = {"file": (name, data, "application/octet-stream")} @@ -53,12 +67,14 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma assert a3["asset_hash"] == a1["asset_hash"] assert a3["id"] != a1["id"] assert a3["id"] != a2["id"] + assert a3.get("loader_path") is None + assert a3.get("display_name") is None def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_base: str): # Seed a small file first name = "fastpath_seed.safetensors" - tags = ["models", "checkpoints", "unit-tests"] + tags = ["input", "unit-tests"] meta = {} files = {"file": (name, b"B" * 1024, "application/octet-stream")} form = {"tags": json.dumps(tags), "name": name, "user_metadata": json.dumps(meta)} @@ -69,9 +85,10 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_ assert b1["hash"] == h # Now POST /api/assets with only hash and no file + hash_only_tags = ["models", "checkpoints", "unit-tests", "hash-labels"] files = [ ("hash", (None, h)), - ("tags", (None, json.dumps(tags))), + ("tags", (None, json.dumps(hash_only_tags))), ("name", (None, "fastpath_copy.safetensors")), ("user_metadata", (None, json.dumps({"purpose": "copy"}))), ] @@ -81,6 +98,53 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_ assert b2["created_new"] is False assert b2["asset_hash"] == h assert b2["hash"] == h + assert "models" in b2["tags"] + assert "checkpoints" in b2["tags"] + assert "uploaded" not in b2["tags"] + assert not any(tag.startswith("model_type:") for tag in b2["tags"]) + assert b2.get("loader_path") is None + assert b2.get("display_name") is None + + rg = http.get(f"{api_base}/api/assets/{b2['id']}", timeout=120) + detail = rg.json() + assert rg.status_code == 200, detail + assert detail.get("loader_path") is None + assert detail.get("display_name") is None + + +def test_create_from_hash_with_model_tags_does_not_synthesize_loader_path( + http: requests.Session, api_base: str +): + seed_name = "from_hash_seed.safetensors" + seed_tags = ["models", "model_type:checkpoints", "unit-tests"] + files = {"file": (seed_name, b"D" * 1024, "application/octet-stream")} + form = { + "tags": json.dumps(seed_tags), + "name": seed_name, + "user_metadata": json.dumps({}), + } + seed_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + seed = seed_r.json() + assert seed_r.status_code == 201, seed + + payload = { + "hash": seed["asset_hash"], + "name": "from_hash_copy.safetensors", + "tags": ["models", "model_type:checkpoints", "unit-tests", "spoofed"], + } + created_r = http.post(api_base + "/api/assets/from-hash", json=payload, timeout=120) + created = created_r.json() + assert created_r.status_code == 201, created + assert created["created_new"] is False + assert created["asset_hash"] == seed["asset_hash"] + assert created.get("loader_path") is None + assert created.get("display_name") is None + + detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120) + detail = detail_r.json() + assert detail_r.status_code == 200, detail + assert detail.get("loader_path") is None + assert detail.get("display_name") is None def test_upload_fastpath_with_known_hash_and_file( @@ -88,7 +152,7 @@ def test_upload_fastpath_with_known_hash_and_file( ): # Seed files = {"file": ("seed.safetensors", b"C" * 128, "application/octet-stream")} - form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})} + form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})} r1 = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) b1 = r1.json() assert r1.status_code == 201, b1 @@ -104,11 +168,49 @@ def test_upload_fastpath_with_known_hash_and_file( assert b2["created_new"] is False assert b2["asset_hash"] == h assert b2["hash"] == h + assert "checkpoints" in b2["tags"] + assert "uploaded" not in b2["tags"] + assert not any(tag == "model_type:checkpoints" for tag in b2["tags"]) + + +def test_duplicate_byte_upload_is_reference_only_and_does_not_need_destination( + http: requests.Session, api_base: str +): + data = b"duplicate-reference-only" * 64 + seed_files = {"file": ("duplicate-seed.bin", data, "application/octet-stream")} + seed_form = { + "tags": json.dumps(["input", "unit-tests", "duplicate-seed"]), + "name": "duplicate-seed.bin", + "user_metadata": json.dumps({}), + } + seed_response = http.post(api_base + "/api/assets", data=seed_form, files=seed_files, timeout=120) + seed = seed_response.json() + assert seed_response.status_code == 201, seed + + duplicate_files = {"file": ("duplicate-copy.bin", data, "application/octet-stream")} + duplicate_form = { + "tags": json.dumps(["not-a-destination", "unit-tests", "duplicate-copy"]), + "name": "duplicate-copy.bin", + "user_metadata": json.dumps({}), + } + duplicate_response = http.post( + api_base + "/api/assets", data=duplicate_form, files=duplicate_files, timeout=120 + ) + duplicate = duplicate_response.json() + + assert duplicate_response.status_code == 200, duplicate + assert duplicate["created_new"] is False + assert duplicate["asset_hash"] == seed["asset_hash"] + assert "not-a-destination" in duplicate["tags"] + assert "uploaded" not in duplicate["tags"] + assert "input" not in duplicate["tags"] + assert duplicate.get("loader_path") is None + assert duplicate.get("display_name") is None def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base: str): data = [ - ("tags", "models,checkpoints"), + ("tags", "models,model_type:checkpoints"), ("tags", json.dumps(["unit-tests", "alpha"])), ("name", "merge.safetensors"), ("user_metadata", json.dumps({"u": 1})), @@ -124,7 +226,71 @@ def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base detail = rg.json() assert rg.status_code == 200, detail tags = set(detail["tags"]) - assert {"models", "checkpoints", "unit-tests", "alpha"}.issubset(tags) + assert {"models", "model_type:checkpoints", "unit-tests", "alpha"}.issubset(tags) + + +@pytest.mark.parametrize( + ( + "tags", + "extension", + "expected_display_prefix", + ), + [ + (["input", "unit-tests"], ".png", ""), + ( + ["models", "model_type:checkpoints", "unit-tests"], + ".safetensors", + "checkpoints/", + ), + ], +) +def test_upload_response_includes_loader_path_and_display_name( + tags: list[str], + extension: str, + expected_display_prefix: str, + http: requests.Session, + api_base: str, + make_asset_bytes, +): + scope = f"response-paths-{uuid.uuid4().hex[:6]}" + scoped_tags = [*tags, scope] + name = f"asset_response_path{extension}" + + files = {"file": (name, make_asset_bytes(name, 1024), "application/octet-stream")} + form = { + "tags": json.dumps(scoped_tags), + "name": name, + "user_metadata": json.dumps({}), + } + created_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + created = created_r.json() + assert created_r.status_code in (200, 201), created + stored_filename = get_asset_filename(created["asset_hash"], extension) + expected_suffix = stored_filename + expected_display_name = f"{expected_display_prefix}{expected_suffix}" + # In-root loader path: model category dropped, no subfolders here -> just the filename. + expected_loader_path = expected_suffix + + assert created["loader_path"] == expected_loader_path + assert created["display_name"] == expected_display_name + assert "logical_path" not in created + + detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120) + detail = detail_r.json() + assert detail_r.status_code == 200, detail + assert detail["loader_path"] == expected_loader_path + assert detail["display_name"] == expected_display_name + + list_r = http.get( + api_base + "/api/assets", + params={"include_tags": f"unit-tests,{scope}", "limit": "50"}, + timeout=120, + ) + listed = list_r.json() + assert list_r.status_code == 200, listed + match = next(a for a in listed["assets"] if a["id"] == created["id"]) + assert match["loader_path"] == expected_loader_path + assert match["display_name"] == expected_display_name @pytest.mark.parametrize("root", ["input", "output"]) @@ -192,16 +358,55 @@ def test_create_from_hash_endpoint_404(http: requests.Session, api_base: str): assert body["error"]["code"] == "ASSET_NOT_FOUND" +def test_create_from_hash_accepts_arbitrary_system_looking_tags( + http: requests.Session, api_base: str +): + files = {"file": ("hash-seed.bin", b"hash-seed" * 64, "application/octet-stream")} + form = { + "tags": json.dumps(["input", "unit-tests", "hash-seed"]), + "name": "hash-seed.bin", + "user_metadata": json.dumps({}), + } + seed_response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + seed = seed_response.json() + assert seed_response.status_code == 201, seed + + response = http.post( + api_base + "/api/assets/from-hash", + json={ + "hash": seed["asset_hash"], + "name": "hash-copy.bin", + "tags": [ + "models", + "model:true", + "models:foo", + "temporary:true", + "unit-tests", + "hash-copy", + ], + }, + timeout=120, + ) + body = response.json() + + assert response.status_code == 201, body + assert "models" in body["tags"] + assert "model:true" in body["tags"] + assert "models:foo" in body["tags"] + assert "temporary:true" in body["tags"] + assert "uploaded" not in body["tags"] + + def test_upload_zero_byte_rejected(http: requests.Session, api_base: str): files = {"file": ("empty.safetensors", b"", "application/octet-stream")} - form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})} + form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) body = r.json() assert r.status_code == 400 assert body["error"]["code"] == "EMPTY_UPLOAD" -def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str): +def test_upload_rejects_arbitrary_labels_without_required_destination_role(http: requests.Session, api_base: str): files = {"file": ("badroot.bin", b"A" * 64, "application/octet-stream")} form = {"tags": json.dumps(["not-a-root", "whatever"]), "name": "badroot.bin", "user_metadata": json.dumps({})} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) @@ -212,7 +417,7 @@ def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str) def test_upload_user_metadata_must_be_json(http: requests.Session, api_base: str): files = {"file": ("badmeta.bin", b"A" * 128, "application/octet-stream")} - form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"} + form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) body = r.json() assert r.status_code == 400 @@ -228,7 +433,7 @@ def test_upload_requires_multipart(http: requests.Session, api_base: str): def test_upload_missing_file_and_hash(http: requests.Session, api_base: str): files = [ - ("tags", (None, json.dumps(["models", "checkpoints", "unit-tests"]))), + ("tags", (None, json.dumps(["models", "model_type:checkpoints", "unit-tests"]))), ("name", (None, "x.safetensors")), ] r = http.post(api_base + "/api/assets", files=files, timeout=120) @@ -237,17 +442,33 @@ def test_upload_missing_file_and_hash(http: requests.Session, api_base: str): assert body["error"]["code"] == "MISSING_FILE" -def test_upload_models_unknown_category(http: requests.Session, api_base: str): +def test_upload_models_unknown_model_type(http: requests.Session, api_base: str): files = {"file": ("m.safetensors", b"A" * 128, "application/octet-stream")} - form = {"tags": json.dumps(["models", "no_such_category", "unit-tests"]), "name": "m.safetensors"} + form = {"tags": json.dumps(["models", "model_type:no_such_category", "unit-tests"]), "name": "m.safetensors"} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) body = r.json() - assert r.status_code == 400 + assert r.status_code == 400, body assert body["error"]["code"] == "INVALID_BODY" - assert body["error"]["message"].startswith("unknown models category") -def test_upload_models_requires_category(http: requests.Session, api_base: str): +@pytest.mark.parametrize("model_type", ["configs", "custom_nodes"]) +def test_upload_models_rejects_non_model_registered_folder( + model_type: str, http: requests.Session, api_base: str +): + files = {"file": ("not-a-model.py", b"A" * 128, "application/octet-stream")} + form = { + "tags": json.dumps(["models", f"model_type:{model_type}", "unit-tests"]), + "name": "not-a-model.py", + } + + response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 400, body + assert body["error"]["code"] == "INVALID_BODY" + + +def test_upload_models_requires_model_type(http: requests.Session, api_base: str): files = {"file": ("nocat.safetensors", b"A" * 64, "application/octet-stream")} form = {"tags": json.dumps(["models"]), "name": "nocat.safetensors", "user_metadata": json.dumps({})} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) @@ -256,13 +477,152 @@ def test_upload_models_requires_category(http: requests.Session, api_base: str): assert body["error"]["code"] == "INVALID_BODY" -def test_upload_tags_traversal_guard(http: requests.Session, api_base: str): +def test_upload_extra_tags_are_labels_not_path_components(http: requests.Session, api_base: str): files = {"file": ("evil.safetensors", b"A" * 256, "application/octet-stream")} - form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"} + form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"} r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) body = r.json() - assert r.status_code == 400 - assert body["error"]["code"] in ("BAD_REQUEST", "INVALID_BODY") + assert r.status_code == 201, body + assert ".." in body["tags"] + assert "zzz" in body["tags"] + assert "models" in body["tags"] + assert "model_type:checkpoints" in body["tags"] + + +@pytest.mark.parametrize( + ("subfolder", "expected_tag", "unexpected_tags"), + [ + ("custom/session", None, {"custom", "session"}), + ("pasted", "pasted", set()), + ], +) +def test_upload_image_accepts_arbitrary_subfolder_but_only_known_values_become_tags( + http: requests.Session, + api_base: str, + comfy_tmp_base_dir: Path, + subfolder: str, + expected_tag: str | None, + unexpected_tags: set[str], +): + name = f"upload-image-{uuid.uuid4().hex}.png" + files = {"image": (name, b"image-upload" * 64, "image/png")} + form = {"type": "input", "subfolder": subfolder} + + response = http.post(api_base + "/upload/image", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 200, body + assert body["subfolder"] == subfolder + assert (comfy_tmp_base_dir / "input" / subfolder / body["name"]).exists() + + asset = body["asset"] + tags = set(asset["tags"]) + assert "input" in tags + assert "uploaded" in tags + if expected_tag: + assert expected_tag in tags + assert tags.isdisjoint(unexpected_tags) + + +def test_multipart_upload_accepts_system_looking_extra_labels( + http: requests.Session, api_base: str +): + files = {"file": ("relaxed-labels.bin", b"relaxed" * 64, "application/octet-stream")} + form = { + "tags": json.dumps( + [ + "input", + "unit-tests", + "model:true", + "models:foo", + "temporary", + "uploaded:true", + ] + ), + "name": "relaxed-labels.bin", + "user_metadata": json.dumps({}), + } + response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 201, body + assert "input" in body["tags"] + assert "model:true" in body["tags"] + assert "models:foo" in body["tags"] + assert "temporary" in body["tags"] + assert "uploaded:true" in body["tags"] + + +def test_multipart_upload_rejects_ambiguous_destination_roles( + http: requests.Session, api_base: str +): + files = {"file": ("ambiguous.bin", b"ambiguous" * 64, "application/octet-stream")} + form = { + "tags": json.dumps(["input", "output", "unit-tests"]), + "name": "ambiguous.bin", + "user_metadata": json.dumps({}), + } + response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 400, body + assert body["error"]["code"] == "INVALID_BODY" + + +def test_multipart_upload_rejects_multiple_model_types_for_models_destination( + http: requests.Session, api_base: str +): + files = {"file": ("ambiguous-model.safetensors", b"ambiguous-model" * 64, "application/octet-stream")} + form = { + "tags": json.dumps( + ["models", "model_type:checkpoints", "model_type:loras", "unit-tests"] + ), + "name": "ambiguous-model.safetensors", + "user_metadata": json.dumps({}), + } + response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 400, body + assert body["error"]["code"] == "INVALID_BODY" + + +@pytest.mark.parametrize( + ("tags", "expected_root", "extension"), + [ + (["input", "unit-tests", "upload-location-input"], "input", ".bin"), + (["output", "unit-tests", "upload-location-output"], "output", ".bin"), + ( + ["models", "model_type:checkpoints", "unit-tests", "upload-location-model"], + "models/checkpoints", + ".safetensors", + ), + ], +) +def test_multipart_upload_role_selects_write_location( + http: requests.Session, + api_base: str, + comfy_tmp_base_dir: Path, + tags: list[str], + expected_root: str, + extension: str, +): + role = next(tag for tag in tags if tag in {"input", "models", "output"}) + name = f"{role}-role-upload{extension}" + files = {"file": (name, f"{role}-role-bytes".encode() * 64, "application/octet-stream")} + form = { + "tags": json.dumps(tags), + "name": name, + "user_metadata": json.dumps({}), + } + + response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120) + body = response.json() + + assert response.status_code == 201, body + stored_name = get_asset_filename(body["asset_hash"], extension) + expected_disk_path = comfy_tmp_base_dir / expected_root / stored_name + assert expected_disk_path.exists() def test_upload_empty_tags_rejected(http: requests.Session, api_base: str): diff --git a/tests-unit/feature_flags_test.py b/tests-unit/feature_flags_test.py index 8ec52a124..a436ab1ec 100644 --- a/tests-unit/feature_flags_test.py +++ b/tests-unit/feature_flags_test.py @@ -29,6 +29,8 @@ class TestFeatureFlags: features = get_server_features() assert "supports_preview_metadata" in features assert features["supports_preview_metadata"] is True + assert "supports_model_type_tags" in features + assert features["supports_model_type_tags"] is True assert "max_upload_size" in features assert isinstance(features["max_upload_size"], (int, float)) diff --git a/tests-unit/websocket_feature_flags_test.py b/tests-unit/websocket_feature_flags_test.py index e93b2e1dd..4950bd9d0 100644 --- a/tests-unit/websocket_feature_flags_test.py +++ b/tests-unit/websocket_feature_flags_test.py @@ -12,6 +12,8 @@ class TestWebSocketFeatureFlags: # Check expected server features assert "supports_preview_metadata" in features assert features["supports_preview_metadata"] is True + assert "supports_model_type_tags" in features + assert features["supports_model_type_tags"] is True assert "max_upload_size" in features assert isinstance(features["max_upload_size"], (int, float)) @@ -75,3 +77,5 @@ class TestWebSocketFeatureFlags: assert server_message["type"] == "feature_flags" assert "supports_preview_metadata" in server_message["data"] assert server_message["data"]["supports_preview_metadata"] is True + assert "supports_model_type_tags" in server_message["data"] + assert server_message["data"]["supports_model_type_tags"] is True From d0008a8958b9170b7eb1e4d5e45bff95249f36ee Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 8 Jul 2026 22:50:25 -0700 Subject: [PATCH 073/211] Fix qwen3vl reference images when used as a text encode models. (#14845) Should not affect use as a text generation model. --- comfy/text_encoders/qwen3vl.py | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/comfy/text_encoders/qwen3vl.py b/comfy/text_encoders/qwen3vl.py index 2082c42e7..7a329d2d6 100644 --- a/comfy/text_encoders/qwen3vl.py +++ b/comfy/text_encoders/qwen3vl.py @@ -90,6 +90,27 @@ class Qwen3VL(BaseLlama, BaseQwen3, BaseGenerate, torch.nn.Module): deepstack = [torch.cat([deepstack[i], ds[i]], dim=0) for i in range(len(ds))] return position_ids, visual_pos_masks, deepstack + def forward(self, input_ids, attention_mask=None, embeds=None, num_tokens=None, intermediate_output=None, final_layer_norm_intermediate=True, dtype=None, embeds_info=[], **kwargs): + position_ids = kwargs.pop("position_ids", None) + visual_pos_masks = kwargs.pop("visual_pos_masks", None) + deepstack_embeds = kwargs.pop("deepstack_embeds", None) + if embeds is not None and position_ids is None: + position_ids, visual_pos_masks, deepstack_embeds = self.build_image_inputs(embeds, embeds_info) + return self.model( + input_ids, + attention_mask=attention_mask, + embeds=embeds, + num_tokens=num_tokens, + intermediate_output=intermediate_output, + final_layer_norm_intermediate=final_layer_norm_intermediate, + dtype=dtype, + position_ids=position_ids, + embeds_info=embeds_info, + visual_pos_masks=visual_pos_masks, + deepstack_embeds=deepstack_embeds, + **kwargs, + ) + def _make_qwen3vl_model(model_type): class Qwen3VL_(Qwen3VL): From b35819712e9e3c1618bb0fa3d9a93c8e3a93d2fb Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 9 Jul 2026 09:20:10 +0300 Subject: [PATCH 074/211] feat: allow --comfy-api-base target ephemeral testenvs (#14569) * feat: allow --comfy-api-base target ephemeral testenvs Signed-off-by: bigcat88 * refactor: name /features data as backend flags, not frontend --------- Signed-off-by: bigcat88 Co-authored-by: guill --- comfy/comfy_api_env.py | 46 ++++++++++++++++++++++ comfy_api_nodes/util/_helpers.py | 3 +- server.py | 7 +++- tests-unit/feature_flags_test.py | 67 ++++++++++++++++++++++++++++++++ 4 files changed, 121 insertions(+), 2 deletions(-) create mode 100644 comfy/comfy_api_env.py diff --git a/comfy/comfy_api_env.py b/comfy/comfy_api_env.py new file mode 100644 index 000000000..17b47933f --- /dev/null +++ b/comfy/comfy_api_env.py @@ -0,0 +1,46 @@ +"""Runtime config the frontend reads from /features to follow --comfy-api-base. + +For a non-prod comfy.org backend (staging or an ephemeral preview env), "/features" exposes the api and +platform base so the frontend talks to it without a rebuild, plus the Firebase environment it should use. +Prod bases are left alone and keep their build-time defaults. +""" + +from typing import Any +from urllib.parse import urlparse + +from comfy.cli_args import args + +_STAGING_API_HOST = "stagingapi.comfy.org" +_TESTENV_HOST_SUFFIX = ".testenvs.comfy.org" +_STAGING_PLATFORM_BASE_URL = "https://stagingplatform.comfy.org" + + +def _is_staging_tier(host: str) -> bool: + return host == _STAGING_API_HOST or host.endswith(_TESTENV_HOST_SUFFIX) + + +def normalize_comfy_api_base(url: str) -> str: + """Rewrite a testenv's friendly main host to its comfy-api '-registry' sibling.""" + parsed = urlparse(url) + host = parsed.hostname or "" + if not host.endswith(_TESTENV_HOST_SUFFIX): + return url + label = host[: -len(_TESTENV_HOST_SUFFIX)] + if label.endswith("-registry"): + return url + return f"{parsed.scheme or 'https'}://{label}-registry{_TESTENV_HOST_SUFFIX}" + + +def environment_overrides_for_base(base_url: str) -> dict[str, Any] | None: + """The /features overrides for a staging-tier base, or None for prod.""" + if not _is_staging_tier(urlparse(base_url).hostname or ""): + return None + return { + "comfy_api_base_url": normalize_comfy_api_base(base_url).rstrip("/"), + "comfy_platform_base_url": _STAGING_PLATFORM_BASE_URL, + "firebase_env": "dev", + } + + +def get_environment_overrides() -> dict[str, Any] | None: + return environment_overrides_for_base(getattr(args, "comfy_api_base", "") or "") diff --git a/comfy_api_nodes/util/_helpers.py b/comfy_api_nodes/util/_helpers.py index 6b8121cab..7eb1ec664 100644 --- a/comfy_api_nodes/util/_helpers.py +++ b/comfy_api_nodes/util/_helpers.py @@ -11,6 +11,7 @@ from io import BytesIO from yarl import URL from comfy.cli_args import args +from comfy.comfy_api_env import normalize_comfy_api_base from comfy.deploy_environment import get_deploy_environment from comfy.model_management import processing_interrupted from comfy_api.latest import IO @@ -63,7 +64,7 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]: def default_base_url() -> str: - return getattr(args, "comfy_api_base", "https://api.comfy.org") + return normalize_comfy_api_base(getattr(args, "comfy_api_base", "https://api.comfy.org")) async def sleep_with_interrupt( diff --git a/server.py b/server.py index 5ab93ddc2..e28fe2d22 100644 --- a/server.py +++ b/server.py @@ -39,6 +39,7 @@ from comfy.deploy_environment import get_deploy_environment import comfy.utils import comfy.model_management from comfy_api import feature_flags +from comfy.comfy_api_env import get_environment_overrides import node_helpers from comfyui_version import __version__ from app.frontend_management import FrontendManager, parse_version @@ -727,7 +728,11 @@ class PromptServer(): @routes.get("/features") async def get_features(request): - return web.json_response(feature_flags.get_server_features()) + features = feature_flags.get_server_features() + overrides = get_environment_overrides() + if overrides: + features.update(overrides) + return web.json_response(features) @routes.get("/prompt") async def get_prompt(request): diff --git a/tests-unit/feature_flags_test.py b/tests-unit/feature_flags_test.py index a436ab1ec..df16df6ab 100644 --- a/tests-unit/feature_flags_test.py +++ b/tests-unit/feature_flags_test.py @@ -11,6 +11,11 @@ from comfy_api.feature_flags import ( _coerce_flag_value, _parse_cli_feature_flags, ) +from comfy.comfy_api_env import ( + environment_overrides_for_base, + get_environment_overrides, + normalize_comfy_api_base, +) class TestFeatureFlags: @@ -183,3 +188,65 @@ class TestCliFeatureFlagRegistry: assert "type" in info, f"{key} missing 'type'" assert "default" in info, f"{key} missing 'default'" assert "description" in info, f"{key} missing 'description'" + + +class TestComfyApiEnv: + """--comfy-api-base staging-tier detection + testenv main-host -> -registry rewrite.""" + + @pytest.mark.parametrize( + "url, expected", + [ + # testenv friendly main host -> comfy-api -registry sibling (slash trimmed) + ("https://pr-4398.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"), + ("https://pr-4398.testenvs.comfy.org/", "https://pr-4398-registry.testenvs.comfy.org"), + ("https://pr-4398-registry.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"), + # staging + everything else -> unchanged (no -registry split) + ("https://stagingapi.comfy.org", "https://stagingapi.comfy.org"), + ("https://api.comfy.org", "https://api.comfy.org"), + ("https://pr-1.testenvs.comfy.org.evil.com", "https://pr-1.testenvs.comfy.org.evil.com"), + ("", ""), + ], + ) + def test_normalize_comfy_api_base(self, url, expected): + assert normalize_comfy_api_base(url) == expected + + def test_config_for_staging_tier_else_none(self): + # ephemeral testenv: friendly main host -> -registry, staging platform, dev Firebase env + eph = environment_overrides_for_base("https://pr-1234.testenvs.comfy.org/") + assert eph["comfy_api_base_url"] == "https://pr-1234-registry.testenvs.comfy.org" + assert eph["comfy_platform_base_url"] == "https://stagingplatform.comfy.org" + assert eph["firebase_env"] == "dev" + # staging api host: emitted as-is + stg = environment_overrides_for_base("https://stagingapi.comfy.org") + assert stg["comfy_api_base_url"] == "https://stagingapi.comfy.org" + assert stg["comfy_platform_base_url"] == "https://stagingplatform.comfy.org" + assert stg["firebase_env"] == "dev" + # prod / unknown: nothing + assert environment_overrides_for_base("https://api.comfy.org") is None + + def test_environment_overrides_only_for_staging_tier(self, monkeypatch): + def set_base(url): + monkeypatch.setattr( + "comfy.comfy_api_env.args", + type("Args", (), {"comfy_api_base": url})(), + ) + + # The overrides merged into the HTTP /features response are present for staging-tier bases... + set_base("https://stagingapi.comfy.org") + assert "comfy_api_base_url" in get_environment_overrides() + set_base("https://pr-7.testenvs.comfy.org") + assert "comfy_api_base_url" in get_environment_overrides() + # ...but never for prod. + set_base("https://api.comfy.org") + assert get_environment_overrides() is None + + def test_server_features_never_carry_env_overrides(self, monkeypatch): + """The WebSocket capability handshake must stay free of routing keys.""" + monkeypatch.setattr( + "comfy.comfy_api_env.args", + type("Args", (), {"comfy_api_base": "https://pr-7.testenvs.comfy.org"})(), + ) + features = get_server_features() + assert "comfy_api_base_url" not in features + assert "comfy_platform_base_url" not in features + assert "firebase_env" not in features From 04a30fb375a6c1312365fbbbdd0a7d0669212e92 Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Thu, 9 Jul 2026 10:42:20 -0400 Subject: [PATCH 075/211] fix: Load3D failing path validation from double path resolution (#14852) --- comfy_extras/nodes_load_3d.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/comfy_extras/nodes_load_3d.py b/comfy_extras/nodes_load_3d.py index 6e3e88471..6ef9a1ca3 100644 --- a/comfy_extras/nodes_load_3d.py +++ b/comfy_extras/nodes_load_3d.py @@ -61,14 +61,10 @@ class Load3D(IO.ComfyNode): @classmethod def execute(cls, model_file, image, **kwargs) -> IO.NodeOutput: - image_path = folder_paths.get_annotated_filepath(image['image']) - mask_path = folder_paths.get_annotated_filepath(image['mask']) - normal_path = folder_paths.get_annotated_filepath(image['normal']) - load_image_node = nodes.LoadImage() - output_image, ignore_mask = load_image_node.load_image(image=image_path) - ignore_image, output_mask = load_image_node.load_image(image=mask_path) - normal_image, ignore_mask2 = load_image_node.load_image(image=normal_path) + output_image, ignore_mask = load_image_node.load_image(image=image['image']) + ignore_image, output_mask = load_image_node.load_image(image=image['mask']) + normal_image, ignore_mask2 = load_image_node.load_image(image=image['normal']) video = None From 412aaab0e27d244f3fe47b2e593fa40cc9bcb4c6 Mon Sep 17 00:00:00 2001 From: Simon Pinfold Date: Fri, 10 Jul 2026 07:59:30 +1200 Subject: [PATCH 076/211] feat(api): expose registered extension filters on /experiment/models (#14797) Each folder in the listing now carries its registered extension allowlist verbatim; an empty array means the folder accepts any extension (match-all), mirroring filter_files_extensions semantics. Gives consumers the filtering rule itself rather than just its output: /models/{folder} lists files by the per-folder rule but the rule is not exposed anywhere, and /experiment/models/{folder} filters everything by the global supported_pt_extensions regardless of registration. Presentation-level filtering of match-all folders (e.g. hiding README/config noise that repository-downloading custom nodes leave in model directories) is deliberately left to the consumer. Co-authored-by: guill --- app/model_manager.py | 6 +++++- openapi.yaml | 8 ++++++++ tests-unit/app_test/model_manager_test.py | 22 ++++++++++++++++++++++ 3 files changed, 35 insertions(+), 1 deletion(-) diff --git a/app/model_manager.py b/app/model_manager.py index b0329ce17..5928781ca 100644 --- a/app/model_manager.py +++ b/app/model_manager.py @@ -35,7 +35,11 @@ class ModelFileManager: for folder in model_types: if folder in folder_black_list: continue - output_folders.append({"name": folder, "folders": folder_paths.get_folder_paths(folder)}) + output_folders.append({ + "name": folder, + "folders": folder_paths.get_folder_paths(folder), + "extensions": sorted(folder_paths.folder_names_and_paths[folder][1]), + }) return web.json_response(output_folders) # NOTE: This is an experiment to replace `/models/{folder}` diff --git a/openapi.yaml b/openapi.yaml index 0cf177815..c09b1eeac 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -775,6 +775,14 @@ components: ModelFolder: description: Represents a folder containing models properties: + extensions: + description: The folder's registered file-extension allowlist. An empty array means the folder accepts any extension (match-all). + example: + - .ckpt + - .safetensors + items: + type: string + type: array folders: description: List of paths where models of this type are stored example: diff --git a/tests-unit/app_test/model_manager_test.py b/tests-unit/app_test/model_manager_test.py index ae59206f6..d7cc20fcd 100644 --- a/tests-unit/app_test/model_manager_test.py +++ b/tests-unit/app_test/model_manager_test.py @@ -24,6 +24,28 @@ def app(model_manager): app.add_routes(routes) return app +async def test_get_model_folders_includes_registered_extensions(aiohttp_client, app, tmp_path): + """Folders expose their registered extension set verbatim; an empty list + means match-all (filter_files_extensions semantics).""" + with patch('folder_paths.folder_names_and_paths', { + 'test_checkpoints': ([str(tmp_path)], {'.safetensors', '.ckpt'}), + 'test_configs': ([str(tmp_path)], ['.yaml']), + 'test_match_all': ([str(tmp_path)], set()), + 'configs': ([str(tmp_path)], ['.yaml']), + }): + client = await aiohttp_client(app) + response = await client.get('/experiment/models') + + assert response.status == 200 + folders = {f['name']: f for f in await response.json()} + + assert 'configs' not in folders # blocklisted + assert folders['test_checkpoints']['folders'] == [str(tmp_path)] + assert folders['test_checkpoints']['extensions'] == ['.ckpt', '.safetensors'] + assert folders['test_configs']['extensions'] == ['.yaml'] + # Match-all registrations are exposed honestly, not substituted. + assert folders['test_match_all']['extensions'] == [] + async def test_get_model_preview_safetensors(aiohttp_client, app, tmp_path): img = Image.new('RGB', (100, 100), 'white') img_byte_arr = BytesIO() From 1ea724339ca6bea1808e6600fefdd08bf42822e6 Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Fri, 10 Jul 2026 05:57:44 +0800 Subject: [PATCH 077/211] Update cla.yml (#14851) --- .github/workflows/cla.yml | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/cla.yml b/.github/workflows/cla.yml index b75397e50..bc0f779cf 100644 --- a/.github/workflows/cla.yml +++ b/.github/workflows/cla.yml @@ -32,9 +32,11 @@ jobs: PR_NUMBER: ${{ github.event.pull_request.number || github.event.issue.number }} PR_AUTHOR: ${{ github.event.pull_request.user.login || github.event.issue.user.login }} BASE_ALLOWLIST: action@github.com,actions-user,ampagent,claude,comfy-pr-bot,GitHub Action,github-actions,github-actions[bot],Glary Bot,Glary-Bot,*[bot] + # For each commit emit the GitHub login when the author/committer email resolves to a GitHub account + # otherwise fall back to the raw git name. run: | others=$(gh api "repos/${{ github.repository }}/pulls/${PR_NUMBER}/commits" --paginate \ - --jq '.[] | (.author.login // empty), (.committer.login // empty)' \ + --jq '.[] | (.author.login // .commit.author.name // empty), (.committer.login // .commit.committer.name // empty)' \ | sort -u | grep -vix "${PR_AUTHOR}" | paste -sd, -) if [ -n "$others" ]; then echo "allowlist=${BASE_ALLOWLIST},${others}" >> "$GITHUB_OUTPUT" @@ -43,7 +45,7 @@ jobs: fi - name: CLA Assistant - # Run on PR events, on "recheck" comment, or when someone posts the exact signing phrase. + # Run on PR events, on "recheck" comment, or when someone posts the signing phrase. # IMPORTANT: this phrase must match `custom-pr-sign-comment` below. if: > github.event_name == 'pull_request_target' || From 73e84d5ec8b943dcb42535229eb94ee7ab3abea1 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 9 Jul 2026 15:57:09 -0700 Subject: [PATCH 078/211] Support convrot int4 models. (#14859) linear_dtype in comfy_quant metadata can be used to set if the int4 op does the matrix multiplication in int8 or int4, the default is int4 on GPUs that support it with fallback to int8 for GPUs that don't. --- comfy/ops.py | 26 +++++++++ comfy/quant_ops.py | 14 +++++ requirements.txt | 2 +- .../comfy_quant/test_mixed_precision.py | 54 ++++++++++++++++++- 4 files changed, 94 insertions(+), 2 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 35a1ee31e..0c6fe4cb4 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1104,6 +1104,21 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat scales["convrot_groupsize"] = int( layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)) ) + elif module.quant_format == "convrot_w4a4": + scale = pop_scale("weight_scale") + if scale is None: + raise ValueError(f"Missing ConvRot W4A4 weight scale for layer {layer_name}") + params_conf = layer_conf.get("params", {}) + if not isinstance(params_conf, dict): + params_conf = {} + scales = { + "scale": scale, + "convrot_groupsize": int( + layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)) + ), + "quant_group_size": 64, + "linear_dtype": layer_conf.get("linear_dtype", params_conf.get("linear_dtype", "int4")), + } else: raise ValueError(f"Unsupported quantization format: {module.quant_format}") @@ -1150,6 +1165,11 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False): quant_conf["convrot"] = True quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256) + elif module.quant_format == "convrot_w4a4": + quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256) + linear_dtype = getattr(params, "linear_dtype", "int4") + if linear_dtype != "int4": + quant_conf["linear_dtype"] = linear_dtype if extra_quant_conf: quant_conf.update(extra_quant_conf) sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8) @@ -1430,6 +1450,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec } if hasattr(params, "block_scale"): # NVFP4 kwargs["block_scale"] = params.block_scale[i] + if hasattr(params, "quant_group_size"): + kwargs["quant_group_size"] = params.quant_group_size + if hasattr(params, "convrot_groupsize"): + kwargs["convrot_groupsize"] = params.convrot_groupsize + if hasattr(params, "linear_dtype"): + kwargs["linear_dtype"] = params.linear_dtype return QuantizedTensor(weight._qdata[i], weight._layout_cls, type(params)(**kwargs)) def state_dict(self, *args, destination=None, prefix="", **kwargs): diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 44f25a97e..53a0cb603 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -10,6 +10,7 @@ try: QuantizedLayout, TensorCoreFP8Layout as _CKFp8Layout, TensorCoreNVFP4Layout as _CKNvfp4Layout, + TensorCoreConvRotW4A4Layout as _CKTensorCoreConvRotW4A4Layout, TensorWiseINT8Layout as _CKTensorWiseINT8Layout, register_layout_op, register_layout_class, @@ -51,6 +52,9 @@ except ImportError as e: class _CKTensorWiseINT8Layout: pass + class _CKTensorCoreConvRotW4A4Layout: + pass + def register_layout_class(name, cls): pass @@ -179,6 +183,7 @@ class TensorCoreFP8E5M2Layout(_TensorCoreFP8LayoutBase): # Backward compatibility alias - default to E4M3 TensorCoreFP8Layout = TensorCoreFP8E4M3Layout TensorWiseINT8Layout = _CKTensorWiseINT8Layout +TensorCoreConvRotW4A4Layout = _CKTensorCoreConvRotW4A4Layout # ============================================================================== @@ -190,6 +195,7 @@ register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout) register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout) register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout) register_layout_class("TensorWiseINT8Layout", _CKTensorWiseINT8Layout) +register_layout_class("TensorCoreConvRotW4A4Layout", _CKTensorCoreConvRotW4A4Layout) if _CK_MXFP8_AVAILABLE: register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout) @@ -227,6 +233,13 @@ QUANT_ALGOS["int8_tensorwise"] = { "quantize_input": False, } +QUANT_ALGOS["convrot_w4a4"] = { + "storage_t": torch.int8, + "parameters": {"weight_scale"}, + "comfy_tensor_layout": "TensorCoreConvRotW4A4Layout", + "quantize_input": False, +} + # ============================================================================== # Re-exports for backward compatibility @@ -239,6 +252,7 @@ __all__ = [ "TensorCoreFP8E4M3Layout", "TensorCoreFP8E5M2Layout", "TensorCoreNVFP4Layout", + "TensorCoreConvRotW4A4Layout", "TensorWiseINT8Layout", "QUANT_ALGOS", "register_layout_op", diff --git a/requirements.txt b/requirements.txt index e72f3045b..a8ea0eace 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.16 +comfy-kitchen==0.2.17 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 diff --git a/tests-unit/comfy_quant/test_mixed_precision.py b/tests-unit/comfy_quant/test_mixed_precision.py index 43b4b7ce9..7bbc96616 100644 --- a/tests-unit/comfy_quant/test_mixed_precision.py +++ b/tests-unit/comfy_quant/test_mixed_precision.py @@ -15,7 +15,7 @@ if not has_gpu(): args.cpu = True from comfy import ops -from comfy.quant_ops import QuantizedTensor +from comfy.quant_ops import QUANT_ALGOS, QuantizedTensor import comfy.utils @@ -283,7 +283,59 @@ class TestMixedPrecisionOps(unittest.TestCase): saved = model.state_dict() saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes()) self.assertTrue(saved_conf["convrot"]) + + def test_convrot_w4a4_loads_into_params(self): + """ConvRot W4A4 checkpoints must load as the dedicated kitchen layout.""" + if "convrot_w4a4" not in QUANT_ALGOS: + self.skipTest("comfy_kitchen does not provide ConvRot W4A4") + + torch.manual_seed(456) + layer_quant_config = { + "layer": { + "format": "convrot_w4a4", + "convrot_groupsize": 256, + "linear_dtype": "int8", + } + } + weight = torch.randn(16, 256, dtype=torch.bfloat16) + bias = torch.randn(16, dtype=torch.bfloat16) + q_weight = QuantizedTensor.from_float( + weight, + "TensorCoreConvRotW4A4Layout", + convrot_groupsize=256, + quant_group_size=64, + ) + state_dict = { + "layer.weight": q_weight._qdata, + "layer.bias": bias, + "layer.weight_scale": q_weight._params.scale, + } + + state_dict, _ = comfy.utils.convert_old_quants( + state_dict, + metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})}, + ) + model = torch.nn.Module() + model.layer = ops.mixed_precision_ops({}).Linear(256, 16, device="cpu", dtype=torch.bfloat16) + model.load_state_dict(state_dict, strict=False) + + self.assertIsInstance(model.layer.weight, QuantizedTensor) + self.assertEqual(model.layer.weight._layout_cls, "TensorCoreConvRotW4A4Layout") + self.assertEqual(model.layer.weight._params.convrot_groupsize, 256) + self.assertEqual(model.layer.weight._params.quant_group_size, 64) + self.assertEqual(model.layer.weight._params.linear_dtype, "int8") + + input_tensor = torch.randn(4, 256, dtype=torch.bfloat16) + loaded_out = model.layer(input_tensor) + ref_out = torch.nn.functional.linear(input_tensor, q_weight, bias) + self.assertTrue(torch.equal(loaded_out, ref_out)) + + saved = model.state_dict() + saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes()) + self.assertEqual(saved_conf["format"], "convrot_w4a4") self.assertEqual(saved_conf["convrot_groupsize"], 256) + self.assertEqual(saved_conf["linear_dtype"], "int8") + self.assertNotIn("quant_group_size", saved_conf) if __name__ == "__main__": unittest.main() From b7a648ca2011489ba40eaacf01a5d6f4e9fab539 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 9 Jul 2026 16:39:01 -0700 Subject: [PATCH 079/211] Try to fix the model reloading issue some people have. (#14822) --- comfy/model_management.py | 11 +++++++++++ comfy_execution/caching.py | 25 +++++++++++++++++-------- execution.py | 17 +++++++++++------ 3 files changed, 39 insertions(+), 14 deletions(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index b15d08ba1..222005b6f 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -616,6 +616,8 @@ PIN_PRESSURE_HYSTERESIS = 256 * 1024 * 1024 #Freeing registerables on pressure does imply a GPU sync, so go big on #the hysteresis so each expensive sync gives us back a good chunk. REGISTERABLE_PIN_HYSTERESIS = 2048 * 1024 * 1024 +WINDOWS_PIN_EVICTION_SWAP_PERCENT = 5.0 +WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE = 512 * 1024 ** 2 def module_size(module): module_mem = 0 @@ -642,6 +644,15 @@ def free_pins(size, evict_active=False): size -= freed return freed_total +def should_free_pins_for_ram_pressure(shortfall): + if shortfall <= 0: + return False + if not WINDOWS: + return True + if psutil.virtual_memory().available < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE: + return True + return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT + def ensure_pin_budget(size, evict_active=False): if args.high_ram: return True diff --git a/comfy_execution/caching.py b/comfy_execution/caching.py index ad75a0e50..6bd99b68f 100644 --- a/comfy_execution/caching.py +++ b/comfy_execution/caching.py @@ -503,6 +503,8 @@ RAM_CACHE_DEFAULT_RAM_USAGE = 0.05 RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER = 1.3 +RAM_CACHE_LARGE_INTERMEDIATE = 512 * 1024 ** 2 + def all_outputs_dynamic(outputs): if outputs is None: @@ -517,7 +519,6 @@ def all_outputs_dynamic(outputs): return True - class RAMPressureCache(LRUCache): def __init__(self, key_class, enable_providers=False): @@ -539,9 +540,9 @@ class RAMPressureCache(LRUCache): self.timestamps[self.cache_key_set.get_data_key(node_id)] = time.time() super().set_local(node_id, value) - def ram_release(self, target, free_active=False): + def ram_release(self, target, free_active=False, min_entry_size=0): if psutil.virtual_memory().available >= target: - return + return 0 clean_list = [] @@ -555,8 +556,9 @@ class RAMPressureCache(LRUCache): oom_score = RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER ** (self.generation - self.used_generation[key]) ram_usage = RAM_CACHE_DEFAULT_RAM_USAGE + oom_ram_usage = ram_usage def scan_list_for_ram_usage(outputs): - nonlocal ram_usage + nonlocal ram_usage, oom_ram_usage if outputs is None: return for output in outputs: @@ -564,19 +566,26 @@ class RAMPressureCache(LRUCache): scan_list_for_ram_usage(output) elif isinstance(output, torch.Tensor) and output.device.type == 'cpu': ram_usage += output.numel() * output.element_size() + oom_ram_usage += output.numel() * output.element_size() elif isinstance(output, ModelPatcher) and self.used_generation[key] != self.generation: #old ModelPatchers are the first to go - ram_usage = 1e30 + oom_ram_usage = 1e30 scan_list_for_ram_usage(cache_entry.outputs) - oom_score *= ram_usage + if ram_usage < min_entry_size: + continue + + oom_score *= oom_ram_usage #In the case where we have no information on the node ram usage at all, #break OOM score ties on the last touch timestamp (pure LRU) - bisect.insort(clean_list, (oom_score, self.timestamps[key], key)) + bisect.insort(clean_list, (oom_score, self.timestamps[key], key, ram_usage)) + freed = 0 while psutil.virtual_memory().available < target and clean_list: - _, _, key = clean_list.pop() + _, _, key, ram_usage = clean_list.pop() del self.cache[key] self.used_generation.pop(key, None) self.timestamps.pop(key, None) self.children.pop(key, None) + freed += ram_usage + return freed diff --git a/execution.py b/execution.py index c45317593..19b8cdd68 100644 --- a/execution.py +++ b/execution.py @@ -29,6 +29,7 @@ from comfy_execution.caching import ( HierarchicalCache, LRUCache, RAMPressureCache, + RAM_CACHE_LARGE_INTERMEDIATE, ) from comfy_execution.graph import ( DynamicPrompt, @@ -794,12 +795,16 @@ class PromptExecutor: if self.cache_type == CacheType.RAM_PRESSURE: ram_release_callback(ram_inactive_headroom) ram_shortfall = ram_headroom - psutil.virtual_memory().available - freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2)) - if freed < ram_shortfall: - if freed > 64 * (1024 ** 2): - # AIMDO MEM_DECOMMIT can outrun psutil.available catching up. - time.sleep(0.05) - ram_release_callback(ram_headroom, free_active=True) + if ram_shortfall > 0: + freed = ram_release_callback(ram_headroom, free_active=True, min_entry_size=RAM_CACHE_LARGE_INTERMEDIATE) + ram_shortfall -= freed + if comfy.model_management.should_free_pins_for_ram_pressure(ram_shortfall): + freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2)) + if freed < ram_shortfall: + if freed > 64 * (1024 ** 2): + # AIMDO MEM_DECOMMIT can outrun psutil.available catching up. + time.sleep(0.05) + ram_release_callback(ram_headroom, free_active=True) else: # Only execute when the while-loop ends without break # Send cached UI for intermediate output nodes that weren't executed From 62e025a4f34d16eeedfc1c93e50a48a69098df7f Mon Sep 17 00:00:00 2001 From: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com> Date: Fri, 10 Jul 2026 10:30:26 +0800 Subject: [PATCH 080/211] Fix FP8 activation quantization for >2D activations in mixed_precision_ops (#14643) mixed_precision_ops.Linear.forward only quantized activations that were 2D, or 3D (reshaped to 2D). Inputs with rank >= 4 (e.g. Anima's MLP activations, which are not reshaped to 3D the way the attention path is) fell through the `input_reshaped.ndim == 2` guard and reached scaled_mm as bf16, silently dispatching a bf16 kernel instead of FP8. Since MLP is roughly half the compute, the FP8 speedup was far below expectation. Generalize the existing 3D->2D reshape to any rank >= 3 (flatten the leading dims, keep the contraction dim) and reshape the output back to the original leading dims. 2D and 3D inputs are handled exactly as before; only rank >= 4 inputs change (now quantized instead of skipped). This matches the rank-agnostic handling already used by the training path (flatten(0, -2) / unflatten). Fixes #14595. --- comfy/ops.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 0c6fe4cb4..13c2604fb 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1257,7 +1257,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec run_every_op() input_shape = input.shape - reshaped_3d = False + reshaped_nd = False #If cast needs to apply lora, it should be done in the compute dtype compute_dtype = input.dtype @@ -1294,12 +1294,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec # Inference path (unchanged) if _use_quantized and quantize_input: - # Reshape 3D tensors to 2D for quantization (needed for NVFP4 and others) - input_reshaped = input.reshape(-1, input_shape[2]) if input.ndim == 3 else input + # Reshape >=3D tensors to 2D for quantization (needed for NVFP4 and others) + input_reshaped = input.reshape(-1, input_shape[-1]) if input.ndim >= 3 else input # Fall back to non-quantized for non-2D tensors if input_reshaped.ndim == 2: - reshaped_3d = input.ndim == 3 + reshaped_nd = input.ndim >= 3 # dtype is now implicit in the layout class scale = getattr(self, 'input_scale', None) if scale is not None: @@ -1314,9 +1314,9 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec weight_only_quant=weight_only_quant, ) - # Reshape output back to 3D if input was 3D - if reshaped_3d: - output = output.reshape((input_shape[0], input_shape[1], self.weight.shape[0])) + # Reshape output back to original rank if input was >2D + if reshaped_nd: + output = output.reshape((*input_shape[:-1], self.weight.shape[0])) return output From 099522f85bcd8586eac4133c02f31c70dafe85d2 Mon Sep 17 00:00:00 2001 From: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com> Date: Fri, 10 Jul 2026 11:11:52 +0800 Subject: [PATCH 081/211] Enable comfy-kitchen Triton backend by default on ROCm/AMD (#14862) On AMD/ROCm the CUDA backend is unavailable, so Triton is the only accelerated comfy-kitchen backend. It was disabled by default (opt-in --enable-triton-backend), leaving AMD on the slow eager path. Enable it by default when torch.version.hip is set AND Triton is >= 3.7 -- older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path, so on Triton < 3.7 it stays disabled with a log line. NVIDIA behavior is unchanged; the explicit --enable-triton-backend flag still works as an override. Fixes #14861 --- comfy/quant_ops.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 53a0cb603..91b3e4fe9 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -25,10 +25,18 @@ try: ck.registry.disable("cuda") logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.") - if args.enable_triton_backend: + # On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated + # comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7: + # older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path. + if args.enable_triton_backend or torch.version.hip is not None: try: import triton - logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__) + triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2]) + if args.enable_triton_backend or triton_version >= (3, 7): + logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__) + else: + logging.info("Triton %s is too old for the ROCm INT8 path (needs >= 3.7); comfy-kitchen triton backend disabled.", triton.__version__) + ck.registry.disable("triton") except ImportError as e: logging.error(f"Failed to import triton, Error: {e}, the comfy-kitchen triton backend will not be available.") ck.registry.disable("triton") From e2a6e30d892402ffcf01d6280c8e2744a4448b9d Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 9 Jul 2026 20:17:06 -0700 Subject: [PATCH 082/211] Fix black image on turing when using int4 models. (#14864) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index a8ea0eace..790ef4940 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.17 +comfy-kitchen==0.2.18 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 8e2e54e2b8fd2623e5bc255b7c90ccf09bcb04bb Mon Sep 17 00:00:00 2001 From: John Pollock Date: Fri, 10 Jul 2026 02:07:42 -0500 Subject: [PATCH 083/211] Add SeedVR2 support (CORE-6) (#14424) --- comfy/latent_formats.py | 4 + comfy/ldm/modules/diffusionmodules/model.py | 6 +- comfy/ldm/seedvr/attention.py | 51 + comfy/ldm/seedvr/color_fix.py | 301 +++ comfy/ldm/seedvr/constants.py | 48 + comfy/ldm/seedvr/model.py | 1361 ++++++++++++++ comfy/ldm/seedvr/vae.py | 1612 +++++++++++++++++ comfy/model_base.py | 12 + comfy/model_detection.py | 45 +- comfy/sd.py | 102 +- comfy/supported_models.py | 35 + comfy/supported_models_base.py | 6 +- comfy/text_encoders/gemma4.py | 2 +- comfy_extras/nodes_seedvr.py | 614 +++++++ nodes.py | 1 + .../test_seedvr2_conditioning.py | 186 ++ .../comfy_extras_test/test_seedvr2_nodes.py | 55 + .../test_seedvr2_post_processing.py | 51 + .../test_seedvr2_temporal_chunk.py | 77 + tests-unit/comfy_test/model_detection_test.py | 83 +- .../comfy_test/seedvr_vae_forward_test.py | 74 + tests-unit/comfy_test/test_seedvr2_dtype.py | 50 + .../comfy_test/test_seedvr2_internals.py | 169 ++ tests-unit/comfy_test/test_seedvr2_model.py | 320 ++++ .../comfy_test/test_seedvr2_vae_decode.py | 94 + .../comfy_test/test_seedvr2_vae_tiled.py | 382 ++++ 26 files changed, 5712 insertions(+), 29 deletions(-) create mode 100644 comfy/ldm/seedvr/attention.py create mode 100644 comfy/ldm/seedvr/color_fix.py create mode 100644 comfy/ldm/seedvr/constants.py create mode 100644 comfy/ldm/seedvr/model.py create mode 100644 comfy/ldm/seedvr/vae.py create mode 100644 comfy_extras/nodes_seedvr.py create mode 100644 tests-unit/comfy_extras_test/test_seedvr2_conditioning.py create mode 100644 tests-unit/comfy_extras_test/test_seedvr2_nodes.py create mode 100644 tests-unit/comfy_extras_test/test_seedvr2_post_processing.py create mode 100644 tests-unit/comfy_extras_test/test_seedvr2_temporal_chunk.py create mode 100644 tests-unit/comfy_test/seedvr_vae_forward_test.py create mode 100644 tests-unit/comfy_test/test_seedvr2_dtype.py create mode 100644 tests-unit/comfy_test/test_seedvr2_internals.py create mode 100644 tests-unit/comfy_test/test_seedvr2_model.py create mode 100644 tests-unit/comfy_test/test_seedvr2_vae_decode.py create mode 100644 tests-unit/comfy_test/test_seedvr2_vae_tiled.py diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index bbdfd4bc2..8a16cfe55 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -779,6 +779,10 @@ class ACEAudio(LatentFormat): latent_channels = 8 latent_dimensions = 2 +class SeedVR2(LatentFormat): + latent_channels = 16 + latent_dimensions = 3 + class ACEAudio15(LatentFormat): latent_channels = 64 latent_dimensions = 1 diff --git a/comfy/ldm/modules/diffusionmodules/model.py b/comfy/ldm/modules/diffusionmodules/model.py index fcbaa074f..e752d0ecb 100644 --- a/comfy/ldm/modules/diffusionmodules/model.py +++ b/comfy/ldm/modules/diffusionmodules/model.py @@ -22,7 +22,7 @@ def torch_cat_if_needed(xl, dim): else: return None -def get_timestep_embedding(timesteps, embedding_dim): +def get_timestep_embedding(timesteps, embedding_dim, flip_sin_to_cos=False, downscale_freq_shift=1): """ This matches the implementation in Denoising Diffusion Probabilistic Models: From Fairseq. @@ -33,11 +33,13 @@ def get_timestep_embedding(timesteps, embedding_dim): assert len(timesteps.shape) == 1 half_dim = embedding_dim // 2 - emb = math.log(10000) / (half_dim - 1) + emb = math.log(10000) / (half_dim - downscale_freq_shift) emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb) emb = emb.to(device=timesteps.device) emb = timesteps.float()[:, None] * emb[None, :] emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) if embedding_dim % 2 == 1: # zero pad emb = torch.nn.functional.pad(emb, (0,1,0,0)) return emb diff --git a/comfy/ldm/seedvr/attention.py b/comfy/ldm/seedvr/attention.py new file mode 100644 index 000000000..11b4c1e4a --- /dev/null +++ b/comfy/ldm/seedvr/attention.py @@ -0,0 +1,51 @@ +import torch + +from comfy.ldm.modules import attention as _attention + + +def _var_attention_qkv(q, k, v, heads, skip_reshape): + if skip_reshape: + return q, k, v, q.shape[-1] + total_tokens, embed_dim = q.shape + head_dim = embed_dim // heads + return ( + q.view(total_tokens, heads, head_dim), + k.view(k.shape[0], heads, head_dim), + v.view(v.shape[0], heads, head_dim), + head_dim, + ) + + +def _var_attention_output(out, heads, head_dim, skip_output_reshape): + if skip_output_reshape: + return out + return out.reshape(-1, heads * head_dim) + + +def var_attention_optimized_split(q, k, v, heads, cu_seqlens_q, cu_seqlens_k, *args, skip_reshape=False, skip_output_reshape=False, **kwargs): + q, k, v, head_dim = _var_attention_qkv(q, k, v, heads, skip_reshape) + + q_split_indices = cu_seqlens_q[1:-1] + k_split_indices = cu_seqlens_k[1:-1] + if k.shape[0] != v.shape[0]: + raise ValueError("cu_seqlens_k does not match v token count") + + q_splits = torch.tensor_split(q, q_split_indices, dim=0) + k_splits = torch.tensor_split(k, k_split_indices, dim=0) + v_splits = torch.tensor_split(v, k_split_indices, dim=0) + if len(q_splits) != len(k_splits) or len(q_splits) != len(v_splits): + raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same sequence count") + + out = [] + for q_i, k_i, v_i in zip(q_splits, k_splits, v_splits): + q_i = q_i.permute(1, 0, 2).unsqueeze(0) + k_i = k_i.permute(1, 0, 2).unsqueeze(0) + v_i = v_i.permute(1, 0, 2).unsqueeze(0) + out_i = _attention.optimized_attention(q_i, k_i, v_i, heads, skip_reshape=True, skip_output_reshape=True) + out.append(out_i.squeeze(0).permute(1, 0, 2)) + + out = torch.cat(out, dim=0) + return _var_attention_output(out, heads, head_dim, skip_output_reshape) + + +optimized_var_attention = var_attention_optimized_split diff --git a/comfy/ldm/seedvr/color_fix.py b/comfy/ldm/seedvr/color_fix.py new file mode 100644 index 000000000..a43cb5270 --- /dev/null +++ b/comfy/ldm/seedvr/color_fix.py @@ -0,0 +1,301 @@ +import torch +import torch.nn.functional as F +from torch import Tensor + +from comfy.ldm.seedvr.constants import ( + CIELAB_DELTA, + CIELAB_KAPPA, + D65_WHITE_X, + D65_WHITE_Z, + WAVELET_DECOMP_LEVELS, +) + + +def wavelet_blur(image: Tensor, radius): + max_safe_radius = max(1, min(image.shape[-2:]) // 8) + if radius > max_safe_radius: + radius = max_safe_radius + + num_channels = image.shape[1] + + kernel_vals = [ + [0.0625, 0.125, 0.0625], + [0.125, 0.25, 0.125], + [0.0625, 0.125, 0.0625], + ] + kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device) + kernel = kernel[None, None].repeat(num_channels, 1, 1, 1) + + image = F.pad(image, (radius, radius, radius, radius), mode='replicate') + output = F.conv2d(image, kernel, groups=num_channels, dilation=radius) + + return output + +def wavelet_decomposition(image: Tensor, levels: int = WAVELET_DECOMP_LEVELS): + high_freq = torch.zeros_like(image) + + for i in range(levels): + radius = 2 ** i + low_freq = wavelet_blur(image, radius) + high_freq.add_(image).sub_(low_freq) + image = low_freq + + return high_freq, low_freq + +def wavelet_reconstruction(content_feat: Tensor, style_feat: Tensor) -> Tensor: + + if content_feat.shape != style_feat.shape: + if len(content_feat.shape) >= 3: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False + ) + + content_high_freq, content_low_freq = wavelet_decomposition(content_feat) + del content_low_freq + + style_high_freq, style_low_freq = wavelet_decomposition(style_feat) + del style_high_freq + + if content_high_freq.shape != style_low_freq.shape: + style_low_freq = F.interpolate( + style_low_freq, + size=content_high_freq.shape[-2:], + mode='bilinear', + align_corners=False + ) + + content_high_freq.add_(style_low_freq) + + return content_high_freq.clamp_(-1.0, 1.0) + +def _histogram_matching_channel(source: Tensor, reference: Tensor) -> Tensor: + original_shape = source.shape + + source_flat = source.flatten() + reference_flat = reference.flatten() + + source_sorted, source_indices = torch.sort(source_flat) + reference_sorted, _ = torch.sort(reference_flat) + del reference_flat + + n_source = len(source_sorted) + n_reference = len(reference_sorted) + + if n_source == n_reference: + matched_sorted = reference_sorted + else: + source_quantiles = torch.linspace(0, 1, n_source, device=source.device) + ref_indices = (source_quantiles * (n_reference - 1)).long() + ref_indices.clamp_(0, n_reference - 1) + matched_sorted = reference_sorted[ref_indices] + del source_quantiles, ref_indices, reference_sorted + + del source_sorted, source_flat + + inverse_indices = torch.argsort(source_indices) + del source_indices + matched_flat = matched_sorted[inverse_indices] + del matched_sorted, inverse_indices + + return matched_flat.reshape(original_shape) + +def _lab_to_rgb_batch(lab: Tensor, matrix_inv: Tensor, epsilon: float, kappa: float) -> Tensor: + L, a, b = lab[:, 0], lab[:, 1], lab[:, 2] + + fy = (L + 16.0) / 116.0 + fx = a.div(500.0).add_(fy) + fz = fy - b / 200.0 + del L, a, b + + x = torch.where( + fx > epsilon, + torch.pow(fx, 3.0), + fx.mul(116.0).sub_(16.0).div_(kappa) + ) + y = torch.where( + fy > epsilon, + torch.pow(fy, 3.0), + fy.mul(116.0).sub_(16.0).div_(kappa) + ) + z = torch.where( + fz > epsilon, + torch.pow(fz, 3.0), + fz.mul(116.0).sub_(16.0).div_(kappa) + ) + del fx, fy, fz + + x.mul_(D65_WHITE_X) + z.mul_(D65_WHITE_Z) + + xyz = torch.stack([x, y, z], dim=1) + del x, y, z + + B, _, H, W = xyz.shape + xyz_flat = xyz.permute(0, 2, 3, 1).reshape(-1, 3) + del xyz + + xyz_flat = xyz_flat.to(dtype=matrix_inv.dtype) + rgb_linear_flat = torch.matmul(xyz_flat, matrix_inv.T) + del xyz_flat + + rgb_linear = rgb_linear_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2) + del rgb_linear_flat + + mask = rgb_linear > 0.0031308 + rgb = torch.where( + mask, + torch.pow(torch.clamp(rgb_linear, min=0.0), 1.0 / 2.4).mul_(1.055).sub_(0.055), + rgb_linear * 12.92 + ) + del mask, rgb_linear + + return torch.clamp(rgb, 0.0, 1.0) + +def _rgb_to_lab_batch(rgb: Tensor, matrix: Tensor, epsilon: float, kappa: float) -> Tensor: + mask = rgb > 0.04045 + rgb_linear = torch.where( + mask, + torch.pow((rgb + 0.055) / 1.055, 2.4), + rgb / 12.92 + ) + del mask + + B, _, H, W = rgb_linear.shape + rgb_flat = rgb_linear.permute(0, 2, 3, 1).reshape(-1, 3) + del rgb_linear + + rgb_flat = rgb_flat.to(dtype=matrix.dtype) + xyz_flat = torch.matmul(rgb_flat, matrix.T) + del rgb_flat + + xyz = xyz_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2) + del xyz_flat + + xyz[:, 0].div_(D65_WHITE_X) + xyz[:, 2].div_(D65_WHITE_Z) + + epsilon_cubed = epsilon ** 3 + mask = xyz > epsilon_cubed + f_xyz = torch.where( + mask, + torch.pow(xyz, 1.0 / 3.0), + xyz.mul(kappa).add_(16.0).div_(116.0) + ) + del xyz, mask + + L = f_xyz[:, 1].mul(116.0).sub_(16.0) + a = (f_xyz[:, 0] - f_xyz[:, 1]).mul_(500.0) + b = (f_xyz[:, 1] - f_xyz[:, 2]).mul_(200.0) + del f_xyz + + return torch.stack([L, a, b], dim=1) + +def lab_color_transfer( + content_feat: Tensor, + style_feat: Tensor, + luminance_weight: float = 0.8 +) -> Tensor: + content_feat = wavelet_reconstruction(content_feat, style_feat) + + if content_feat.shape != style_feat.shape: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False + ) + + device = content_feat.device + original_dtype = content_feat.dtype + content_feat = content_feat.float() + style_feat = style_feat.float() + + rgb_to_xyz_matrix = torch.tensor([ + [0.4124564, 0.3575761, 0.1804375], + [0.2126729, 0.7151522, 0.0721750], + [0.0193339, 0.1191920, 0.9503041] + ], dtype=torch.float32, device=device) + + xyz_to_rgb_matrix = torch.tensor([ + [ 3.2404542, -1.5371385, -0.4985314], + [-0.9692660, 1.8760108, 0.0415560], + [ 0.0556434, -0.2040259, 1.0572252] + ], dtype=torch.float32, device=device) + + epsilon = CIELAB_DELTA + kappa = CIELAB_KAPPA + + content_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) + style_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) + + content_lab = _rgb_to_lab_batch(content_feat, rgb_to_xyz_matrix, epsilon, kappa) + del content_feat + + style_lab = _rgb_to_lab_batch(style_feat, rgb_to_xyz_matrix, epsilon, kappa) + del style_feat, rgb_to_xyz_matrix + + matched_a = _histogram_matching_channel(content_lab[:, 1], style_lab[:, 1]) + matched_b = _histogram_matching_channel(content_lab[:, 2], style_lab[:, 2]) + + if luminance_weight < 1.0: + matched_L = _histogram_matching_channel(content_lab[:, 0], style_lab[:, 0]) + result_L = content_lab[:, 0].mul(luminance_weight).add_(matched_L.mul(1.0 - luminance_weight)) + del matched_L + else: + result_L = content_lab[:, 0] + + del content_lab, style_lab + + result_lab = torch.stack([result_L, matched_a, matched_b], dim=1) + del result_L, matched_a, matched_b + + result_rgb = _lab_to_rgb_batch(result_lab, xyz_to_rgb_matrix, epsilon, kappa) + del result_lab, xyz_to_rgb_matrix + + result = result_rgb.mul_(2.0).sub_(1.0) + del result_rgb + + result = result.to(original_dtype) + + return result + + +def wavelet_color_transfer(content_feat: Tensor, style_feat: Tensor) -> Tensor: + return wavelet_reconstruction(content_feat, style_feat) + + +def adain_color_transfer(content_feat: Tensor, style_feat: Tensor, eps: float = 1e-5) -> Tensor: + if content_feat.shape != style_feat.shape: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False, + ) + + original_dtype = content_feat.dtype + content_feat = content_feat.float() + style_feat = style_feat.float() + + b, c = content_feat.shape[:2] + content_flat = content_feat.reshape(b, c, -1) + style_flat = style_feat.reshape(b, c, -1) + + content_mean = content_flat.mean(dim=2).reshape(b, c, 1, 1) + content_std = (content_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1) + style_mean = style_flat.mean(dim=2).reshape(b, c, 1, 1) + style_std = (style_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1) + del content_flat, style_flat + + normalized = (content_feat - content_mean) / content_std + del content_mean, content_std + result = normalized * style_std + style_mean + del normalized, style_mean, style_std + + result = result.clamp_(-1.0, 1.0) + if result.dtype != original_dtype: + result = result.to(original_dtype) + return result diff --git a/comfy/ldm/seedvr/constants.py b/comfy/ldm/seedvr/constants.py new file mode 100644 index 000000000..12c4b4bef --- /dev/null +++ b/comfy/ldm/seedvr/constants.py @@ -0,0 +1,48 @@ +"""SeedVR2 constants.""" + +# Temporal chunk-size law: the sampler's activation wall is linear in +# T_latent * pixel area (17-cell resolution sweep + T bisection, RTX 5090, 3b fp16): +# max_latent_frames = (free_GiB - RESERVED - K*SIGMA) / (GIB_PER_MPX_FRAME * megapixels) +# RESERVED covers model staging plus fixed CUDA/torch overhead; SIGMA is the measured +# run-to-run spread of the wall; K=4 trades ~10% smaller chunks for ~1e-5 OOM odds. +SEEDVR2_CHUNK_GIB_PER_MPX_FRAME = 0.55 +SEEDVR2_CHUNK_RESERVED_GIB = 8.5 +SEEDVR2_CHUNK_SIGMA_GIB = 0.55 +SEEDVR2_CHUNK_SIGMA_K = 4 + +SEEDVR2_7B_VID_DIM = 3072 +SEEDVR2_OOM_BACKOFF_DIVISOR = 2 +SEEDVR2_DTYPE_BYTES_FLOOR = 4 +SEEDVR2_7B_MLP_CHUNK = 8192 +SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS = 4096 # partial-RoPE application token-chunk. +SEEDVR2_LATENT_CHANNELS = 16 + +SEEDVR2_COLOR_MEM_HEADROOM = 0.75 +SEEDVR2_LAB_SCALE_MULTIPLIER = 13 +SEEDVR2_WAVELET_SCALE_MULTIPLIER = 10 # per-frame byte multiplier, wavelet path. +SEEDVR2_ADAIN_SCALE_MULTIPLIER = 6 + +BYTEDANCE_VAE_SCALING_FACTOR = 0.9152 # configs_3b/main.yaml:57. +BYTEDANCE_VAE_SHIFTING_FACTOR = 0.0 +BYTEDANCE_VAE_CONV_MEM_GIB = 0.5 +BYTEDANCE_VAE_NORM_MEM_GIB = 0.5 +BYTEDANCE_LOGVAR_CLAMP_MIN = -30.0 # video_vae_v3/modules/types.py:28. +BYTEDANCE_LOGVAR_CLAMP_MAX = 20.0 # video_vae_v3/modules/types.py:28. +BYTEDANCE_GN_CHUNKS_FP16 = 4 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp16). +BYTEDANCE_GN_CHUNKS_FP32 = 2 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp32). +BYTEDANCE_BLOCK_OUT_CHANNELS = (128, 256, 512, 512) # s8_c16_t4_inflation_sd3.yaml:7-11. +BYTEDANCE_SLICING_SAMPLE_MIN = 4 # s8_c16_t4_inflation_sd3.yaml:22 (slicing_sample_min_size). +BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE = 4 # infer.py:230 (temporal_downsample_factor); the 4n+1 factor. +BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE = 8 # infer.py:231 (spatial_downsample_factor). +BYTEDANCE_720P_REF_AREA = 45 * 80 # dit_v2/window.py:32 (720p reference area for window scaling). +BYTEDANCE_MAX_TEMPORAL_WINDOW = 30 # dit_v2/window.py:35 (max temporal window frames). +BYTEDANCE_ROPE_MAX_FREQ = 256 # dit_v2/rope.py:31 (pixel-RoPE max frequency). +BYTEDANCE_SINUSOIDAL_DIM = 256 # dit_3b/nadit.py:120 (timestep sinusoidal embed dim). + +ROPE_THETA = 10000 # RoPE base; Su et al., "RoFormer", arXiv:2104.09864. + +CIELAB_DELTA = 6.0 / 29.0 # CIE 15 (delta). +CIELAB_KAPPA = (29.0 / 3.0) ** 3 # CIE 15 (kappa). +D65_WHITE_X = 0.95047 # CIE D65 standard illuminant Xn (Yn = 1). +D65_WHITE_Z = 1.08883 # CIE D65 standard illuminant Zn. +WAVELET_DECOMP_LEVELS = 5 # wavelet color-fix decomposition depth (GIMP/Krita; StableSR). diff --git a/comfy/ldm/seedvr/model.py b/comfy/ldm/seedvr/model.py new file mode 100644 index 000000000..a978698d5 --- /dev/null +++ b/comfy/ldm/seedvr/model.py @@ -0,0 +1,1361 @@ +from dataclasses import dataclass +from typing import Optional, Tuple, Union, List, Dict, Any, Callable +import torch.nn.functional as F +from math import ceil, pi +import torch +from itertools import accumulate, chain +from comfy.ldm.modules.diffusionmodules.model import get_timestep_embedding +from comfy.ldm.seedvr.attention import optimized_var_attention +from torch.nn.modules.utils import _triple +from torch import nn +import math +from comfy.ldm.flux.math import apply_rope1 +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_720P_REF_AREA, + BYTEDANCE_MAX_TEMPORAL_WINDOW, + BYTEDANCE_ROPE_MAX_FREQ, + BYTEDANCE_SINUSOIDAL_DIM, + ROPE_THETA, + SEEDVR2_7B_MLP_CHUNK, + SEEDVR2_7B_VID_DIM, + SEEDVR2_LATENT_CHANNELS, + SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS, +) +import comfy.model_management +import comfy.ops + +class Cache: + def __init__(self, disable=False, prefix="", cache=None): + self.cache = cache if cache is not None else {} + self.disable = disable + self.prefix = prefix + + def __call__(self, key: str, fn: Callable): + if self.disable: + return fn() + + key = self.prefix + key + if key not in self.cache: + result = fn() + self.cache[key] = result + return self.cache[key] + + def namespace(self, namespace: str): + return Cache( + disable=self.disable, + prefix=self.prefix + namespace + ".", + cache=self.cache, + ) + +def repeat_concat( + vid: torch.FloatTensor, # (VL ... c) + txt: torch.FloatTensor, # (TL ... c) + vid_len: torch.LongTensor, # (n*b) + txt_len: torch.LongTensor, # (b) + txt_repeat: List, # (n) +) -> torch.FloatTensor: # (L ... c) + vid = torch.split(vid, vid_len.tolist()) + txt = torch.split(txt, txt_len.tolist()) + txt = [[x] * n for x, n in zip(txt, txt_repeat)] + txt = list(chain(*txt)) + return torch.cat(list(chain(*zip(vid, txt)))) + +def repeat_concat_idx( + vid_len: torch.LongTensor, # (n*b) + txt_len: torch.LongTensor, # (b) + txt_repeat: torch.LongTensor, # (n) +) -> Tuple[ + Callable, + Callable, +]: + device = vid_len.device + vid_idx = torch.arange(vid_len.sum(), device=device) + txt_idx = torch.arange(len(vid_idx), len(vid_idx) + txt_len.sum(), device=device) + txt_repeat_list = txt_repeat.tolist() + tgt_idx = repeat_concat(vid_idx, txt_idx, vid_len, txt_len, txt_repeat_list) + src_idx = torch.argsort(tgt_idx) + txt_idx_len = len(tgt_idx) - len(vid_idx) + repeat_txt_len = (txt_len * txt_repeat).tolist() + + def unconcat_coalesce(all): + vid_out, txt_out = all[src_idx].split([len(vid_idx), txt_idx_len]) + txt_out_coalesced = [] + for txt, repeat_time in zip(txt_out.split(repeat_txt_len), txt_repeat_list): + txt = txt.reshape(-1, repeat_time, *txt.shape[1:]).mean(1) + txt_out_coalesced.append(txt) + return vid_out, torch.cat(txt_out_coalesced) + + return ( + lambda vid, txt: torch.cat([vid, txt])[tgt_idx], + lambda all: unconcat_coalesce(all), + ) + +def cumulative_lengths(lengths): + return [0, *accumulate(lengths)] + + +@dataclass +class MMArg: + vid: Any + txt: Any + +def get_args(key: str, args: List[Any]) -> List[Any]: + return [getattr(v, key) if isinstance(v, MMArg) else v for v in args] + + +def get_kwargs(key: str, kwargs: Dict[str, Any]) -> Dict[str, Any]: + return {k: getattr(v, key) if isinstance(v, MMArg) else v for k, v in kwargs.items()} + + +def get_window_op(name: str): + if name == "720pwin_by_size_bysize": + return make_720Pwindows_bysize + if name == "720pswin_by_size_bysize": + return make_shifted_720Pwindows_bysize + raise ValueError(f"Unknown windowing method: {name}") + + +def make_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]): + t, h, w = size + resized_nt, resized_nh, resized_nw = num_windows + scale = math.sqrt(BYTEDANCE_720P_REF_AREA / (h * w)) + resized_h, resized_w = round(h * scale), round(w * scale) + wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) + wt = ceil(min(t, BYTEDANCE_MAX_TEMPORAL_WINDOW) / resized_nt) + nt, nh, nw = ceil(t / wt), ceil(h / wh), ceil(w / ww) + return [ + ( + slice(it * wt, min((it + 1) * wt, t)), + slice(ih * wh, min((ih + 1) * wh, h)), + slice(iw * ww, min((iw + 1) * ww, w)), + ) + for iw in range(nw) + if min((iw + 1) * ww, w) > iw * ww + for ih in range(nh) + if min((ih + 1) * wh, h) > ih * wh + for it in range(nt) + if min((it + 1) * wt, t) > it * wt + ] + +def make_shifted_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]): + t, h, w = size + resized_nt, resized_nh, resized_nw = num_windows + scale = math.sqrt(BYTEDANCE_720P_REF_AREA / (h * w)) + resized_h, resized_w = round(h * scale), round(w * scale) + wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) + wt = ceil(min(t, BYTEDANCE_MAX_TEMPORAL_WINDOW) / resized_nt) + + st, sh, sw = ( + 0.5 if wt < t else 0, + 0.5 if wh < h else 0, + 0.5 if ww < w else 0, + ) + nt, nh, nw = ceil((t - st) / wt), ceil((h - sh) / wh), ceil((w - sw) / ww) + nt, nh, nw = ( + nt + 1 if st > 0 else 1, + nh + 1 if sh > 0 else 1, + nw + 1 if sw > 0 else 1, + ) + return [ + ( + slice(max(int((it - st) * wt), 0), min(int((it - st + 1) * wt), t)), + slice(max(int((ih - sh) * wh), 0), min(int((ih - sh + 1) * wh), h)), + slice(max(int((iw - sw) * ww), 0), min(int((iw - sw + 1) * ww), w)), + ) + for iw in range(nw) + if min(int((iw - sw + 1) * ww), w) > max(int((iw - sw) * ww), 0) + for ih in range(nh) + if min(int((ih - sh + 1) * wh), h) > max(int((ih - sh) * wh), 0) + for it in range(nt) + if min(int((it - st + 1) * wt), t) > max(int((it - st) * wt), 0) + ] + +class RotaryEmbedding(nn.Module): + def __init__( + self, + dim, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + ): + super().__init__() + + self.freqs_for = freqs_for + + if freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + else: + raise ValueError(f"Unknown rotary frequency type: {freqs_for}") + + self.register_buffer("freqs", freqs) + + @property + def device(self): + return self.freqs.device + + def get_axial_freqs( + self, + *dims, + offsets = None + ): + Colon = slice(None) + all_freqs = [] + + if exists(offsets): + if len(offsets) != len(dims): + raise ValueError(f"SeedVR2 rotary offsets length must match dims length, got {len(offsets)} and {len(dims)}.") + + for ind, dim in enumerate(dims): + + offset = 0 + if exists(offsets): + offset = offsets[ind] + + if self.freqs_for == 'pixel': + pos = torch.linspace(-1, 1, steps = dim, device = self.device) + else: + pos = torch.arange(dim, device = self.device) + + pos = pos + offset + + freqs = self.forward(pos) + + all_axis = [None] * len(dims) + all_axis[ind] = Colon + + new_axis_slice = (Ellipsis, *all_axis, Colon) + all_freqs.append(freqs[new_axis_slice]) + + all_freqs = torch.broadcast_tensors(*all_freqs) + return torch.cat(all_freqs, dim = -1) + + def forward( + self, + t, + ): + freqs = self.freqs + + freqs = torch.einsum('..., f -> ... f', t.type(freqs.dtype), freqs) + freqs = freqs.unsqueeze(-1).expand(*freqs.shape, 2).flatten(-2) + + return freqs + +class RotaryEmbeddingBase(nn.Module): + def __init__(self, dim: int, rope_dim: int): + super().__init__() + self.rope = RotaryEmbedding( + dim=dim // rope_dim, + freqs_for="pixel", + max_freq=BYTEDANCE_ROPE_MAX_FREQ, + ) + + def get_axial_freqs(self, *dims): + return self.rope.get_axial_freqs(*dims) + + +class RotaryEmbedding3d(RotaryEmbeddingBase): + def __init__(self, dim: int): + super().__init__(dim, rope_dim=3) + self.mm = False + + +class NaRotaryEmbedding3d(RotaryEmbedding3d): + def forward( + self, + q: torch.FloatTensor, + k: torch.FloatTensor, + shape: torch.LongTensor, + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + freqs = cache("rope_freqs_3d", lambda: self.get_freqs(shape)) + freqs = freqs.to(device=q.device) + q = q.transpose(0, 1) + k = k.transpose(0, 1) + q = _apply_seedvr2_rotary_emb(freqs, q.float()).to(q.dtype) + k = _apply_seedvr2_rotary_emb(freqs, k.float()).to(k.dtype) + q = q.transpose(0, 1) + k = k.transpose(0, 1) + return q, k + + @torch._dynamo.disable + def get_freqs( + self, + shape: torch.LongTensor, + ) -> torch.Tensor: + # Primary provenance: ByteDance-Seed/SeedVR models/dit/rope.py builds + # 7B pixel RoPE with the interleaved-angle convention, not Comfy's + # Flux freqs_cis matrix. + plain_rope = RotaryEmbedding( + dim=self.rope.freqs.numel() * 2, + freqs_for="pixel", + max_freq=BYTEDANCE_ROPE_MAX_FREQ, + ) + plain_rope = plain_rope.to(self.rope.device) + freq_list = [] + for f, h, w in shape.tolist(): + freqs = plain_rope.get_axial_freqs(f, h, w) + freq_list.append(freqs.view(-1, freqs.size(-1))) + return torch.cat(freq_list, dim=0) + + +class MMRotaryEmbeddingBase(RotaryEmbeddingBase): + def __init__(self, dim: int, rope_dim: int): + super().__init__(dim, rope_dim) + self.rope = RotaryEmbedding( + dim=dim // rope_dim, + freqs_for="lang", + theta=ROPE_THETA, + ) + self.mm = True + +def slice_at_dim(t, dim_slice: slice, *, dim): + dim += (t.ndim if dim < 0 else 0) + colons = [slice(None)] * t.ndim + colons[dim] = dim_slice + return t[tuple(colons)] + +def rotate_half(x): + x = x.reshape(*x.shape[:-1], x.shape[-1] // 2, 2) + x1, x2 = x.unbind(dim = -1) + x = torch.stack((-x2, x1), dim = -1) + return x.flatten(-2) +def exists(val): + return val is not None + +def _apply_seedvr2_rotary_emb( + freqs: torch.Tensor, + t: torch.Tensor, + start_index: int = 0, + scale: float = 1.0, + seq_dim: int = -2, + freqs_seq_dim: int | None = None, +) -> torch.Tensor: + dtype = t.dtype + if freqs_seq_dim is None and (freqs.ndim == 2 or t.ndim == 3): + freqs_seq_dim = 0 + + if t.ndim == 3 or freqs_seq_dim is not None: + seq_len = t.shape[seq_dim] + freqs = slice_at_dim(freqs, slice(-seq_len, None), dim=freqs_seq_dim) + + rot_feats = freqs.shape[-1] + end_index = start_index + rot_feats + + t_left = t[..., :start_index] + t_middle = t[..., start_index:end_index] + t_right = t[..., end_index:] + + freqs = freqs.to(device=t_middle.device, dtype=t_middle.dtype) + cos = freqs.cos() * scale + sin = freqs.sin() * scale + t_middle = (t_middle * cos) + (rotate_half(t_middle) * sin) + return torch.cat((t_left, t_middle, t_right), dim=-1).to(dtype) + +def _to_flux_freqs_cis(freqs_interleaved: torch.Tensor) -> torch.Tensor: + angles = freqs_interleaved[..., ::2].float() + cos = torch.cos(angles) + sin = torch.sin(angles) + out = torch.stack([cos, -sin, sin, cos], dim=-1) + return out.reshape(*out.shape[:-1], 2, 2) + + +def _apply_rope1_partial(t: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: + out = t.clone() if t.requires_grad or comfy.model_management.in_training else t + rot_d = 2 * freqs_cis.shape[-3] + seq_len = out.shape[-2] + for start in range(0, seq_len, SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS): + end = min(start + SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS, seq_len) + freqs_chunk = freqs_cis[start:end] + if rot_d == out.shape[-1]: + out[..., start:end, :] = apply_rope1(out[..., start:end, :], freqs_chunk).to(out.dtype) + else: + out[..., start:end, :rot_d] = apply_rope1(out[..., start:end, :rot_d], freqs_chunk).to(out.dtype) + return out + + +class NaMMRotaryEmbedding3d(MMRotaryEmbeddingBase): + def __init__(self, dim: int): + super().__init__(dim, rope_dim=3) + + def forward( + self, + vid_q: torch.FloatTensor, # L h d + vid_k: torch.FloatTensor, # L h d + vid_shape: torch.LongTensor, # B 3 + txt_q: torch.FloatTensor, # L h d + txt_k: torch.FloatTensor, # L h d + txt_shape: torch.LongTensor, # B 1 + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_freqs, txt_freqs = cache( + "mmrope_freqs_3d", + lambda: self.get_freqs(vid_shape, txt_shape), + ) + target_device = vid_q.device + if vid_freqs.device != target_device: + vid_freqs = vid_freqs.to(target_device) + if txt_freqs.device != target_device: + txt_freqs = txt_freqs.to(target_device) + vid_q = vid_q.transpose(0, 1) + vid_k = vid_k.transpose(0, 1) + vid_q = _apply_rope1_partial(vid_q, vid_freqs) + vid_k = _apply_rope1_partial(vid_k, vid_freqs) + vid_q = vid_q.transpose(0, 1) + vid_k = vid_k.transpose(0, 1) + + txt_q = txt_q.transpose(0, 1) + txt_k = txt_k.transpose(0, 1) + txt_q = _apply_rope1_partial(txt_q, txt_freqs) + txt_k = _apply_rope1_partial(txt_k, txt_freqs) + txt_q = txt_q.transpose(0, 1) + txt_k = txt_k.transpose(0, 1) + return vid_q, vid_k, txt_q, txt_k + + @torch._dynamo.disable # Disable compilation: .tolist() is data-dependent and causes graph breaks + def get_freqs( + self, + vid_shape: torch.LongTensor, + txt_shape: torch.LongTensor, + ) -> Tuple[ + torch.Tensor, + torch.Tensor, + ]: + + max_temporal = 0 + max_height = 0 + max_width = 0 + max_txt_len = 0 + + for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()): + max_temporal = max(max_temporal, l + f) + max_height = max(max_height, h) + max_width = max(max_width, w) + max_txt_len = max(max_txt_len, l) + + autocast_device = "cuda" if torch.cuda.is_available() else "cpu" + with torch.amp.autocast(autocast_device, enabled=False): + vid_freqs = self.get_axial_freqs( + max_temporal + 16, + max_height + 4, + max_width + 4, + ).float() + txt_freqs = self.get_axial_freqs(max_txt_len + 16) + + vid_freq_list, txt_freq_list = [], [] + for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()): + vid_freq = vid_freqs[l : l + f, :h, :w].reshape(-1, vid_freqs.size(-1)) + txt_freq = txt_freqs[:l].repeat(1, 3).reshape(-1, vid_freqs.size(-1)) + vid_freq_list.append(vid_freq) + txt_freq_list.append(txt_freq) + vid_freqs_interleaved = torch.cat(vid_freq_list, dim=0) + txt_freqs_interleaved = torch.cat(txt_freq_list, dim=0) + + return _to_flux_freqs_cis(vid_freqs_interleaved), _to_flux_freqs_cis(txt_freqs_interleaved) + +class MMModule(nn.Module): + def __init__( + self, + module: Callable[..., nn.Module], + *args, + shared_weights: bool = False, + vid_only: bool = False, + **kwargs, + ): + super().__init__() + self.shared_weights = shared_weights + self.vid_only = vid_only + if self.shared_weights: + if get_args("vid", args) != get_args("txt", args): + raise ValueError("SeedVR2 shared MMModule requires matching vid/txt args.") + if get_kwargs("vid", kwargs) != get_kwargs("txt", kwargs): + raise ValueError("SeedVR2 shared MMModule requires matching vid/txt kwargs.") + self.all = module(*get_args("vid", args), **get_kwargs("vid", kwargs)) + else: + self.vid = module(*get_args("vid", args), **get_kwargs("vid", kwargs)) + self.txt = ( + module(*get_args("txt", args), **get_kwargs("txt", kwargs)) + if not vid_only + else None + ) + + def forward( + self, + vid: torch.FloatTensor, + txt: torch.FloatTensor, + *args, + **kwargs, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_module = self.vid if not self.shared_weights else self.all + vid = vid_module(vid, *get_args("vid", args), **get_kwargs("vid", kwargs)) + if not self.vid_only: + txt_module = self.txt if not self.shared_weights else self.all + txt = txt.to(device=vid.device, dtype=vid.dtype) + txt = txt_module(txt, *get_args("txt", args), **get_kwargs("txt", kwargs)) + return vid, txt + +def get_na_rope(rope_type: Optional[str], dim: int): + if rope_type is None: + return None + if rope_type == "rope3d": + return NaRotaryEmbedding3d(dim=dim) + if rope_type == "mmrope3d": + return NaMMRotaryEmbedding3d(dim=dim) + raise ValueError(f"Unknown SeedVR2 rope type: {rope_type}") + +class NaMMAttention(nn.Module): + def __init__( + self, + vid_dim: int, + txt_dim: int, + heads: int, + head_dim: int, + qk_bias: bool, + qk_norm, + qk_norm_eps: float, + rope_type: Optional[str], + rope_dim: int, + shared_weights: bool, + device, dtype, operations, + ): + super().__init__() + dim = MMArg(vid_dim, txt_dim) + self.heads = heads + inner_dim = heads * head_dim + qkv_dim = inner_dim * 3 + self.head_dim = head_dim + self.proj_qkv = MMModule( + operations.Linear, dim, qkv_dim, bias=qk_bias, shared_weights=shared_weights, device=device, dtype=dtype + ) + self.proj_out = MMModule(operations.Linear, inner_dim, dim, shared_weights=shared_weights, device=device, dtype=dtype) + self.norm_q = MMModule( + qk_norm, + normalized_shape=head_dim, + eps=qk_norm_eps, + elementwise_affine=True, + shared_weights=shared_weights, + device=device, dtype=dtype + ) + self.norm_k = MMModule( + qk_norm, + normalized_shape=head_dim, + eps=qk_norm_eps, + elementwise_affine=True, + shared_weights=shared_weights, + device=device, dtype=dtype + ) + + + self.rope = get_na_rope(rope_type=rope_type, dim=rope_dim) + +def window( + hid: torch.FloatTensor, # (L c) + hid_shape: torch.LongTensor, # (b n) + window_fn: Callable[[torch.Tensor], List[torch.Tensor]], +): + hid = unflatten(hid, hid_shape) + hid = list(map(window_fn, hid)) + hid_windows_list = [len(x) for x in hid] + hid_windows = torch.as_tensor(hid_windows_list, device=hid_shape.device) + hid = list(chain(*hid)) + hid_len_list = [math.prod(x.shape[:-1]) for x in hid] + hid, hid_shape = flatten(hid) + return hid, hid_shape, hid_windows, hid_len_list, hid_windows_list + +def window_idx( + hid_shape: torch.LongTensor, # (b n) + window_fn: Callable[[torch.Tensor], List[torch.Tensor]], +): + hid_idx = torch.arange(hid_shape.prod(-1).sum(), device=hid_shape.device).unsqueeze(-1) + tgt_idx, tgt_shape, tgt_windows, tgt_len_list, tgt_windows_list = window(hid_idx, hid_shape, window_fn) + tgt_idx = tgt_idx.squeeze(-1) + src_idx = torch.argsort(tgt_idx) + return ( + lambda hid: torch.index_select(hid, 0, tgt_idx), + lambda hid: torch.index_select(hid, 0, src_idx), + tgt_shape, + tgt_windows, + tgt_len_list, + tgt_windows_list, + ) + +class NaSwinAttention(NaMMAttention): + def __init__( + self, + *args, + window: Union[int, Tuple[int, int, int]], + window_method: str, + version: bool = False, + **kwargs, + ): + super().__init__(*args, **kwargs) + self.version_7b = version + self.window = _triple(window) + self.window_method = window_method + if not all(isinstance(v, int) and v >= 0 for v in self.window): + raise ValueError(f"SeedVR2 window must contain non-negative integers, got {self.window}.") + + self.window_op = get_window_op(window_method) + + def forward( + self, + vid: torch.FloatTensor, # l c + txt: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, # b 3 + txt_shape: torch.LongTensor, # b 1 + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + + vid_qkv, txt_qkv = self.proj_qkv(vid, txt) + + cache_win = cache.namespace(f"{self.window_method}_{self.window}_sd3") + + def make_window(x: torch.Tensor): + t, h, w, _ = x.shape + window_slices = self.window_op((t, h, w), self.window) + return [x[st, sh, sw] for (st, sh, sw) in window_slices] + + window_partition, window_reverse, window_shape, window_count, vid_len_win_list, window_count_list = cache_win( + "win_transform", + lambda: window_idx(vid_shape, make_window), + ) + vid_qkv_win = window_partition(vid_qkv) + + vid_qkv_win = vid_qkv_win.reshape(vid_qkv_win.shape[0], 3, self.heads, self.head_dim) + txt_qkv = txt_qkv.reshape(txt_qkv.shape[0], 3, self.heads, self.head_dim) + + vid_q, vid_k, vid_v = vid_qkv_win.unbind(1) + txt_q, txt_k, txt_v = txt_qkv.unbind(1) + + vid_q, txt_q = self.norm_q(vid_q, txt_q) + vid_k, txt_k = self.norm_k(vid_k, txt_k) + + txt_len = cache("txt_len", lambda: txt_shape.prod(-1)) + + vid_len_win = cache_win("vid_len", lambda: window_shape.prod(-1)) + txt_len = txt_len.to(window_count.device) + + if self.rope: + if self.version_7b: + vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + elif self.rope.mm: + _, num_h, _ = txt_q.shape + txt_q_repeat = txt_q.flatten(1, 2) + txt_q_repeat = unflatten(txt_q_repeat, txt_shape) + txt_q_repeat = [[x] * n for x, n in zip(txt_q_repeat, window_count_list)] + txt_q_repeat = list(chain(*txt_q_repeat)) + txt_q_repeat, txt_shape_repeat = flatten(txt_q_repeat) + txt_q_repeat = txt_q_repeat.reshape(txt_q_repeat.shape[0], num_h, self.head_dim) + + txt_k_repeat = txt_k.flatten(1, 2) + txt_k_repeat = unflatten(txt_k_repeat, txt_shape) + txt_k_repeat = [[x] * n for x, n in zip(txt_k_repeat, window_count_list)] + txt_k_repeat = list(chain(*txt_k_repeat)) + txt_k_repeat, _ = flatten(txt_k_repeat) + txt_k_repeat = txt_k_repeat.reshape(txt_k_repeat.shape[0], num_h, self.head_dim) + + vid_q, vid_k, txt_q, txt_k = self.rope( + vid_q, vid_k, window_shape, txt_q_repeat, txt_k_repeat, txt_shape_repeat, cache_win + ) + else: + vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + + txt_len_win_list = cache_win( + "txt_len_list", + lambda: [txt_len for txt_len, window_count in zip(txt_len.tolist(), window_count_list) for _ in range(window_count)], + ) + all_len_win = cache_win("all_len", lambda: [vid_len + txt_len for vid_len, txt_len in zip(vid_len_win_list, txt_len_win_list)]) + concat_win, unconcat_win = cache_win( + "mm_pnp", lambda: repeat_concat_idx(vid_len_win, txt_len, window_count) + ) + out = optimized_var_attention( + q=concat_win(vid_q, txt_q), + k=concat_win(vid_k, txt_k), + v=concat_win(vid_v, txt_v), + heads=self.heads, skip_reshape=True, skip_output_reshape=True, + cu_seqlens_q=cache_win("vid_seqlens_q", lambda: cumulative_lengths(all_len_win)), + cu_seqlens_k=cache_win("vid_seqlens_k", lambda: cumulative_lengths(all_len_win)), + ) + vid_out, txt_out = unconcat_win(out) + + vid_out = vid_out.flatten(1, 2) + txt_out = txt_out.flatten(1, 2) + vid_out = window_reverse(vid_out) + + vid_out, txt_out = self.proj_out(vid_out, txt_out) + + return vid_out, txt_out + +class MLP(nn.Module): + def __init__( + self, + dim: int, + expand_ratio: int, + device, dtype, operations + ): + super().__init__() + self.proj_in = operations.Linear(dim, dim * expand_ratio, device=device, dtype=dtype) + self.act = nn.GELU("tanh") + self.proj_out = operations.Linear(dim * expand_ratio, dim, device=device, dtype=dtype) + + def forward(self, x: torch.FloatTensor) -> torch.FloatTensor: + x = self.proj_in(x) + x = self.act(x) + x = self.proj_out(x) + return x + + +class SwiGLUMLP(nn.Module): + def __init__( + self, + dim: int, + expand_ratio: int, + multiple_of: int = 256, + device=None, dtype=None, operations=None + ): + super().__init__() + hidden_dim = int(2 * dim * expand_ratio / 3) + hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of) + self.proj_in_gate = operations.Linear(dim, hidden_dim, bias=False, device=device, dtype=dtype) + self.proj_out = operations.Linear(hidden_dim, dim, bias=False, device=device, dtype=dtype) + self.proj_in = operations.Linear(dim, hidden_dim, bias=False, device=device, dtype=dtype) + + def forward(self, x: torch.FloatTensor) -> torch.FloatTensor: + return self.proj_out(F.silu(self.proj_in_gate(x)) * self.proj_in(x)) + +def get_mlp(mlp_type: Optional[str] = "normal"): + if mlp_type == "normal": + return MLP + if mlp_type == "swiglu": + return SwiGLUMLP + raise ValueError(f"Unknown SeedVR2 MLP type: {mlp_type}") + +class NaMMSRTransformerBlock(nn.Module): + def __init__( + self, + *, + vid_dim: int, + txt_dim: int, + emb_dim: int, + heads: int, + head_dim: int, + expand_ratio: int, + norm, + norm_eps: float, + ada, + qk_bias: bool, + qk_norm, + mlp_type: str, + shared_weights: bool, + rope_type: str, + rope_dim: int, + is_last_layer: bool, + window: Union[int, Tuple[int, int, int]], + window_method: str, + version: bool, + device, dtype, operations, + ): + super().__init__() + dim = MMArg(vid_dim, txt_dim) + self.attn_norm = MMModule(norm, normalized_shape=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights, device=device, dtype=dtype) + + self.attn = NaSwinAttention( + vid_dim=vid_dim, + txt_dim=txt_dim, + heads=heads, + head_dim=head_dim, + qk_bias=qk_bias, + qk_norm=qk_norm, + qk_norm_eps=norm_eps, + rope_type=rope_type, + rope_dim=rope_dim, + shared_weights=shared_weights, + window=window, + window_method=window_method, + version=version, + device=device, dtype=dtype, operations=operations + ) + + self.mlp_norm = MMModule(norm, normalized_shape=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights, vid_only=is_last_layer, device=device, dtype=dtype) + self.mlp = MMModule( + get_mlp(mlp_type), + dim=dim, + expand_ratio=expand_ratio, + shared_weights=shared_weights, + vid_only=is_last_layer, + device=device, dtype=dtype, operations=operations + ) + self.ada = MMModule(ada, dim=dim, emb_dim=emb_dim, layers=["attn", "mlp"], shared_weights=shared_weights, vid_only=is_last_layer, device=device, dtype=dtype) + self.is_last_layer = is_last_layer + self.version = version + + def _seedvr2_7b_mlp( + self, + vid: torch.FloatTensor, + txt: torch.FloatTensor, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_module = self.mlp.vid if not self.mlp.shared_weights else self.mlp.all + if comfy.model_management.in_training or vid.requires_grad: + vid = torch.cat([vid_module(chunk) for chunk in vid.split(SEEDVR2_7B_MLP_CHUNK, dim=0)], dim=0) + else: + vid_out = None + offset = 0 + for chunk in vid.split(SEEDVR2_7B_MLP_CHUNK, dim=0): + chunk_out = vid_module(chunk) + if vid_out is None: + vid_out = chunk_out.new_empty((vid.shape[0], *chunk_out.shape[1:])) + vid_out[offset:offset + chunk_out.shape[0]] = chunk_out + offset += chunk_out.shape[0] + vid = vid_out + if not self.mlp.vid_only: + txt_module = self.mlp.txt if not self.mlp.shared_weights else self.mlp.all + txt = txt.to(device=vid.device, dtype=vid.dtype) + txt = txt_module(txt) + return vid, txt + + def forward( + self, + vid: torch.FloatTensor, # l c + txt: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, # b 3 + txt_shape: torch.LongTensor, # b 1 + emb: torch.FloatTensor, + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + torch.LongTensor, + torch.LongTensor, + ]: + hid_len = MMArg( + cache("vid_len", lambda: vid_shape.prod(-1)), + cache("txt_len", lambda: txt_shape.prod(-1)), + ) + ada_kwargs = { + "emb": emb, + "hid_len": hid_len, + "cache": cache, + "branch_tag": MMArg("vid", "txt"), + } + + vid_attn, txt_attn = self.attn_norm(vid, txt) + vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="in", **ada_kwargs) + vid_attn, txt_attn = self.attn(vid_attn, txt_attn, vid_shape, txt_shape, cache) + vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="out", **ada_kwargs) + vid_attn, txt_attn = (vid_attn + vid), (txt_attn + txt) + + vid_mlp, txt_mlp = self.mlp_norm(vid_attn, txt_attn) + vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="in", **ada_kwargs) + if self.version: + vid_mlp, txt_mlp = self._seedvr2_7b_mlp(vid_mlp, txt_mlp) + else: + vid_mlp, txt_mlp = self.mlp(vid_mlp, txt_mlp) + vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="out", **ada_kwargs) + vid_mlp, txt_mlp = (vid_mlp + vid_attn), (txt_mlp + txt_attn) + + return vid_mlp, txt_mlp, vid_shape, txt_shape + +class PatchOut(nn.Module): + def __init__( + self, + out_channels: int, + patch_size: Union[int, Tuple[int, int, int]], + dim: int, + device, dtype, operations + ): + super().__init__() + t, h, w = _triple(patch_size) + self.patch_size = t, h, w + self.proj = operations.Linear(dim, out_channels * t * h * w, device=device, dtype=dtype) + + def forward( + self, + vid: torch.Tensor, + ) -> torch.Tensor: + t, h, w = self.patch_size + vid = self.proj(vid) + b, T, H, W, channels = vid.shape + c = channels // (t * h * w) + vid = vid.view(b, T, H, W, t, h, w, c).permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(b, c, T * t, H * h, W * w) + if t > 1: + vid = vid[:, :, (t - 1) :] + return vid + +class NaPatchOut(PatchOut): + def forward( + self, + vid: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, + cache: Optional[Cache] = None, + vid_shape_before_patchify = None + ) -> Tuple[ + torch.FloatTensor, + torch.LongTensor, + ]: + if cache is None: + cache = Cache(disable=True) + + t, h, w = self.patch_size + vid = self.proj(vid) + + if not (t == h == w == 1): + vid = unflatten(vid, vid_shape) + for i in range(len(vid)): + T, H, W, channels = vid[i].shape + c = channels // (t * h * w) + vid[i] = vid[i].view(T, H, W, t, h, w, c).permute(0, 3, 1, 4, 2, 5, 6).reshape(T * t, H * h, W * w, c) + if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: + vid[i] = vid[i][(t - vid_shape_before_patchify[i, 0] % t) :] + vid, vid_shape = flatten(vid) + + return vid, vid_shape + +class PatchIn(nn.Module): + def __init__( + self, + in_channels: int, + patch_size: Union[int, Tuple[int, int, int]], + dim: int, + device, dtype, operations + ): + super().__init__() + t, h, w = _triple(patch_size) + self.patch_size = t, h, w + self.proj = operations.Linear(in_channels * t * h * w, dim, device=device, dtype=dtype) + + def forward( + self, + vid: torch.Tensor, + ) -> torch.Tensor: + t, h, w = self.patch_size + if t > 1: + if vid.size(2) % t != 1: + raise ValueError( + f"SeedVR2 patch input temporal size must satisfy T % {t} == 1, got {vid.size(2)}." + ) + vid = torch.cat([vid[:, :, :1]] * (t - 1) + [vid], dim=2) + b, c, Tt, Hh, Ww = vid.shape + vid = vid.view(b, c, Tt // t, t, Hh // h, h, Ww // w, w).permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(b, Tt // t, Hh // h, Ww // w, t * h * w * c) + vid = self.proj(vid) + return vid + +class NaPatchIn(PatchIn): + def forward( + self, + vid: torch.Tensor, # l c + vid_shape: torch.LongTensor, + cache: Optional[Cache] = None, + ) -> torch.Tensor: + if cache is None: + cache = Cache(disable=True) + cache = cache.namespace("patch") + vid_shape_before_patchify = cache("vid_shape_before_patchify", lambda: vid_shape) + t, h, w = self.patch_size + if not (t == h == w == 1): + vid = unflatten(vid, vid_shape) + for i in range(len(vid)): + if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: + vid[i] = torch.cat([vid[i][:1]] * (t - vid[i].size(0) % t) + [vid[i]], dim=0) + Tt, Hh, Ww, c = vid[i].shape + vid[i] = vid[i].view(Tt // t, t, Hh // h, h, Ww // w, w, c).permute(0, 2, 4, 1, 3, 5, 6).reshape(Tt // t, Hh // h, Ww // w, t * h * w * c) + vid, vid_shape = flatten(vid) + + vid = self.proj(vid) + return vid, vid_shape + +def expand_dims(x: torch.Tensor, dim: int, ndim: int): + shape = x.shape + shape = shape[:dim] + (1,) * (ndim - len(shape)) + shape[dim:] + return x.reshape(shape) + + +class AdaSingle(nn.Module): + def __init__( + self, + dim: int, + emb_dim: int, + layers: List[str], + modes: Tuple[str, ...] = ("in", "out"), + device = None, dtype = None, + ): + if emb_dim != 6 * dim: + raise ValueError(f"SeedVR2 AdaSingle requires emb_dim == 6 * dim, got emb_dim={emb_dim}, dim={dim}.") + super().__init__() + self.dim = dim + self.emb_dim = emb_dim + self.layers = layers + + param_kwargs = {"device": device, "dtype": dtype} + + for l in layers: + if "in" in modes: + self.register_parameter(f"{l}_shift", nn.Parameter(torch.empty(dim, **param_kwargs))) + self.register_parameter(f"{l}_scale", nn.Parameter(torch.empty(dim, **param_kwargs))) + if "out" in modes: + self.register_parameter(f"{l}_gate", nn.Parameter(torch.empty(dim, **param_kwargs))) + + def forward( + self, + hid: torch.FloatTensor, # b ... c + emb: torch.FloatTensor, # b d + layer: str, + mode: str, + cache: Optional[Cache] = None, + branch_tag: str = "", + hid_len: Optional[torch.LongTensor] = None, # b + ) -> torch.FloatTensor: + if cache is None: + cache = Cache(disable=True) + idx = self.layers.index(layer) + emb = emb.reshape(emb.shape[0], -1, len(self.layers), 3)[:, :, idx, :] + emb = expand_dims(emb, 1, hid.ndim + 1) + + if hid_len is not None: + emb = cache( + f"emb_repeat_{idx}_{branch_tag}", + lambda: torch.repeat_interleave(emb, hid_len, dim=0), + ) + + shiftA, scaleA, gateA = emb.unbind(-1) + shiftB, scaleB, gateB = ( + getattr(self, f"{layer}_shift", None), + getattr(self, f"{layer}_scale", None), + getattr(self, f"{layer}_gate", None), + ) + + if mode == "in": + shiftB = comfy.ops.cast_to_input(shiftB, hid) + scaleB = comfy.ops.cast_to_input(scaleB, hid) + return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) + if mode == "out": + if gateB is not None: + gateB = comfy.ops.cast_to_input(gateB, hid) + return hid.mul_(gateA + gateB) + else: + return hid.mul_(gateA) + + raise ValueError(f"Unknown AdaSingle mode: {mode}") + + +class TimeEmbedding(nn.Module): + def __init__( + self, + sinusoidal_dim: int, + hidden_dim: int, + output_dim: int, + device, dtype, operations + ): + super().__init__() + self.sinusoidal_dim = sinusoidal_dim + self.proj_in = operations.Linear(sinusoidal_dim, hidden_dim, device=device, dtype=dtype) + self.proj_hid = operations.Linear(hidden_dim, hidden_dim, device=device, dtype=dtype) + self.proj_out = operations.Linear(hidden_dim, output_dim, device=device, dtype=dtype) + self.act = nn.SiLU() + + def forward( + self, + timestep: Union[int, float, torch.IntTensor, torch.FloatTensor], + device: torch.device, + dtype: torch.dtype, + ) -> torch.FloatTensor: + if not torch.is_tensor(timestep): + timestep = torch.tensor([timestep], device=device, dtype=dtype) + if timestep.ndim == 0: + timestep = timestep[None] + + emb = get_timestep_embedding( + timesteps=timestep, + embedding_dim=self.sinusoidal_dim, + flip_sin_to_cos=False, + downscale_freq_shift=0, + ).to(dtype) + emb = self.proj_in(emb) + emb = self.act(emb) + emb = self.proj_hid(emb) + emb = self.act(emb) + emb = self.proj_out(emb) + return emb + +def flatten( + hid: List[torch.FloatTensor], # List of (*** c) +) -> Tuple[ + torch.FloatTensor, # (L c) + torch.LongTensor, # (b n) +]: + if len(hid) == 0: + raise ValueError("SeedVR2 flatten requires at least one tensor.") + shape = torch.as_tensor([x.shape[:-1] for x in hid], device=hid[0].device) + hid = torch.cat([x.flatten(0, -2) for x in hid]) + return hid, shape + + +def unflatten( + hid: torch.FloatTensor, # (L c) or (L ... c) + hid_shape: torch.LongTensor, # (b n) +) -> List[torch.Tensor]: # List of (*** c) or (*** ... c) + hid_len = hid_shape.prod(-1) + hid = hid.split(hid_len.tolist()) + hid = [x.unflatten(0, s.tolist()) for x, s in zip(hid, hid_shape)] + return hid + +class NaDiT(nn.Module): + + def __init__( + self, + norm_eps, + num_layers, + mlp_type, + vid_in_channels = 33, + vid_out_channels = SEEDVR2_LATENT_CHANNELS, + vid_dim = 2560, + txt_in_dim = 5120, + heads = 20, + head_dim = 128, + mm_layers = 10, + expand_ratio = 4, + qk_bias = False, + patch_size = (1, 2, 2), + rope_dim = 128, + rope_type = "mmrope3d", + vid_out_norm: Optional[str] = None, + image_model = None, + device = None, + dtype = None, + operations = None, + ): + if image_model not in (None, "seedvr2"): + raise ValueError(f"SeedVR2 NaDiT expected image_model='seedvr2', got {image_model!r}.") + self._7b_version = vid_dim == SEEDVR2_7B_VID_DIM + if self._7b_version: + rope_type = "rope3d" + self.dtype = dtype + factory_kwargs = {"device": device, "dtype": dtype} + window_method = num_layers // 2 * ["720pwin_by_size_bysize","720pswin_by_size_bysize"] + txt_dim = vid_dim + emb_dim = vid_dim * 6 + window = num_layers * [(4,3,3)] + ada = AdaSingle + norm = operations.RMSNorm + qk_norm = operations.RMSNorm + super().__init__() + self.register_buffer("positive_conditioning", torch.empty((58, 5120), device=device, dtype=dtype)) + self.register_buffer("negative_conditioning", torch.empty((64, 5120), device=device, dtype=dtype)) + self.vid_in = NaPatchIn( + in_channels=vid_in_channels, + patch_size=patch_size, + dim=vid_dim, + device=device, dtype=dtype, operations=operations + ) + self.txt_in = ( + operations.Linear(txt_in_dim, txt_dim, **factory_kwargs) + if txt_in_dim and txt_in_dim != txt_dim + else nn.Identity() + ) + self.emb_in = TimeEmbedding( + sinusoidal_dim=BYTEDANCE_SINUSOIDAL_DIM, + hidden_dim=max(vid_dim, txt_dim), + output_dim=emb_dim, + device=device, dtype=dtype, operations=operations + ) + + if window is None or isinstance(window[0], int): + window = [window] * num_layers + + rope_dim = rope_dim if rope_dim is not None else head_dim // 2 + self.blocks = nn.ModuleList( + [ + NaMMSRTransformerBlock( + vid_dim=vid_dim, + txt_dim=txt_dim, + emb_dim=emb_dim, + heads=heads, + head_dim=head_dim, + expand_ratio=expand_ratio, + norm=norm, + norm_eps=norm_eps, + ada=ada, + qk_bias=qk_bias, + qk_norm=qk_norm, + mlp_type=mlp_type, + rope_dim = rope_dim, + window=window[i], + window_method=window_method[i], + version = self._7b_version, + is_last_layer=(i == num_layers - 1) and not self._7b_version, + rope_type = rope_type, + shared_weights=not ( + (i < mm_layers) if isinstance(mm_layers, int) else mm_layers[i] + ), + operations = operations, + **factory_kwargs + ) + for i in range(num_layers) + ] + ) + self.vid_out = NaPatchOut( + out_channels=vid_out_channels, + patch_size=patch_size, + dim=vid_dim, + device=device, dtype=dtype, operations=operations + ) + + self.vid_out_norm = None + if vid_out_norm is not None: + self.vid_out_norm = operations.RMSNorm( + normalized_shape=vid_dim, + eps=norm_eps, + elementwise_affine=True, + device=device, dtype=dtype + ) + self.vid_out_ada = ada( + dim=vid_dim, + emb_dim=emb_dim, + layers=["out"], + modes=["in"], + device=device, dtype=dtype + ) + + def _resolve_text_conditioning(self, context, cond_or_uncond=None): + if context is None or context.numel() == 0: + context = self.positive_conditioning + return flatten([context]) + if NaDiT._seedvr2_is_single_conditioning_branch(cond_or_uncond): + if context.shape[0] == 1: + context = context.squeeze(0) + return flatten([context]) + return flatten(context.unbind(0)) + if context.shape[0] % 2 != 0: + raise ValueError(f"SeedVR2 expected an even text-conditioning batch, got shape {tuple(context.shape)}") + neg_cond, pos_cond = context.chunk(2, dim=0) + if pos_cond.shape[0] == 1: + pos_cond, neg_cond = pos_cond.squeeze(0), neg_cond.squeeze(0) + return flatten([pos_cond, neg_cond]) + return flatten((*pos_cond.unbind(0), *neg_cond.unbind(0))) + + @staticmethod + def _seedvr2_is_single_conditioning_branch(cond_or_uncond): + if cond_or_uncond is None or len(cond_or_uncond) == 0: + return False + first = cond_or_uncond[0] + return all(entry == first for entry in cond_or_uncond) + + @staticmethod + def _check_seedvr2_video_latent(x, channels, name): + if x.ndim != 5: + raise ValueError(f"SeedVR2 expected {name} to be 5-D native latent, got shape {tuple(x.shape)}.") + if x.shape[1] != channels: + raise ValueError(f"SeedVR2 expected {name} channels to be {channels}, got shape {tuple(x.shape)}.") + return x + + def _swap_pos_neg_halves(self, out, cond_or_uncond=None): + if NaDiT._seedvr2_is_single_conditioning_branch(cond_or_uncond): + return out + pos, neg = out.chunk(2, dim=0) + return torch.cat([neg, pos], dim=0) + + def forward( + self, + x, + timestep, + context, # l c + disable_cache: bool = False, + **kwargs + ): + transformer_options = kwargs.get("transformer_options", {}) + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + conditions = kwargs.get("condition") + if conditions is None: + raise ValueError("SeedVR2 requires conditioning latents from the SeedVR2Conditioning node.") + x = self._check_seedvr2_video_latent(x, SEEDVR2_LATENT_CHANNELS, "latent") + conditions = self._check_seedvr2_video_latent(conditions, SEEDVR2_LATENT_CHANNELS + 1, "conditioning") + b, _, t, h, w = x.shape + if conditions.shape[0] != b or conditions.shape[2:] != (t, h, w): + raise ValueError( + f"SeedVR2 conditioning shape must match latent batch/temporal/spatial dimensions; got latent {tuple(x.shape)} and conditioning {tuple(conditions.shape)}." + ) + x = x.movedim(1, -1) + conditions = conditions.movedim(1, -1) + cache = Cache(disable=disable_cache) + + txt, txt_shape = self._resolve_text_conditioning(context, transformer_options.get("cond_or_uncond")) + + vid, vid_shape = flatten(x) + cond_latent, _ = flatten(conditions) + + vid = torch.cat([vid, cond_latent], dim=-1) + + txt = self.txt_in(txt) + + vid_shape_before_patchify = vid_shape + vid, vid_shape = self.vid_in(vid, vid_shape, cache=cache) + + emb = self.emb_in(timestep, device=vid.device, dtype=vid.dtype) + + for i, block in enumerate(self.blocks): + if ("block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["vid"], out["txt"], out["vid_shape"], out["txt_shape"] = block( + vid=args["vid"], + txt=args["txt"], + vid_shape=args["vid_shape"], + txt_shape=args["txt_shape"], + emb=args["emb"], + cache=args["cache"], + ) + return out + out = blocks_replace[("block", i)]({ + "vid":vid, + "txt":txt, + "vid_shape":vid_shape, + "txt_shape":txt_shape, + "emb":emb, + "cache":cache, + }, {"original_block": block_wrap}) + vid, txt, vid_shape, txt_shape = out["vid"], out["txt"], out["vid_shape"], out["txt_shape"] + else: + vid, txt, vid_shape, txt_shape = block( + vid=vid, + txt=txt, + vid_shape=vid_shape, + txt_shape=txt_shape, + emb=emb, + cache=cache, + ) + + if self.vid_out_norm: + vid = self.vid_out_norm(vid) + vid = self.vid_out_ada( + vid, + emb=emb, + layer="out", + mode="in", + hid_len=cache("vid_len", lambda: vid_shape.prod(-1)), + cache=cache, + branch_tag="vid", + ) + + vid, vid_shape = self.vid_out(vid, vid_shape, cache, vid_shape_before_patchify = vid_shape_before_patchify) + vid = unflatten(vid, vid_shape) + out = torch.stack(vid) + out = out.movedim(-1, 1) + return self._swap_pos_neg_halves(out, transformer_options.get("cond_or_uncond")) diff --git a/comfy/ldm/seedvr/vae.py b/comfy/ldm/seedvr/vae.py new file mode 100644 index 000000000..c9f430184 --- /dev/null +++ b/comfy/ldm/seedvr/vae.py @@ -0,0 +1,1612 @@ +from typing import Literal, Optional, Tuple +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from contextlib import contextmanager +from comfy.utils import ProgressBar + +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_BLOCK_OUT_CHANNELS, + BYTEDANCE_GN_CHUNKS_FP16, + BYTEDANCE_GN_CHUNKS_FP32, + BYTEDANCE_LOGVAR_CLAMP_MAX, + BYTEDANCE_LOGVAR_CLAMP_MIN, + BYTEDANCE_SLICING_SAMPLE_MIN, + BYTEDANCE_VAE_CONV_MEM_GIB, + BYTEDANCE_VAE_NORM_MEM_GIB, + BYTEDANCE_VAE_SCALING_FACTOR, + BYTEDANCE_VAE_SHIFTING_FACTOR, + BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE, + BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE, + SEEDVR2_LATENT_CHANNELS, +) +from comfy.ldm.modules.attention import optimized_attention +from comfy.ldm.modules.diffusionmodules.model import vae_attention + +import math +from enum import Enum + +import logging +import comfy.model_management +import comfy.ops +ops = comfy.ops.disable_weight_init + + +def _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap, temporal_scale=1): + if temporal_size is None: + return None + + temporal_size = int(temporal_size) + if temporal_size <= 0: + return None + + temporal_overlap = max(0, int(temporal_overlap or 0)) + temporal_overlap = min(temporal_overlap, temporal_size - 1) + temporal_step = temporal_size - temporal_overlap + temporal_scale = max(1, int(temporal_scale)) + return max(1, math.ceil(temporal_step / temporal_scale)) + + +def _seedvr2_clamped_spatial_overlap(overlap, tile_size): + overlap = max(0, int(overlap)) + tile_size = max(1, int(tile_size)) + return min(overlap, tile_size - 1) + + +def tiled_vae( + x, + vae_model, + tile_size=(512, 512), + tile_overlap=(64, 64), + temporal_size=16, + temporal_overlap=0, + encode=True, +): + if x.ndim != 5: + x = x.unsqueeze(2) + + _, _, d, h, w = x.shape + + sf_s = getattr(vae_model, "spatial_downsample_factor", BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE) + sf_t = getattr(vae_model, "temporal_downsample_factor", BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE) + if encode: + slicing_attr = "slicing_sample_min_size" + slicing_min_size = _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap) + else: + slicing_attr = "slicing_latent_min_size" + slicing_min_size = _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap, sf_t) + if encode: + ti_h, ti_w = tile_size + ov_h = _seedvr2_clamped_spatial_overlap(tile_overlap[0], ti_h) + ov_w = _seedvr2_clamped_spatial_overlap(tile_overlap[1], ti_w) + blend_ov_h = max(0, ov_h // sf_s) + blend_ov_w = max(0, ov_w // sf_s) + target_d = (d + sf_t - 1) // sf_t + target_h = (h + sf_s - 1) // sf_s + target_w = (w + sf_s - 1) // sf_s + else: + ti_h = max(1, tile_size[0] // sf_s) + ti_w = max(1, tile_size[1] // sf_s) + ov_h = _seedvr2_clamped_spatial_overlap(tile_overlap[0] // sf_s, ti_h) + ov_w = _seedvr2_clamped_spatial_overlap(tile_overlap[1] // sf_s, ti_w) + blend_ov_h = ov_h * sf_s + blend_ov_w = ov_w * sf_s + + target_d = max(1, d * sf_t - (sf_t - 1)) + target_h = h * sf_s + target_w = w * sf_s + + stride_h = max(1, ti_h - ov_h) + stride_w = max(1, ti_w - ov_w) + + storage_device = vae_model.device + result = None + count = None + def run_temporal_chunks(spatial_tile, model=vae_model, device=storage_device): + device = torch.device(device) + t_chunk = spatial_tile.to(device=device, dtype=next(model.parameters()).dtype, non_blocking=True).contiguous() + old_device = getattr(model, "device", None) + model.device = device + old_slicing_min_size = getattr(model, slicing_attr, None) + if old_slicing_min_size is not None and slicing_min_size is not None: + if slicing_min_size <= 0: + setattr(model, slicing_attr, t_chunk.shape[2]) + else: + setattr(model, slicing_attr, slicing_min_size) + try: + if encode: + out = model.encode(t_chunk) + else: + out = model.decode_(t_chunk) + finally: + if old_slicing_min_size is not None and slicing_min_size is not None: + setattr(model, slicing_attr, old_slicing_min_size) + if old_device is not None: + model.device = old_device + if out.ndim == 4: + out = out.unsqueeze(2) + return out.to(storage_device) + + ramp_cache = {} + def get_ramp(steps): + if steps not in ramp_cache: + t = torch.linspace(0, 1, steps=steps, device=storage_device, dtype=torch.float32) + ramp_cache[steps] = 0.5 - 0.5 * torch.cos(t * torch.pi) + return ramp_cache[steps] + + tile_ranges = [] + for y_idx in range(0, h, stride_h): + y_end = min(y_idx + ti_h, h) + if y_idx > 0 and (y_end - y_idx) <= ov_h: + continue + for x_idx in range(0, w, stride_w): + x_end = min(x_idx + ti_w, w) + if x_idx > 0 and (x_end - x_idx) <= ov_w: + continue + tile_ranges.append((y_idx, y_end, x_idx, x_end)) + + total_tiles = len(tile_ranges) + bar = ProgressBar(total_tiles) + single_spatial_tile = h <= ti_h and w <= ti_w + + def run_tile(tile_index, tile_range): + y_idx, y_end, x_idx, x_end = tile_range + tile_x = x[:, :, :, y_idx:y_end, x_idx:x_end] + tile_out = run_temporal_chunks(tile_x) + return tile_index, y_idx, y_end, x_idx, x_end, tile_out + + ordered_tile_outputs = ( + run_tile(tile_index, tile_range) + for tile_index, tile_range in enumerate(tile_ranges) + ) + + for _, y_idx, y_end, x_idx, x_end, tile_out in ordered_tile_outputs: + + if single_spatial_tile: + result = tile_out[:, :, :target_d, :target_h, :target_w] + if result.device != x.device or result.dtype != x.dtype: + result = result.to(device=x.device, dtype=x.dtype) + if x.shape[2] == 1 and sf_t == 1: + result = result.squeeze(2) + bar.update(1) + return result + + if result is None: + b_out, c_out = tile_out.shape[0], tile_out.shape[1] + result = torch.zeros((b_out, c_out, target_d, target_h, target_w), device=storage_device, dtype=torch.float32) + count = torch.zeros((1, 1, 1, target_h, target_w), device=storage_device, dtype=torch.float32) + + if encode: + ys, ye = y_idx // sf_s, (y_idx // sf_s) + tile_out.shape[3] + xs, xe = x_idx // sf_s, (x_idx // sf_s) + tile_out.shape[4] + cur_ov_h = max(0, min(blend_ov_h, tile_out.shape[3] // 2)) + cur_ov_w = max(0, min(blend_ov_w, tile_out.shape[4] // 2)) + else: + ys, ye = y_idx * sf_s, (y_idx * sf_s) + tile_out.shape[3] + xs, xe = x_idx * sf_s, (x_idx * sf_s) + tile_out.shape[4] + cur_ov_h = max(0, min(blend_ov_h, tile_out.shape[3] // 2)) + cur_ov_w = max(0, min(blend_ov_w, tile_out.shape[4] // 2)) + + w_h = torch.ones((tile_out.shape[3],), device=storage_device) + w_w = torch.ones((tile_out.shape[4],), device=storage_device) + + if cur_ov_h > 0: + r = get_ramp(cur_ov_h) + if y_idx > 0: + w_h[:cur_ov_h] = r + if y_end < h: + w_h[-cur_ov_h:] = 1.0 - r + + if cur_ov_w > 0: + r = get_ramp(cur_ov_w) + if x_idx > 0: + w_w[:cur_ov_w] = r + if x_end < w: + w_w[-cur_ov_w:] = 1.0 - r + + final_weight = w_h.view(1,1,1,-1,1) * w_w.view(1,1,1,1,-1) + + valid_d = min(tile_out.shape[2], result.shape[2]) + tile_out = tile_out[:, :, :valid_d, :, :] + + tile_out.mul_(final_weight) + + result[:, :, :valid_d, ys:ye, xs:xe] += tile_out + count[:, :, :, ys:ye, xs:xe] += final_weight + + del tile_out, final_weight, w_h, w_w + bar.update(1) + + result.div_(count.clamp(min=1e-6)) + + if result.device != x.device or result.dtype != x.dtype: + result = result.to(device=x.device, dtype=x.dtype) + + if x.shape[2] == 1 and sf_t == 1: + result = result.squeeze(2) + + return result + +_NORM_LIMIT = float("inf") +def get_norm_limit(): + return _NORM_LIMIT + + +def set_norm_limit(value: Optional[float] = None): + global _NORM_LIMIT + if value is None: + value = float("inf") + _NORM_LIMIT = value + +@contextmanager +def ignore_padding(model): + orig_padding = model.padding + model.padding = (0, 0, 0) + try: + yield + finally: + model.padding = orig_padding + +class MemoryState(Enum): + DISABLED = 0 + INITIALIZING = 1 + ACTIVE = 2 + UNSET = 3 + +def get_cache_size(conv_module, input_len, pad_len, dim=0): + dilated_kernel_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1 + output_len = (input_len + pad_len - dilated_kernel_size) // conv_module.stride[dim] + 1 + remain_len = ( + input_len + pad_len - ((output_len - 1) * conv_module.stride[dim] + dilated_kernel_size) + ) + overlap_len = dilated_kernel_size - conv_module.stride[dim] + cache_len = overlap_len + remain_len + + if output_len <= 0: + raise ValueError( + f"SeedVR2 VAE cache input is too short for convolution: input_len={input_len}, pad_len={pad_len}." + ) + return cache_len + +class DiagonalGaussianDistribution(object): + def __init__(self, parameters: torch.Tensor): + self.parameters = parameters + self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) + self.logvar = torch.clamp(self.logvar, BYTEDANCE_LOGVAR_CLAMP_MIN, BYTEDANCE_LOGVAR_CLAMP_MAX) + + def mode(self): + return self.mean + +class SpatialNorm(nn.Module): + def __init__( + self, + f_channels: int, + zq_channels: int, + ): + super().__init__() + self.norm_layer = ops.GroupNorm(num_channels=f_channels, num_groups=32, eps=1e-6, affine=True) + self.conv_y = ops.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) + self.conv_b = ops.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) + + def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor: + f_size = f.shape[-2:] + zq = F.interpolate(zq, size=f_size, mode="nearest") + norm_f = self.norm_layer(f) + new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) + return new_f + +class Attention(nn.Module): + def __init__( + self, + query_dim: int, + heads: int = 8, + dim_head: int = 64, + bias: bool = False, + norm_num_groups: Optional[int] = None, + spatial_norm_dim: Optional[int] = None, + out_bias: bool = True, + eps: float = 1e-5, + rescale_output_factor: float = 1.0, + residual_connection: bool = False, + ): + super().__init__() + + self.inner_dim = dim_head * heads + self.rescale_output_factor = rescale_output_factor + self.residual_connection = residual_connection + self.out_dim = query_dim + self.heads = heads + + if norm_num_groups is not None: + self.group_norm = ops.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True) + else: + self.group_norm = None + + if spatial_norm_dim is not None: + self.spatial_norm = SpatialNorm(f_channels=query_dim, zq_channels=spatial_norm_dim) + else: + self.spatial_norm = None + + self.to_q = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_k = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_v = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_out = nn.ModuleList([]) + self.to_out.append(ops.Linear(self.inner_dim, self.out_dim, bias=out_bias)) + self.to_out.append(nn.Identity()) + + self.optimized_vae_attention = vae_attention() + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + + residual = hidden_states + if self.spatial_norm is not None: + hidden_states = self.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size = hidden_states.shape[0] + + if self.group_norm is not None: + hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = self.to_q(hidden_states) + key = self.to_k(hidden_states) + value = self.to_v(hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // self.heads + + query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + + if input_ndim == 4 and self.heads == 1: + query = query.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + key = key.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + value = value.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + hidden_states = self.optimized_vae_attention(query, key, value).reshape(batch_size, self.heads, head_dim, height * width).transpose(2, 3) + else: + hidden_states = optimized_attention(query, key, value, heads = self.heads, skip_reshape=True, skip_output_reshape=True) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + hidden_states = self.to_out[0](hidden_states) + hidden_states = self.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if self.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / self.rescale_output_factor + + return hidden_states + + +def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: + input_dtype = x.dtype + if isinstance(norm_layer, (ops.LayerNorm, ops.RMSNorm)): + if x.ndim == 4: + x = x.permute(0, 2, 3, 1) + x = norm_layer(x) + x = x.permute(0, 3, 1, 2) + return x.to(input_dtype) + if x.ndim == 5: + x = x.permute(0, 2, 3, 4, 1) + x = norm_layer(x) + x = x.permute(0, 4, 1, 2, 3) + return x.to(input_dtype) + if isinstance(norm_layer, (ops.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): + if x.ndim <= 4: + return norm_layer(x).to(input_dtype) + if x.ndim == 5: + b, c, t, h, w = x.shape + x = x.transpose(1, 2).reshape(b * t, c, h, w) + memory_occupy = x.numel() * x.element_size() / 1024**3 + if isinstance(norm_layer, ops.GroupNorm) and memory_occupy > get_norm_limit(): + num_chunks = min(BYTEDANCE_GN_CHUNKS_FP16 if x.element_size() == 2 else BYTEDANCE_GN_CHUNKS_FP32, norm_layer.num_groups) + if norm_layer.num_groups % num_chunks != 0: + raise ValueError( + f"SeedVR2 VAE GroupNorm groups must divide chunks: groups={norm_layer.num_groups}, chunks={num_chunks}." + ) + num_groups_per_chunk = norm_layer.num_groups // num_chunks + + x = list(x.chunk(num_chunks, dim=1)) + weights = norm_layer.weight.chunk(num_chunks, dim=0) + biases = norm_layer.bias.chunk(num_chunks, dim=0) + for i, (w, bias) in enumerate(zip(weights, biases)): + x[i] = F.group_norm(x[i], num_groups_per_chunk, w, bias, norm_layer.eps) + x[i] = x[i].to(input_dtype) + x = torch.cat(x, dim=1) + else: + x = norm_layer(x) + x = x.reshape((b, t, x.size(1), x.size(2), x.size(3))).transpose(1, 2) + return x.to(input_dtype) + raise TypeError(f"SeedVR2 VAE unsupported norm layer type: {type(norm_layer).__name__}") + +_receptive_field_t = Literal["half", "full"] + +def extend_head(tensor, times: int = 2, memory = None): + if memory is not None: + return torch.cat((memory.to(tensor), tensor), dim=2) + if times < 0: + raise ValueError(f"SeedVR2 VAE extend_head expected times >= 0, got {times}.") + if times == 0: + return tensor + else: + tile_repeat = [1] * tensor.ndim + tile_repeat[2] = times + return torch.cat(tensors=(torch.tile(tensor[:, :, :1], tile_repeat), tensor), dim=2) + +def cache_send_recv(tensor, cache_size, times, memory=None): + recv_buffer = None + + if memory is not None: + recv_buffer = memory.to(tensor[0]) + elif times > 0: + tile_repeat = [1] * tensor[0].ndim + tile_repeat[2] = times + recv_buffer = torch.tile(tensor[0][:, :, :1], tile_repeat) + + return recv_buffer + +class InflatedCausalConv3d(ops.Conv3d): + def __init__( + self, + *args, + inflation_mode, + **kwargs, + ): + self.inflation_mode = inflation_mode + super().__init__(*args, **kwargs) + self.temporal_padding = self.padding[0] + self.padding = (0, *self.padding[1:]) + self.memory_limit = float("inf") + self.logged_once = False + + def set_memory_limit(self, value: float): + self.memory_limit = value + + def _conv_forward(self, input, weight, bias, *args, **kwargs): + try: + return super()._conv_forward(input, weight, bias, *args, **kwargs) + except NotImplementedError: + # for: Could not run 'aten::cudnn_convolution' with arguments from the 'CPU' backend + if not self.logged_once: + logging.warning("VAE is on CPU for decoding. This is most likely due to not enough memory") + self.logged_once = True + return F.conv3d(input, weight, bias, *args, **kwargs) + + def memory_limit_conv( + self, + x, + *, + split_dim=3, + padding=(0, 0, 0, 0, 0, 0), + prev_cache=None, + ): + if math.isinf(self.memory_limit): + if prev_cache is not None: + x = torch.cat([prev_cache, x], dim=split_dim - 1) + return super().forward(x) + + shape = list(x.size()) + if prev_cache is not None: + shape[split_dim - 1] += prev_cache.size(split_dim - 1) + for i, pad_sum in enumerate((padding[4] + padding[5], padding[2] + padding[3], padding[0] + padding[1])): + shape[-3 + i] += pad_sum + memory_occupy = math.prod(shape) * x.element_size() / 1024**3 # GiB + if memory_occupy < self.memory_limit or split_dim == x.ndim: + x_concat = x + if prev_cache is not None: + x_concat = torch.cat([prev_cache, x], dim=split_dim - 1) + + def pad_and_forward(): + padded = F.pad(x_concat, padding, mode='constant', value=0.0) + if not padded.is_contiguous(): + padded = padded.contiguous() + with ignore_padding(self): + return torch.nn.Conv3d.forward(self, padded) + + return pad_and_forward() + + num_splits = math.ceil(memory_occupy / self.memory_limit) + size_per_split = x.size(split_dim) // num_splits + split_sizes = [size_per_split] * (num_splits - 1) + split_sizes += [x.size(split_dim) - sum(split_sizes)] + + x = list(x.split(split_sizes, dim=split_dim)) + if prev_cache is not None: + prev_cache = list(prev_cache.split(split_sizes, dim=split_dim)) + cache = None + for idx in range(len(x)): + if prev_cache is not None: + x[idx] = torch.cat([prev_cache[idx], x[idx]], dim=split_dim - 1) + + lpad_dim = (x[idx].ndim - split_dim - 1) * 2 + rpad_dim = lpad_dim + 1 + padding = list(padding) + padding[lpad_dim] = self.padding[split_dim - 2] if idx == 0 else 0 + padding[rpad_dim] = self.padding[split_dim - 2] if idx == len(x) - 1 else 0 + pad_len = padding[lpad_dim] + padding[rpad_dim] + padding = tuple(padding) + + next_cache = None + cache_len = cache.size(split_dim) if cache is not None else 0 + next_cache_size = get_cache_size( + conv_module=self, + input_len=x[idx].size(split_dim) + cache_len, + pad_len=pad_len, + dim=split_dim - 2, + ) + if next_cache_size != 0: + if next_cache_size > x[idx].size(split_dim): + raise ValueError( + f"SeedVR2 VAE cache size {next_cache_size} exceeds split size {x[idx].size(split_dim)}." + ) + next_cache = ( + x[idx].transpose(0, split_dim)[-next_cache_size:].transpose(0, split_dim) + ) + + x[idx] = self.memory_limit_conv( + x[idx], + split_dim=split_dim + 1, + padding=padding, + prev_cache=cache + ) + + cache = next_cache + + output = torch.cat(x, dim=split_dim) + return output + + def forward( + self, + input, + memory_state: MemoryState = MemoryState.UNSET, + memory_cache = None, + ) -> Tensor: + if memory_state == MemoryState.UNSET: + raise ValueError("SeedVR2 VAE convolution requires an explicit MemoryState.") + if memory_cache is None: + memory_cache = {} + if memory_state != MemoryState.ACTIVE: + memory_cache.pop(self, None) + if ( + math.isinf(self.memory_limit) + and torch.is_tensor(input) + ): + return self.basic_forward(input, memory_state, memory_cache) + return self.slicing_forward(input, memory_state, memory_cache) + + def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET, memory_cache = None): + mem_size = self.stride[0] - self.kernel_size[0] + memory = memory_cache.get(self) if memory_cache is not None else None + if (memory is not None) and (memory_state == MemoryState.ACTIVE): + input = extend_head(input, memory=memory, times=-1) + else: + input = extend_head(input, times=self.temporal_padding * 2) + next_memory = ( + input[:, :, mem_size:].detach() + if (mem_size != 0 and memory_state != MemoryState.DISABLED) + else None + ) + if memory_cache is not None and memory_state != MemoryState.DISABLED: + if next_memory is None: + memory_cache.pop(self, None) + else: + memory_cache[self] = next_memory + return super().forward(input) + + def slicing_forward( + self, + input, + memory_state: MemoryState = MemoryState.UNSET, + memory_cache = None, + ) -> Tensor: + if memory_cache is None: + memory_cache = {} + squeeze_out = False + if torch.is_tensor(input): + input = [input] + squeeze_out = True + + cache_size = self.kernel_size[0] - self.stride[0] + memory = memory_cache.get(self) if memory_cache is not None else None + cache = cache_send_recv( + input, cache_size=cache_size, memory=memory, times=self.temporal_padding * 2 + ) + + if ( + memory_state in [MemoryState.INITIALIZING, MemoryState.ACTIVE] + and cache_size != 0 + ): + if cache_size > input[-1].size(2) and cache is not None and len(input) == 1: + input[0] = torch.cat([cache, input[0]], dim=2) + cache = None + if cache_size <= input[-1].size(2): + memory_cache[self] = input[-1][:, :, -cache_size:].detach().contiguous() + + padding = tuple(x for x in reversed(self.padding) for _ in range(2)) + for i in range(len(input)): + next_cache = None + cache_size = 0 + if i < len(input) - 1: + cache_len = cache.size(2) if cache is not None else 0 + cache_size = get_cache_size(self, input[i].size(2) + cache_len, pad_len=0) + if cache_size != 0: + if cache_size > input[i].size(2) and cache is not None: + input[i] = torch.cat([cache, input[i]], dim=2) + cache = None + if cache_size > input[i].size(2): + raise ValueError(f"SeedVR2 VAE cache size {cache_size} exceeds input length {input[i].size(2)}.") + next_cache = input[i][:, :, -cache_size:] + + input[i] = self.memory_limit_conv( + input[i], + padding=padding, + prev_cache=cache + ) + + cache = next_cache + + return input[0] if squeeze_out else input + +def remove_head(tensor: Tensor, times: int = 1) -> Tensor: + if times == 0: + return tensor + return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2) + +class Upsample3D(nn.Module): + + def __init__( + self, + channels, + out_channels = None, + inflation_mode = "tail", + temporal_up: bool = False, + spatial_up: bool = True, + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + + conv = InflatedCausalConv3d( + self.channels, + self.out_channels, + 3, + padding=1, + inflation_mode=inflation_mode, + ) + + self.temporal_up = temporal_up + self.spatial_up = spatial_up + self.temporal_ratio = 2 if temporal_up else 1 + self.spatial_ratio = 2 if spatial_up else 1 + + upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio + self.upscale_conv = ops.Conv3d( + self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0 + ) + + self.conv = conv + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state=None, + memory_cache=None, + ) -> torch.FloatTensor: + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 upsample expected {self.channels} channels, got {hidden_states.shape[1]}.") + + hidden_states = self.upscale_conv(hidden_states) + b, channels, f, h, w = hidden_states.shape + c = channels // (self.spatial_ratio * self.spatial_ratio * self.temporal_ratio) + hidden_states = hidden_states.view(b, self.spatial_ratio, self.spatial_ratio, self.temporal_ratio, c, f, h, w) + hidden_states = hidden_states.permute(0, 4, 5, 3, 6, 1, 7, 2).reshape( + b, + c, + f * self.temporal_ratio, + h * self.spatial_ratio, + w * self.spatial_ratio, + ) + + if self.temporal_up and memory_state != MemoryState.ACTIVE: + hidden_states = remove_head(hidden_states) + + hidden_states = self.conv(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class Downsample3D(nn.Module): + def __init__( + self, + channels, + out_channels = None, + inflation_mode = "tail", + spatial_down: bool = False, + temporal_down: bool = False, + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.temporal_down = temporal_down + self.spatial_down = spatial_down + + self.temporal_ratio = 2 if temporal_down else 1 + self.spatial_ratio = 2 if spatial_down else 1 + + self.temporal_kernel = 3 if temporal_down else 1 + self.spatial_kernel = 3 if spatial_down else 1 + + self.conv = InflatedCausalConv3d( + self.channels, + self.out_channels, + kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel), + stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + padding=(1 if self.temporal_down else 0, 0, 0), + inflation_mode=inflation_mode, + ) + + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 downsample expected {self.channels} channels, got {hidden_states.shape[1]}.") + + if self.spatial_down: + pad = (0, 1, 0, 1) + hidden_states = F.pad(hidden_states, pad, mode="constant", value=0) + + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 downsample expected {self.channels} channels after padding, got {hidden_states.shape[1]}.") + + hidden_states = self.conv(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class ResnetBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: Optional[int] = None, + temb_channels: int = 512, + groups: int = 32, + groups_out: Optional[int] = None, + eps: float = 1e-6, + output_scale_factor: float = 1.0, + skip_time_act: bool = False, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = in_channels if out_channels is None else out_channels + self.output_scale_factor = output_scale_factor + self.skip_time_act = skip_time_act + self.nonlinearity = nn.SiLU() + if temb_channels is not None: + self.time_emb_proj = ops.Linear(temb_channels, self.out_channels) + else: + self.time_emb_proj = None + self.norm1 = ops.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True) + if groups_out is None: + groups_out = groups + self.norm2 = ops.GroupNorm(num_groups=groups_out, num_channels=self.out_channels, eps=eps, affine=True) + self.use_in_shortcut = self.in_channels != self.out_channels + self.conv1 = InflatedCausalConv3d( + self.in_channels, + self.out_channels, + kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3), + stride=1, + padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1), + inflation_mode=inflation_mode, + ) + + self.conv2 = InflatedCausalConv3d( + self.out_channels, + self.out_channels, + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.conv_shortcut = None + if self.use_in_shortcut: + self.conv_shortcut = InflatedCausalConv3d( + self.in_channels, + self.out_channels, + kernel_size=1, + stride=1, + padding=0, + bias=True, + inflation_mode=inflation_mode, + ) + + def forward(self, input_tensor, temb, memory_state = None, memory_cache = None): + hidden_states = input_tensor + + hidden_states = causal_norm_wrapper(self.norm1, hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + + hidden_states = self.conv1(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + if self.time_emb_proj is not None: + if not self.skip_time_act: + temb = self.nonlinearity(temb) + temb = self.time_emb_proj(temb)[:, :, None, None] + + if temb is not None: + hidden_states = hidden_states + temb + + hidden_states = causal_norm_wrapper(self.norm2, hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + + hidden_states = self.conv2(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + if self.conv_shortcut is not None: + input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state, memory_cache=memory_cache) + + output_tensor = (input_tensor + hidden_states) / self.output_scale_factor + + return output_tensor + + +class DownEncoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_groups: int = 32, + output_scale_factor: float = 1.0, + add_downsample: bool = True, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_down: bool = True, + spatial_down: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=None, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample3D( + out_channels, + out_channels=out_channels, + temporal_down=temporal_down, + spatial_down=spatial_down, + inflation_mode=inflation_mode, + ) + ] + ) + else: + self.downsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, memory_cache=memory_cache) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class UpDecoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_groups: int = 32, + output_scale_factor: float = 1.0, + add_upsample: bool = True, + temb_channels: Optional[int] = None, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up: bool = True, + spatial_up: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + + resnets.append( + ResnetBlock3D( + in_channels=input_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_upsample: + self.upsamplers = nn.ModuleList( + [ + Upsample3D( + out_channels, + out_channels=out_channels, + temporal_up=temporal_up, + spatial_up=spatial_up, + inflation_mode=inflation_mode, + ) + ] + ) + else: + self.upsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + temb: Optional[torch.FloatTensor] = None, + memory_state=None, + memory_cache=None, + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, memory_cache=memory_cache) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class UNetMidBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + temb_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", # default, spatial + resnet_groups: int = 32, + add_attention: bool = True, + attention_head_dim: int = 1, + output_scale_factor: float = 1.0, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) + self.add_attention = add_attention + + resnets = [ + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ] + attentions = [] + + if attention_head_dim is None: + attention_head_dim = in_channels + + for _ in range(num_layers): + if self.add_attention: + attentions.append( + Attention( + in_channels, + heads=in_channels // attention_head_dim, + dim_head=attention_head_dim, + rescale_output_factor=output_scale_factor, + eps=resnet_eps, + norm_num_groups=( + resnet_groups if resnet_time_scale_shift == "default" else None + ), + spatial_norm_dim=( + temb_channels if resnet_time_scale_shift == "spatial" else None + ), + residual_connection=True, + bias=True, + ) + ) + else: + attentions.append(None) + + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + def forward(self, hidden_states, temb=None, memory_state=None, memory_cache=None): + video_length = hidden_states.size(2) + hidden_states = self.resnets[0](hidden_states, temb, memory_state=memory_state, memory_cache=memory_cache) + for attn, resnet in zip(self.attentions, self.resnets[1:]): + if attn is not None: + b, c, f, h, w = hidden_states.shape + hidden_states = hidden_states.transpose(1, 2).reshape(b * f, c, h, w) + hidden_states = attn(hidden_states, temb=temb) + hidden_states = hidden_states.reshape(b, video_length, c, h, w).transpose(1, 2) + hidden_states = resnet(hidden_states, temb, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class Encoder3D(nn.Module): + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + down_block_types: Tuple[str, ...] = ("DownEncoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + mid_block_add_attention=True, + temporal_down_num: int = 2, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_down_num = temporal_down_num + + self.conv_in = InflatedCausalConv3d( + in_channels, + block_out_channels[0], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.down_blocks = nn.ModuleList([]) + + output_channel = block_out_channels[0] + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1 + + if down_block_type != "DownEncoderBlock3D": + raise ValueError(f"SeedVR2 encoder only supports DownEncoderBlock3D, got {down_block_type}.") + + down_block = DownEncoderBlock3D( + num_layers=self.layers_per_block, + in_channels=input_channel, + out_channels=output_channel, + add_downsample=not is_final_block, + resnet_eps=1e-6, + resnet_groups=norm_num_groups, + temporal_down=is_temporal_down_block, + spatial_down=True, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.down_blocks.append(down_block) + + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + output_scale_factor=1, + resnet_time_scale_shift="default", + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=None, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.conv_norm_out = ops.GroupNorm( + num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + + conv_out_channels = 2 * out_channels + self.conv_out = InflatedCausalConv3d( + block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + + def forward( + self, + sample: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + sample = sample.to(next(self.parameters()).device) + sample = self.conv_in(sample, memory_state=memory_state, memory_cache=memory_cache) + for down_block in self.down_blocks: + sample = down_block(sample, memory_state=memory_state, memory_cache=memory_cache) + + sample = self.mid_block(sample, memory_state=memory_state, memory_cache=memory_cache) + + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state, memory_cache=memory_cache) + + return sample + + +class Decoder3D(nn.Module): + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + up_block_types: Tuple[str, ...] = ("UpDecoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + mid_block_add_attention=True, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up_num: int = 2, + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_up_num = temporal_up_num + + self.conv_in = InflatedCausalConv3d( + in_channels, + block_out_channels[-1], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.up_blocks = nn.ModuleList([]) + + temb_channels = None + + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + output_scale_factor=1, + resnet_time_scale_shift="default", + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=temb_channels, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + reversed_block_out_channels = list(reversed(block_out_channels)) + output_channel = reversed_block_out_channels[0] + for i, up_block_type in enumerate(up_block_types): + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + + is_final_block = i == len(block_out_channels) - 1 + is_temporal_up_block = i < self.temporal_up_num + if up_block_type != "UpDecoderBlock3D": + raise ValueError(f"SeedVR2 decoder only supports UpDecoderBlock3D, got {up_block_type}.") + up_block = UpDecoderBlock3D( + num_layers=self.layers_per_block + 1, + in_channels=prev_output_channel, + out_channels=output_channel, + add_upsample=not is_final_block, + resnet_eps=1e-6, + resnet_groups=norm_num_groups, + temb_channels=temb_channels, + temporal_up=is_temporal_up_block, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.up_blocks.append(up_block) + prev_output_channel = output_channel + + self.conv_norm_out = ops.GroupNorm( + num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + self.conv_out = InflatedCausalConv3d( + block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + + def forward( + self, + sample: torch.FloatTensor, + latent_embeds: Optional[torch.FloatTensor] = None, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + + sample = sample.to(next(self.parameters()).device) + sample = self.conv_in(sample, memory_state=memory_state, memory_cache=memory_cache) + + upscale_dtype = next(iter(self.up_blocks.parameters())).dtype + sample = self.mid_block(sample, latent_embeds, memory_state=memory_state, memory_cache=memory_cache) + sample = sample.to(upscale_dtype) + + for up_block in self.up_blocks: + sample = up_block(sample, latent_embeds, memory_state=memory_state, memory_cache=memory_cache) + + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state, memory_cache=memory_cache) + + return sample + +class VideoAutoencoderKL(nn.Module): + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + layers_per_block: int = 2, + latent_channels: int = SEEDVR2_LATENT_CHANNELS, + norm_num_groups: int = 32, + temporal_scale_num: int = 2, + inflation_mode = "pad", + time_receptive_field: _receptive_field_t = "full", + slicing_sample_min_size = BYTEDANCE_SLICING_SAMPLE_MIN, + ): + self.slicing_sample_min_size = slicing_sample_min_size + self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num) + block_out_channels = BYTEDANCE_BLOCK_OUT_CHANNELS + down_block_types = ("DownEncoderBlock3D",) * 4 + up_block_types = ("UpDecoderBlock3D",) * 4 + super().__init__() + + self.encoder = Encoder3D( + in_channels=in_channels, + out_channels=latent_channels, + down_block_types=down_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + norm_num_groups=norm_num_groups, + temporal_down_num=temporal_scale_num, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.decoder = Decoder3D( + in_channels=latent_channels, + out_channels=out_channels, + up_block_types=up_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + norm_num_groups=norm_num_groups, + temporal_up_num=temporal_scale_num, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.use_slicing = True + + def encode(self, x: torch.FloatTensor, return_dict: bool = True): + h = self.slicing_encode(x) + posterior = DiagonalGaussianDistribution(h).mode() + + if not return_dict: + return (posterior,) + + return posterior + + def decode_( + self, z: torch.Tensor, return_dict: bool = True + ): + decoded = self.slicing_decode(z) + + if not return_dict: + return (decoded,) + + return decoded + + def _encode( + self, x, memory_state = MemoryState.DISABLED, memory_cache = None + ) -> torch.Tensor: + _x = x.to(self.device) + h = self.encoder(_x, memory_state=memory_state, memory_cache=memory_cache) + return h.to(x.device) + + def _decode( + self, z, memory_state = MemoryState.DISABLED, memory_cache = None + ) -> torch.Tensor: + _z = z.to(self.device) + output = self.decoder(_z, memory_state=memory_state, memory_cache=memory_cache) + return output.to(z.device) + + def slicing_encode(self, x: torch.Tensor) -> torch.Tensor: + if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size: + memory_cache = {} + split_size = max( + self.slicing_sample_min_size, + getattr(self, "temporal_downsample_factor", 1), + ) + x_slices = list(x[:, :, 1:].split(split_size=split_size, dim=2)) + min_active_len = getattr(self, "temporal_downsample_factor", 1) + if len(x_slices) > 1 and x_slices[-1].shape[2] < min_active_len: + x_slices[-2] = torch.cat((x_slices[-2], x_slices[-1]), dim=2) + x_slices.pop() + encoded_slices = [ + self._encode( + torch.cat((x[:, :, :1], x_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + memory_cache=memory_cache, + ) + ] + for x_idx in range(1, len(x_slices)): + encoded_slices.append( + self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE, memory_cache=memory_cache) + ) + out = torch.cat(encoded_slices, dim=2) + return out + else: + return self._encode(x) + + def slicing_decode(self, z: torch.Tensor) -> torch.Tensor: + if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size: + memory_cache = {} + z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size, dim=2) + decoded_slices = [ + self._decode( + torch.cat((z[:, :, :1], z_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + memory_cache=memory_cache, + ) + ] + for z_idx in range(1, len(z_slices)): + decoded_slices.append( + self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE, memory_cache=memory_cache) + ) + out = torch.cat(decoded_slices, dim=2) + return out + else: + return self._decode(z) + + def forward(self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all"): + def _unwrap(value): + return value[0] if isinstance(value, tuple) else value + + if mode == "encode": + return _unwrap(self.encode(x)) + if mode == "decode": + return _unwrap(self.decode_(x)) + if mode == "all": + latent = _unwrap(self.encode(x)) + return _unwrap(self.decode_(latent)) + raise ValueError(f"Unknown SeedVR2 VAE forward mode: {mode}") + +class VideoAutoencoderKLWrapper(VideoAutoencoderKL): + def __init__( + self, + spatial_downsample_factor = 8, + temporal_downsample_factor = 4, + ): + self.spatial_downsample_factor = spatial_downsample_factor + self.temporal_downsample_factor = temporal_downsample_factor + super().__init__() + self.set_memory_limit(BYTEDANCE_VAE_CONV_MEM_GIB, BYTEDANCE_VAE_NORM_MEM_GIB) + + def forward(self, x: torch.FloatTensor): + z, p = self._encode_with_raw_latent(x) + x = self.decode(z) + return x, z, p + + def _encode_with_raw_latent(self, x): + if x.ndim == 4: + x = x.unsqueeze(2) + x = x.to(dtype=next(self.parameters()).dtype) + self.device = x.device + p = super().encode(x) + z = p.squeeze(2) + return z, p + + def encode(self, x): + z, _ = self._encode_with_raw_latent(x) + return z + + def decode(self, z, seedvr2_tiling=None): + seedvr2_tiling = {} if seedvr2_tiling is None else seedvr2_tiling + if not isinstance(seedvr2_tiling, dict): + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: `seedvr2_tiling` must be a dict; " + f"got {type(seedvr2_tiling).__name__} with value {seedvr2_tiling!r}." + ) + + if z.ndim == 5: + _, c, _, _, _ = z.shape + if c != SEEDVR2_LATENT_CHANNELS: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: 5-D latent input must " + f"have {SEEDVR2_LATENT_CHANNELS} channels; got shape {tuple(z.shape)}." + ) + latent = z + elif z.ndim == 4: + b, tc, h, w = z.shape + if tc % SEEDVR2_LATENT_CHANNELS != 0: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: 4-D latent input must " + f"use collapsed channel layout (B, {SEEDVR2_LATENT_CHANNELS}*T, H, W); " + f"got shape {tuple(z.shape)}." + ) + latent = z.reshape(b, SEEDVR2_LATENT_CHANNELS, -1, h, w) + else: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: latent input must be " + f"4-D collapsed (B, {SEEDVR2_LATENT_CHANNELS}*T, H, W) or " + f"5-D (B, {SEEDVR2_LATENT_CHANNELS}, T, H, W); " + f"got shape {tuple(z.shape)}." + ) + scale = BYTEDANCE_VAE_SCALING_FACTOR + shift = BYTEDANCE_VAE_SHIFTING_FACTOR + latent = latent / scale + shift + + self.device = latent.device + enable_tiling = seedvr2_tiling.get("enable_tiling", False) + + if enable_tiling: + decode_seedvr2_args = dict(seedvr2_tiling) + decode_seedvr2_args.pop("enable_tiling", None) + tile_h, tile_w = decode_seedvr2_args.get("tile_size", (512, 512)) + ov_h, ov_w = decode_seedvr2_args.get("tile_overlap", (64, 64)) + decode_seedvr2_args["tile_overlap"] = ( + min(ov_h, max(0, tile_h - 8)), + min(ov_w, max(0, tile_w - 8)), + ) + x = tiled_vae(latent, self, **decode_seedvr2_args, encode=False) + if x.ndim == 4: + # tiled_vae squeezes the temporal axis when + # temporal_downsample_factor == 1 AND latent T == 1 + # (see tiled_vae line 179-180); re-add it so the post-decode + # pipeline can keep batch and time distinct on the tiled path. + x = x.unsqueeze(2) + else: + x = super().decode_(latent) + + h, w = x.shape[-2:] + w2 = w - (w % 2) + h2 = h - (h % 2) + x = x[..., :h2, :w2] + + return x + + def decode_tiled(self, z, tile_x=32, tile_y=32, overlap=8, tile_t=None, overlap_t=None): + # SeedVR2's causal VAE owns temporal via the MemoryState cache; external + # temporal tiling breaks that continuity, so only spatial tiling is applied. + sf = self.spatial_downsample_factor + seedvr2_tiling = { + "enable_tiling": True, + "tile_size": (tile_y * sf, tile_x * sf), + "tile_overlap": (overlap * sf, overlap * sf), + "temporal_size": None, + "temporal_overlap": None, + } + return self.decode(z, seedvr2_tiling=seedvr2_tiling) + + def encode_tiled(self, x, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + # External temporal tiling knobs are discarded; the causal VAE keeps its + # own internal MemoryState slicing. + if tile_y is None: + tile_y = 512 + if tile_x is None: + tile_x = 512 + if overlap is None: + overlap_y = 64 + overlap_x = 64 + else: + overlap_y = overlap + overlap_x = overlap + overlap_y = min(overlap_y, max(0, tile_y - 8)) + overlap_x = min(overlap_x, max(0, tile_x - 8)) + self.device = x.device + return tiled_vae( + x, + self, + tile_size=(tile_y, tile_x), + tile_overlap=(overlap_y, overlap_x), + temporal_size=None, + temporal_overlap=None, + encode=True, + ) + + def comfy_format_encoded(self, samples): + if samples.ndim == 4: + samples = samples.unsqueeze(2) + samples = samples.contiguous() + samples = samples * BYTEDANCE_VAE_SCALING_FACTOR + return samples + + def comfy_memory_used_decode(self, shape): + bytes_per_output_pixel = 160 + + def output_pixels(latent_t, latent_h, latent_w): + output_t = max(1, (latent_t - 1) * 4 + 1) + return output_t * latent_h * 8 * latent_w * 8 + + # SeedVR2 decode performs full-frame LAB histogram matching: fp32 channels + # plus int64 sort indices dominate peak memory, not the VAE weight dtype. + if len(shape) == 5: + candidates = [] + if shape[1] == SEEDVR2_LATENT_CHANNELS: + candidates.append((shape[2], shape[3], shape[4])) + if shape[-1] == SEEDVR2_LATENT_CHANNELS: + candidates.append((shape[1], shape[2], shape[3])) + if len(candidates) == 0: + candidates.append((shape[2], shape[3], shape[4])) + pixels = max(output_pixels(*candidate) for candidate in candidates) + elif len(shape) == 4: + latent_t = max(1, (shape[1] + SEEDVR2_LATENT_CHANNELS - 1) // SEEDVR2_LATENT_CHANNELS) + pixels = output_pixels(latent_t, shape[2], shape[3]) + else: + pixels = output_pixels(1, shape[-2], shape[-1]) + return pixels * bytes_per_output_pixel + + def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]): + set_norm_limit(norm_max_mem) + for m in self.modules(): + if isinstance(m, InflatedCausalConv3d): + m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf")) diff --git a/comfy/model_base.py b/comfy/model_base.py index dcfa555dc..786a7c127 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -55,6 +55,7 @@ import comfy.ldm.pixeldit.model import comfy.ldm.pixeldit.pid import comfy.ldm.ace.model import comfy.ldm.omnigen.omnigen2 +import comfy.ldm.seedvr.model import comfy.ldm.boogu.model import comfy.ldm.qwen_image.model import comfy.ldm.ideogram4.model @@ -932,6 +933,17 @@ class HunyuanDiT(BaseModel): out['image_meta_size'] = comfy.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]])) return out +class SeedVR2(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.seedvr.model.NaDiT) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + condition = kwargs.get("condition", None) + if condition is not None: + out["condition"] = comfy.conds.CONDRegular(condition) + return out + class PixArt(BaseModel): def __init__(self, model_config, model_type=ModelType.EPS, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.pixart.pixartms.PixArtMS) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index e53d848c9..174bc77cc 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -598,6 +598,44 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): return dit_config + seedvr2_7b_separate_key = "{}blocks.35.mlp.vid.proj_out.weight".format(key_prefix) + if seedvr2_7b_separate_key in state_dict_keys and state_dict[seedvr2_7b_separate_key].shape[0] == 3072: # seedvr2 7b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 3072 + dit_config["heads"] = 24 + dit_config["num_layers"] = 36 + # This checkpoint uses separate vid/txt MMModule keys in every block. + dit_config["mm_layers"] = 36 + dit_config["norm_eps"] = 1e-5 + dit_config["rope_type"] = "rope3d" + dit_config["rope_dim"] = 64 + dit_config["mlp_type"] = "normal" + return dit_config + if "{}blocks.35.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 7b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 3072 + dit_config["heads"] = 24 + dit_config["num_layers"] = 36 + # This checkpoint uses shared all.* MMModule keys after the initial blocks. + dit_config["mm_layers"] = 10 + dit_config["norm_eps"] = 1e-5 + dit_config["rope_type"] = "rope3d" + dit_config["rope_dim"] = 64 + dit_config["mlp_type"] = "swiglu" + return dit_config + if "{}blocks.31.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 3b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 2560 + dit_config["heads"] = 20 + dit_config["num_layers"] = 32 + dit_config["norm_eps"] = 1.0e-05 + dit_config["mlp_type"] = "swiglu" + dit_config["vid_out_norm"] = True + return dit_config + if '{}head.modulation'.format(key_prefix) in state_dict_keys: # Wan 2.1 dit_config = {} dit_config["image_model"] = "wan2.1" @@ -1119,9 +1157,10 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): return unet_config -def model_config_from_unet_config(unet_config, state_dict=None): + +def model_config_from_unet_config(unet_config, state_dict=None, unet_key_prefix=""): for model_config in comfy.supported_models.models: - if model_config.matches(unet_config, state_dict): + if model_config.matches(unet_config, state_dict, unet_key_prefix=unet_key_prefix): return model_config(unet_config) logging.error("no match {}".format(unet_config)) @@ -1131,7 +1170,7 @@ def model_config_from_unet(state_dict, unet_key_prefix, use_base_if_no_match=Fal unet_config = detect_unet_config(state_dict, unet_key_prefix, metadata=metadata) if unet_config is None: return None - model_config = model_config_from_unet_config(unet_config, state_dict) + model_config = model_config_from_unet_config(unet_config, state_dict, unet_key_prefix) if model_config is None and use_base_if_no_match: model_config = comfy.supported_models_base.BASE(unet_config) diff --git a/comfy/sd.py b/comfy/sd.py index 071a3102a..4a0742e7a 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -16,6 +16,7 @@ import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae import comfy.ldm.wan.vae2_2 import comfy.ldm.hunyuan3d.vae +import comfy.ldm.seedvr.vae import comfy.ldm.triposplat.vae import comfy.ldm.ace.vae.music_dcae_pipeline import comfy.ldm.cogvideo.vae @@ -473,7 +474,8 @@ class CLIP: class VAE: def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None): - if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format + is_seedvr2_vae = "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd + if not is_seedvr2_vae and 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format sd = diffusers_convert.convert_vae_state_dict(sd) if model_management.is_amd(): @@ -500,6 +502,8 @@ class VAE: self.upscale_index_formula = None self.extra_1d_channel = None self.crop_input = True + self.handles_tiling = False + self.format_encoded = None self.audio_sample_rate = 44100 @@ -546,6 +550,22 @@ class VAE: self.first_stage_model = StageC_coder() self.downscale_ratio = 32 self.latent_channels = 16 + elif "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd: # seedvr2 + self.first_stage_model = comfy.ldm.seedvr.vae.VideoAutoencoderKLWrapper() + self.latent_channels = comfy.ldm.seedvr.vae.SEEDVR2_LATENT_CHANNELS + self.latent_dim = 3 + self.disable_offload = True + self.memory_used_decode = lambda shape, dtype: self.first_stage_model.comfy_memory_used_decode(shape) + self.memory_used_encode = lambda shape, dtype: (max(shape[2], 5) * shape[3] * shape[4] * 64) * model_management.dtype_size(dtype) + self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32] + self.handles_tiling = True + self.format_encoded = self.first_stage_model.comfy_format_encoded + self.downscale_ratio = (lambda a: max(0, math.floor((a + 3) / 4)), 8, 8) + self.downscale_index_formula = (4, 8, 8) + self.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8) + self.upscale_index_formula = (4, 8, 8) + self.process_input = lambda image: image * 2.0 - 1.0 + self.crop_input = False elif "decoder.conv_in.weight" in sd: if sd['decoder.conv_in.weight'].shape[1] == 64: ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True} @@ -1012,6 +1032,10 @@ class VAE: decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device)) + def _decode_tiled_owned(self, samples, **kwargs): + out = self.first_stage_model.decode_tiled(samples.to(self.vae_dtype).to(self.device), **kwargs) + return self.process_output(out.to(device=self.output_device, dtype=self.vae_output_dtype(), copy=True)) + def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64): steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap) steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap) @@ -1048,6 +1072,25 @@ class VAE: encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device) + def _encode_tiled_owned(self, pixel_samples, **kwargs): + x = self.process_input(pixel_samples).to(self.vae_dtype).to(self.device) + out = self.first_stage_model.encode_tiled(x, **kwargs) + return out.to(device=self.output_device, dtype=self.vae_output_dtype()) + + def _owned_tiled_args(self, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + args = {} + if tile_x is not None: + args["tile_x"] = tile_x + if tile_y is not None: + args["tile_y"] = tile_y + if overlap is not None: + args["overlap"] = overlap + if tile_t is not None: + args["tile_t"] = tile_t + if overlap_t is not None: + args["overlap_t"] = overlap_t + return args + def decode(self, samples_in, vae_options={}): self.throw_exception_if_invalid() pixel_samples = None @@ -1095,11 +1138,19 @@ class VAE: if dims == 1 or self.extra_1d_channel is not None: pixel_samples = self.decode_tiled_1d(samples_in) elif dims == 2: - pixel_samples = self.decode_tiled_(samples_in) + if self.handles_tiling: + tile = 256 // self.spacial_compression_decode() + overlap = tile // 4 + pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) + else: + pixel_samples = self.decode_tiled_(samples_in) elif dims == 3: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 - pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + if self.handles_tiling: + pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) + else: + pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples @@ -1118,7 +1169,9 @@ class VAE: args["overlap"] = overlap with model_management.cuda_device_context(self.device): - if dims == 1 or self.extra_1d_channel is not None: + if self.handles_tiling and dims in (2, 3): + output = self._decode_tiled_owned(samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) + elif dims == 1 or self.extra_1d_channel is not None: args.pop("tile_y") output = self.decode_tiled_1d(samples, **args) elif dims == 2: @@ -1179,12 +1232,17 @@ class VAE: if self.latent_dim == 3: tile = 256 overlap = tile // 4 - samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + if self.handles_tiling: + samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap) + else: + samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) elif self.latent_dim == 1 or self.extra_1d_channel is not None: samples = self.encode_tiled_1d(pixel_samples) else: samples = self.encode_tiled_(pixel_samples) + if self.format_encoded is not None: + samples = self.format_encoded(samples) return samples def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): @@ -1192,7 +1250,7 @@ class VAE: pixel_samples = self.vae_encode_crop_pixels(pixel_samples) dims = self.latent_dim pixel_samples = pixel_samples.movedim(-1, 1) - if dims == 3: + if dims == 3 and pixel_samples.ndim < 5: if not self.not_video: pixel_samples = pixel_samples.movedim(1, 0).unsqueeze(0) else: @@ -1216,21 +1274,27 @@ class VAE: elif dims == 2: samples = self.encode_tiled_(pixel_samples, **args) elif dims == 3: - if tile_t is not None: - tile_t_latent = max(2, self.downscale_ratio[0](tile_t)) + if self.handles_tiling: + samples = self._encode_tiled_owned(pixel_samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) else: - tile_t_latent = 9999 - args["tile_t"] = self.upscale_ratio[0](tile_t_latent) + if tile_t is not None: + tile_t_latent = max(2, self.downscale_ratio[0](tile_t)) + else: + tile_t_latent = 9999 + args["tile_t"] = self.upscale_ratio[0](tile_t_latent) - if overlap_t is None: - args["overlap"] = (1, overlap, overlap) - else: - args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), overlap, overlap) - maximum = pixel_samples.shape[2] - maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum)) + spatial_overlap = overlap if overlap is not None else 64 + if overlap_t is None: + args["overlap"] = (1, spatial_overlap, spatial_overlap) + else: + args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), spatial_overlap, spatial_overlap) + maximum = pixel_samples.shape[2] + maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum)) - samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args) + samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args) + if self.format_encoded is not None: + samples = self.format_encoded(samples) return samples def get_sd(self): @@ -1898,7 +1962,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes) else: manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) + model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device) if model_config.clip_vision_prefix is not None: if output_clipvision: @@ -2039,7 +2103,7 @@ def load_diffusion_model_state_dict(sd, model_options={}, metadata=None, disable manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes) else: manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) + model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device) if custom_operations is not None: model_config.custom_operations = custom_operations diff --git a/comfy/supported_models.py b/comfy/supported_models.py index afb66e6f3..b82e4178f 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -1685,6 +1685,40 @@ class Chroma(supported_models_base.BASE): t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.pixart_t5.PixArtTokenizer, comfy.text_encoders.pixart_t5.pixart_te(**t5_detect)) +class SeedVR2(supported_models_base.BASE): + unet_config = { + "image_model": "seedvr2" + } + unet_extra_config = {} + required_keys = { + "{}positive_conditioning", + "{}negative_conditioning", + } + latent_format = comfy.latent_formats.SeedVR2 + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32] + sampling_settings = { + "shift": 1.0, + } + + def set_inference_dtype(self, dtype, manual_cast_dtype, device=None): + if ( + dtype == torch.float16 + and manual_cast_dtype is None + and comfy.model_management.should_use_bf16(device) + ): + manual_cast_dtype = torch.bfloat16 + super().set_inference_dtype(dtype, manual_cast_dtype, device=device) + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.SeedVR2(self, device=device) + return out + + def clip_target(self, state_dict={}): + return None + class ChromaRadiance(Chroma): unet_config = { "image_model": "chroma_radiance", @@ -2348,6 +2382,7 @@ models = [ HiDream, HiDreamO1, Chroma, + SeedVR2, ChromaRadiance, ACEStep, ACEStep15, diff --git a/comfy/supported_models_base.py b/comfy/supported_models_base.py index 0e7a829ba..e3a8e131f 100644 --- a/comfy/supported_models_base.py +++ b/comfy/supported_models_base.py @@ -54,13 +54,13 @@ class BASE: optimizations = {"fp8": False} @classmethod - def matches(s, unet_config, state_dict=None): + def matches(s, unet_config, state_dict=None, unet_key_prefix=""): for k in s.unet_config: if k not in unet_config or s.unet_config[k] != unet_config[k]: return False if state_dict is not None: for k in s.required_keys: - if k not in state_dict: + if k.format(unet_key_prefix) not in state_dict: return False return True @@ -115,7 +115,7 @@ class BASE: replace_prefix = {"": self.vae_key_prefix[0]} return utils.state_dict_prefix_replace(state_dict, replace_prefix) - def set_inference_dtype(self, dtype, manual_cast_dtype): + def set_inference_dtype(self, dtype, manual_cast_dtype, device=None): self.unet_config['dtype'] = dtype self.manual_cast_dtype = manual_cast_dtype diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index f050061ed..0bba8341b 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1088,7 +1088,7 @@ class Gemma4_Tokenizer(): h, w = samples.shape[2], samples.shape[3] patch_size = 16 pooling_k = 3 - max_soft_tokens = 70 if is_video else 280 # video uses smaller token budget per frame + max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280) max_patches = max_soft_tokens * pooling_k * pooling_k target_px = max_patches * patch_size * patch_size factor = (target_px / (h * w)) ** 0.5 diff --git a/comfy_extras/nodes_seedvr.py b/comfy_extras/nodes_seedvr.py new file mode 100644 index 000000000..c4ca3b55c --- /dev/null +++ b/comfy_extras/nodes_seedvr.py @@ -0,0 +1,614 @@ +import logging + +from typing_extensions import override +from comfy_api.latest import ComfyExtension, io +import torch + +import comfy.model_management +from comfy.ldm.seedvr.color_fix import ( + adain_color_transfer, + lab_color_transfer, + wavelet_color_transfer, +) +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE, + SEEDVR2_ADAIN_SCALE_MULTIPLIER, + SEEDVR2_CHUNK_GIB_PER_MPX_FRAME, + SEEDVR2_CHUNK_RESERVED_GIB, + SEEDVR2_CHUNK_SIGMA_GIB, + SEEDVR2_CHUNK_SIGMA_K, + SEEDVR2_COLOR_MEM_HEADROOM, + SEEDVR2_DTYPE_BYTES_FLOOR, + SEEDVR2_LAB_SCALE_MULTIPLIER, + SEEDVR2_LATENT_CHANNELS, + SEEDVR2_OOM_BACKOFF_DIVISOR, + SEEDVR2_WAVELET_SCALE_MULTIPLIER, +) + +from torchvision.transforms import functional as TVF +from torchvision.transforms.functional import InterpolationMode + + +_SEEDVR2_INVALID_MODEL_MSG_PREFIX = "SeedVR2Conditioning: model object does not match expected SeedVR2 structure" +_ATTR_MISSING = object() + + +def _resolve_seedvr2_diffusion_model(model): + inner = getattr(model, "model", _ATTR_MISSING) + if inner is _ATTR_MISSING: + raise RuntimeError( + f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input has no 'model' attribute " + f"(got type {type(model).__name__})." + ) + if inner is None: + raise RuntimeError( + f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input.model is None " + f"(input type {type(model).__name__})." + ) + diffusion_model = getattr(inner, "diffusion_model", _ATTR_MISSING) + if diffusion_model is _ATTR_MISSING: + raise RuntimeError( + f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model' has no " + f"'diffusion_model' attribute (got type {type(inner).__name__})." + ) + if diffusion_model is None: + raise RuntimeError( + f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model.diffusion_model' " + f"is None (model.model type {type(inner).__name__})." + ) + return diffusion_model + + +def div_pad(image, factor): + height_factor, width_factor = factor + height, width = image.shape[-2:] + + pad_height = (height_factor - (height % height_factor)) % height_factor + pad_width = (width_factor - (width % width_factor)) % width_factor + + if pad_height == 0 and pad_width == 0: + return image + + padding = (0, pad_width, 0, pad_height) + return torch.nn.functional.pad(image, padding, mode='constant', value=0.0) + +def cut_videos(videos): + t = videos.size(1) + if t < 1: + raise ValueError("SeedVR2Preprocess expected at least one frame.") + if t == 1: + return videos + if t <= 4: + padding = videos[:, -1:].repeat(1, 4 - t + 1, 1, 1, 1) + return torch.cat([videos, padding], dim=1) + if (t - 1) % 4 == 0: + return videos + padding = videos[:, -1:].repeat(1, 4 - ((t - 1) % 4), 1, 1, 1) + videos = torch.cat([videos, padding], dim=1) + if (videos.size(1) - 1) % 4 != 0: + raise ValueError(f"SeedVR2Preprocess failed to pad video length to 4n+1; got {videos.size(1)} frames.") + return videos + +def _seedvr2_input_shorter_edge(images, node_name): + if images.dim() == 4: + return min(images.shape[1], images.shape[2]) + if images.dim() == 5: + return min(images.shape[2], images.shape[3]) + raise ValueError( + f"{node_name}: expected 4-D or 5-D IMAGE tensor, " + f"got shape {tuple(images.shape)}" + ) + + +def _seedvr2_pad(images, upscaled_shorter_edge, node_name): + if upscaled_shorter_edge < 2: + raise ValueError( + f"{node_name}: input shorter edge must be at least 2 pixels; " + f"got {upscaled_shorter_edge}." + ) + if images.shape[-1] > 3: + images = images[..., :3] + if images.dim() == 4: + # Comfy video components arrive as a 4-D IMAGE frame sequence: + # (frames, H, W, C). SeedVR2 consumes that as one video. + images = images.unsqueeze(0) + elif images.dim() != 5: + raise ValueError( + f"{node_name}: expected 4-D or 5-D IMAGE tensor, " + f"got shape {tuple(images.shape)}" + ) + images = images.permute(0, 1, 4, 2, 3) + + b, t, c, h, w = images.shape + images = images.reshape(b * t, c, h, w) + + images = torch.clamp(images, 0.0, 1.0) + images = div_pad(images, (16, 16)) + _, _, new_h, new_w = images.shape + + images = images.reshape(b, t, c, new_h, new_w) + images = cut_videos(images) + images_bthwc = images.permute(0, 1, 3, 4, 2).contiguous() + + return io.NodeOutput(images_bthwc) + + +class SeedVR2Preprocess(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedVR2Preprocess", + display_name="Pre-Process SeedVR2 Input", + category="image/pre-processors", + description="Pad a resized image for SeedVR2 model. Alpha channel is dropped. The node Post-Process SeedVR2 Output re-applies it from the original resized image.", + search_aliases=["seedvr2", "upscale", "video upscale", "pad", "preprocess"], + inputs=[ + io.Image.Input("resized_images", tooltip="The resized image to process."), + ], + outputs=[ + io.Image.Output("images", tooltip="The padded image for VAE encoding."), + ] + ) + + @classmethod + def execute(cls, resized_images): + upscaled_shorter_edge = _seedvr2_input_shorter_edge(resized_images, "SeedVR2Preprocess") + return _seedvr2_pad( + resized_images, upscaled_shorter_edge, "SeedVR2Preprocess", + ) + + +class SeedVR2PostProcessing(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedVR2PostProcessing", + display_name="Post-Process SeedVR2 Output", + category="image/post-processors", + description="Align the generated image with the original resized image and apply color correction.", + search_aliases=["seedvr2", "upscale", "color correction", "color match", "postprocess"], + inputs=[ + io.Image.Input("images", tooltip="The generated image to process."), + io.Image.Input("original_resized_images", tooltip="The original resized image before pre-processing, used as reference."), + io.Combo.Input("color_correction_method", options=["lab", "wavelet", "adain", "none"], default="lab", tooltip="Method to match the generated image colors to the original image. lab: transfer color in CIELAB space, preserving detail (most faithful). wavelet: transfer low-frequency color, keeping upscaled high-frequency detail. adain: match per-channel mean/std (fastest, global tint). none: skip color transfer (geometry alignment only)."), + ], + outputs=[io.Image.Output(display_name="images", tooltip="The aligned, color-corrected image.")], + ) + + @classmethod + def execute(cls, images, original_resized_images, color_correction_method): + alpha_input = None + if original_resized_images.shape[-1] == 4: + alpha_input = original_resized_images[..., 3:4] + original_resized_images = original_resized_images[..., :3] + decoded_5d, decoded_was_4d = cls._as_bthwc(images) + reference_full, _ = cls._as_bthwc(original_resized_images) + decoded_5d = cls._restore_reference_batch_time(decoded_5d, reference_full) + + b = min(decoded_5d.shape[0], reference_full.shape[0]) + t = min(decoded_5d.shape[1], reference_full.shape[1]) + reference_h = reference_full.shape[2] + reference_w = reference_full.shape[3] + + decoded_5d = decoded_5d[:b, :t, :, :, :] + target_h = min(decoded_5d.shape[2], reference_h) + target_w = min(decoded_5d.shape[3], reference_w) + decoded_5d = decoded_5d[:, :, :target_h, :target_w, :] + if color_correction_method in ("lab", "wavelet", "adain"): + reference_5d = reference_full[:b, :t, :, :, :] + reference_5d = cls._resize_reference(reference_5d, target_h, target_w) + output_device = decoded_5d.device + decoded_raw = cls._to_seedvr2_raw(decoded_5d) + reference_raw = cls._to_seedvr2_raw(reference_5d) + decoded_flat = decoded_raw.permute(0, 1, 4, 2, 3).reshape(b * t, decoded_raw.shape[4], target_h, target_w) + reference_flat = reference_raw.permute(0, 1, 4, 2, 3).reshape(b * t, reference_raw.shape[4], target_h, target_w) + output = cls._color_transfer_chunked( + decoded_flat, reference_flat, output_device, color_correction_method, + ) + output = output.reshape(b, t, output.shape[1], output.shape[2], output.shape[3]).permute(0, 1, 3, 4, 2) + output = output.add(1.0).div(2.0).clamp(0.0, 1.0) + elif color_correction_method == "none": + output = decoded_5d + else: + raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}") + + if alpha_input is not None: + alpha_5d, _ = cls._as_bthwc(alpha_input) + alpha_5d = alpha_5d[:output.shape[0], :output.shape[1], :output.shape[2], :output.shape[3], :] + output = torch.cat([output, alpha_5d.to(dtype=output.dtype, device=output.device)], dim=-1) + h2 = output.shape[-3] - (output.shape[-3] % 2) + w2 = output.shape[-2] - (output.shape[-2] % 2) + output = output[:, :, :h2, :w2, :] + if decoded_was_4d: + output = output.reshape(-1, output.shape[-3], output.shape[-2], output.shape[-1]) + return io.NodeOutput(output) + + @staticmethod + def _as_bthwc(images): + if images.ndim == 4: + return images.unsqueeze(0), True + if images.ndim == 5: + return images, False + raise ValueError( + f"SeedVR2PostProcessing: expected 4-D or 5-D IMAGE tensor, got shape {tuple(images.shape)}" + ) + + @staticmethod + def _restore_reference_batch_time(decoded, reference): + if decoded.shape[0] != 1: + return decoded + ref_b, ref_t = reference.shape[:2] + if ref_b < 1 or decoded.shape[1] % ref_b != 0: + return decoded + decoded_t = decoded.shape[1] // ref_b + if decoded_t < ref_t: + return decoded + return decoded.reshape(ref_b, decoded_t, decoded.shape[2], decoded.shape[3], decoded.shape[4]) + + @staticmethod + def _to_seedvr2_raw(images): + return images.mul(2.0).sub(1.0) + + @staticmethod + def _color_transfer_on_vae_device(decoded_flat, reference_flat, output_device, transfer_fn): + color_device = comfy.model_management.vae_device() + decoded_flat = decoded_flat.to(device=color_device) + reference_flat = reference_flat.to(device=color_device) + output = transfer_fn(decoded_flat, reference_flat) + return output.to(device=output_device) + + @staticmethod + def _lab_color_transfer_on_vae_device(decoded_flat, reference_flat, output_device): + color_device = comfy.model_management.vae_device() + result = None + for start in range(decoded_flat.shape[0]): + decoded_frame = decoded_flat[start:start + 1].to(device=color_device).clone() + reference_frame = reference_flat[start:start + 1].to(device=color_device).clone() + output = lab_color_transfer(decoded_frame, reference_frame).to(device=output_device) + if result is None: + result = torch.empty( + (decoded_flat.shape[0],) + tuple(output.shape[1:]), + device=output_device, + dtype=output.dtype, + ) + result[start:start + 1].copy_(output) + if result is None: + raise ValueError("SeedVR2PostProcessing: LAB color correction requires at least one frame.") + return result + + @classmethod + def _color_transfer_chunked(cls, decoded_flat, reference_flat, output_device, color_correction_method): + chunk_size = cls._estimate_color_correction_chunk_size(decoded_flat, color_correction_method) + while True: + try: + return cls._run_color_transfer_chunks( + decoded_flat, reference_flat, output_device, color_correction_method, chunk_size, + ) + except Exception as e: + comfy.model_management.raise_non_oom(e) + if chunk_size <= 1: + raise RuntimeError( + "SeedVR2PostProcessing: color correction OOM at one frame; " + f"color_correction_method={color_correction_method}, shape={tuple(decoded_flat.shape)}." + ) from e + chunk_size = max(1, chunk_size // SEEDVR2_OOM_BACKOFF_DIVISOR) + + @classmethod + def _run_color_transfer_chunks(cls, decoded_flat, reference_flat, output_device, color_correction_method, chunk_size): + result = None + for start in range(0, decoded_flat.shape[0], chunk_size): + end = min(start + chunk_size, decoded_flat.shape[0]) + decoded_chunk = decoded_flat[start:end] + reference_chunk = reference_flat[start:end] + if color_correction_method == "lab": + output = cls._lab_color_transfer_on_vae_device(decoded_chunk, reference_chunk, output_device) + elif color_correction_method == "wavelet": + output = cls._color_transfer_on_vae_device( + decoded_chunk, reference_chunk, output_device, wavelet_color_transfer, + ) + else: + output = cls._color_transfer_on_vae_device( + decoded_chunk, reference_chunk, output_device, adain_color_transfer, + ) + if result is None: + result = torch.empty( + (decoded_flat.shape[0],) + tuple(output.shape[1:]), + device=output_device, + dtype=output.dtype, + ) + result[start:end].copy_(output) + if result is None: + raise ValueError("SeedVR2PostProcessing: color correction requires at least one frame.") + return result + + @classmethod + def _estimate_color_correction_chunk_size(cls, decoded_flat, color_correction_method): + multiplier = cls._color_correction_memory_multiplier(color_correction_method) + frames = decoded_flat.shape[0] + _, channels, height, width = decoded_flat.shape + dtype_bytes = max(decoded_flat.element_size(), SEEDVR2_DTYPE_BYTES_FLOOR) + bytes_per_frame = height * width * channels * dtype_bytes * multiplier + if bytes_per_frame <= 0: + return frames + color_device = comfy.model_management.vae_device() + free_memory = comfy.model_management.get_free_memory(color_device) + chunk_size = int((free_memory * SEEDVR2_COLOR_MEM_HEADROOM) // bytes_per_frame) + return max(1, min(frames, chunk_size)) + + @staticmethod + def _color_correction_memory_multiplier(color_correction_method): + if color_correction_method == "lab": + return SEEDVR2_LAB_SCALE_MULTIPLIER + if color_correction_method == "wavelet": + return SEEDVR2_WAVELET_SCALE_MULTIPLIER + if color_correction_method == "adain": + return SEEDVR2_ADAIN_SCALE_MULTIPLIER + raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}") + + @staticmethod + def _resize_reference(reference, height, width): + if reference.shape[2] == height and reference.shape[3] == width: + return reference + b, t = reference.shape[:2] + reference_flat = reference.permute(0, 1, 4, 2, 3).reshape(b * t, reference.shape[4], reference.shape[2], reference.shape[3]) + resized = TVF.resize( + reference_flat, + size=(height, width), + interpolation=InterpolationMode.BICUBIC, + antialias=not (isinstance(reference_flat, torch.Tensor) and reference_flat.device.type == "mps"), + ) + return resized.reshape(b, t, resized.shape[1], height, width).permute(0, 1, 3, 4, 2) + + +class SeedVR2Conditioning(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedVR2Conditioning", + display_name="Apply SeedVR2 Conditioning", + category="model/conditioning", + description="Build SeedVR2 positive/negative conditioning from a VAE latent.", + search_aliases=["seedvr2", "upscale", "conditioning"], + inputs=[ + io.Model.Input("model", tooltip="The SeedVR2 model."), + io.Latent.Input("vae_conditioning", display_name="latent"), + ], + outputs=[ + io.Conditioning.Output(display_name="positive", tooltip="The positive conditioning for sampling."), + io.Conditioning.Output(display_name="negative", tooltip="The negative conditioning for sampling."), + ], + ) + + @classmethod + def execute(cls, model, vae_conditioning) -> io.NodeOutput: + + vae_conditioning = vae_conditioning["samples"] + if vae_conditioning.ndim != 5: + raise ValueError( + "SeedVR2Conditioning expects a 5-D VAE latent in Comfy " + f"channel-first layout; got shape {tuple(vae_conditioning.shape)}." + ) + if vae_conditioning.shape[1] != SEEDVR2_LATENT_CHANNELS: + if vae_conditioning.shape[-1] == SEEDVR2_LATENT_CHANNELS: + raise ValueError( + "SeedVR2Conditioning expects SeedVR2 VAE latents in Comfy " + f"channel-first layout (B, {SEEDVR2_LATENT_CHANNELS}, T, H, W); " + f"got channel-last shape {tuple(vae_conditioning.shape)}." + ) + raise ValueError( + "SeedVR2Conditioning expects SeedVR2 VAE latents with " + f"{SEEDVR2_LATENT_CHANNELS} channels; got shape {tuple(vae_conditioning.shape)}." + ) + vae_conditioning = vae_conditioning.movedim(1, -1).contiguous() + model = _resolve_seedvr2_diffusion_model(model) + pos_cond = model.positive_conditioning + neg_cond = model.negative_conditioning + + mask = vae_conditioning.new_ones(vae_conditioning.shape[:-1] + (1,)) + condition = torch.cat((vae_conditioning, mask), dim=-1) + condition = condition.movedim(-1, 1) + + negative = [[neg_cond.unsqueeze(0), {"condition": condition}]] + positive = [[pos_cond.unsqueeze(0), {"condition": condition}]] + + return io.NodeOutput(positive, negative) + +def _seedvr2_chunk_crossfade_weights(overlap, device, dtype): + """Descending previous-chunk weights across the overlap (next chunk gets ``1 - w``): a Hann fade over the middle third, flat shoulders on the outer thirds.""" + ramp = torch.linspace(0.0, 1.0, steps=overlap, device=device, dtype=dtype) + ramp = ((ramp - 1.0 / 3.0) / (1.0 / 3.0)).clamp(0.0, 1.0) + return 0.5 + 0.5 * torch.cos(torch.pi * ramp) + + +class SeedVR2TemporalChunk(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedVR2TemporalChunk", + display_name="Split SeedVR2 Latent", + category="model/latent/batch", + description="Split a SeedVR2 video latent into overlapping temporal chunks small enough to sample one at a time within VRAM, wiring latents outputs to both Apply SeedVR2 Conditioning and the sampler latent input before recombining with Merge SeedVR2 Latents.", + search_aliases=["seedvr2", "split", "chunk", "temporal", "video upscale", "rebatch"], + inputs=[ + io.Latent.Input("latent", tooltip="The VAE-encoded SeedVR2 latent to split."), + io.Int.Input("temporal_overlap", default=0, min=0, max=16384, + tooltip="Latent frames shared between adjacent chunks and crossfaded at merge; 0 = no overlap."), + io.DynamicCombo.Input("chunking_mode", + tooltip="manual = use frames_per_chunk exactly; auto = predict the largest chunk that fits free VRAM.", + options=[ + io.DynamicCombo.Option("auto", []), + io.DynamicCombo.Option("manual", [ + io.Int.Input("frames_per_chunk", default=21, min=1, max=16384, step=4, + tooltip="Pixel frames per temporal chunk (4n+1: 1, 5, 9, 13, ...)."), + ]), + ]), + ], + outputs=[ + io.Latent.Output(display_name="latents", is_output_list=True, + tooltip="The temporal chunks in sequence order."), + io.Int.Output(display_name="temporal_overlap", + tooltip="The effective latent-frame overlap between adjacent chunks, for Merge SeedVR2 Latents."), + ], + ) + + @classmethod + def execute(cls, latent, temporal_overlap, chunking_mode) -> io.NodeOutput: + samples = latent["samples"] + if samples.ndim != 5: + raise ValueError( + f"SeedVR2TemporalChunk: expected a 5-D video latent (B, C, T, H, W); " + f"got shape {tuple(samples.shape)}." + ) + if samples.shape[1] != SEEDVR2_LATENT_CHANNELS: + raise ValueError( + f"SeedVR2TemporalChunk: expected {SEEDVR2_LATENT_CHANNELS} latent channels; " + f"got shape {tuple(samples.shape)}." + ) + if temporal_overlap < 0: + raise ValueError( + f"SeedVR2TemporalChunk: temporal_overlap must be >= 0; got {temporal_overlap}." + ) + mode = chunking_mode["chunking_mode"] + if mode not in ("auto", "manual"): + raise ValueError( + f"SeedVR2TemporalChunk: chunking_mode must be 'auto' or 'manual'; " + f"got {mode!r}." + ) + t_latent = samples.shape[2] + t_pixel = 4 * (t_latent - 1) + 1 + + if mode == "auto": + free_gb = comfy.model_management.get_free_memory( + comfy.model_management.get_torch_device()) / (1024 ** 3) + mpx_per_frame = (samples.shape[0] * samples.shape[3] * samples.shape[4]) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6 + budget_gb = free_gb - SEEDVR2_CHUNK_RESERVED_GIB - SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB + chunk_latent_max = max(1, int(budget_gb / (SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame))) + frames_per_chunk = min(4 * (chunk_latent_max - 1) + 1, t_pixel) + logging.info( + "SeedVR2TemporalChunk auto: free=%.2fGiB, %.2fMpx -> frames_per_chunk=%d (t_pixel=%d).", + free_gb, mpx_per_frame, frames_per_chunk, t_pixel, + ) + else: + frames_per_chunk = chunking_mode["frames_per_chunk"] + if frames_per_chunk < 1 or (frames_per_chunk - 1) % 4 != 0: + raise ValueError( + f"SeedVR2TemporalChunk: frames_per_chunk must be a 4n+1 pixel-frame count " + f"(1, 5, 9, 13, 17, 21, ...); got {frames_per_chunk}." + ) + + if t_pixel <= frames_per_chunk: + return io.NodeOutput([latent], 0) + + chunk_latent = (frames_per_chunk - 1) // 4 + 1 + temporal_overlap = min(temporal_overlap, chunk_latent - 1) + step = chunk_latent - temporal_overlap + + chunks = [] + for start in range(0, t_latent, step): + end = min(start + chunk_latent, t_latent) + chunk = latent.copy() + chunk["samples"] = samples[:, :, start:end].contiguous() + chunks.append(chunk) + if end >= t_latent: + break + return io.NodeOutput(chunks, temporal_overlap) + + +class SeedVR2TemporalMerge(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SeedVR2TemporalMerge", + display_name="Merge SeedVR2 Latents", + category="model/latent/batch", + is_input_list=True, + description="Recombine sampled SeedVR2 latent temporal chunks into one latent, crossfading each overlap with a Hann window sized by the temporal_overlap wired from Split SeedVR2 Latent.", + search_aliases=["seedvr2", "merge", "temporal", "hann", "crossfade"], + inputs=[ + io.Latent.Input("latents", tooltip="The sampled temporal chunks in sequence order."), + io.Int.Input("temporal_overlap", default=0, min=0, max=16384, force_input=True, + tooltip="The temporal_overlap output of Split SeedVR2 Latent. 0 = plain concatenation."), + ], + outputs=[ + io.Latent.Output(display_name="latent", tooltip="The recombined full-length latent."), + ], + ) + + @classmethod + def execute(cls, latents, temporal_overlap) -> io.NodeOutput: + temporal_overlap = temporal_overlap[0] + if temporal_overlap < 0: + raise ValueError( + f"SeedVR2TemporalMerge: temporal_overlap must be >= 0; got {temporal_overlap}." + ) + chunks = [entry["samples"] for entry in latents] + first = chunks[0] + if first.ndim != 5: + raise ValueError( + f"SeedVR2TemporalMerge: expected 5-D video latents (B, C, T, H, W); " + f"chunk 0 has shape {tuple(first.shape)}." + ) + for i, chunk in enumerate(chunks[1:], start=1): + if chunk.shape[:2] != first.shape[:2] or chunk.shape[3:] != first.shape[3:]: + raise ValueError( + f"SeedVR2TemporalMerge: chunk {i} shape {tuple(chunk.shape)} does not " + f"match chunk 0 shape {tuple(first.shape)} outside the temporal axis." + ) + if i < len(chunks) - 1 and chunk.shape[2] != first.shape[2]: + raise ValueError( + f"SeedVR2TemporalMerge: chunk {i} has {chunk.shape[2]} latent frames but " + f"chunk 0 has {first.shape[2]}; only the final chunk may be shorter." + ) + + out = latents[0].copy() + out.pop("noise_mask", None) + + if len(chunks) == 1: + out["samples"] = first + return io.NodeOutput(out) + if temporal_overlap == 0: + out["samples"] = torch.cat(chunks, dim=2) + return io.NodeOutput(out) + + chunk_latent = first.shape[2] + step = chunk_latent - min(temporal_overlap, chunk_latent - 1) + t_total = step * (len(chunks) - 1) + chunks[-1].shape[2] + b, c, _, h, w = first.shape + merged = torch.empty((b, c, t_total, h, w), device=first.device, dtype=first.dtype) + + merged[:, :, :chunk_latent] = first + filled = chunk_latent + for i, chunk in enumerate(chunks[1:], start=1): + start = i * step + end = start + chunk.shape[2] + # Crossfade width is bounded by the previous fill frontier and by a runt + # final chunk shorter than the configured overlap. + fade = min(filled - start, chunk.shape[2]) + if fade > 0: + w_prev = _seedvr2_chunk_crossfade_weights( + fade, chunk.device, chunk.dtype).view(1, 1, fade, 1, 1) + merged[:, :, start:start + fade] = ( + merged[:, :, start:start + fade] * w_prev + chunk[:, :, :fade] * (1.0 - w_prev) + ) + merged[:, :, start + fade:end] = chunk[:, :, fade:] + else: + merged[:, :, start:end] = chunk + filled = end + + out["samples"] = merged + return io.NodeOutput(out) + + +class SeedVRExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + SeedVR2Conditioning, + SeedVR2Preprocess, + SeedVR2PostProcessing, + SeedVR2TemporalChunk, + SeedVR2TemporalMerge, + ] + +async def comfy_entrypoint() -> SeedVRExtension: + return SeedVRExtension() diff --git a/nodes.py b/nodes.py index e126576fe..474e188fe 100644 --- a/nodes.py +++ b/nodes.py @@ -2458,6 +2458,7 @@ async def init_builtin_extra_nodes(): "nodes_camera_trajectory.py", "nodes_edit_model.py", "nodes_tcfg.py", + "nodes_seedvr.py", "nodes_context_windows.py", "nodes_qwen.py", "nodes_boogu.py", diff --git a/tests-unit/comfy_extras_test/test_seedvr2_conditioning.py b/tests-unit/comfy_extras_test/test_seedvr2_conditioning.py new file mode 100644 index 000000000..045502b5b --- /dev/null +++ b/tests-unit/comfy_extras_test/test_seedvr2_conditioning.py @@ -0,0 +1,186 @@ +"""SeedVR2 conditioning node regression tests.""" + +import importlib +import sys +from unittest.mock import MagicMock + +import pytest +import torch +import torch.nn as nn + +from comfy.cli_args import args as cli_args +from comfy.ldm.seedvr.constants import SEEDVR2_LATENT_CHANNELS + +if not torch.cuda.is_available(): + cli_args.cpu = True + + +_SENTINEL = object() +_TARGETS = ( + ("comfy.model_management", "comfy"), + ("comfy_extras.nodes_seedvr", "comfy_extras"), +) + + +def _import_nodes_seedvr_isolated(): + """Import comfy_extras.nodes_seedvr with comfy.model_management mocked.""" + priors = [] + for mod_name, parent_name in _TARGETS: + prior_mod = sys.modules.get(mod_name, _SENTINEL) + parent = sys.modules.get(parent_name) + attr = mod_name.split(".")[-1] + prior_attr = ( + getattr(parent, attr, _SENTINEL) if parent is not None else _SENTINEL + ) + priors.append((mod_name, parent_name, attr, prior_mod, prior_attr)) + + mock_mm = MagicMock() + for fn in ( + "xformers_enabled", "xformers_enabled_vae", + "pytorch_attention_enabled", "pytorch_attention_enabled_vae", + "sage_attention_enabled", "flash_attention_enabled", + "is_intel_xpu", + ): + getattr(mock_mm, fn).return_value = False + tv = torch.version.__version__.split(".") + mock_mm.torch_version_numeric = (int(tv[0]), int(tv[1])) + mock_mm.WINDOWS = False + sys.modules["comfy.model_management"] = mock_mm + if sys.modules.get("comfy") is None: + importlib.import_module("comfy") + comfy_pkg = sys.modules.get("comfy") + if comfy_pkg is not None: + setattr(comfy_pkg, "model_management", mock_mm) + nodes_seedvr = sys.modules.get("comfy_extras.nodes_seedvr") or ( + importlib.import_module("comfy_extras.nodes_seedvr") + ) + + def _restore(): + for mod_name, parent_name, attr, prior_mod, prior_attr in priors: + if prior_mod is _SENTINEL: + sys.modules.pop(mod_name, None) + else: + sys.modules[mod_name] = prior_mod + parent = sys.modules.get(parent_name) + if parent is None: + continue + if prior_attr is _SENTINEL: + if hasattr(parent, attr): + delattr(parent, attr) + else: + setattr(parent, attr, prior_attr) + + return nodes_seedvr, _restore + + +class _Rope(nn.Module): + def __init__(self): + super().__init__() + self.freqs = nn.Parameter(torch.zeros(4)) + + +class _Block(nn.Module): + def __init__(self): + super().__init__() + self.rope = _Rope() + + +class _DiffusionModel(nn.Module): + def __init__(self, n_blocks=3, conditioning_dtype=torch.float32): + super().__init__() + self.blocks = nn.ModuleList([_Block() for _ in range(n_blocks)]) + self.register_buffer("positive_conditioning", torch.ones((2, 4), dtype=conditioning_dtype)) + self.register_buffer("negative_conditioning", torch.zeros((3, 4), dtype=conditioning_dtype)) + + +class _ModelInner: + def __init__(self, diffusion_model): + self.diffusion_model = diffusion_model + + +class _ModelPatcher: + def __init__(self, diffusion_model): + self.model = _ModelInner(diffusion_model) + + +def test_seedvr2_conditioning_schema_exposes_conditioning_outputs(): + nodes_seedvr, restore = _import_nodes_seedvr_isolated() + try: + schema = nodes_seedvr.SeedVR2Conditioning.define_schema() + assert [input_item.id for input_item in schema.inputs] == [ + "model", + "vae_conditioning", + ] + assert schema.inputs[1].display_name == "latent" + assert [output.display_name for output in schema.outputs] == [ + "positive", + "negative", + ] + finally: + restore() + + +def test_seedvr2_conditioning_rejects_wrong_latent_channels(): + nodes_seedvr, restore = _import_nodes_seedvr_isolated() + try: + patcher = _ModelPatcher(_DiffusionModel()) + vae_conditioning = {"samples": torch.zeros(1, 8, 2, 2, 2)} + + with pytest.raises(ValueError, match=f"{SEEDVR2_LATENT_CHANNELS} channels"): + nodes_seedvr.SeedVR2Conditioning.execute(patcher, vae_conditioning) + finally: + restore() + + +def test_seedvr2_conditioning_returns_conditioning_deterministically(): + nodes_seedvr, restore = _import_nodes_seedvr_isolated() + try: + diffusion_model = _DiffusionModel() + patcher = _ModelPatcher(diffusion_model) + samples = torch.arange( + 1, + 1 + SEEDVR2_LATENT_CHANNELS * 3 * 2 * 2, + dtype=torch.float32, + ).reshape(1, SEEDVR2_LATENT_CHANNELS, 3, 2, 2) + vae_conditioning = {"samples": samples} + + first_positive, first_negative = ( + nodes_seedvr.SeedVR2Conditioning.execute( + patcher, + vae_conditioning, + ) + ) + second_positive, second_negative = ( + nodes_seedvr.SeedVR2Conditioning.execute( + patcher, + vae_conditioning, + ) + ) + + channel_last = samples.movedim(1, -1).contiguous() + expected_condition = torch.cat( + [ + channel_last, + torch.ones((*channel_last.shape[:-1], 1)), + ], + dim=-1, + ).movedim(-1, 1) + + assert torch.equal( + first_positive[0][1]["condition"], + expected_condition, + ) + assert torch.equal( + second_positive[0][1]["condition"], + expected_condition, + ) + assert torch.equal( + first_negative[0][1]["condition"], + expected_condition, + ) + assert torch.equal( + second_negative[0][1]["condition"], + expected_condition, + ) + finally: + restore() diff --git a/tests-unit/comfy_extras_test/test_seedvr2_nodes.py b/tests-unit/comfy_extras_test/test_seedvr2_nodes.py new file mode 100644 index 000000000..1c5d20ac9 --- /dev/null +++ b/tests-unit/comfy_extras_test/test_seedvr2_nodes.py @@ -0,0 +1,55 @@ +import importlib +import inspect +import sys +from unittest.mock import MagicMock, patch + +import torch + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + + +def test_seedvr_node_signature_matches_schema(): + mock_mm = MagicMock() + mock_mm.xformers_enabled.return_value = False + mock_mm.xformers_enabled_vae.return_value = False + mock_mm.sage_attention_enabled.return_value = False + mock_mm.flash_attention_enabled.return_value = False + + sentinel = object() + prior_cpu = cli_args.cpu + cli_args.cpu = True + prior_module = sys.modules.get("comfy_extras.nodes_seedvr", sentinel) + comfy_pkg = sys.modules.get("comfy") + prior_mm_attr = getattr(comfy_pkg, "model_management", sentinel) if comfy_pkg else sentinel + + with patch.dict(sys.modules, {"comfy.model_management": mock_mm}): + if comfy_pkg is not None: + setattr(comfy_pkg, "model_management", mock_mm) + sys.modules.pop("comfy_extras.nodes_seedvr", None) + try: + nodes_seedvr = importlib.import_module("comfy_extras.nodes_seedvr") + for node_cls in (nodes_seedvr.SeedVR2Preprocess, nodes_seedvr.SeedVR2PostProcessing, nodes_seedvr.SeedVR2Conditioning): + schema_ids = [i.id for i in node_cls.define_schema().inputs] + exec_params = [ + p for p in inspect.signature(node_cls.execute).parameters.keys() + if p != "cls" + ] + assert schema_ids == exec_params, ( + f"{node_cls.__name__} schema/execute drift: " + f"schema_ids={schema_ids}, exec_params={exec_params}" + ) + finally: + cli_args.cpu = prior_cpu + if prior_module is sentinel: + sys.modules.pop("comfy_extras.nodes_seedvr", None) + else: + sys.modules["comfy_extras.nodes_seedvr"] = prior_module + if comfy_pkg is not None: + if prior_mm_attr is sentinel: + if hasattr(comfy_pkg, "model_management"): + delattr(comfy_pkg, "model_management") + else: + setattr(comfy_pkg, "model_management", prior_mm_attr) diff --git a/tests-unit/comfy_extras_test/test_seedvr2_post_processing.py b/tests-unit/comfy_extras_test/test_seedvr2_post_processing.py new file mode 100644 index 000000000..6c821136d --- /dev/null +++ b/tests-unit/comfy_extras_test/test_seedvr2_post_processing.py @@ -0,0 +1,51 @@ +from unittest.mock import patch + +import pytest +import torch + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +from comfy_extras import nodes_seedvr # noqa: E402 + + +def _schema_ids(items): + return [item.id for item in items] + + +def test_seedvr2_post_processing_schema(): + schema = nodes_seedvr.SeedVR2PostProcessing.define_schema() + + assert _schema_ids(schema.inputs) == ["images", "original_resized_images", "color_correction_method"] + assert schema.inputs[2].options == ["lab", "wavelet", "adain", "none"] + assert schema.inputs[2].default == "lab" + assert schema.outputs[0].get_io_type() == "IMAGE" + + +def test_seedvr2_post_processing_oom_error_uses_color_correction_method(monkeypatch): + decoded = torch.full((1, 3, 4, 4), 0.25) + reference = torch.full((1, 3, 4, 4), 0.75) + + def _lab(content, style): + raise torch.cuda.OutOfMemoryError("CUDA out of memory") + + monkeypatch.setattr(nodes_seedvr.comfy.model_management, "vae_device", lambda: torch.device("cpu")) + monkeypatch.setattr(nodes_seedvr.comfy.model_management, "get_free_memory", lambda device: 1_000_000) + + with patch.object(nodes_seedvr, "lab_color_transfer", _lab): + with pytest.raises(RuntimeError) as excinfo: + nodes_seedvr.SeedVR2PostProcessing._color_transfer_chunked( + decoded, reference, torch.device("cpu"), "lab", + ) + assert "color_correction_method=lab" in str(excinfo.value) + assert " method=lab" not in str(excinfo.value) + + +def test_seedvr2_post_processing_unknown_color_correction_method_raises(): + decoded = torch.zeros(1, 2, 4, 4, 3) + original = torch.zeros(1, 2, 4, 4, 3) + with pytest.raises(ValueError) as excinfo: + nodes_seedvr.SeedVR2PostProcessing.execute(decoded, original, "bogus") + assert "color_correction_method" in str(excinfo.value) diff --git a/tests-unit/comfy_extras_test/test_seedvr2_temporal_chunk.py b/tests-unit/comfy_extras_test/test_seedvr2_temporal_chunk.py new file mode 100644 index 000000000..328355b49 --- /dev/null +++ b/tests-unit/comfy_extras_test/test_seedvr2_temporal_chunk.py @@ -0,0 +1,77 @@ +"""SeedVR2 temporal chunk/merge node regression tests.""" + +import pytest +import torch + +from comfy.cli_args import args as cli_args +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE, + SEEDVR2_CHUNK_GIB_PER_MPX_FRAME, + SEEDVR2_CHUNK_RESERVED_GIB, + SEEDVR2_CHUNK_SIGMA_GIB, + SEEDVR2_CHUNK_SIGMA_K, + SEEDVR2_LATENT_CHANNELS, +) + +if not torch.cuda.is_available(): + cli_args.cpu = True + +import comfy.model_management # noqa: E402 +from comfy_extras.nodes_seedvr import SeedVR2TemporalChunk, SeedVR2TemporalMerge, _seedvr2_chunk_crossfade_weights # noqa: E402 + +def _latent(t_latent, h=8, w=8, b=1): + g = torch.Generator().manual_seed(7) + return {"samples": torch.randn(b, SEEDVR2_LATENT_CHANNELS, t_latent, h, w, generator=g)} + +def _split(latent, frames_per_chunk, temporal_overlap, chunking_mode="manual"): + combo = {"chunking_mode": chunking_mode} + if chunking_mode != "auto": + combo["frames_per_chunk"] = frames_per_chunk + return SeedVR2TemporalChunk.execute(latent, temporal_overlap, combo).args + +def _merge(chunks, temporal_overlap): + return SeedVR2TemporalMerge.execute(chunks, [temporal_overlap]).args[0] + +def test_chunk_temporal_windows_and_validation(): + with pytest.raises(ValueError, match="4n\\+1"): + _split(_latent(9), 20, 0) + with pytest.raises(ValueError, match="5-D"): + _split({"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS * 9, 8, 8)}, 21, 0) + with pytest.raises(ValueError, match="chunking_mode"): + _split(_latent(13), 21, 0, "adaptive") + latent = _latent(13) + chunks, overlap = _split(latent, 21, 2) # chunk_latent=6, step=4 -> [0:6], [4:10], [8:13] + assert overlap == 2 and [c["samples"].shape[2] for c in chunks] == [6, 6, 5] + assert all(torch.equal(c["samples"], latent["samples"][:, :, s:e]) for c, (s, e) in zip(chunks, [(0, 6), (4, 10), (8, 13)])) + assert len(_split(_latent(13), 21, 999)[0]) == 8 # overlap clamps to chunk_latent-1 -> step=1 + assert (r := _split(_latent(5), 21, 3)) and len(r[0]) == 1 and r[1] == 0 # t_pixel <= 21: passthrough + +def test_chunk_auto_mode_applies_vram_law(monkeypatch): + mpx_per_frame = (32 * 32) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6 + free_gb = ( + SEEDVR2_CHUNK_RESERVED_GIB + + SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB + + 5.1 * SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame + ) + monkeypatch.setattr(comfy.model_management, "get_free_memory", lambda dev=None: free_gb * (1024 ** 3)) + assert [c["samples"].shape[2] for c in _split(_latent(13, h=32, w=32), 1, 0, "auto")[0]] == [5, 5, 3] + assert _split(_latent(13, h=32, w=32, b=2), 1, 0, "auto")[0][0]["samples"].shape[2] == 2 # batch halves the chunk + +def test_merge_crossfade_and_reassembly(): + latent = _latent(13) + latent["noise_mask"] = torch.rand(1, 1, 13, 8, 8) + latent["batch_index"] = [0] + merged = _merge(_split(latent, 21, 0)[0], 0) + assert torch.equal(merged["samples"], latent["samples"]) + assert "noise_mask" not in merged and merged["batch_index"] == [0] + assert torch.allclose(_merge(_split(latent, 21, 3)[0], 3)["samples"], latent["samples"], atol=1e-6) + w = _seedvr2_chunk_crossfade_weights(3, merged["samples"].device, merged["samples"].dtype) + assert w[0] == 1.0 and w[-1] == 0.0 and torch.all(w[:-1] >= w[1:]) + ones, zeros = {"samples": torch.ones(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}, {"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)} + fused = _merge([ones, zeros], 3)["samples"] # overlap equals w: prev fades out, next fades in + assert torch.equal(fused[:, :, 3:6], w.view(1, 1, 3, 1, 1).expand(1, SEEDVR2_LATENT_CHANNELS, 3, 8, 8)) + assert torch.equal(fused[:, :, :3], ones["samples"][:, :, :3]) and torch.equal(fused[:, :, 6:], zeros["samples"][:, :, :3]) + short = _split(latent, 21, 2)[0] + short[0]["samples"] = short[0]["samples"][:, :, :4] + with pytest.raises(ValueError, match="only the final chunk may be shorter"): + _merge(short, 2) diff --git a/tests-unit/comfy_test/model_detection_test.py b/tests-unit/comfy_test/model_detection_test.py index 4e9350602..6e7d71f79 100644 --- a/tests-unit/comfy_test/model_detection_test.py +++ b/tests-unit/comfy_test/model_detection_test.py @@ -2,7 +2,7 @@ from collections import defaultdict import torch -from comfy.model_detection import detect_unet_config, model_config_from_unet_config +from comfy.model_detection import detect_unet_config, model_config_from_unet, model_config_from_unet_config import comfy.supported_models @@ -73,6 +73,34 @@ def _make_flux_schnell_comfyui_sd(): return sd +def _make_seedvr2_7b_separate_mm_sd(): + return { + "blocks.35.mlp.vid.proj_out.weight": torch.empty(3072, 1), + "positive_conditioning": torch.empty(58, 5120), + "negative_conditioning": torch.empty(64, 5120), + } + + +def _make_seedvr2_7b_shared_mm_sd(): + return { + "blocks.35.mlp.all.proj_in_gate.weight": torch.empty(1, 1), + "positive_conditioning": torch.empty(58, 5120), + "negative_conditioning": torch.empty(64, 5120), + } + + +def _make_seedvr2_3b_shared_mm_sd(): + return { + "blocks.31.mlp.all.proj_in_gate.weight": torch.empty(1, 1), + "positive_conditioning": torch.empty(58, 5120), + "negative_conditioning": torch.empty(64, 5120), + } + + +def _add_model_diffusion_prefix(sd): + return {f"model.diffusion_model.{k}": v for k, v in sd.items()} + + class TestModelDetection: """Verify that first-match model detection selects the correct model based on list ordering and unet_config specificity.""" @@ -125,6 +153,59 @@ class TestModelDetection: assert model_config is not None assert type(model_config).__name__ == "FluxSchnell" + def test_seedvr2_7b_separate_mm_detection_config(self): + sd = _make_seedvr2_7b_separate_mm_sd() + unet_config = detect_unet_config(sd, "") + + assert unet_config is not None + assert unet_config["image_model"] == "seedvr2" + assert unet_config["vid_dim"] == 3072 + assert unet_config["heads"] == 24 + assert unet_config["num_layers"] == 36 + assert unet_config["mm_layers"] == 36 + assert unet_config["mlp_type"] == "normal" + assert unet_config["rope_type"] == "rope3d" + assert unet_config["rope_dim"] == 64 + + def test_seedvr2_7b_shared_mm_detection_config(self): + sd = _make_seedvr2_7b_shared_mm_sd() + unet_config = detect_unet_config(sd, "") + + assert unet_config is not None + assert unet_config["image_model"] == "seedvr2" + assert unet_config["vid_dim"] == 3072 + assert unet_config["heads"] == 24 + assert unet_config["num_layers"] == 36 + assert unet_config["mm_layers"] == 10 + assert unet_config["mlp_type"] == "swiglu" + assert unet_config["rope_type"] == "rope3d" + assert unet_config["rope_dim"] == 64 + + def test_seedvr2_3b_shared_mm_detection_config(self): + sd = _make_seedvr2_3b_shared_mm_sd() + unet_config = detect_unet_config(sd, "") + + assert unet_config is not None + assert unet_config["image_model"] == "seedvr2" + assert unet_config["vid_dim"] == 2560 + assert unet_config["heads"] == 20 + assert unet_config["num_layers"] == 32 + assert unet_config["mlp_type"] == "swiglu" + + def test_seedvr2_model_match_requires_conditioning_tensors(self): + sd = _make_seedvr2_7b_shared_mm_sd() + unet_config = detect_unet_config(sd, "") + + assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "SeedVR2" + + del sd["positive_conditioning"] + assert model_config_from_unet_config(unet_config, sd) is None + + def test_seedvr2_model_match_accepts_full_checkpoint_prefix(self): + sd = _add_model_diffusion_prefix(_make_seedvr2_7b_shared_mm_sd()) + + assert type(model_config_from_unet(sd, "model.diffusion_model.")).__name__ == "SeedVR2" + def test_unet_config_and_required_keys_combination_is_unique(self): """Each model in the registry must have a unique combination of ``unet_config`` and ``required_keys``. If two models share the same diff --git a/tests-unit/comfy_test/seedvr_vae_forward_test.py b/tests-unit/comfy_test/seedvr_vae_forward_test.py new file mode 100644 index 000000000..7ea7a143e --- /dev/null +++ b/tests-unit/comfy_test/seedvr_vae_forward_test.py @@ -0,0 +1,74 @@ +"""Regression tests for the SeedVR2 VAE forward return contract.""" + +import pytest +import torch +import torch.nn as nn + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +from comfy.ldm.seedvr.vae import SEEDVR2_LATENT_CHANNELS, VideoAutoencoderKL # noqa: E402 + + +_LATENT_SHAPE = (1, SEEDVR2_LATENT_CHANNELS, 2, 2, 2) +_DECODED_SHAPE = (1, 3, 5, 16, 16) +_INPUT_ENCODE_SHAPE = (1, 3, 5, 16, 16) +_INPUT_DECODE_SHAPE = _LATENT_SHAPE + + +class _StubVAE(VideoAutoencoderKL): + def __init__(self): + nn.Module.__init__(self) + self._encode_out = torch.zeros(*_LATENT_SHAPE) + self._decode_out = torch.zeros(*_DECODED_SHAPE) + + def encode(self, x, return_dict=True): + return self._encode_out + + def decode_(self, z, return_dict=True): + return self._decode_out + + +def test_forward_encode_returns_tensor(): + vae = _StubVAE() + x = torch.zeros(*_INPUT_ENCODE_SHAPE) + result = vae.forward(x, mode="encode") + assert type(result) is torch.Tensor + assert result.shape == torch.Size(_LATENT_SHAPE) + + +def test_forward_decode_returns_tensor(): + vae = _StubVAE() + z = torch.zeros(*_INPUT_DECODE_SHAPE) + result = vae.forward(z, mode="decode") + assert type(result) is torch.Tensor + assert result.shape == torch.Size(_DECODED_SHAPE) + + +class _TupleReturningStubVAE(VideoAutoencoderKL): + def __init__(self): + nn.Module.__init__(self) + self._encode_tensor = torch.zeros(*_LATENT_SHAPE) + self._decode_tensor = torch.zeros(*_DECODED_SHAPE) + + def encode(self, x, return_dict=True): + return (self._encode_tensor,) + + def decode_(self, z, return_dict=True): + return (self._decode_tensor,) + + +def test_forward_all_unwraps_one_tuple_at_each_step(): + vae = _TupleReturningStubVAE() + x = torch.zeros(*_INPUT_ENCODE_SHAPE) + result = vae.forward(x, mode="all") + assert type(result) is torch.Tensor + assert result.shape == torch.Size(_DECODED_SHAPE) + + +def test_forward_rejects_unknown_mode(): + vae = _StubVAE() + with pytest.raises(ValueError, match="Unknown SeedVR2 VAE forward mode"): + vae.forward(torch.zeros(*_INPUT_ENCODE_SHAPE), mode="bogus") diff --git a/tests-unit/comfy_test/test_seedvr2_dtype.py b/tests-unit/comfy_test/test_seedvr2_dtype.py new file mode 100644 index 000000000..8e08b6dde --- /dev/null +++ b/tests-unit/comfy_test/test_seedvr2_dtype.py @@ -0,0 +1,50 @@ +import torch + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +import comfy.sd +import comfy.supported_models +import comfy.ldm.seedvr.model as seedvr_model +import comfy.ldm.seedvr.vae as seedvr_vae + + +def test_seedvr2_fp16_manual_cast_only_for_bf16_device(monkeypatch): + bf16_device = object() + fp16_device = object() + + monkeypatch.setattr( + comfy.supported_models.comfy.model_management, + "should_use_bf16", + lambda device=None: device is bf16_device, + ) + + bf16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"}) + bf16_config.set_inference_dtype(torch.float16, None, device=bf16_device) + assert bf16_config.manual_cast_dtype is torch.bfloat16 + + fp16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"}) + fp16_config.set_inference_dtype(torch.float16, None, device=fp16_device) + assert fp16_config.manual_cast_dtype is None + + +def test_seedvr2_text_conditioning_accepts_cfg1_single_branch(): + context = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2) + + txt, txt_shape = seedvr_model.NaDiT._resolve_text_conditioning(object(), context, [0]) + + torch.testing.assert_close(txt, context.squeeze(0)) + torch.testing.assert_close(txt_shape, torch.tensor([[3]], device=context.device)) + + +def test_seedvr2_vae_decode_memory_covers_full_frame_lab_transfer(): + wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper) + latent_channels = seedvr_vae.SEEDVR2_LATENT_CHANNELS + estimate = wrapper.comfy_memory_used_decode((1, latent_channels, 26, 120, 160)) + old_estimate = latent_channels * 120 * 160 * (4 * 8 * 8) * 2 + + assert estimate == 101 * 960 * 1280 * 160 + assert estimate > 15 * 1024 ** 3 + assert estimate > old_estimate * 100 diff --git a/tests-unit/comfy_test/test_seedvr2_internals.py b/tests-unit/comfy_test/test_seedvr2_internals.py new file mode 100644 index 000000000..fe4bde1c4 --- /dev/null +++ b/tests-unit/comfy_test/test_seedvr2_internals.py @@ -0,0 +1,169 @@ +"""SeedVR2 internals regression tests.""" + +from __future__ import annotations + +from unittest.mock import patch + +import pytest +import torch + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +import comfy.ldm.seedvr.model as seedvr_model # noqa: E402 +import comfy.ldm.seedvr.vae as vae_mod # noqa: E402 +import comfy.ldm.modules.attention as attention # noqa: E402 +import comfy.ops as comfy_ops # noqa: E402 +from comfy.ldm.seedvr.vae import ( # noqa: E402 + causal_norm_wrapper, + set_norm_limit, +) +from comfy.ldm.seedvr.attention import var_attention_optimized_split # noqa: E402 + + +_NUM_CHANNELS = 8 +_NUM_GROUPS = 4 +_TENSOR_SHAPE = (1, 8, 2, 4, 4) + +_GROUPNORM_SUBCLASSES = [ + pytest.param(comfy_ops.disable_weight_init.GroupNorm, id="disable_weight_init"), + pytest.param(comfy_ops.manual_cast.GroupNorm, id="manual_cast"), +] + + +@pytest.mark.parametrize("groupnorm_cls", _GROUPNORM_SUBCLASSES) +def test_seedvr_groupnorm_low_limit_uses_chunked_groupnorm_path(groupnorm_cls): + real_group_norm = vae_mod.F.group_norm + set_norm_limit(1e-9) + try: + gn = groupnorm_cls(num_channels=_NUM_CHANNELS, num_groups=_NUM_GROUPS) + gn.eval() + + forward_hook_calls = [] + + def _hook(module, inputs, output): + forward_hook_calls.append(tuple(inputs[0].shape)) + + spy_calls = [] + + def _group_norm_spy(input_tensor, num_groups_arg, *args, **kwargs): + spy_calls.append({"num_groups": int(num_groups_arg)}) + return real_group_norm(input_tensor, num_groups_arg, *args, **kwargs) + + handle = gn.register_forward_hook(_hook) + try: + with patch.object(vae_mod.F, "group_norm", side_effect=_group_norm_spy): + out_tensor = causal_norm_wrapper(gn, torch.randn(*_TENSOR_SHAPE)) + finally: + handle.remove() + + full_calls = len(forward_hook_calls) + chunked_calls = sum(1 for entry in spy_calls if entry["num_groups"] < _NUM_GROUPS) + + assert tuple(int(s) for s in out_tensor.shape) == _TENSOR_SHAPE + assert full_calls == 0, ( + f"low-limit GroupNorm gate must NOT take the full-forward path; got full_calls={full_calls}" + ) + assert chunked_calls > 0, ( + f"low-limit GroupNorm gate must take the chunked path; got chunked_calls={chunked_calls}" + ) + finally: + set_norm_limit(None) + + +def test_seedvr2_7b_swin_attention_forward_uses_optimized_var_attention(monkeypatch): + dim = 8 + heads = 2 + head_dim = 4 + attn = seedvr_model.NaSwinAttention( + vid_dim=dim, + txt_dim=dim, + heads=heads, + head_dim=head_dim, + qk_bias=False, + qk_norm=comfy_ops.disable_weight_init.RMSNorm, + qk_norm_eps=1e-6, + rope_type=None, + rope_dim=head_dim, + shared_weights=False, + window=(2, 1, 1), + window_method="720pwin_by_size_bysize", + version=True, + device="cpu", + dtype=torch.float32, + operations=comfy_ops.disable_weight_init, + ) + generator = torch.Generator(device="cpu").manual_seed(11) + vid = torch.randn(8, dim, generator=generator) + txt = torch.randn(3, dim, generator=generator) + vid_shape = torch.tensor([[2, 2, 2]], dtype=torch.long) + txt_shape = torch.tensor([[3]], dtype=torch.long) + calls = [] + + def fake_optimized_var_attention(**kwargs): + calls.append(kwargs) + return kwargs["q"] + + monkeypatch.setattr(seedvr_model, "optimized_var_attention", fake_optimized_var_attention) + + vid_out, txt_out = attn(vid, txt, vid_shape, txt_shape, seedvr_model.Cache(disable=True)) + + assert tuple(vid_out.shape) == (8, dim) + assert tuple(txt_out.shape) == (3, dim) + assert len(calls) == 1 + call = calls[0] + assert tuple(call["q"].shape) == (14, heads, head_dim) + assert tuple(call["k"].shape) == (14, heads, head_dim) + assert tuple(call["v"].shape) == (14, heads, head_dim) + assert call["heads"] == heads + assert call["skip_reshape"] is True + assert call["skip_output_reshape"] is True + assert call["cu_seqlens_q"] == [0, 7, 14] + assert call["cu_seqlens_k"] == [0, 7, 14] + + +def test_var_attention_optimized_split_calls_dense_backend_per_window(monkeypatch): + heads = 2 + head_dim = 3 + q = torch.arange(30, dtype=torch.float32).reshape(5, heads, head_dim) + k = q + 100 + v = q + 200 + cu = [0, 2, 5] + calls = [] + + def fake_optimized_attention(q_arg, k_arg, v_arg, heads_arg, **kwargs): + calls.append( + { + "q_shape": tuple(q_arg.shape), + "k_shape": tuple(k_arg.shape), + "v_shape": tuple(v_arg.shape), + "heads": heads_arg, + "kwargs": kwargs, + } + ) + return q_arg + v_arg + + monkeypatch.setattr(attention, "optimized_attention", fake_optimized_attention) + + out = var_attention_optimized_split( + q, + k, + v, + heads, + cu, + cu, + skip_reshape=True, + skip_output_reshape=True, + ) + + assert tuple(out.shape) == (5, heads, head_dim) + assert len(calls) == 2 + assert calls[0]["q_shape"] == (1, heads, 2, head_dim) + assert calls[1]["q_shape"] == (1, heads, 3, head_dim) + assert all(call["heads"] == heads for call in calls) + assert all(call["kwargs"]["skip_reshape"] is True for call in calls) + assert all(call["kwargs"]["skip_output_reshape"] is True for call in calls) + torch.testing.assert_close(out, q + v, rtol=0, atol=0) + diff --git a/tests-unit/comfy_test/test_seedvr2_model.py b/tests-unit/comfy_test/test_seedvr2_model.py new file mode 100644 index 000000000..1d454aaf1 --- /dev/null +++ b/tests-unit/comfy_test/test_seedvr2_model.py @@ -0,0 +1,320 @@ +"""SeedVR2 model, latent-format, and VAE graph regression tests.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +import torch +from torch import nn + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +import comfy # noqa: E402 +import comfy.latent_formats # noqa: E402 +import comfy.ldm.seedvr.model as seedvr_model # noqa: E402 +import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402 +import comfy.model_management # noqa: E402 +import comfy.ops as comfy_ops # noqa: E402 +import comfy.sample # noqa: E402 +import comfy.sd as sd_mod # noqa: E402 +import nodes as nodes_mod # noqa: E402 +from comfy.ldm.seedvr.model import NaDiT # noqa: E402 + + +_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS + + +def _make_standin(positive_conditioning): + class _StandIn(torch.nn.Module): + def __init__(self): + super().__init__() + self.register_buffer( + "positive_conditioning", positive_conditioning + ) + + _resolve_text_conditioning = NaDiT._resolve_text_conditioning + + return _StandIn() + + +class _StubModule(nn.Module): + def __init__(self, *args, **kwargs): + super().__init__() + + +def _capture_last_layer_flags(monkeypatch, vid_dim: int, txt_in_dim: int) -> list[bool]: + flags = [] + + class _Block(_StubModule): + def __init__(self, *args, **kwargs): + flags.append(kwargs["is_last_layer"]) + super().__init__() + + monkeypatch.setattr(seedvr_model, "NaPatchIn", _StubModule) + monkeypatch.setattr(seedvr_model, "NaPatchOut", _StubModule) + monkeypatch.setattr(seedvr_model, "TimeEmbedding", _StubModule) + monkeypatch.setattr(seedvr_model, "NaMMSRTransformerBlock", _Block) + + seedvr_model.NaDiT( + norm_eps=1e-5, + num_layers=4, + mlp_type="normal", + vid_dim=vid_dim, + txt_in_dim=txt_in_dim, + heads=24, + mm_layers=3, + operations=comfy_ops.disable_weight_init, + ) + + return flags + + +class _Model: + def __init__(self, latent_format): + self._latent_format = latent_format + + def get_model_object(self, name): + assert name == "latent_format" + return self._latent_format + + +class _Patcher: + def get_free_memory(self, device): + return 1024 * 1024 * 1024 + + +class _EncodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper): + def __init__(self, encoded): + nn.Module.__init__(self) + self.encoded = encoded + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.seen = [] + + def encode(self, x): + self.seen.append(tuple(x.shape)) + return self.encoded.to(device=x.device, dtype=x.dtype) + + +class _DecodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper): + def __init__(self): + nn.Module.__init__(self) + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.calls = [] + + def decode(self, z, seedvr2_tiling=None): + self.calls.append({"shape": tuple(z.shape), "seedvr2_tiling": seedvr2_tiling}) + if z.ndim == 4: + b, tc, h, w = z.shape + t = tc // _LATENT_CHANNELS + else: + b, _, t, h, w = z.shape + return torch.zeros(b, 3, t, h * 8, w * 8, dtype=z.dtype, device=z.device) + + +def test_seedvr2_wrapper_public_encode_returns_tensor(monkeypatch): + raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 2.0) + seen_shapes = [] + + def base_encode(self, x): + seen_shapes.append(tuple(x.shape)) + return raw_latent.to(device=x.device, dtype=x.dtype) + + monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode) + + vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper) + nn.Module.__init__(vae) + vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32)) + + latent = vae.encode(torch.zeros(1, 3, 32, 40)) + + assert type(latent) is torch.Tensor + assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5) + assert seen_shapes == [(1, 3, 1, 32, 40)] + + +def test_seedvr2_wrapper_private_encode_helper_keeps_raw_latent(monkeypatch): + raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 3.0) + + def base_encode(self, x): + return raw_latent.to(device=x.device, dtype=x.dtype) + + monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode) + + vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper) + nn.Module.__init__(vae) + vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32)) + + latent, raw = vae._encode_with_raw_latent(torch.zeros(1, 3, 32, 40)) + + assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5) + assert tuple(raw.shape) == (1, _LATENT_CHANNELS, 1, 4, 5) + assert torch.equal(raw, raw_latent) + + +def _make_vae(wrapper): + vae = sd_mod.VAE.__new__(sd_mod.VAE) + vae.first_stage_model = wrapper + vae.device = torch.device("cpu") + vae.output_device = torch.device("cpu") + vae.vae_dtype = torch.float32 + vae.latent_channels = _LATENT_CHANNELS + vae.latent_dim = 3 + vae.downscale_ratio = (lambda a: max(0, (a + 3) // 4), 8, 8) + vae.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8) + vae.output_channels = 3 + vae.disable_offload = True + vae.extra_1d_channel = None + vae.crop_input = False + vae.not_video = False + vae.handles_tiling = isinstance(wrapper, seedvr_vae_mod.VideoAutoencoderKLWrapper) + vae.format_encoded = wrapper.comfy_format_encoded + vae.patcher = _Patcher() + vae.process_input = lambda image: image + vae.process_output = lambda image: image.add(1.0).div(2.0).clamp(0.0, 1.0) + vae.vae_output_dtype = lambda: torch.float32 + vae.memory_used_encode = lambda shape, dtype: 1 + vae.memory_used_decode = lambda shape, dtype: 1 + vae.throw_exception_if_invalid = lambda: None + vae.vae_encode_crop_pixels = lambda pixels: pixels + vae.spacial_compression_decode = lambda: 8 + vae.temporal_compression_decode = lambda: 4 + return vae + + +def test_missing_context_falls_back_to_positive_buffer(): + pos_buffer = torch.full((58, 5120), 7.0) + standin = _make_standin(pos_buffer) + txt, txt_shape = standin._resolve_text_conditioning(None) + assert txt.shape == (58, 5120) + assert (txt == 7.0).all(), ( + "fallback path must use the positive_conditioning buffer " + "verbatim, not a zero tensor" + ) + assert txt_shape.shape == (1, 1) + assert txt_shape[0, 0].item() == 58 + + +def test_seedvr2_7b_keeps_final_block_text_path(monkeypatch): + assert _capture_last_layer_flags(monkeypatch, vid_dim=3072, txt_in_dim=3072) == [ + False, + False, + False, + False, + ] + + +def test_seedvr2_7b_rope3d_matches_wrapper_oracle(): + rope = seedvr_model.get_na_rope("rope3d", dim=64) + generator = torch.Generator(device="cpu").manual_seed(0) + q = torch.randn(4, 2, 128, generator=generator) + k = torch.randn(4, 2, 128, generator=generator) + shape = torch.tensor([[1, 2, 2]], dtype=torch.long) + freqs = rope.get_axial_freqs(1, 2, 2).reshape(4, -1) + + expected_q = seedvr_model._apply_seedvr2_rotary_emb( + freqs, + q.permute(1, 0, 2).float(), + ).to(q.dtype).permute(1, 0, 2) + expected_k = seedvr_model._apply_seedvr2_rotary_emb( + freqs, + k.permute(1, 0, 2).float(), + ).to(k.dtype).permute(1, 0, 2) + + actual_q, actual_k = rope(q.clone(), k.clone(), shape, seedvr_model.Cache(disable=True)) + + torch.testing.assert_close(actual_q, expected_q, rtol=0, atol=0) + torch.testing.assert_close(actual_k, expected_k, rtol=0, atol=0) + + +def test_seedvr2_forward_requires_conditioning_latents(): + model = NaDiT.__new__(NaDiT) + x = torch.zeros(1, _LATENT_CHANNELS, 1, 4, 5) + + with pytest.raises(ValueError, match="requires conditioning latents"): + NaDiT.forward(model, x, timestep=torch.tensor([1.0]), context=None) + + +def test_seedvr2_latent_format_uses_native_video_latent_shape(): + latent_format = comfy.latent_formats.SeedVR2() + latent_image = torch.zeros(1, 1, 4, 5) + + fixed = comfy.sample.fix_empty_latent_channels(_Model(latent_format), latent_image) + + assert latent_format.latent_channels == _LATENT_CHANNELS + assert latent_format.latent_dimensions == 3 + assert fixed.shape == (1, _LATENT_CHANNELS, 1, 4, 5) + + +def test_seedvr2_model_requires_native_5d_latent(): + latent = torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5) + assert NaDiT._check_seedvr2_video_latent(latent, _LATENT_CHANNELS, "latent") is latent + + with pytest.raises(ValueError, match="5-D native latent"): + NaDiT._check_seedvr2_video_latent(torch.zeros(1, _LATENT_CHANNELS * 2, 4, 5), _LATENT_CHANNELS, "latent") + + +def test_seedvr2_encode_and_encode_tiled_preserve_native_latent_contract(monkeypatch): + monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None) + + encoded = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 2.0) + vae = _make_vae(_EncodeWrapper(encoded)) + pixels = torch.zeros(1, 5, 32, 40, 3) + + node_output = nodes_mod.VAEEncode().encode(vae, pixels)[0] + node_latent = node_output["samples"] + assert set(node_output) == {"samples"} + assert tuple(node_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5) + assert node_latent.dtype == torch.float32 + assert node_latent.stride()[-1] == 1 + assert torch.equal(node_latent, torch.full_like(node_latent, 2.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR)) + + tiled = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 3.0) + monkeypatch.setattr(seedvr_vae_mod, "tiled_vae", MagicMock(return_value=tiled)) + tiled_output = nodes_mod.VAEEncodeTiled().encode( + vae, + pixels, + tile_size=512, + overlap=64, + temporal_size=16, + temporal_overlap=4, + )[0] + tiled_latent = tiled_output["samples"] + assert set(tiled_output) == {"samples"} + assert tuple(tiled_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5) + assert tiled_latent.dtype == torch.float32 + assert torch.equal(tiled_latent, torch.full_like(tiled_latent, 3.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR)) + + +def test_vaedecode_tiled_spatial_applies_temporal_discarded(monkeypatch): + monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None) + vae = _make_vae(_DecodeWrapper()) + + nodes_mod.VAEDecodeTiled().decode( + vae, + {"samples": torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)}, + tile_size=512, + overlap=64, + temporal_size=16, + temporal_overlap=4, + ) + + # Spatial inputs flow through; temporal inputs are discarded as public tiling + # knobs, but SeedVR2's internal MemoryState causal slicing is left intact. + assert vae.first_stage_model.calls == [ + { + "shape": (1, _LATENT_CHANNELS, 2, 4, 5), + "seedvr2_tiling": { + "enable_tiling": True, + "tile_size": (512, 512), + "tile_overlap": (64, 64), + "temporal_size": None, + "temporal_overlap": None, + }, + } + ] diff --git a/tests-unit/comfy_test/test_seedvr2_vae_decode.py b/tests-unit/comfy_test/test_seedvr2_vae_decode.py new file mode 100644 index 000000000..c486b9195 --- /dev/null +++ b/tests-unit/comfy_test/test_seedvr2_vae_decode.py @@ -0,0 +1,94 @@ +from unittest.mock import patch + +import pytest +import torch +import torch.nn as nn + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +import comfy.ldm.seedvr.vae as vae_mod # noqa: E402 +from comfy_extras import nodes_seedvr # noqa: E402 + + +_LATENT_CHANNELS = vae_mod.SEEDVR2_LATENT_CHANNELS + + +def _make_wrapper() -> vae_mod.VideoAutoencoderKLWrapper: + wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__( + vae_mod.VideoAutoencoderKLWrapper + ) + nn.Module.__init__(wrapper) + return wrapper + + +def _fingerprint_decode_(self, z, return_dict=True): + b = int(z.shape[0]) + t = int(z.shape[2]) + h = int(z.shape[3]) + w = int(z.shape[4]) + out = torch.empty(b, 3, t, h * 8, w * 8) + for batch_idx in range(b): + out[batch_idx].fill_(float(batch_idx + 1)) + return out + + +def _decode_with_patches(wrapper, z): + with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _fingerprint_decode_): + return wrapper.decode(z) + + +def test_decode_b2_t3_multi_frame_batch_unchanged(): + wrapper = _make_wrapper() + + out = _decode_with_patches(wrapper, torch.zeros(2, _LATENT_CHANNELS * 3, 2, 2)) + + assert tuple(out.shape) == (2, 3, 3, 16, 16) + + +class _Wrapper(vae_mod.VideoAutoencoderKLWrapper): + def __init__(self): + nn.Module.__init__(self) + self.calls = [] + + def parameters(self): + return iter([torch.nn.Parameter(torch.zeros(()))]) + +def _decode_stub(self, latent): + self.calls.append(tuple(latent.shape)) + return torch.zeros(latent.shape[0], 3, latent.shape[2], latent.shape[3] * 8, latent.shape[4] * 8) + + +def test_seedvr2_wrapper_decode_accepts_5d_channel_first_latents_without_preprocessor_state(): + wrapper = _Wrapper() + + with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _decode_stub): + out = wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)) + + assert tuple(out.shape) == (1, 3, 2, 32, 40) + assert wrapper.calls == [(1, _LATENT_CHANNELS, 2, 4, 5)] + + +def test_seedvr2_wrapper_decode_rejects_wrong_rank_latents(): + wrapper = _Wrapper() + + with pytest.raises(RuntimeError, match=r"latent input must be 4-D collapsed .* or 5-D"): + wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 4)) + + +def _t_padded(t_in: int) -> int: + if t_in == 1: + return 1 + if t_in <= 4: + return 5 + if (t_in - 1) % 4 == 0: + return t_in + return t_in + (4 - ((t_in - 1) % 4)) + + +@pytest.mark.parametrize("t_in", [1, 5, 9]) +def test_t_padded_matches_cut_videos(t_in): + dummy = torch.zeros(1, t_in, 1, 1, 1) + assert nodes_seedvr.cut_videos(dummy).shape[1] == _t_padded(t_in) diff --git a/tests-unit/comfy_test/test_seedvr2_vae_tiled.py b/tests-unit/comfy_test/test_seedvr2_vae_tiled.py new file mode 100644 index 000000000..a2866b609 --- /dev/null +++ b/tests-unit/comfy_test/test_seedvr2_vae_tiled.py @@ -0,0 +1,382 @@ +from contextlib import ExitStack +from unittest.mock import MagicMock, patch + +import pytest +import torch +import torch.nn as nn + +from comfy.cli_args import args as cli_args + +if not torch.cuda.is_available(): + cli_args.cpu = True + +import comfy.ldm.seedvr.vae as vae_mod # noqa: E402 +import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402 +import comfy.sd as sd_mod # noqa: E402 +from comfy.ldm.seedvr.vae import MemoryState, tiled_vae # noqa: E402 + + +_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS + + +def test_runtime_decode_zero_temporal_size_preserves_model_slicing(): + class StubVAEModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.slicing_latent_min_size = 2 + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.device = torch.device("cpu") + self.use_slicing = True + self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32)) + self.decode_min_sizes = [] + self.memory_states = [] + + def decode_(self, t_chunk): + self.decode_min_sizes.append(self.slicing_latent_min_size) + return vae_mod.VideoAutoencoderKL.slicing_decode(self, t_chunk) + + def _decode(self, z, memory_state=MemoryState.DISABLED, memory_cache=None): + self.memory_states.append(memory_state) + b, c, d, h, w = z.shape + return torch.zeros((b, 3, d, h * 8, w * 8), dtype=z.dtype) + + vae = StubVAEModel() + z = torch.zeros((1, _LATENT_CHANNELS, 5, 8, 8), dtype=torch.float32) + + tiled_vae( + z, + vae, + tile_size=(64, 64), + tile_overlap=(0, 0), + temporal_size=0, + temporal_overlap=0, + encode=False, + ) + + assert vae.decode_min_sizes == [2] + assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE] + assert vae.slicing_latent_min_size == 2 + + +def test_zero_temporal_size_preserves_min_size_when_encode_raises(): + class RaisingVAEModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.slicing_sample_min_size = 4 + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.device = torch.device("cpu") + self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32)) + + def encode(self, t_chunk): + raise RuntimeError("simulated encode failure") + + vae = RaisingVAEModel() + x = torch.zeros((1, 3, 12, 64, 64), dtype=torch.float32) + + with pytest.raises(RuntimeError, match="simulated encode failure"): + tiled_vae( + x, + vae, + tile_size=(64, 64), + tile_overlap=(0, 0), + temporal_size=0, + temporal_overlap=0, + encode=True, + ) + + assert vae.slicing_sample_min_size == 4 + + +def test_tiled_vae_encode_uses_tensor_return_without_indexing(): + class TensorEncodeVAEModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.slicing_sample_min_size = 4 + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.device = torch.device("cpu") + self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32)) + self.calls = [] + + def encode(self, t_chunk): + self.calls.append(tuple(t_chunk.shape)) + b, _, _, h, w = t_chunk.shape + return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype) + + vae = TensorEncodeVAEModel() + x = torch.zeros((2, 3, 1, 64, 64), dtype=torch.float32) + + out = tiled_vae( + x, + vae, + tile_size=(64, 64), + tile_overlap=(0, 0), + temporal_size=0, + temporal_overlap=0, + encode=True, + ) + + assert vae.calls == [(2, 3, 1, 64, 64)] + assert tuple(out.shape) == (2, _LATENT_CHANNELS, 1, 8, 8) + + +def test_tiled_vae_preserves_input_dtype_on_single_tile(): + class FloatOutputVAEModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.slicing_sample_min_size = 4 + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.device = torch.device("cpu") + self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32)) + + def encode(self, t_chunk): + b, _, _, h, w = t_chunk.shape + return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=torch.float32) + + out = tiled_vae( + torch.zeros((1, 3, 1, 64, 64), dtype=torch.float16), + FloatOutputVAEModel(), + tile_size=(64, 64), + tile_overlap=(0, 0), + temporal_size=0, + temporal_overlap=0, + encode=True, + ) + + assert out.dtype == torch.float16 + + +class _SlicingDecodeVAE(nn.Module): + def __init__(self, slicing_latent_min_size): + super().__init__() + self.slicing_latent_min_size = slicing_latent_min_size + self.spatial_downsample_factor = 8 + self.temporal_downsample_factor = 4 + self.device = torch.device("cpu") + self.use_slicing = True + self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float32)) + self.decode_min_sizes = [] + self.memory_states = [] + + def decode_(self, z): + self.decode_min_sizes.append(self.slicing_latent_min_size) + return vae_mod.VideoAutoencoderKL.slicing_decode(self, z) + + def _decode(self, z, memory_state=MemoryState.DISABLED, memory_cache=None): + self.memory_states.append(memory_state) + x = z[:, :1].repeat( + 1, + 3, + 1, + self.spatial_downsample_factor, + self.spatial_downsample_factor, + ) + return x + + +def test_decode_tiled_vae_maps_temporal_args_to_latent_slicing_min_size(): + vae = _SlicingDecodeVAE(slicing_latent_min_size=2) + z = torch.arange( + _LATENT_CHANNELS * 5 * 8 * 8, + dtype=torch.float32, + ).reshape(1, _LATENT_CHANNELS, 5, 8, 8) + + tiled_vae( + z, + vae, + tile_size=(64, 64), + tile_overlap=(0, 0), + temporal_size=12, + temporal_overlap=4, + encode=False, + ) + + assert vae.decode_min_sizes == [2] + assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE] + assert vae.slicing_latent_min_size == 2 + + wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__( + vae_mod.VideoAutoencoderKLWrapper + ) + nn.Module.__init__(wrapper) + seedvr2_tiling = { + "enable_tiling": True, + "tile_size": (64, 64), + "tile_overlap": (0, 0), + "temporal_size": 8, + "temporal_overlap": 7, + } + + captured = {} + + def _fake_tiled_vae(latent, model, **kwargs): + captured.update(kwargs) + return torch.zeros(1, 3, 1, 16, 16) + + with patch.object(vae_mod, "tiled_vae", side_effect=_fake_tiled_vae): + wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 2, 2), seedvr2_tiling=seedvr2_tiling) + + assert captured["temporal_overlap"] == 7 + + +def _force_oom(*a, **k): + raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test") + + +def _make_vae(first_stage_model, latent_channels, latent_dim): + vae = sd_mod.VAE.__new__(sd_mod.VAE) + vae.first_stage_model = first_stage_model + vae.patcher = MagicMock() + vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024) + vae.device = vae.output_device = torch.device("cpu") + vae.vae_dtype = torch.float32 + vae.disable_offload = True + vae.extra_1d_channel = None + vae.upscale_ratio = vae.downscale_ratio = 8 + vae.upscale_index_formula = vae.downscale_index_formula = None + vae.output_channels = 3 + vae.latent_channels = latent_channels + vae.latent_dim = latent_dim + vae.vae_output_dtype = lambda: torch.float32 + vae.spacial_compression_decode = lambda: 8 + vae.handles_tiling = isinstance(first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper) + vae.format_encoded = None + vae.process_input = lambda x: x + vae.process_output = lambda x: x + vae.throw_exception_if_invalid = lambda: None + vae.memory_used_decode = lambda *a, **k: 1 + return vae + + +def _dispatch(vae, samples, seedvr2_call, generic_call, patch_wrapper_decode): + mm = sd_mod.model_management + with ExitStack() as stack: + stack.enter_context(patch.object(mm, "raise_non_oom", lambda e: None)) + stack.enter_context(patch.object(mm, "load_models_gpu", lambda *a, **k: None)) + stack.enter_context(patch.object(mm, "soft_empty_cache", lambda: None)) + stack.enter_context(patch.object(sd_mod.VAE, "_decode_tiled_owned", seedvr2_call)) + stack.enter_context(patch.object(sd_mod.VAE, "decode_tiled_", generic_call)) + if patch_wrapper_decode: + stack.enter_context(patch.object( + seedvr_vae_mod.VideoAutoencoderKLWrapper, "decode", + side_effect=_force_oom)) + vae.decode(samples) + + +def test_4d_seedvr2_latent_routes_to_owned_decode_tiled(): + wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__( + seedvr_vae_mod.VideoAutoencoderKLWrapper) + vae = _make_vae(wrapper, latent_channels=_LATENT_CHANNELS, latent_dim=3) + seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64)) + generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64)) + _dispatch(vae, torch.zeros(1, _LATENT_CHANNELS * 3, 8, 8), seedvr2_call, generic_call, True) + assert seedvr2_call.call_count == 1 + assert generic_call.call_count == 0 + + +def test_4d_non_seedvr2_latent_still_routes_to_generic_decode_tiled(): + first_stage = MagicMock() + first_stage.decode = MagicMock(side_effect=_force_oom) + vae = _make_vae(first_stage, latent_channels=4, latent_dim=2) + seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64)) + generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64)) + _dispatch(vae, torch.zeros(1, 4, 8, 8), seedvr2_call, generic_call, False) + assert generic_call.call_count == 1 + assert seedvr2_call.call_count == 0 + + +def _populate_common_vae_attrs_fallback(vae): + vae.patcher = MagicMock() + vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024) + vae.device = torch.device("cpu") + vae.output_device = torch.device("cpu") + vae.vae_dtype = torch.float32 + vae.disable_offload = True + vae.extra_1d_channel = None + vae.upscale_ratio = 8 + vae.upscale_index_formula = None + vae.output_channels = 3 + vae.latent_channels = _LATENT_CHANNELS + vae.latent_dim = 3 + vae.downscale_ratio = 8 + vae.downscale_index_formula = None + vae.not_video = False + vae.crop_input = False + vae.pad_channel_value = None + vae.handles_tiling = isinstance(vae.first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper) + vae.format_encoded = None + + vae.vae_output_dtype = lambda: torch.float32 + vae.spacial_compression_encode = lambda: 8 + vae.process_input = lambda x: x + vae.process_output = lambda x: x + vae.throw_exception_if_invalid = lambda: None + vae.memory_used_encode = lambda *a, **k: 1 + + +def _make_seedvr2_vae_fallback(): + vae = sd_mod.VAE.__new__(sd_mod.VAE) + wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__( + seedvr_vae_mod.VideoAutoencoderKLWrapper + ) + vae.first_stage_model = wrapper + _populate_common_vae_attrs_fallback(vae) + return vae + + +def _make_non_seedvr2_vae_fallback(): + vae = sd_mod.VAE.__new__(sd_mod.VAE) + vae.first_stage_model = MagicMock() + _populate_common_vae_attrs_fallback(vae) + return vae + + +def _force_regular_encode_oom(*args, **kwargs): + raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test") + + +def test_seedvr2_3d_routes_to_owned_encode_tiled_on_oom(): + vae = _make_seedvr2_vae_fallback() + pixel_samples = torch.zeros((1, 8, 64, 64, 3)) + + seedvr2_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8)) + generic_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8)) + + with patch.object(sd_mod.model_management, "raise_non_oom", + lambda e: None), \ + patch.object(sd_mod.model_management, "load_models_gpu", + lambda *a, **k: None), \ + patch.object(sd_mod.model_management, "soft_empty_cache", + lambda: None), \ + patch.object(seedvr_vae_mod.VideoAutoencoderKLWrapper, "encode", + side_effect=_force_regular_encode_oom), \ + patch.object(sd_mod.VAE, "_encode_tiled_owned", seedvr2_call), \ + patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call): + vae.encode(pixel_samples) + + assert seedvr2_call.call_count == 1, ( + f"Expected _encode_tiled_owned to be called once for a SeedVR2 3D " + f"input under OOM fallback; got {seedvr2_call.call_count} calls." + ) + assert generic_call.call_count == 0, ( + f"encode_tiled_3d must NOT be called for a SeedVR2 input; got " + f"{generic_call.call_count} calls." + ) + + +def test_non_seedvr2_encode_tiled_3d_default_overlap_is_concrete(): + vae = _make_non_seedvr2_vae_fallback() + vae.downscale_ratio = (lambda a: max(1, a // 4), 8, 8) + vae.upscale_ratio = (lambda a: a * 4, 8, 8) + generic_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8)) + pixel_samples = torch.zeros((1, 8, 64, 64, 3)) + + with patch.object(sd_mod.model_management, "load_models_gpu", + lambda *a, **k: None), \ + patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call): + vae.encode_tiled(pixel_samples) + + assert generic_call.call_args.kwargs["overlap"] == (1, 64, 64) From 89ecc5cf8c1992230ddab3ee67ef836b7c884442 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 10 Jul 2026 11:58:22 +0300 Subject: [PATCH 084/211] [Partner Nodes] feat(Seedream): add widget to disable thinking (#14853) Signed-off-by: bigcat88 Co-authored-by: Daxiong (Lin) --- comfy_api_nodes/apis/bytedance.py | 5 +++++ comfy_api_nodes/nodes_bytedance.py | 21 +++++++++++++++++++++ 2 files changed, 26 insertions(+) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 76573304b..515e124ca 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -17,6 +17,10 @@ class Seedream4Options(BaseModel): max_images: int = Field(15) +class Seedream5OptimizePromptOptions(BaseModel): + thinking: Literal["auto", "enabled", "disabled"] = Field(...) + + class Seedream4TaskCreationRequest(BaseModel): model: str = Field(...) prompt: str = Field(...) @@ -28,6 +32,7 @@ class Seedream4TaskCreationRequest(BaseModel): sequential_image_generation_options: Seedream4Options | None = Field(Seedream4Options(max_images=15)) watermark: bool = Field(False) output_format: str | None = None + optimize_prompt_options: Seedream5OptimizePromptOptions | None = None class ImageTaskCreationResponse(BaseModel): diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 043bc9526..a84399ad3 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -34,6 +34,7 @@ from comfy_api_nodes.apis.bytedance import ( SeedanceVirtualLibraryCreateAssetRequest, Seedream4Options, Seedream4TaskCreationRequest, + Seedream5OptimizePromptOptions, TaskAudioContent, TaskAudioContentUrl, TaskCreationResponse, @@ -875,6 +876,17 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): tooltip='Whether to add an "AI generated" watermark to the image.', advanced=True, ), + IO.Boolean.Input( + "thinking", + default=True, + tooltip=( + "Enable the model's prompt-optimization reasoning ('thinking') for better adherence. " + "Can substantially increase generation time — notably on Seedream 5.0 Pro. " + "Can only be disabled for text-to-image (not when reference images are provided)." + ), + optional=True, + advanced=True, + ), ], outputs=[ IO.Image.Output(), @@ -920,6 +932,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): model: dict, seed: int = 0, watermark: bool = False, + thinking: bool = True, ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1) model_id = SEEDREAM_MODELS[model["model"]] @@ -979,6 +992,10 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): raise ValueError( "The maximum number of generated images plus the number of reference images cannot exceed 15." ) + if not thinking and n_input_images > 0: + raise ValueError( + "'thinking' can only be disabled for text-to-image; enable it when using reference images." + ) reference_images_urls: list[str] = [] if image_tensors: @@ -992,6 +1009,9 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): wait_label="Uploading reference images", ) + optimize_prompt_options = None + if n_input_images == 0: + optimize_prompt_options = Seedream5OptimizePromptOptions(thinking="enabled" if thinking else "disabled") response = await sync_op( cls, ApiEndpoint(path=BYTEPLUS_IMAGE_ENDPOINT, method="POST"), @@ -1005,6 +1025,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): sequential_image_generation=None if is_pro else sequential_image_generation, sequential_image_generation_options=None if is_pro else Seedream4Options(max_images=max_images), watermark=watermark, + optimize_prompt_options=optimize_prompt_options, ), ) if len(response.data) == 1: From 206b9245dcc4e497bf18808adf7585b8ae7595ec Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 10 Jul 2026 12:33:32 +0300 Subject: [PATCH 085/211] [Partner Nodes] fix(Tencent): restore Tencent3DPartNode FBX output via staged generation (#14867) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/hunyuan3d.py | 1 + comfy_api_nodes/nodes_hunyuan3d.py | 1 + 2 files changed, 2 insertions(+) diff --git a/comfy_api_nodes/apis/hunyuan3d.py b/comfy_api_nodes/apis/hunyuan3d.py index dad9bc2fa..91f630e81 100644 --- a/comfy_api_nodes/apis/hunyuan3d.py +++ b/comfy_api_nodes/apis/hunyuan3d.py @@ -77,6 +77,7 @@ class To3DUVTaskRequest(BaseModel): class To3DPartTaskRequest(BaseModel): File: TaskFile3DInput = Field(...) + EnableStagedGeneration: bool | None = Field(None) class TextureEditImageInfo(BaseModel): diff --git a/comfy_api_nodes/nodes_hunyuan3d.py b/comfy_api_nodes/nodes_hunyuan3d.py index fcd27b7fb..a9942476c 100644 --- a/comfy_api_nodes/nodes_hunyuan3d.py +++ b/comfy_api_nodes/nodes_hunyuan3d.py @@ -642,6 +642,7 @@ class Tencent3DPartNode(IO.ComfyNode): response_model=To3DProTaskCreateResponse, data=To3DPartTaskRequest( File=TaskFile3DInput(Type=file_format.upper(), Url=model_url), + EnableStagedGeneration=True, ), is_rate_limited=_is_tencent_rate_limited, ) From 1377a2f72925ed7a5518c1900ff71c6740217b0d Mon Sep 17 00:00:00 2001 From: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com> Date: Fri, 10 Jul 2026 18:31:20 +0800 Subject: [PATCH 086/211] Only auto-enable the ROCm comfy-kitchen Triton backend on matrix-core GPUs (#14869) #14862 auto-enables the comfy-kitchen Triton backend whenever torch.version.hip is set and Triton >= 3.7. The INT8 matmul kernels compile tl.dot to matrix-core instructions (WMMA on RDNA3+/gfx11xx-gfx12xx, MFMA on CDNA/gfx9xx); RDNA1/RDNA2 (gfx10xx) have neither, so the auto-enabled INT8 path hangs the GPU there (reported on RDNA2 + triton-windows 3.7.1: native and custom-node INT8 freeze until reset). Gate the automatic ROCm default on GPU architecture as well as Triton version so RDNA1/RDNA2 stay on the working eager fallback. Add --disable-triton-backend as an explicit override; --enable-triton-backend still force-enables on any arch. --- comfy/cli_args.py | 1 + comfy/quant_ops.py | 24 ++++++++++++++++++++++-- 2 files changed, 23 insertions(+), 2 deletions(-) diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 0d7df5e13..e2e0d97ec 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -92,6 +92,7 @@ parser.add_argument("--directml", type=int, nargs="?", metavar="DIRECTML_DEVICE" parser.add_argument("--oneapi-device-selector", type=str, default=None, metavar="SELECTOR_STRING", help="Sets the oneAPI device(s) this instance will use.") parser.add_argument("--supports-fp8-compute", action="store_true", help="ComfyUI will act like if the device supports fp8 compute.") parser.add_argument("--enable-triton-backend", action="store_true", help="ComfyUI will enable the use of Triton backend in comfy-kitchen. Is disabled at launch by default.") +parser.add_argument("--disable-triton-backend", action="store_true", help="Force-disable the comfy-kitchen Triton backend, overriding the automatic ROCm/AMD default and --enable-triton-backend.") class LatentPreviewMethod(enum.Enum): NoPreviews = "none" diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 91b3e4fe9..b1aabdc93 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -3,6 +3,22 @@ import logging from comfy.cli_args import args + +def _rocm_kitchen_arch_supported(): + """comfy-kitchen's INT8 Triton kernels compile tl.dot to matrix-core instructions. + RDNA3/3.5/4 (gfx11xx/gfx12xx) have WMMA and CDNA (gfx9xx) has MFMA; RDNA1/RDNA2 + (gfx10xx) have neither, so the INT8 path hangs the GPU there. Gates the automatic + ROCm default so those cards stay on the eager fallback (an explicit + --enable-triton-backend still forces it on any arch).""" + try: + arch = torch.cuda.get_device_properties(torch.cuda.current_device()).gcnArchName.split(":")[0] + except Exception: + return False + if arch.startswith(("gfx11", "gfx12")): + return True + return arch in ("gfx908", "gfx90a", "gfx940", "gfx941", "gfx942", "gfx950") + + try: import comfy_kitchen as ck from comfy_kitchen.tensor import ( @@ -26,9 +42,13 @@ try: logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.") # On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated - # comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7: + # comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7 AND a + # matrix-core GPU (RDNA3+ WMMA gfx11xx/gfx12xx, CDNA MFMA gfx9xx). RDNA1/RDNA2 + # (gfx10xx) have no WMMA -> the INT8 tl.dot path hangs the GPU, so they stay eager. # older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path. - if args.enable_triton_backend or torch.version.hip is not None: + if args.disable_triton_backend: + ck.registry.disable("triton") + elif args.enable_triton_backend or (torch.version.hip is not None and _rocm_kitchen_arch_supported()): try: import triton triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2]) From 94fa08223e611f1e95693fab90f8c00adf353ccc Mon Sep 17 00:00:00 2001 From: "Yousef R. Gamaleldin" <81116377+yousef-rafat@users.noreply.github.com> Date: Fri, 10 Jul 2026 22:54:56 +0300 Subject: [PATCH 087/211] Save Text Node (CORE-176) (#14102) --- comfy_execution/jobs.py | 15 ++++++-- comfy_extras/nodes_text.py | 71 ++++++++++++++++++++++++++++++++++++++ nodes.py | 1 + 3 files changed, 84 insertions(+), 3 deletions(-) create mode 100644 comfy_extras/nodes_text.py diff --git a/comfy_execution/jobs.py b/comfy_execution/jobs.py index fa3ab0faf..f0ad59f86 100644 --- a/comfy_execution/jobs.py +++ b/comfy_execution/jobs.py @@ -56,6 +56,9 @@ PREVIEWABLE_MEDIA_TYPES = frozenset({'images', 'video', 'audio', '3d', 'text'}) # 3D file extensions for preview fallback (no dedicated media_type exists) THREE_D_EXTENSIONS = frozenset({'.obj', '.fbx', '.gltf', '.glb', '.usdz'}) +# Text file extensions for preview fallback (the formats SaveText can produce) +TEXT_EXTENSIONS = frozenset({'.txt', '.md', '.json'}) + def has_3d_extension(filename: str) -> bool: lower = filename.lower() @@ -143,9 +146,10 @@ def is_previewable(media_type: str, item: dict) -> bool: Maintains backwards compatibility with existing logic. Priority: - 1. media_type is 'images', 'video', 'audio', or '3d' + 1. media_type is 'images', 'video', 'audio', '3d', or 'text' 2. format field starts with 'video/' or 'audio/' 3. filename has a 3D extension (.obj, .fbx, .gltf, .glb, .usdz) + 4. filename has a text extension (.txt, .md, .json, ...) """ if media_type in PREVIEWABLE_MEDIA_TYPES: return True @@ -156,10 +160,12 @@ def is_previewable(media_type: str, item: dict) -> bool: if fmt and (fmt.startswith('video/') or fmt.startswith('audio/')): return True - # Check for 3D files by extension + # Check for 3D and text files by extension filename = item.get('filename', '').lower() if any(filename.endswith(ext) for ext in THREE_D_EXTENSIONS): return True + if any(filename.endswith(ext) for ext in TEXT_EXTENSIONS): + return True return False @@ -255,6 +261,10 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: Preview priority (matching frontend): 1. type="output" with previewable media 2. Any previewable media + + Text content entries (strings under 'text') are preview-only metadata, + matching the frontend's METADATA_KEYS: they can serve as the fallback + preview but are not counted as outputs. """ count = 0 preview_output = None @@ -275,7 +285,6 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: if normalized is None: # Not a 3D file string — check for text preview if media_type == 'text': - count += 1 if preview_output is None: if isinstance(item, tuple): text_value = item[0] if item else '' diff --git a/comfy_extras/nodes_text.py b/comfy_extras/nodes_text.py new file mode 100644 index 000000000..a485f5df8 --- /dev/null +++ b/comfy_extras/nodes_text.py @@ -0,0 +1,71 @@ +import os +import json +from typing_extensions import override +from comfy_api.latest import io, ComfyExtension, ui +import folder_paths + + +class SaveTextNode(io.ComfyNode): + """Save text content to .txt, .md, or .json.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SaveText", + search_aliases=["save text", "write text", "export text"], + display_name="Save Text", + category="text", + description="Save text content to a file in the output directory.", + inputs=[ + io.String.Input("text", force_input=True), + io.String.Input("filename_prefix", default="ComfyUI"), + io.Combo.Input("format", options=["txt", "md", "json"], default="txt"), + ], + outputs=[io.String.Output(display_name="text")], + is_output_node=True, + ) + + @classmethod + def execute(cls, text, filename_prefix, format): + full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( + filename_prefix, + folder_paths.get_output_directory(), + 1, + 1, + ) + + file = f"{filename}_{counter:05}.{format}" + filepath = os.path.join(full_output_folder, file) + + if format == "json": + # tries to pretty print otherwise saves normally + try: + data = json.loads(text) + with open(filepath, "w", encoding="utf-8") as f: + json.dump(data, f, indent=2, ensure_ascii=False) + except json.JSONDecodeError: + with open(filepath, "w", encoding="utf-8") as f: + f.write(text) + else: + with open(filepath, "w", encoding="utf-8") as f: + f.write(text) + + return io.NodeOutput( + text, + ui={ + "text": (text,), + "files": [ + ui.SavedResult(file, subfolder, io.FolderType.output) + ] + } + ) + +class TextExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + SaveTextNode + ] + +async def comfy_entrypoint() -> TextExtension: + return TextExtension() diff --git a/nodes.py b/nodes.py index 474e188fe..31602e582 100644 --- a/nodes.py +++ b/nodes.py @@ -2504,6 +2504,7 @@ async def init_builtin_extra_nodes(): "nodes_triposplat.py", "nodes_depth_anything_3.py", "nodes_seed.py", + "nodes_text.py", ] import_failed = [] From 8310b0e0dbc3361c69c70985ee63b73f1970449a Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Fri, 10 Jul 2026 15:58:03 -0400 Subject: [PATCH 088/211] feat: add bboxes input to Create Bounding Boxes node (#14724) --- comfy_extras/nodes_bounding_boxes.py | 132 ++++++++++++++++++++++++++- 1 file changed, 129 insertions(+), 3 deletions(-) diff --git a/comfy_extras/nodes_bounding_boxes.py b/comfy_extras/nodes_bounding_boxes.py index 77cbf8649..de3709b91 100644 --- a/comfy_extras/nodes_bounding_boxes.py +++ b/comfy_extras/nodes_bounding_boxes.py @@ -1,3 +1,5 @@ +import json + import numpy as np import torch from PIL import Image, ImageDraw, ImageEnhance, ImageFont @@ -166,6 +168,111 @@ def boxes_to_regions(boxes, width: int, height: int) -> list: return regions +def normalize_incoming_boxes(bboxes) -> list: + if isinstance(bboxes, dict): + frame = [bboxes] + elif not isinstance(bboxes, list) or not bboxes: + frame = [] + elif isinstance(bboxes[0], dict): + frame = bboxes + else: + frame = bboxes[0] if isinstance(bboxes[0], list) else [] + boxes = [] + for box in frame: + if not isinstance(box, dict): + continue + norm = { + "x": box.get("x", 0), + "y": box.get("y", 0), + "width": box.get("width", 0), + "height": box.get("height", 0), + } + meta = box.get("metadata") + if isinstance(meta, dict): + norm["metadata"] = meta + boxes.append(norm) + return boxes + + +def _looks_like_element(box: dict) -> bool: + bbox = box.get("bbox") + return isinstance(bbox, (list, tuple)) and len(bbox) == 4 + + +def _looks_like_bbox(box: dict) -> bool: + return all(key in box for key in ("x", "y", "width", "height")) + + +def elements_to_boxes(elements: list, width: int, height: int) -> list: + boxes = [] + for element in elements: + if not isinstance(element, dict): + continue + bbox = element.get("bbox") + if not (isinstance(bbox, (list, tuple)) and len(bbox) == 4): + raise ValueError("bboxes element is missing a valid 'bbox' [ymin, xmin, ymax, xmax]") + try: + ymin, xmin, ymax, xmax = (float(v) / 1000.0 for v in bbox) + except (TypeError, ValueError): + raise ValueError("bboxes element 'bbox' must contain four numbers") + etype = "text" if element.get("type") == "text" else "obj" + boxes.append({ + "x": round(min(xmin, xmax) * width), + "y": round(min(ymin, ymax) * height), + "width": round(abs(xmax - xmin) * width), + "height": round(abs(ymax - ymin) * height), + "metadata": { + "type": etype, + "text": element.get("text", "") if etype == "text" else "", + "desc": element.get("desc", ""), + "palette": element.get("color_palette", []) or [], + }, + }) + return boxes + + +def boxes_from_input(data, width: int, height: int) -> list: + if data is None: + return [] + if isinstance(data, str): + text = data.strip() + if not text: + return [] + try: + data = json.loads(text) + except (ValueError, TypeError) as exc: + raise ValueError(f"bboxes string input is not valid JSON: {exc}") from exc + if isinstance(data, dict): + if _looks_like_element(data): + return elements_to_boxes([data], width, height) + if _looks_like_bbox(data): + return normalize_incoming_boxes(data) + raise ValueError( + "bboxes dict must be a bounding box (x, y, width, height) or an element (with a 'bbox')" + ) + if not isinstance(data, list): + raise ValueError( + "bboxes input must be bounding boxes, elements, or a JSON string, " + f"got {type(data).__name__}" + ) + if not data: + return [] + first = data[0] + if isinstance(first, list): + return normalize_incoming_boxes(data) + if isinstance(first, dict): + if _looks_like_element(first): + return elements_to_boxes(data, width, height) + if _looks_like_bbox(first): + return normalize_incoming_boxes(data) + raise ValueError( + "bboxes items must be bounding boxes (x, y, width, height) or elements (with a 'bbox')" + ) + raise ValueError( + f"bboxes list must contain bounding boxes or elements, got {type(first).__name__}" + ) + + def _norm_bbox(region: dict) -> list[int]: def grid(value: float) -> int: return max(0, min(1000, round(value * 1000))) @@ -217,29 +324,48 @@ class CreateBoundingBoxes(io.ComfyNode): optional=True, tooltip="Optional image used as background in the canvas and preview.", ), + io.MultiType.Input( + "bboxes", + [io.BoundingBox, io.Array, io.String], + optional=True, + tooltip="Bounding boxes, elements, or a JSON string to initialize the canvas. A new upstream value initializes the canvas; edits made on the canvas take priority and are kept until the upstream value changes again.", + ), io.Int.Input("width", default=1024, min=64, max=16384, step=16, tooltip="Width of the canvas and the pixel grid for the bounding boxes."), io.Int.Input("height", default=1024, min=64, max=16384, step=16, tooltip="Height of the canvas and the pixel grid for the bounding boxes."), editor_state, + io.BoundingBoxes.Input( + "last_incoming", + optional=True, + tooltip="Internal state managed by the canvas: the upstream bboxes value that last initialized it. Leave empty to re-initialize the canvas from the bboxes input on the next run.", + ), ], outputs=[ io.Image.Output(display_name="preview"), io.BoundingBox.Output(display_name="bboxes"), io.Array.Output(display_name="elements"), ], + is_output_node=True, is_experimental=True, ) @classmethod - def execute(cls, width, height, editor_state=None, background=None) -> io.NodeOutput: - regions = boxes_to_regions(editor_state, width, height) + def execute(cls, width, height, editor_state=None, last_incoming=None, background=None, bboxes=None) -> io.NodeOutput: + incoming = boxes_from_input(bboxes, width, height) + applied = last_incoming if isinstance(last_incoming, list) else [] + upstream_changed = bool(incoming) and incoming != applied + source = incoming if upstream_changed else (editor_state or []) + regions = boxes_to_regions(source, width, height) preview = render_preview(regions, width, height, _bg_from_image(background)) + ui = {"dims": [width, height]} + if incoming: + ui["input_bboxes"] = incoming return io.NodeOutput( preview, fractions_to_bbox_frame(regions, width, height), build_elements(regions), - ui={"dims": [width, height]}, + ui=ui, ) From 328144ce24c6ce4b979dd850027f95d4cfa8449a Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Fri, 10 Jul 2026 16:03:34 -0400 Subject: [PATCH 089/211] CORE-329 feat: add Save 3D (Advanced) node family (#14701) --- comfy_extras/nodes_save_3d.py | 158 +++++++++++++++++++++++++++++++++- 1 file changed, 156 insertions(+), 2 deletions(-) diff --git a/comfy_extras/nodes_save_3d.py b/comfy_extras/nodes_save_3d.py index 1b6592bb2..7c524caa1 100644 --- a/comfy_extras/nodes_save_3d.py +++ b/comfy_extras/nodes_save_3d.py @@ -13,7 +13,7 @@ from typing_extensions import override import folder_paths from comfy.cli_args import args -from comfy_api.latest import ComfyExtension, IO, Types +from comfy_api.latest import ComfyExtension, IO, Types, UI def pack_variable_mesh_batch(vertices, faces, colors=None, uvs=None, texture=None, unlit=False): @@ -406,10 +406,164 @@ class SaveGLB(IO.ComfyNode): return IO.NodeOutput(ui={"3d": results}) +def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str: + full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( + filename_prefix, folder_paths.get_output_directory() + ) + ext = model_3d.format or "glb" + saved_filename = f"{filename}_{counter:05}.{ext}" + model_3d.save_to(os.path.join(full_output_folder, saved_filename)) + return f"{subfolder}/{saved_filename}" if subfolder else saved_filename + + +def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput: + model_file = _save_file3d_to_output(model_3d, filename_prefix) + camera_info_input = kwargs.get("camera_info", None) + camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + model_3d_info_input = kwargs.get("model_3d_info", None) + model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) + return IO.NodeOutput( + model_3d, + model_3d_info, + camera_info, + width, + height, + ui=UI.PreviewUI3DAdvanced(model_file, camera_info, model_3d_info), + ) + + +class Save3DAdvanced(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="Save3DAdvanced", + display_name="Save 3D (Advanced)", + search_aliases=["save 3d", "export 3d model", "save mesh advanced"], + category="3d", + is_experimental=True, + is_output_node=True, + inputs=[ + IO.MultiType.Input( + "model_3d", + types=[ + IO.File3DGLB, + IO.File3DGLTF, + IO.File3DFBX, + IO.File3DOBJ, + IO.File3DSTL, + IO.File3DUSDZ, + IO.File3DAny, + ], + tooltip="3D model file from an upstream 3D node.", + ), + IO.String.Input("filename_prefix", default="3d/ComfyUI"), + IO.Load3D.Input("viewport_state"), + IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True), + IO.Load3DCamera.Input("camera_info", optional=True, advanced=True), + IO.Int.Input("width", default=1024, min=1, max=4096, step=1), + IO.Int.Input("height", default=1024, min=1, max=4096, step=1), + ], + outputs=[ + IO.File3DAny.Output(display_name="model_3d"), + IO.Load3DModelInfo.Output(display_name="model_3d_info"), + IO.Load3DCamera.Output(display_name="camera_info"), + IO.Int.Output(display_name="width"), + IO.Int.Output(display_name="height"), + ], + ) + + @classmethod + def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput: + return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) + + +class SaveGaussianSplat(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="SaveGaussianSplat", + display_name="Save Splat", + search_aliases=["save splat", "save gaussian splat", "export gaussian", "export splat"], + category="3d", + is_experimental=True, + is_output_node=True, + inputs=[ + IO.MultiType.Input( + "model_3d", + types=[ + IO.File3DSplatAny, + IO.File3DPLY, + IO.File3DSPLAT, + IO.File3DSPZ, + IO.File3DKSPLAT, + ], + tooltip="A gaussian splat 3D file.", + ), + IO.String.Input("filename_prefix", default="3d/ComfyUI"), + IO.Load3D.Input("viewport_state"), + IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True), + IO.Load3DCamera.Input("camera_info", optional=True, advanced=True), + IO.Int.Input("width", default=1024, min=1, max=4096, step=1), + IO.Int.Input("height", default=1024, min=1, max=4096, step=1), + ], + outputs=[ + IO.File3DSplatAny.Output(display_name="model_3d"), + IO.Load3DModelInfo.Output(display_name="model_3d_info"), + IO.Load3DCamera.Output(display_name="camera_info"), + IO.Int.Output(display_name="width"), + IO.Int.Output(display_name="height"), + ], + ) + + @classmethod + def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput: + return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) + + +class SavePointCloud(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="SavePointCloud", + display_name="Save Point Cloud", + search_aliases=["save point cloud", "save pointcloud", "export point cloud"], + category="3d", + is_experimental=True, + is_output_node=True, + inputs=[ + IO.MultiType.Input( + "model_3d", + types=[ + IO.File3DPointCloudAny, + IO.File3DPLY, + ], + tooltip="Point cloud file (.ply)", + ), + IO.String.Input("filename_prefix", default="3d/ComfyUI"), + IO.Load3D.Input("viewport_state"), + IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True), + IO.Load3DCamera.Input("camera_info", optional=True, advanced=True), + IO.Int.Input("width", default=1024, min=1, max=4096, step=1), + IO.Int.Input("height", default=1024, min=1, max=4096, step=1), + ], + outputs=[ + IO.File3DPointCloudAny.Output(display_name="model_3d"), + IO.Load3DModelInfo.Output(display_name="model_3d_info"), + IO.Load3DCamera.Output(display_name="camera_info"), + IO.Int.Output(display_name="width"), + IO.Int.Output(display_name="height"), + ], + ) + + @classmethod + def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput: + return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) + + class Save3DExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: - return [SaveGLB] + return [SaveGLB, Save3DAdvanced, SaveGaussianSplat, SavePointCloud] async def comfy_entrypoint() -> Save3DExtension: From 5976ee37cd0ff1a5d28d14228daf8e5710390836 Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Sat, 11 Jul 2026 07:31:39 +0800 Subject: [PATCH 090/211] Bringing back the text node (#14870) --- comfy_extras/nodes_primitive.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/comfy_extras/nodes_primitive.py b/comfy_extras/nodes_primitive.py index 7f90daf14..35761863f 100644 --- a/comfy_extras/nodes_primitive.py +++ b/comfy_extras/nodes_primitive.py @@ -10,11 +10,10 @@ class String(io.ComfyNode): return io.Schema( node_id="PrimitiveString", search_aliases=["text", "string", "text box", "prompt"], - display_name="Text String (DEPRECATED)", + display_name="Text", category="utilities/primitive", inputs=[io.String.Input("value")], - outputs=[io.String.Output()], - is_deprecated=True + outputs=[io.String.Output()] ) @classmethod @@ -28,7 +27,7 @@ class StringMultiline(io.ComfyNode): return io.Schema( node_id="PrimitiveStringMultiline", search_aliases=["text", "string", "text multiline", "string multiline", "text box", "prompt"], - display_name="Input Text", + display_name="Text (Multiline)", category="utilities/primitive", essentials_category="Basics", inputs=[io.String.Input("value", multiline=True)], From 1f51e146a884462a28b2071c2268ce7c72550ce9 Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Sat, 11 Jul 2026 07:32:53 +0800 Subject: [PATCH 091/211] chore: Update preview nodes (#14871) --- comfy_extras/nodes_audio.py | 1 + comfy_extras/nodes_load_3d.py | 4 ++++ comfy_extras/nodes_mask.py | 5 +++-- comfy_extras/nodes_preview_any.py | 1 + nodes.py | 1 + 5 files changed, 10 insertions(+), 2 deletions(-) diff --git a/comfy_extras/nodes_audio.py b/comfy_extras/nodes_audio.py index 6adcc95fa..4ac5ced53 100644 --- a/comfy_extras/nodes_audio.py +++ b/comfy_extras/nodes_audio.py @@ -298,6 +298,7 @@ class PreviewAudio(IO.ComfyNode): search_aliases=["play audio"], display_name="Preview Audio", category="audio", + description="Preview the audio without saving it to the ComfyUI output directory.", inputs=[ IO.Audio.Input("audio"), ], diff --git a/comfy_extras/nodes_load_3d.py b/comfy_extras/nodes_load_3d.py index 6ef9a1ca3..a9df557c2 100644 --- a/comfy_extras/nodes_load_3d.py +++ b/comfy_extras/nodes_load_3d.py @@ -92,6 +92,7 @@ class Preview3D(IO.ComfyNode): search_aliases=["view mesh", "3d viewer"], display_name="Preview 3D & Animation", category="3d", + description="Preview a 3D model file without saving it to the ComfyUI output directory.", is_experimental=True, is_output_node=True, inputs=[ @@ -136,6 +137,7 @@ class Preview3DAdvanced(IO.ComfyNode): display_name="Preview 3D (Advanced)", search_aliases=["preview 3d", "3d viewer", "view mesh", "frame 3d", "3d camera output"], category="3d", + description="Preview a 3D model file without saving it to the ComfyUI output directory.", is_experimental=True, is_output_node=True, inputs=[ @@ -193,6 +195,7 @@ class PreviewGaussianSplat(IO.ComfyNode): node_id="PreviewGaussianSplat", display_name="Preview Splat", category="3d", + description="Preview a gaussian splat 3D file without saving it to the ComfyUI output directory.", is_experimental=True, is_output_node=True, search_aliases=[ @@ -261,6 +264,7 @@ class PreviewPointCloud(IO.ComfyNode): node_id="PreviewPointCloud", display_name="Preview Point Cloud", category="3d", + description="Preview a point cloud 3D file without saving it to the ComfyUI output directory.", is_experimental=True, is_output_node=True, search_aliases=[ diff --git a/comfy_extras/nodes_mask.py b/comfy_extras/nodes_mask.py index 76af338de..3fae7221f 100644 --- a/comfy_extras/nodes_mask.py +++ b/comfy_extras/nodes_mask.py @@ -419,17 +419,18 @@ class MaskPreview(IO.ComfyNode): search_aliases=["show mask", "view mask", "inspect mask", "debug mask"], display_name="Preview Mask", category="image/mask", - description="Saves the input images to your ComfyUI output directory.", + description="Preview the masks without saving them to the ComfyUI output directory.", inputs=[ IO.Mask.Input("mask"), ], hidden=[IO.Hidden.prompt, IO.Hidden.extra_pnginfo], is_output_node=True, + outputs=[IO.Mask.Output(display_name="mask")] ) @classmethod def execute(cls, mask, filename_prefix="ComfyUI") -> IO.NodeOutput: - return IO.NodeOutput(ui=UI.PreviewMask(mask)) + return IO.NodeOutput(mask, ui=UI.PreviewMask(mask)) class MaskExtension(ComfyExtension): diff --git a/comfy_extras/nodes_preview_any.py b/comfy_extras/nodes_preview_any.py index 1070a69d0..d985f3287 100644 --- a/comfy_extras/nodes_preview_any.py +++ b/comfy_extras/nodes_preview_any.py @@ -18,6 +18,7 @@ class PreviewAny(): CATEGORY = "utilities" SEARCH_ALIASES = ["show output", "inspect", "debug", "print value", "show text"] + DESCRIPTION = "Preview any input value as text." def main(self, source=None): torch.set_printoptions(edgeitems=6) diff --git a/nodes.py b/nodes.py index 31602e582..883258bd1 100644 --- a/nodes.py +++ b/nodes.py @@ -1709,6 +1709,7 @@ class PreviewImage(SaveImage): self.compress_level = 1 SEARCH_ALIASES = ["preview", "preview image", "show image", "view image", "display image", "image viewer"] + DESCRIPTION = "Preview the images without saving them to the ComfyUI output directory." @classmethod def INPUT_TYPES(s): From 92ddf07ba14711cb579ab090846e0d51289c0619 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 10 Jul 2026 16:54:28 -0700 Subject: [PATCH 092/211] Try to fix some issues with the seedvr VAE. (#14877) --- comfy/ldm/seedvr/vae.py | 20 ++++++------- tests-unit/comfy_test/test_seedvr2_dtype.py | 29 +++++++++++++++++++ .../comfy_test/test_seedvr2_vae_tiled.py | 25 ++++++++++++++++ 3 files changed, 63 insertions(+), 11 deletions(-) diff --git a/comfy/ldm/seedvr/vae.py b/comfy/ldm/seedvr/vae.py index c9f430184..7a8070b65 100644 --- a/comfy/ldm/seedvr/vae.py +++ b/comfy/ldm/seedvr/vae.py @@ -30,7 +30,7 @@ from enum import Enum import logging import comfy.model_management import comfy.ops -ops = comfy.ops.disable_weight_init +ops = comfy.ops.manual_cast def _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap, temporal_scale=1): @@ -103,11 +103,10 @@ def tiled_vae( storage_device = vae_model.device result = None count = None - def run_temporal_chunks(spatial_tile, model=vae_model, device=storage_device): - device = torch.device(device) - t_chunk = spatial_tile.to(device=device, dtype=next(model.parameters()).dtype, non_blocking=True).contiguous() + def run_temporal_chunks(spatial_tile, model=vae_model): + t_chunk = spatial_tile.contiguous() old_device = getattr(model, "device", None) - model.device = device + model.device = t_chunk.device old_slicing_min_size = getattr(model, slicing_attr, None) if old_slicing_min_size is not None and slicing_min_size is not None: if slicing_min_size <= 0: @@ -397,7 +396,7 @@ class Attention(nn.Module): def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: input_dtype = x.dtype - if isinstance(norm_layer, (ops.LayerNorm, ops.RMSNorm)): + if isinstance(norm_layer, (nn.LayerNorm, nn.RMSNorm)): if x.ndim == 4: x = x.permute(0, 2, 3, 1) x = norm_layer(x) @@ -408,14 +407,14 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: x = norm_layer(x) x = x.permute(0, 4, 1, 2, 3) return x.to(input_dtype) - if isinstance(norm_layer, (ops.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): + if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): if x.ndim <= 4: return norm_layer(x).to(input_dtype) if x.ndim == 5: b, c, t, h, w = x.shape x = x.transpose(1, 2).reshape(b * t, c, h, w) memory_occupy = x.numel() * x.element_size() / 1024**3 - if isinstance(norm_layer, ops.GroupNorm) and memory_occupy > get_norm_limit(): + if isinstance(norm_layer, nn.GroupNorm) and memory_occupy > get_norm_limit(): num_chunks = min(BYTEDANCE_GN_CHUNKS_FP16 if x.element_size() == 2 else BYTEDANCE_GN_CHUNKS_FP32, norm_layer.num_groups) if norm_layer.num_groups % num_chunks != 0: raise ValueError( @@ -423,9 +422,9 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: ) num_groups_per_chunk = norm_layer.num_groups // num_chunks + weights = comfy.ops.cast_to_input(norm_layer.weight, x).chunk(num_chunks, dim=0) + biases = comfy.ops.cast_to_input(norm_layer.bias, x).chunk(num_chunks, dim=0) x = list(x.chunk(num_chunks, dim=1)) - weights = norm_layer.weight.chunk(num_chunks, dim=0) - biases = norm_layer.bias.chunk(num_chunks, dim=0) for i, (w, bias) in enumerate(zip(weights, biases)): x[i] = F.group_norm(x[i], num_groups_per_chunk, w, bias, norm_layer.eps) x[i] = x[i].to(input_dtype) @@ -1459,7 +1458,6 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL): def _encode_with_raw_latent(self, x): if x.ndim == 4: x = x.unsqueeze(2) - x = x.to(dtype=next(self.parameters()).dtype) self.device = x.device p = super().encode(x) z = p.squeeze(2) diff --git a/tests-unit/comfy_test/test_seedvr2_dtype.py b/tests-unit/comfy_test/test_seedvr2_dtype.py index 8e08b6dde..d743cc848 100644 --- a/tests-unit/comfy_test/test_seedvr2_dtype.py +++ b/tests-unit/comfy_test/test_seedvr2_dtype.py @@ -1,4 +1,5 @@ import torch +import torch.nn as nn from comfy.cli_args import args as cli_args @@ -48,3 +49,31 @@ def test_seedvr2_vae_decode_memory_covers_full_frame_lab_transfer(): assert estimate == 101 * 960 * 1280 * 160 assert estimate > 15 * 1024 ** 3 assert estimate > old_estimate * 100 + + +def test_seedvr2_vae_encode_preserves_compute_dtype(monkeypatch): + wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper) + nn.Module.__init__(wrapper) + wrapper._dummy = nn.Parameter(torch.empty(1, dtype=torch.float16)) + input_dtype = None + + def encode(self, x): + nonlocal input_dtype + input_dtype = x.dtype + return x + + monkeypatch.setattr(seedvr_vae.VideoAutoencoderKL, "encode", encode) + + x = torch.zeros((1, 3, 1, 8, 8), dtype=torch.float32) + wrapper._encode_with_raw_latent(x) + + assert input_dtype == torch.float32 + + +def test_seedvr2_vae_ops_cast_weights_to_compute_dtype(): + attention = seedvr_vae.Attention(query_dim=4, heads=1, dim_head=4).to(torch.float16) + hidden_states = torch.zeros((1, 2, 4), dtype=torch.float32) + + output = attention(hidden_states) + + assert output.dtype == torch.float32 diff --git a/tests-unit/comfy_test/test_seedvr2_vae_tiled.py b/tests-unit/comfy_test/test_seedvr2_vae_tiled.py index a2866b609..d64f51918 100644 --- a/tests-unit/comfy_test/test_seedvr2_vae_tiled.py +++ b/tests-unit/comfy_test/test_seedvr2_vae_tiled.py @@ -122,6 +122,31 @@ def test_tiled_vae_encode_uses_tensor_return_without_indexing(): assert tuple(out.shape) == (2, _LATENT_CHANNELS, 1, 8, 8) +def test_tiled_vae_preserves_compute_dtype_with_different_parameter_dtype(): + class DummyVAE(nn.Module): + spatial_downsample_factor = 8 + temporal_downsample_factor = 4 + slicing_sample_min_size = 8 + + def __init__(self): + super().__init__() + self.device = torch.device("cpu") + self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float16)) + self.input_dtype = None + + def encode(self, t_chunk): + self.input_dtype = t_chunk.dtype + b, _, _, h, w = t_chunk.shape + return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype) + + vae = DummyVAE() + x = torch.zeros((1, 3, 1, 64, 64), dtype=torch.float32) + + tiled_vae(x, vae, tile_size=(64, 64), tile_overlap=(16, 16), encode=True) + + assert vae.input_dtype == torch.float32 + + def test_tiled_vae_preserves_input_dtype_on_single_tile(): class FloatOutputVAEModel(torch.nn.Module): def __init__(self): From f3a36e74844893f32f77f22d249d08862805d8f4 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 10 Jul 2026 18:37:59 -0700 Subject: [PATCH 093/211] Temporarily disable auto enabling triton by default on AMD. (#14878) I get freezing issues on my test machine. --- comfy/quant_ops.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index b1aabdc93..15f9b1fdb 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -48,7 +48,7 @@ try: # older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path. if args.disable_triton_backend: ck.registry.disable("triton") - elif args.enable_triton_backend or (torch.version.hip is not None and _rocm_kitchen_arch_supported()): + elif args.enable_triton_backend: # or (torch.version.hip is not None and _rocm_kitchen_arch_supported()): try: import triton triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2]) From 69ea58697bb2f05124f5dc7e00ad111f7cfff645 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 11 Jul 2026 17:16:40 -0700 Subject: [PATCH 094/211] Try to fix flash attention related issue on AMD. (#14880) --- comfy/ldm/modules/attention.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 2411aff5c..e6500cff4 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -709,7 +709,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape return out try: - @torch.library.custom_op("flash_attention::flash_attn", mutates_args=()) + @torch.library.custom_op("comfy::flash_attn", mutates_args=()) def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor: softmax_scale_arg = None if softmax_scale == -1.0 else softmax_scale From 8b099de36acd81acd1afa3b5442951dc847e0a52 Mon Sep 17 00:00:00 2001 From: Gustavo Schneiter Date: Sun, 12 Jul 2026 01:58:25 -0300 Subject: [PATCH 095/211] Fix SaveVideo description: says images, saves video (#14885) --- comfy_extras/nodes_video.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy_extras/nodes_video.py b/comfy_extras/nodes_video.py index d3acc9ad0..3bfd00be4 100644 --- a/comfy_extras/nodes_video.py +++ b/comfy_extras/nodes_video.py @@ -81,7 +81,7 @@ class SaveVideo(io.ComfyNode): display_name="Save Video", category="video", essentials_category="Basics", - description="Saves the input images to your ComfyUI output directory.", + description="Saves the input videos to your ComfyUI output directory.", inputs=[ io.Video.Input("video", tooltip="The video to save."), io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."), From 917faef771a2fd2f14f44af94f17da3d0b2803a3 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sun, 12 Jul 2026 09:43:30 -0700 Subject: [PATCH 096/211] Support PID 1.5 models. (#14894) --- comfy/ldm/pixeldit/model.py | 4 ++ comfy/ldm/pixeldit/pid.py | 64 +++++++++++++++---- comfy/model_detection.py | 37 ++++++++++- tests-unit/comfy_test/model_detection_test.py | 52 +++++++++++++++ 4 files changed, 140 insertions(+), 17 deletions(-) diff --git a/comfy/ldm/pixeldit/model.py b/comfy/ldm/pixeldit/model.py index b044b9b29..3b30b9226 100644 --- a/comfy/ldm/pixeldit/model.py +++ b/comfy/ldm/pixeldit/model.py @@ -197,6 +197,9 @@ class PixDiT_T2I(nn.Module): """Hook for subclasses to inject per-block state into the patch stream (e.g. PiD's LQ gate).""" return s + def _pre_pixel_blocks(self, s, **kwargs): + return s + def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, **kwargs): H_orig, W_orig = x.shape[2], x.shape[3] x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size)) @@ -226,6 +229,7 @@ class PixDiT_T2I(nn.Module): s, y_emb = blk(s, y_emb, condition, pos_img, pos_txt, None, transformer_options=transformer_options) s = F.silu(t_emb + s) + s = self._pre_pixel_blocks(s, **kwargs) s_cond = s.view(B * L, self.hidden_size) x_pixels = self.pixel_embedder(x, patch_size=self.patch_size) for blk in self.pixel_blocks: diff --git a/comfy/ldm/pixeldit/pid.py b/comfy/ldm/pixeldit/pid.py index 21b73907a..8590408d9 100644 --- a/comfy/ldm/pixeldit/pid.py +++ b/comfy/ldm/pixeldit/pid.py @@ -13,15 +13,15 @@ from .model import PixDiT_T2I from .modules import precompute_freqs_cis_2d -class SigmaAwareGatePerTokenPerDim(nn.Module): +class SigmaAwareGate(nn.Module): """gate = sigmoid(content_proj(cat[x, lq]) - exp(log_alpha) * sigma); out = x + gate * lq. Trained init gives ~0.88 gate at sigma=0, ~0.05 at sigma=1. """ - def __init__(self, dim: int, dtype=None, device=None, operations=None): + def __init__(self, dim: int, per_token: bool = False, dtype=None, device=None, operations=None): super().__init__() - self.content_proj = operations.Linear(dim * 2, dim, dtype=dtype, device=device) + self.content_proj = operations.Linear(dim * 2, 1 if per_token else dim, dtype=dtype, device=device) self.log_alpha = nn.Parameter(torch.empty((), dtype=dtype, device=device)) def forward(self, x: torch.Tensor, lq: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: @@ -36,15 +36,15 @@ class SigmaAwareGatePerTokenPerDim(nn.Module): class ResBlock(nn.Module): """Pre-activation ResNet block: GN -> SiLU -> Conv -> GN -> SiLU -> Conv + skip.""" - def __init__(self, channels: int, num_groups: int = 4, dtype=None, device=None, operations=None): + def __init__(self, channels: int, num_groups: int = 4, conv_padding_mode: str = "zeros", dtype=None, device=None, operations=None): super().__init__() self.block = nn.Sequential( operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), ) def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -62,9 +62,13 @@ class LQProjection2D(nn.Module): patch_size: int = 16, sr_scale: int = 4, latent_spatial_down_factor: int = 8, + latent_unpatchify_factor: int = 1, num_res_blocks: int = 4, num_outputs: int = 7, interval: int = 2, + conv_padding_mode: str = "zeros", + gate_per_token: bool = False, + pit_output: bool = False, dtype=None, device=None, operations=None, ): super().__init__() @@ -74,34 +78,38 @@ class LQProjection2D(nn.Module): self.patch_size = patch_size self.sr_scale = sr_scale self.latent_spatial_down_factor = latent_spatial_down_factor + self.latent_unpatchify_factor = latent_unpatchify_factor self.num_outputs = num_outputs self.interval = interval - z_to_patch_ratio = (sr_scale * latent_spatial_down_factor) / patch_size + effective_latent_channels = latent_channels // (latent_unpatchify_factor * latent_unpatchify_factor) + effective_spatial_down_factor = latent_spatial_down_factor // latent_unpatchify_factor + z_to_patch_ratio = (sr_scale * effective_spatial_down_factor) / patch_size self.z_to_patch_ratio = z_to_patch_ratio if z_to_patch_ratio >= 1: self.latent_fold_factor = 0 - latent_proj_in_ch = latent_channels + latent_proj_in_ch = effective_latent_channels else: fold_factor = int(1 / z_to_patch_ratio) assert fold_factor * z_to_patch_ratio == 1.0 self.latent_fold_factor = fold_factor - latent_proj_in_ch = latent_channels * fold_factor * fold_factor + latent_proj_in_ch = effective_latent_channels * fold_factor * fold_factor layers = [ - operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), ] for _ in range(num_res_blocks): - layers.append(ResBlock(hidden_dim, dtype=dtype, device=device, operations=operations)) + layers.append(ResBlock(hidden_dim, conv_padding_mode=conv_padding_mode, dtype=dtype, device=device, operations=operations)) self.latent_proj = nn.Sequential(*layers) self.output_heads = nn.ModuleList( [operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) for _ in range(num_outputs)] ) + self.pit_head = operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) if pit_output else None self.gate_modules = nn.ModuleList( - [SigmaAwareGatePerTokenPerDim(out_dim, dtype=dtype, device=device, operations=operations) + [SigmaAwareGate(out_dim, per_token=gate_per_token, dtype=dtype, device=device, operations=operations) for _ in range(num_outputs)] ) @@ -115,6 +123,11 @@ class LQProjection2D(nn.Module): return self.gate_modules[out_idx](x, lq_feature, sigma) def _align_latent_to_patch_grid(self, lq_latent: torch.Tensor, pH: int, pW: int) -> torch.Tensor: + f = self.latent_unpatchify_factor + if f > 1: + B, C, H, W = lq_latent.shape + lq_latent = lq_latent.reshape(B, C // (f * f), f, f, H, W) + lq_latent = lq_latent.permute(0, 1, 4, 2, 5, 3).reshape(B, C // (f * f), H * f, W * f) B, z_dim = lq_latent.shape[:2] if self.z_to_patch_ratio >= 1: if lq_latent.shape[2] != pH or lq_latent.shape[3] != pW: @@ -134,7 +147,10 @@ class LQProjection2D(nn.Module): feat = self._align_latent_to_patch_grid(lq_latent, target_pH, target_pW) B, C, H, W = feat.shape tokens = feat.permute(0, 2, 3, 1).contiguous().view(B, H * W, C) - return [head(tokens) for head in self.output_heads] + outputs = [head(tokens) for head in self.output_heads] + if self.pit_head is not None: + outputs.append(self.pit_head(tokens)) + return outputs class PidNet(PixDiT_T2I): @@ -148,6 +164,10 @@ class PidNet(PixDiT_T2I): lq_interval: int = 2, sr_scale: int = 4, latent_spatial_down_factor: int = 8, + lq_latent_unpatchify_factor: int = 1, + lq_conv_padding_mode: str = "zeros", + lq_gate_per_token: bool = False, + pit_lq_inject: bool = False, rope_ref_h: int = 1024, # NTK ref resolution in PIXEL units: 1024px / patch=16 -> grid_ref=64. rope_ref_w: int = 1024, image_model=None, @@ -165,6 +185,8 @@ class PidNet(PixDiT_T2I): for blk in self.pixel_blocks: blk._rope_fn = _pit_rope_fn + self.pit_lq_inject = pit_lq_inject + num_lq_outputs = (self.patch_depth + lq_interval - 1) // lq_interval self.lq_proj = LQProjection2D( latent_channels=lq_latent_channels, @@ -173,13 +195,20 @@ class PidNet(PixDiT_T2I): patch_size=self.patch_size, sr_scale=sr_scale, latent_spatial_down_factor=latent_spatial_down_factor, + latent_unpatchify_factor=lq_latent_unpatchify_factor, num_res_blocks=lq_num_res_blocks, num_outputs=num_lq_outputs, interval=lq_interval, + conv_padding_mode=lq_conv_padding_mode, + gate_per_token=lq_gate_per_token, + pit_output=pit_lq_inject, dtype=dtype, device=device, operations=operations, ) + self.pit_lq_gate = SigmaAwareGate( + self.hidden_size, per_token=lq_gate_per_token, dtype=dtype, device=device, operations=operations + ) if pit_lq_inject else None def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts): return precompute_freqs_cis_2d( @@ -197,6 +226,11 @@ class PidNet(PixDiT_T2I): return s return self.lq_proj.gate(s, pid_lq_features[out_idx], pid_degrade_sigma, out_idx) + def _pre_pixel_blocks(self, s, pid_pit_lq_feature=None, pid_degrade_sigma=None, **kwargs): + if pid_pit_lq_feature is None: + return s + return self.pit_lq_gate(s, pid_pit_lq_feature, pid_degrade_sigma) + def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, lq_latent=None, degrade_sigma=None, **kwargs): if lq_latent is None: raise ValueError("PidNet requires lq_latent — attach via PiDConditioning") @@ -216,12 +250,14 @@ class PidNet(PixDiT_T2I): degrade_sigma = degrade_sigma.expand(B).contiguous() lq_features = self.lq_proj(lq_latent=lq_latent.to(x), target_pH=Hs, target_pW=Ws) + pit_lq_feature = lq_features.pop() if self.pit_lq_inject else None return super()._forward( x, timesteps, context=context, attention_mask=attention_mask, transformer_options=transformer_options, pid_lq_features=lq_features, + pid_pit_lq_feature=pit_lq_feature, pid_degrade_sigma=degrade_sigma, **kwargs, ) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 174bc77cc..70c8625e3 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -470,15 +470,46 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): # PiD (Pixel Diffusion Decoder). Must check BEFORE plain PixelDiT_T2I. _lq_w_key = '{}lq_proj.latent_proj.0.weight'.format(key_prefix) if _lq_w_key in state_dict_keys: - in_ch = int(state_dict[_lq_w_key].shape[1]) + latent_proj_in_channels = int(state_dict[_lq_w_key].shape[1]) + hidden_dim = int(state_dict[_lq_w_key].shape[0]) _gate_prefix = '{}lq_proj.gate_modules.'.format(key_prefix) num_gates = len({k[len(_gate_prefix):].split('.')[0] for k in state_dict_keys if k.startswith(_gate_prefix)}) + pid_v1_5 = '{}lq_proj.pit_head.weight'.format(key_prefix) in state_dict_keys dit_config = {"image_model": "pid", - "lq_latent_channels": in_ch, - "latent_spatial_down_factor": 16 if in_ch >= 64 else 8} + "lq_hidden_dim": hidden_dim} if num_gates > 0: dit_config["lq_interval"] = (14 + num_gates - 1) // num_gates + if pid_v1_5: + pid_v1_5_variants = { + 16: { # Flux and QwenImage + "lq_latent_channels": 16, + "latent_spatial_down_factor": 8, + "lq_latent_unpatchify_factor": 1, + }, + 32: { # Flux2 after 2x latent unpatchify + "lq_latent_channels": 128, + "latent_spatial_down_factor": 16, + "lq_latent_unpatchify_factor": 2, + }, + } + variant = pid_v1_5_variants.get(latent_proj_in_channels) + if variant is None: + raise ValueError(f"Unsupported PiD v1.5 latent projection with {latent_proj_in_channels} input channels") + gate_weight = state_dict['{}lq_proj.gate_modules.0.content_proj.weight'.format(key_prefix)] + dit_config.update(variant) + dit_config.update({ + "lq_conv_padding_mode": "replicate", + "lq_gate_per_token": gate_weight.shape[0] == 1, + "pit_lq_inject": True, + "rope_ref_h": 2048, + "rope_ref_w": 2048, + }) + else: + dit_config.update({ + "lq_latent_channels": latent_proj_in_channels, + "latent_spatial_down_factor": 16 if latent_proj_in_channels >= 64 else 8, + }) return dit_config if '{}core.pixel_embedder.proj.weight'.format(key_prefix) in state_dict_keys: # PixelDiT T2I diff --git a/tests-unit/comfy_test/model_detection_test.py b/tests-unit/comfy_test/model_detection_test.py index 6e7d71f79..7c5b271c5 100644 --- a/tests-unit/comfy_test/model_detection_test.py +++ b/tests-unit/comfy_test/model_detection_test.py @@ -97,6 +97,21 @@ def _make_seedvr2_3b_shared_mm_sd(): } +def _make_pid_v1_5_sd(latent_proj_channels=16): + sd = { + "pixel_embedder.proj.weight": torch.empty(16, 3, device="meta"), + "lq_proj.latent_proj.0.weight": torch.empty(1024, latent_proj_channels, 3, 3, device="meta"), + "lq_proj.pit_head.weight": torch.empty(1536, 1024, device="meta"), + "lq_proj.gate_modules.0.content_proj.weight": torch.empty(1, 3072, device="meta"), + "pixel_blocks.0.attn.q_norm.weight": torch.empty(72, device="meta"), + "pixel_blocks.0.adaLN_modulation.0.weight": torch.empty(24576, 1536, device="meta"), + "pixel_blocks.0.adaLN_modulation.0.bias": torch.empty(24576, device="meta"), + } + for i in range(7): + sd[f"lq_proj.gate_modules.{i}.log_alpha"] = torch.empty((), device="meta") + return sd + + def _add_model_diffusion_prefix(sd): return {f"model.diffusion_model.{k}": v for k, v in sd.items()} @@ -206,6 +221,43 @@ class TestModelDetection: assert type(model_config_from_unet(sd, "model.diffusion_model.")).__name__ == "SeedVR2" + def test_pid_v1_5_detection(self): + sd = _make_pid_v1_5_sd() + unet_config = detect_unet_config(sd, "") + + assert unet_config == { + "image_model": "pid", + "lq_latent_channels": 16, + "lq_hidden_dim": 1024, + "latent_spatial_down_factor": 8, + "lq_interval": 2, + "lq_latent_unpatchify_factor": 1, + "lq_conv_padding_mode": "replicate", + "lq_gate_per_token": True, + "pit_lq_inject": True, + "rope_ref_h": 2048, + "rope_ref_w": 2048, + } + assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "PiD" + + def test_pid_v1_5_flux2_detection(self): + unet_config = detect_unet_config(_make_pid_v1_5_sd(latent_proj_channels=32), "") + + assert unet_config["lq_latent_channels"] == 128 + assert unet_config["latent_spatial_down_factor"] == 16 + assert unet_config["lq_latent_unpatchify_factor"] == 2 + + def test_pid_v1_5_pixel_adaln_conversion(self): + sd = _make_pid_v1_5_sd() + model_config = model_config_from_unet_config(detect_unet_config(sd, ""), sd) + processed = model_config.process_unet_state_dict(sd) + + assert processed["pixel_blocks.0.attn.q_norm.weight"].shape == (72,) + assert processed["pixel_blocks.0.adaLN_modulation_msa.weight"].shape == (12288, 1536) + assert processed["pixel_blocks.0.adaLN_modulation_mlp.weight"].shape == (12288, 1536) + assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,) + assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,) + def test_unet_config_and_required_keys_combination_is_unique(self): """Each model in the registry must have a unique combination of ``unet_config`` and ``required_keys``. If two models share the same From b58f829b570d52fbbd41ebb00c6fe02b8755ec75 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 13 Jul 2026 10:38:35 +0300 Subject: [PATCH 097/211] [Partner Nodes] feat(client): send ComfyUI Core version in request headers (#14910) Signed-off-by: bigcat88 --- comfy_api_nodes/util/_helpers.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/comfy_api_nodes/util/_helpers.py b/comfy_api_nodes/util/_helpers.py index 7eb1ec664..acab10d95 100644 --- a/comfy_api_nodes/util/_helpers.py +++ b/comfy_api_nodes/util/_helpers.py @@ -15,6 +15,7 @@ from comfy.comfy_api_env import normalize_comfy_api_base from comfy.deploy_environment import get_deploy_environment from comfy.model_management import processing_interrupted from comfy_api.latest import IO +from comfyui_version import __version__ as comfyui_version from .common_exceptions import ProcessingInterrupted @@ -60,6 +61,7 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]: **get_auth_header(node_cls), "Comfy-Env": get_deploy_environment(), "Comfy-Usage-Source": get_usage_source(node_cls), + "Comfy-Core-Version": comfyui_version, } From 8deaa4d911497f93bbd434a3821efab396f6981f Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 13 Jul 2026 10:53:37 +0300 Subject: [PATCH 098/211] fix(image): correct HLG inverse-OETF clamp in hlg_to_linear (#14762) Signed-off-by: bigcat88 Co-authored-by: Alexis Rolland --- comfy_extras/nodes_images.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/comfy_extras/nodes_images.py b/comfy_extras/nodes_images.py index fe1937ba5..4d7b37200 100644 --- a/comfy_extras/nodes_images.py +++ b/comfy_extras/nodes_images.py @@ -891,10 +891,11 @@ def hlg_to_linear(t: torch.Tensor) -> torch.Tensor: return torch.cat([hlg_to_linear(rgb), alpha], dim=-1) # Piecewise: sqrt branch below 0.5, log branch above. - # Clamp inside the log branch so negative / out-of-range values don't blow up; + # Clamp the log branch at the 0.5 branch point (not above it) so the + # unselected lane stays finite in exp() without altering selected values; # values above 1.0 are allowed and extrapolate naturally. low = (t ** 2) / 3.0 - high = (torch.exp((t.clamp(min=_HLG_C) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0 + high = (torch.exp((t.clamp(min=0.5) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0 return torch.where(t <= 0.5, low, high) From ec0e8b3447d5aa5a91a5a846b7fd94c88318fef7 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 13 Jul 2026 11:38:50 +0300 Subject: [PATCH 099/211] fix(image): support single-channel images in Save Image (Advanced) (#14761) Signed-off-by: bigcat88 Co-authored-by: Alexis Rolland --- comfy_extras/nodes_images.py | 32 +++++++++++++++++++++----------- 1 file changed, 21 insertions(+), 11 deletions(-) diff --git a/comfy_extras/nodes_images.py b/comfy_extras/nodes_images.py index 4d7b37200..7011d9c13 100644 --- a/comfy_extras/nodes_images.py +++ b/comfy_extras/nodes_images.py @@ -844,15 +844,18 @@ class ImageMergeTileList(IO.ComfyNode): # Format specifications # --------------------------------------------------------------------------- -# Maps (file_format, bit_depth, has_alpha) -> (numpy dtype scale, av pixel format, -# stream pix_fmt). Keeps the encode path declarative instead of branchy. +# Maps (file_format, bit_depth, num_channels) -> (quantization scale, numpy dtype, +# av frame pix_fmt, stream pix_fmt). Keeps the encode path declarative instead of branchy. _FORMAT_SPECS = { - ("png", "8-bit", False): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"}, - ("png", "8-bit", True): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"}, - ("png", "16-bit", False): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"}, - ("png", "16-bit", True): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"}, - ("exr", "32-bit float", False): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"}, - ("exr", "32-bit float", True): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"}, + ("png", "8-bit", 1): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "gray", "stream_fmt": "gray"}, + ("png", "8-bit", 3): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"}, + ("png", "8-bit", 4): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"}, + ("png", "16-bit", 1): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "gray16le", "stream_fmt": "gray16be"}, + ("png", "16-bit", 3): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"}, + ("png", "16-bit", 4): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"}, + ("exr", "32-bit float", 1): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "grayf32le", "stream_fmt": "grayf32le"}, + ("exr", "32-bit float", 3): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"}, + ("exr", "32-bit float", 4): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"}, } @@ -1088,7 +1091,8 @@ def _encode_image( bit_depth: str, colorspace: str, ) -> bytes: - """Encode a single HxWxC tensor to PNG or EXR bytes in memory. + """Encode a single HxWxC (or channel-less HxW grayscale) tensor to PNG or + EXR bytes in memory. Grayscale is written as single-channel PNG / Y-only EXR. For EXR the input is interpreted according to `colorspace` and converted to scene-linear (EXR's convention) before writing: @@ -1102,10 +1106,16 @@ def _encode_image( For PNG, colorspace selection does not modify pixels — PNG is delivered sRGB-encoded and there is no PNG path for wide-gamut HDR in this node. """ + if img_tensor.ndim == 2: + img_tensor = img_tensor.unsqueeze(-1) # Some nodes emit grayscale as (H, W) with no channel dim, mask-style. height, width, num_channels = img_tensor.shape - has_alpha = num_channels == 4 - spec = _FORMAT_SPECS[(file_format, bit_depth, has_alpha)] + spec = _FORMAT_SPECS.get((file_format, bit_depth, num_channels)) + if spec is None: + raise ValueError( + f"No {file_format}/{bit_depth} encoder for {num_channels}-channel images: " + "supported channel counts are 1 (grayscale), 3 (RGB) and 4 (RGBA)." + ) if spec["dtype"] == np.float32: # EXR path: preserve full range, no clamp. From 5697b970173bc0c16a05c30d509d0911f2b84822 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 13 Jul 2026 14:18:09 +0300 Subject: [PATCH 100/211] [Partner Nodes] chore(Google): reroute Gemini Image preview models to release versions (#14917) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_gemini.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index aa992802d..a8eb0a797 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -1133,7 +1133,9 @@ class GeminiImage2(IO.ComfyNode): ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1) if model == "Nano Banana 2 (Gemini 3.1 Flash Image)": - model = "gemini-3.1-flash-image-preview" + model = "gemini-3.1-flash-image" + elif model == "gemini-3-pro-image-preview": + model = "gemini-3-pro-image" parts: list[GeminiPart] = [GeminiPart(text=prompt)] if images is not None: @@ -1507,7 +1509,7 @@ class GeminiNanoBanana2V2(IO.ComfyNode): validate_string(prompt, strip_whitespace=True, min_length=1) model_choice = model["model"] if model_choice == "Nano Banana 2 (Gemini 3.1 Flash Image)": - model_id = "gemini-3.1-flash-image-preview" + model_id = "gemini-3.1-flash-image" elif model_choice == "Nano Banana 2 Lite": model_id = "gemini-3.1-flash-lite-image" else: From 5bb831a3f565dbcd517fc157284ce69af528ec2e Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Mon, 13 Jul 2026 22:13:23 +0800 Subject: [PATCH 101/211] chore: update embedded docs to v0.5.8 (#14920) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 790ef4940..b27de8987 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ comfyui-frontend-package==1.45.20 comfyui-workflow-templates==0.11.6 -comfyui-embedded-docs==0.5.7 +comfyui-embedded-docs==0.5.8 torch torchsde torchvision From 5658a68a875e4c1210d72c71d61eca95740adaf0 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Tue, 14 Jul 2026 02:20:58 +0900 Subject: [PATCH 102/211] chore(openapi): sync shared API contract from cloud@bcb8f5f (#14815) Co-authored-by: mattmillerai <7741082+mattmillerai@users.noreply.github.com> Co-authored-by: Alexis Rolland Co-authored-by: Matt Miller --- openapi.yaml | 45 ++++++++++++++++++++++++--------------------- 1 file changed, 24 insertions(+), 21 deletions(-) diff --git a/openapi.yaml b/openapi.yaml index c09b1eeac..e00643bad 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -7,18 +7,18 @@ components: description: Timestamp when the asset was created format: date-time type: string + display_name: + description: Display name of the asset. Mirrors name for backwards compatibility. + nullable: true + type: string + file_path: + description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors") + nullable: true + type: string hash: description: Blake3 hash of the asset content. pattern: ^blake3:[a-f0-9]{64}$ type: string - loader_path: - description: The value a loader consumes to load this asset. Null when no loader can resolve the file. - nullable: true - type: string - display_name: - description: Human-facing label for the asset. Not unique. - nullable: true - type: string id: description: Unique identifier for the asset format: uuid @@ -144,6 +144,14 @@ components: AssetUpdated: description: Response returned when an existing asset is successfully updated. properties: + display_name: + description: Display name of the asset. Mirrors name for backwards compatibility. + nullable: true + type: string + file_path: + description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors") + nullable: true + type: string hash: description: Blake3 hash of the asset content. pattern: ^blake3:[a-f0-9]{64}$ @@ -775,14 +783,6 @@ components: ModelFolder: description: Represents a folder containing models properties: - extensions: - description: The folder's registered file-extension allowlist. An empty array means the folder accepts any extension (match-all). - example: - - .ckpt - - .safetensors - items: - type: string - type: array folders: description: List of paths where models of this type are stored example: @@ -1644,7 +1644,7 @@ paths: format: uuid type: string tags: - description: JSON-encoded array of tag strings. For new byte uploads, include exactly one destination role (`input`, `output`, or `models`); `models` uploads also require exactly one `model_type:` tag. Extra tags are stored as labels and do not create path components. + description: JSON-encoded array of freeform tag strings, e.g. '["models","checkpoint"]'. Common types include "models", "input", "output", and "temp", but any tag can be used in any order. type: string user_metadata: description: Custom JSON metadata as a string @@ -1829,7 +1829,7 @@ paths: content: application/json: schema: - $ref: '#/components/schemas/Asset' + $ref: '#/components/schemas/AssetUpdated' description: Asset updated successfully "400": content: @@ -2470,9 +2470,6 @@ paths: supports_preview_metadata: description: Whether the server supports preview metadata type: boolean - supports_model_type_tags: - description: Whether the server supports namespaced model type asset tags - type: boolean type: object description: Success headers: @@ -3300,6 +3297,12 @@ paths: schema: $ref: '#/components/schemas/ErrorResponse' description: Invalid request parameters + "401": + content: + application/json: + schema: + $ref: '#/components/schemas/ErrorResponse' + description: Unauthorized - Authentication required "500": content: application/json: From da2608926eaf68fd532bba4e1ace3402c5d21399 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Tue, 14 Jul 2026 03:15:34 +0800 Subject: [PATCH 103/211] Update workflow templates to v0.11.9 (#14924) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index b27de8987..e7e7ba747 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.20 -comfyui-workflow-templates==0.11.6 +comfyui-workflow-templates==0.11.9 comfyui-embedded-docs==0.5.8 torch torchsde From c35a622acdabf99f60bba737f07977fef1ef97f4 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 13 Jul 2026 12:52:28 -0700 Subject: [PATCH 104/211] Fix hidream o1 regression. (#14923) --- comfy/ldm/hidream_o1/attention.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/comfy/ldm/hidream_o1/attention.py b/comfy/ldm/hidream_o1/attention.py index 1b68f1771..afb2be9b8 100644 --- a/comfy/ldm/hidream_o1/attention.py +++ b/comfy/ldm/hidream_o1/attention.py @@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None): The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes. """ - def two_pass_attention(q, k, v, heads, **kwargs): + def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs): B, H, T, D = q.shape if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) elif ar_len >= T: - out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True) + out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa) elif ar_len <= 0: - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) else: out_ar = comfy.ops.scaled_dot_product_attention( q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len], - attn_mask=None, dropout_p=0.0, is_causal=True, + attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa, ) out_gen = optimized_attention( q[:, :, ar_len:], k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, - transformer_options=transformer_options, + transformer_options=transformer_options, enable_gqa=enable_gqa, ) out = torch.cat([out_ar, out_gen], dim=2) From 80acfcf0fe6ea006f0d6176be8301ade1c0c4836 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 13 Jul 2026 13:03:29 -0700 Subject: [PATCH 105/211] More optimized int8 and int4 on turing. (#14927) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index e7e7ba747..e1458ca34 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.18 +comfy-kitchen==0.2.19 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 0aecac867d7840b56ad790aa76c5e76e33c74c3d Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Mon, 13 Jul 2026 22:07:23 -0400 Subject: [PATCH 106/211] Fix 3D advanced nodes crashing in API mode (no UI) (#14930) --- comfy_extras/nodes_load_3d.py | 12 ++++++++---- comfy_extras/nodes_save_3d.py | 3 ++- 2 files changed, 10 insertions(+), 5 deletions(-) diff --git a/comfy_extras/nodes_load_3d.py b/comfy_extras/nodes_load_3d.py index a9df557c2..106b01f9d 100644 --- a/comfy_extras/nodes_load_3d.py +++ b/comfy_extras/nodes_load_3d.py @@ -174,8 +174,9 @@ class Preview3DAdvanced(IO.ComfyNode): filename = f"preview3d_advanced_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -243,8 +244,9 @@ class PreviewGaussianSplat(IO.ComfyNode): filename = f"preview_splat_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -303,8 +305,9 @@ class PreviewPointCloud(IO.ComfyNode): filename = f"preview_pointcloud_{uuid.uuid4().hex}.{model_3d.format}" model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( @@ -375,8 +378,9 @@ class Load3DAdvanced(IO.ComfyNode): file_3d = None if model_file and model_file != "none": file_3d = Types.File3D(folder_paths.get_annotated_filepath(model_file)) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} model_3d_info = viewport_state.get('model_3d_info', []) - return IO.NodeOutput(file_3d, model_3d_info, viewport_state['camera_info'], width, height) + return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height) class Load3DExtension(ComfyExtension): diff --git a/comfy_extras/nodes_save_3d.py b/comfy_extras/nodes_save_3d.py index 7c524caa1..e9fd07326 100644 --- a/comfy_extras/nodes_save_3d.py +++ b/comfy_extras/nodes_save_3d.py @@ -418,8 +418,9 @@ def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str: def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput: model_file = _save_file3d_to_output(model_3d, filename_prefix) + viewport_state = viewport_state if isinstance(viewport_state, dict) else {} camera_info_input = kwargs.get("camera_info", None) - camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info'] + camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info') model_3d_info_input = kwargs.get("model_3d_info", None) model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', []) return IO.NodeOutput( From dff0b18fff04fda6558d0c1dd2ad1dca43fc18bc Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 14 Jul 2026 13:13:08 -0700 Subject: [PATCH 107/211] Fix int8 performance regression on 16xx series. (#14941) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index e1458ca34..eb40caa6b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.19 +comfy-kitchen==0.2.20 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 26537080cb1da5a6efb05bfc5a4792569c845913 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 14 Jul 2026 23:21:20 +0300 Subject: [PATCH 108/211] [Partner Nodes] feat(sync.so): add support for "sync-3" model (#14928) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/sync_so.py | 49 ++++ comfy_api_nodes/nodes_sync_so.py | 391 +++++++++++++++++++++++++++++++ 2 files changed, 440 insertions(+) create mode 100644 comfy_api_nodes/apis/sync_so.py create mode 100644 comfy_api_nodes/nodes_sync_so.py diff --git a/comfy_api_nodes/apis/sync_so.py b/comfy_api_nodes/apis/sync_so.py new file mode 100644 index 000000000..af9419580 --- /dev/null +++ b/comfy_api_nodes/apis/sync_so.py @@ -0,0 +1,49 @@ +from pydantic import BaseModel, Field + + +class SyncInputItem(BaseModel): + type: str = Field(..., description="Input kind: 'video', 'image' or 'audio'.") + url: str = Field(...) + + +class SyncActiveSpeakerDetection(BaseModel): + auto_detect: bool | None = Field( + None, description="Detect the active speaker automatically. Video input only; rejected for images." + ) + frame_number: int | None = Field( + None, description="Frame used for manual speaker selection. Must be 0 for image inputs." + ) + coordinates: list[int] | None = Field( + None, description="Pixel [x, y] of the speaker's face in the frame selected by frame_number." + ) + + +class SyncGenerationOptions(BaseModel): + sync_mode: str | None = Field( + None, + description="How to resolve an audio/video duration mismatch: " + "cut_off, bounce, loop, silence or remap. Ignored for image inputs.", + ) + i2v_prompt: str | None = Field( + None, description="Motion prompt for image-to-video generation. Image input only." + ) + active_speaker_detection: SyncActiveSpeakerDetection | None = Field(None) + + +class SyncGenerationRequest(BaseModel): + model: str = Field(..., description="Generation model, e.g. 'sync-3'.") + input: list[SyncInputItem] = Field( + ..., description="Exactly one visual input (video or image) plus one audio input." + ) + options: SyncGenerationOptions | None = Field(None) + + +class SyncGeneration(BaseModel): + """Subset of the Generation object returned by POST /v2/generate and GET /v2/generate/{id}.""" + + id: str = Field(...) + status: str = Field(..., description="PENDING | PROCESSING | COMPLETED | FAILED | REJECTED") + outputUrl: str | None = Field(None) + outputDuration: float | None = Field(None) + error: str | None = Field(None, description="Human-readable failure message.") + errorCode: str | None = Field(None, description="Stable machine-readable code from the GET /v2/errors catalog.") diff --git a/comfy_api_nodes/nodes_sync_so.py b/comfy_api_nodes/nodes_sync_so.py new file mode 100644 index 000000000..27382b399 --- /dev/null +++ b/comfy_api_nodes/nodes_sync_so.py @@ -0,0 +1,391 @@ +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension, Input +from comfy_api_nodes.apis.sync_so import ( + SyncActiveSpeakerDetection, + SyncGeneration, + SyncGenerationOptions, + SyncGenerationRequest, + SyncInputItem, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + download_url_to_video_output, + downscale_image_tensor, + downscale_image_tensor_by_max_side, + get_image_dimensions, + get_number_of_images, + poll_op, + sync_op, + upload_audio_to_comfyapi, + upload_image_to_comfyapi, + upload_video_to_comfyapi, + validate_audio_duration, +) + + +class SyncLipSyncNode(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="SyncLipSyncNode", + display_name="sync.so Lip Sync", + category="partner/video/sync.so", + description=( + "Re-sync mouth movement in a video to new speech audio using sync.so. " + "Handles close-ups, profiles and obstructions automatically while preserving " + "the speaker's expression. Cost scales with output duration." + ), + inputs=[ + IO.Video.Input( + "video", + tooltip="Footage of the speaker to re-sync. Up to 4K (4096x2160); " + "a constant frame rate of 24/25/30 fps works best.", + ), + IO.Audio.Input( + "audio", + tooltip="Speech audio to sync the mouth to.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "sync-3", + [ + IO.Combo.Input( + "sync_mode", + options=["bounce", "cut_off", "loop", "silence", "remap"], + default="bounce", + tooltip=( + "How to handle a duration mismatch between video and audio; " + "this also sets the output length. " + "bounce: video plays forward then backward until the audio ends " + "(output = audio length). " + "loop: video restarts until the audio ends (output = audio length). " + "remap: video is time-stretched to match the audio (output = audio length). " + "cut_off: the longer track is trimmed (output = shorter length). " + "silence: nothing is trimmed; the shorter track is padded " + "(output = longer length)." + ), + ), + IO.Combo.Input( + "speaker_selection", + options=["default", "auto-detect", "coordinates"], + default="default", + tooltip=( + "Which face to lipsync when several people are visible. " + "default: let the model decide. " + "auto-detect: detect and follow the active speaker. " + "coordinates: target the face at pixel (speaker_x, speaker_y) " + "in the frame chosen by speaker_frame." + ), + ), + IO.Int.Input( + "speaker_frame", + default=0, + min=0, + max=1_000_000, + advanced=True, + tooltip="Video frame used to locate the speaker. " + "Only used when speaker_selection is 'coordinates'.", + ), + IO.Int.Input( + "speaker_x", + default=0, + min=0, + max=4096, + advanced=True, + tooltip="X pixel coordinate of the speaker's face. " + "Only used when speaker_selection is 'coordinates'.", + ), + IO.Int.Input( + "speaker_y", + default=0, + min=0, + max=4096, + advanced=True, + tooltip="Y pixel coordinate of the speaker's face. " + "Only used when speaker_selection is 'coordinates'.", + ), + ], + ) + ], + tooltip="sync.so generation model.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + video: Input.Video, + audio: Input.Audio, + seed: int, + model: dict, + ) -> IO.NodeOutput: + try: + width, height = video.get_dimensions() + except Exception: + width = height = None + if width and height and (max(width, height) > 4096 or width * height > 4096 * 2160): + raise ValueError( + f"sync.so rejects videos above 4K (4096x2160); got {width}x{height}. Downscale the video first." + ) + validate_audio_duration(audio, max_duration=600) + + if model["speaker_selection"] == "auto-detect": + speaker_detection = SyncActiveSpeakerDetection(auto_detect=True) + elif model["speaker_selection"] == "coordinates": + speaker_detection = SyncActiveSpeakerDetection( + frame_number=model["speaker_frame"], + coordinates=[model["speaker_x"], model["speaker_y"]], + ) + else: + speaker_detection = None + + video_url = await upload_video_to_comfyapi(cls, video, max_duration=600) + audio_url = await upload_audio_to_comfyapi(cls, audio) + + generation = await sync_op( + cls, + ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"), + response_model=SyncGeneration, + data=SyncGenerationRequest( + model=model["model"], + input=[ + SyncInputItem(type="video", url=video_url), + SyncInputItem(type="audio", url=audio_url), + ], + options=SyncGenerationOptions( + sync_mode=model["sync_mode"], + active_speaker_detection=speaker_detection, + ), + ), + ) + generation = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"), + response_model=SyncGeneration, + status_extractor=lambda g: g.status, + completed_statuses=["COMPLETED", "FAILED", "REJECTED"], + failed_statuses=[], + queued_statuses=["PENDING"], + poll_interval=10.0, + ) + if generation.status != "COMPLETED": + code = f" [{generation.errorCode}]" if generation.errorCode else "" + raise ValueError( + f"sync.so generation {generation.status.lower()}{code}: " + f"{generation.error or 'no error details provided'}" + ) + if not generation.outputUrl: + raise ValueError("sync.so generation completed but no output URL was returned.") + return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl)) + + +class SyncTalkingImageNode(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="SyncTalkingImageNode", + display_name="sync.so Talking Image", + category="partner/video/sync.so", + description=( + "Animate a still portrait into a talking video driven by speech audio, " + "using sync.so's sync-3 model. The output duration matches the audio. " + "Cost scales with output duration." + ), + inputs=[ + IO.Image.Input( + "image", + tooltip="A single image with a clearly visible face, up to 4K (4096x2160).", + ), + IO.Audio.Input( + "audio", + tooltip="Speech audio driving the talking video; the output duration matches it. " + "Chain any TTS node here to drive the animation from text.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Optional guidance for how the portrait comes to life, e.g. " + "'make the subject smile and look at the camera'. " + "Leave empty for natural talking motion.", + ), + IO.Int.Input( + "seed", + default=0, + min=0, + max=2147483647, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "sync-3", + [ + IO.Combo.Input( + "speaker_selection", + options=["default", "coordinates"], + default="default", + tooltip=( + "Which face to animate when several people are visible. " + "default: let the model decide. " + "coordinates: target the face at pixel (speaker_x, speaker_y) " + "in the image. Auto-detection is not supported for images." + ), + ), + IO.Int.Input( + "speaker_x", + default=0, + min=0, + max=4096, + advanced=True, + tooltip="X pixel coordinate of the speaker's face. " + "Only used when speaker_selection is 'coordinates'.", + ), + IO.Int.Input( + "speaker_y", + default=0, + min=0, + max=4096, + advanced=True, + tooltip="Y pixel coordinate of the speaker's face. " + "Only used when speaker_selection is 'coordinates'.", + ), + IO.Boolean.Input( + "auto_downscale", + default=True, + advanced=True, + tooltip="Automatically downscale the image if it exceeds the 4K " + "(4096x2160) input limit; speaker coordinates are scaled to match. " + "When disabled, an oversized image raises an error instead.", + ), + ], + ) + ], + tooltip="sync.so generation model. Image input is exclusive to sync-3.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + audio: Input.Audio, + prompt: str, + seed: int, + model: dict, + ) -> IO.NodeOutput: + if get_number_of_images(image) != 1: + raise ValueError("Exactly one image is required; got a batch. Pick one frame first.") + validate_audio_duration(audio, max_duration=600) + + height, width = get_image_dimensions(image) + speaker_x, speaker_y = model["speaker_x"], model["speaker_y"] + if max(width, height) > 4096 or width * height > 4096 * 2160: + if not model["auto_downscale"]: + raise ValueError( + f"sync.so rejects images above 4K (4096x2160); got {width}x{height}. " + "Downscale the image first or enable auto_downscale." + ) + image = downscale_image_tensor(image, total_pixels=4096 * 2160) + image = downscale_image_tensor_by_max_side(image, max_side=4096) + new_height, new_width = get_image_dimensions(image) + # speaker coordinates are given in the original image's pixel space + speaker_x = min(new_width - 1, round(speaker_x * new_width / width)) + speaker_y = min(new_height - 1, round(speaker_y * new_height / height)) + + if model["speaker_selection"] == "coordinates": + speaker_detection = SyncActiveSpeakerDetection( + frame_number=0, # images have a single frame; auto_detect is rejected by the API + coordinates=[speaker_x, speaker_y], + ) + else: + speaker_detection = None + + image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None) + audio_url = await upload_audio_to_comfyapi(cls, audio) + + generation = await sync_op( + cls, + ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"), + response_model=SyncGeneration, + data=SyncGenerationRequest( + model=model["model"], + input=[ + SyncInputItem(type="image", url=image_url), + SyncInputItem(type="audio", url=audio_url), + ], + options=SyncGenerationOptions( + i2v_prompt=prompt.strip() or None, + active_speaker_detection=speaker_detection, + ), + ), + ) + generation = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"), + response_model=SyncGeneration, + status_extractor=lambda g: g.status, + completed_statuses=["COMPLETED", "FAILED", "REJECTED"], + failed_statuses=[], + queued_statuses=["PENDING"], + poll_interval=10.0, + ) + if generation.status != "COMPLETED": + code = f" [{generation.errorCode}]" if generation.errorCode else "" + raise ValueError( + f"sync.so generation {generation.status.lower()}{code}: " + f"{generation.error or 'no error details provided'}" + ) + if not generation.outputUrl: + raise ValueError("sync.so generation completed but no output URL was returned.") + return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl)) + + +class SyncExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + SyncLipSyncNode, + SyncTalkingImageNode, + ] + + +async def comfy_entrypoint() -> SyncExtension: + return SyncExtension() From 1701cce8dc105e47bc1e6f5a7790c495bbcf331c Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 15 Jul 2026 00:29:01 +0300 Subject: [PATCH 109/211] Fix cached outputs missing from job results when prompt has no client_id (#14939) When a prompt is submitted without client_id and its output nodes are served from cache, _send_cached_ui returned early before recording the cached UI outputs, so /api/jobs/{job_id} (and /history) reported success with empty outputs. Record the outputs before the client_id check. --- execution.py | 4 ++-- tests/execution/test_execution.py | 24 ++++++++++++++++++++++++ 2 files changed, 26 insertions(+), 2 deletions(-) diff --git a/execution.py b/execution.py index 19b8cdd68..387772629 100644 --- a/execution.py +++ b/execution.py @@ -426,12 +426,12 @@ def _is_intermediate_output(dynprompt, node_id): def _send_cached_ui(server, node_id, display_node_id, cached, prompt_id, ui_outputs): + if cached.ui is not None: + ui_outputs[node_id] = cached.ui if server.client_id is None: return cached_ui = cached.ui or {} server.send_sync("executed", { "node": node_id, "display_node": display_node_id, "output": cached_ui.get("output", None), "prompt_id": prompt_id }, server.client_id) - if cached.ui is not None: - ui_outputs[node_id] = cached.ui async def execute(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, pending_subgraph_results, pending_async_nodes, ui_outputs): unique_id = current_item diff --git a/tests/execution/test_execution.py b/tests/execution/test_execution.py index 15e2304fc..c914d2feb 100644 --- a/tests/execution/test_execution.py +++ b/tests/execution/test_execution.py @@ -818,6 +818,30 @@ class TestExecution: except urllib.error.HTTPError: pass # Expected behavior + def test_cached_outputs_in_job_without_client_id(self, client: ComfyClient, builder: GraphBuilder): + g = builder + image = g.node("StubImage", content="BLACK", height=32, width=32, batch_size=1) + output = g.node("SaveImage", images=image.out(0)) + + # Prime the cache with a normal run. + client.run(g) + + # Resubmit anonymously (no client_id) so output nodes are cache hits with no websocket client. + data = json.dumps({"prompt": g.finalize()}).encode('utf-8') + req = urllib.request.Request(f"http://{client.server_address}/prompt", data=data) + prompt_id = json.loads(urllib.request.urlopen(req).read())['prompt_id'] + + for _ in range(100): + job = client.get_job(prompt_id) + if job is not None and job['status'] not in ('pending', 'in_progress'): + break + time.sleep(0.1) + else: + raise AssertionError("Prompt did not complete in time") + + assert job['status'] == 'completed' + assert output.id in job['outputs'], "Cached outputs must appear in job outputs without a client_id" + def _create_history_item(self, client, builder): g = GraphBuilder(prefix="offset_test") input_node = g.node( From 3cd13eb4245a111bda4a46b9129cef19c8cac1d5 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Wed, 15 Jul 2026 14:29:24 +0900 Subject: [PATCH 110/211] Bump comfyui-frontend-package to 1.45.21 (#14944) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index eb40caa6b..e7d301576 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -comfyui-frontend-package==1.45.20 +comfyui-frontend-package==1.45.21 comfyui-workflow-templates==0.11.9 comfyui-embedded-docs==0.5.8 torch From 700821e1364eaab0e8f21c538a2131719fec57bf Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Wed, 15 Jul 2026 01:43:14 -0400 Subject: [PATCH 111/211] ComfyUI v0.28.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index 8e9967f1b..dcc0fee96 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.27.0" +__version__ = "0.28.0" diff --git a/pyproject.toml b/pyproject.toml index 8c17e410e..73de2990f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.27.0" +version = "0.28.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From cc6b3525110bede7c4850b1f40403880b34d1ad8 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 15 Jul 2026 10:23:43 +0300 Subject: [PATCH 112/211] fix(Video): stream the video transcode instead of buffering every frame in RAM (CORE-353) (CORE-351) (#14813) --- comfy_api/latest/_input_impl/video_types.py | 394 +++++++++++-- tests-unit/comfy_api_test/video_types_test.py | 526 +++++++++++++++++- 2 files changed, 868 insertions(+), 52 deletions(-) diff --git a/comfy_api/latest/_input_impl/video_types.py b/comfy_api/latest/_input_impl/video_types.py index bc95a5b99..f5af41973 100644 --- a/comfy_api/latest/_input_impl/video_types.py +++ b/comfy_api/latest/_input_impl/video_types.py @@ -1,5 +1,6 @@ from av.container import InputContainer from av.subtitles.stream import SubtitleStream +from av.video.reformatter import ColorRange from fractions import Fraction from typing import Optional from .._input import AudioInput, VideoInput @@ -9,6 +10,7 @@ import itertools import json import numpy as np import math +import os import torch from .._util import VideoContainer, VideoCodec, VideoComponents import logging @@ -58,6 +60,57 @@ def video_stream_bit_depth(stream) -> int: return max(component.bits for component in stream.format.components) +def last_decodable_audio_stream(container: InputContainer): + """Streams FFmpeg has no decoder for have no codec context, and decoding their + packets crashes the process (e.g. APAC spatial-audio track in iPhone).""" + stream = next( + (s for s in reversed(container.streams.audio) if s.codec_context is not None), + None, + ) + if stream is None and len(container.streams.audio): + logging.warning("No decodable audio stream found in video; ignoring audio.") + return stream + + +def probe_audio_params(container: InputContainer, audio_stream, max_packets: int = 200): + """Containers probed only up to a window (mpegts) leave audio codec parameters unset when + audio starts beyond it; learn them by decoding ahead. The caller must seek back afterwards. + Returns (sample_rate, channels), zeros when the stream never yields a decodable frame.""" + for i, packet in enumerate(container.demux(audio_stream)): + try: + frames = packet.decode() + except av.error.FFmpegError: + frames = () + if frames: + return frames[0].sample_rate, frames[0].layout.nb_channels + if i >= max_packets: + break + return 0, 0 + + +def write_output_metadata(container: InputContainer, output, metadata: dict | None): + """Copy the source container's metadata, then overlay the caller's tags.""" + for key, value in container.metadata.items(): + if metadata is None or key not in metadata: + output.metadata[key] = value + if metadata is not None: + for key, value in metadata.items(): + output.metadata[key] = value if isinstance(value, str) else json.dumps(value) + + +def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> dict: + if format != VideoContainer.AUTO and format != VideoContainer.MP4: + raise ValueError("Only MP4 format is supported for now") + if codec != VideoCodec.AUTO and codec != VideoCodec.H264: + raise ValueError("Only H264 codec is supported for now") + open_kwargs = {"mode": "w", "options": {"movflags": "use_metadata_tags"}} + if isinstance(format, VideoContainer) and format != VideoContainer.AUTO: + open_kwargs["format"] = format.value + elif isinstance(path, io.BytesIO): + open_kwargs["format"] = "mp4" # no file extension to infer the format from + return open_kwargs + + class VideoFromFile(VideoInput): """ Class representing video input from a file. @@ -192,13 +245,10 @@ class VideoFromFile(VideoInput): return estimated_frames # 3. Last resort: decode frames and count them (streaming) - if self.__start_time < 0: - start_time = max(self._get_raw_duration() + self.__start_time, 0) - else: - start_time = self.__start_time + start_time, duration = self.get_active_trim_window() frame_count = 1 start_pts = int(start_time / video_stream.time_base) - end_pts = int((start_time + self.__duration) / video_stream.time_base) + end_pts = int((start_time + duration) / video_stream.time_base) container.seek(start_pts, stream=video_stream) frame_iterator = ( container.decode(video_stream) @@ -253,17 +303,14 @@ class VideoFromFile(VideoInput): def get_components_internal(self, container: InputContainer) -> VideoComponents: video_stream = self._get_first_video_stream(container) - if self.__start_time < 0: - start_time = max(self._get_raw_duration() + self.__start_time, 0) - else: - start_time = self.__start_time + start_time, duration = self.get_active_trim_window() # Get video frames frames = [] audio_frames = [] alphas = None start_pts = int(start_time / video_stream.time_base) - end_pts = int((start_time + self.__duration) / video_stream.time_base) + end_pts = int((start_time + duration) / video_stream.time_base) if start_pts != 0: container.seek(start_pts, stream=video_stream) @@ -281,18 +328,11 @@ class VideoFromFile(VideoInput): video_done = False audio_done = True - # Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context, - # and decoding their packets crashes the process. (e.g. APAC spatial-audio track in iPhone) - audio_stream = next( - (s for s in reversed(container.streams.audio) if s.codec_context is not None), - None, - ) + audio_stream = last_decodable_audio_stream(container) if audio_stream is not None: streams += [audio_stream] resampler = av.audio.resampler.AudioResampler(format='fltp') audio_done = False - elif len(container.streams.audio): - logging.warning("No decodable audio stream found in video; ignoring audio.") for packet in container.demux(*streams): if video_done and audio_done: @@ -305,7 +345,7 @@ class VideoFromFile(VideoInput): for frame in packet.decode(): if frame.pts < start_pts: continue - if self.__duration and frame.pts >= end_pts: + if duration and frame.pts >= end_pts: video_done = True break @@ -372,7 +412,7 @@ class VideoFromFile(VideoInput): map(resampler.resample, packet.decode()) ) for frame in aframes: - if self.__duration and frame.time > start_time + self.__duration: + if duration and frame.time > start_time + duration: audio_done = True break @@ -394,8 +434,8 @@ class VideoFromFile(VideoInput): if len(audio_frames) > 0: audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples) - if self.__duration: - audio_data = audio_data[..., :int(self.__duration * audio_stream.sample_rate)] + if duration: + audio_data = audio_data[..., :int(duration * audio_stream.sample_rate)] audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples) audio = AudioInput({ @@ -441,28 +481,14 @@ class VideoFromFile(VideoInput): if not reuse_streams: if bit_depth is None: bit_depth = source_bit_depth - components = self.get_components_internal(container) - video = VideoFromComponents(components) - return video.save_to( - path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth, - ) + return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth) streams = container.streams open_kwargs = get_open_write_kwargs(path, container_format, format) with av.open(path, **open_kwargs) as output_container: - # Copy over the original metadata - for key, value in container.metadata.items(): - if metadata is None or key not in metadata: - output_container.metadata[key] = value - - # Add our new metadata - if metadata is not None: - for key, value in metadata.items(): - if isinstance(value, str): - output_container.metadata[key] = value - else: - output_container.metadata[key] = json.dumps(value) + # Add metadata before writing any streams + write_output_metadata(container, output_container, metadata) # Add streams to the new container. Streams with no codec context cannot be used as an output template. stream_map = {} @@ -480,6 +506,282 @@ class VideoFromFile(VideoInput): packet.stream = stream_map[packet.stream] output_container.mux(packet) + def _save_transcoded( + self, + container: InputContainer, + path: str | io.BytesIO, + format: VideoContainer, + codec: VideoCodec, + metadata: dict | None, + bit_depth: int, + ): + """Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length.""" + open_kwargs = mp4_output_open_kwargs(path, format, codec) + video_stream = self._get_first_video_stream(container) + start_time, duration = self.get_active_trim_window() + start_pts = int(start_time / video_stream.time_base) + end_pts = int((start_time + duration) / video_stream.time_base) if duration else None + stream_end_pts = None + if video_stream.duration is not None: + stream_end_pts = (video_stream.start_time or 0) + video_stream.duration + output_end_pts = end_pts + if stream_end_pts is not None and (output_end_pts is None or stream_end_pts < output_end_pts): + output_end_pts = stream_end_pts + if start_pts != 0: + container.seek(start_pts, stream=video_stream) + + audio_stream = last_decodable_audio_stream(container) + pix_fmt = "yuv420p10le" if bit_depth >= 10 else "yuv420p" + rate = Fraction(video_stream.average_rate) if video_stream.average_rate else Fraction(1) + + resampler = None + sample_rate = 0 + audio_time_base = None + duration_cap = None + if audio_stream is not None: + sample_rate = audio_stream.codec_context.sample_rate + channels = audio_stream.codec_context.channels + if not sample_rate: + sample_rate, channels = probe_audio_params(container, audio_stream) + container.seek(start_pts, stream=video_stream) + if sample_rate: + audio_stream.codec_context.flush_buffers() + else: + logging.warning("Audio stream parameters could not be determined; ignoring audio.") + audio_stream = None + if audio_stream is not None: + audio_time_base = Fraction(1, sample_rate) + layout = {1: "mono", 2: "stereo", 6: "5.1"}.get(channels, "stereo") + resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=sample_rate) + if duration: + duration_cap = math.ceil(duration * sample_rate) + + streams = [video_stream] if audio_stream is None else [video_stream, audio_stream] + pts_step = max(1, int(round((1 / rate) / video_stream.time_base))) + video_done = False + audio_done = audio_stream is None + video_pts_offset = None + last_video_pts = None + last_video_end = None + # rebased pts -> true display duration: the mp4 muxer pads the last sample with 1/rate otherwise + video_frame_durations = {} + source_size = None + rotation_k = 0 + rotation_filter = None + audio_started = False + samples_written = 0 + pending_audio = [] + # The output opens lazily on the first kept frame: it decides the geometry (90/270 rotation swaps dims), + # and never seeking back keeps webm/mkv leading audio intact. + output = None + out_video = None + out_audio = None + + def audio_frame_from_ndarray(nd_planar): + frame = av.AudioFrame.from_ndarray(np.ascontiguousarray(nd_planar), format="fltp", layout=layout) + frame.sample_rate = sample_rate + return frame + + def drain_audio(final=False): + # Audio may cover the pts span of the video written so far, capped by the requested duration + nonlocal samples_written, audio_done + if last_video_end is None: + cap = 0 + else: + cap = math.ceil(last_video_end * video_stream.time_base * sample_rate) + if duration_cap is not None: + cap = min(cap, duration_cap) + while pending_audio and not audio_done: + frame = pending_audio[0] + if samples_written + frame.samples <= cap: + frame.pts = samples_written + frame.time_base = audio_time_base + output.mux(out_audio.encode(frame)) + samples_written += frame.samples + pending_audio.pop(0) + continue + if final: + keep = frame.to_ndarray()[..., :cap - samples_written] + if keep.shape[-1] > 0: + tail = audio_frame_from_ndarray(keep) + tail.pts = samples_written + tail.time_base = audio_time_base + output.mux(out_audio.encode(tail)) + samples_written += keep.shape[-1] + pending_audio.clear() + break + if duration_cap is not None and samples_written >= duration_cap: + audio_done = True + return cap + + try: + for packet in container.demux(*streams): + if video_done and audio_done: + break + + if packet.stream == video_stream and not video_done: + try: + frames = packet.decode() + except av.error.InvalidDataError: + logging.info("pyav decode error") + continue + for frame in frames: + if frame.pts is not None and frame.pts < start_pts: + continue + if end_pts is not None and frame.pts is not None and frame.pts >= end_pts: + video_done = True + if last_video_pts is not None: + # the source continues past the window: hold the last kept frame to the window end + end_offset = video_pts_offset if video_pts_offset is not None else start_pts + last_video_end = max(last_video_end, end_pts - end_offset) + break + # the source's true display duration of this frame; average_rate is not a + # frame duration (sparse/VFR sources), so it is only the fallback + frame_duration = frame.duration if frame.duration else pts_step + if end_pts is not None and frame.pts is not None: + frame_duration = min(frame_duration, end_pts - frame.pts) + if output is None: + rotation_k = int(round(frame.rotation // 90)) % 4 if frame.rotation else 0 + if rotation_k % 2: + out_width, out_height = frame.height, frame.width + else: + out_width, out_height = frame.width, frame.height + if out_width % 2 or out_height % 2: + raise ValueError(f"H.264 output requires even dimensions, got {out_width}x{out_height}") + source_size = (frame.width, frame.height) + output = av.open(path, **open_kwargs) + # Add metadata before writing any streams + write_output_metadata(container, output, metadata) + out_video = output.add_stream("h264", rate=rate) + # no B-frames: reordering makes mp4 sample durations follow decode order, + # so irregular-VFR spans and trim windows land wrong + out_video.codec_context.max_b_frames = 0 + out_video.width = out_width + out_video.height = out_height + out_video.pix_fmt = pix_fmt + # source pts pass through (rebased to 0), so variable frame rate survives + out_video.codec_context.time_base = video_stream.time_base + if audio_stream is not None: + out_audio = output.add_stream("aac", rate=sample_rate, layout=layout) + if (frame.width, frame.height) != source_size: + # encoding would silently rescale the new geometry into the old one + raise ValueError( + f"Video resolution changes mid-stream " + f"({source_size[0]}x{source_size[1]} -> {frame.width}x{frame.height}); cannot transcode" + ) + if rotation_k: + if rotation_filter is None: + g = av.filter.Graph() + g_src = g.add_buffer(width=frame.width, height=frame.height, + format=frame.format.name, time_base=video_stream.time_base) + tail = g_src + for filter_name, filter_args in {1: [("transpose", "cclock")], + 2: [("hflip", None), ("vflip", None)], + 3: [("transpose", "clock")]}[rotation_k]: + step = g.add(filter_name, filter_args) + tail.link_to(step) + tail = step + g_sink = g.add("buffersink") + tail.link_to(g_sink) + g.configure() + rotation_filter = (g_src, g_sink) + rotation_filter[0].push(frame) + frame = rotation_filter[1].pull() + if frame.color_range == ColorRange.JPEG: + # compress full-range sources (yuvj/MJPEG) to limited range + frame = frame.reformat(format=pix_fmt, src_color_range="JPEG", dst_color_range="MPEG") + else: + frame = frame.reformat(format=pix_fmt) + frame_output_end = None + if frame.pts is not None: + if video_pts_offset is None: + video_pts_offset = frame.pts + frame.pts -= video_pts_offset + if output_end_pts is not None: + frame_output_end = output_end_pts - video_pts_offset + if frame.pts + frame_duration > frame_output_end: + clamped_pts = frame_output_end - frame_duration + if clamped_pts >= 0 and (last_video_pts is None or clamped_pts > last_video_pts): + frame.pts = min(frame.pts, clamped_pts) + elif frame.pts < frame_output_end: + frame_duration = frame_output_end - frame.pts + else: + continue + if frame.pts is None or (last_video_pts is not None and frame.pts <= last_video_pts): + # broken sources emit missing/backward timestamps mid-stream, which the + # muxer rejects; nudge them forward by one nominal frame interval + frame.pts = 0 if last_video_pts is None else last_video_pts + pts_step + if frame_output_end is not None and frame.pts + frame_duration > frame_output_end: + if frame.pts >= frame_output_end: + continue + frame_duration = frame_output_end - frame.pts + last_video_pts = frame.pts + last_video_end = frame.pts + frame_duration + video_frame_durations[frame.pts] = frame_duration + # the decoded pict_type would force x264's frame types (intra-only + # sources like MJPEG/ProRes would come out all-keyframe) + frame.pict_type = 0 + for out_packet in out_video.encode(frame): + out_packet.duration = video_frame_durations.pop(out_packet.pts, 0) + output.mux(out_packet) + drain_audio() + + elif packet.stream == audio_stream and not audio_done: + for resampled in itertools.chain.from_iterable(map(resampler.resample, packet.decode())): + frame_start = None + if resampled.pts is not None: + # passthrough frames keep the source stream's time base + tb = resampled.time_base if resampled.time_base else audio_time_base + frame_start = float(resampled.pts * tb) + if duration and not audio_started and frame_start >= start_time + duration: + audio_done = True + break + if not audio_started: + if frame_start is None: + frame_start = 0.0 + to_skip = max(0, int((start_time - frame_start) * sample_rate)) + if to_skip >= resampled.samples: + continue + audio_started = True + if duration and frame_start > start_time: + duration_cap = min(duration_cap, math.ceil((start_time + duration - frame_start) * sample_rate)) + if to_skip: + pending_audio.append(audio_frame_from_ndarray(resampled.to_ndarray()[..., to_skip:])) + continue + pending_audio.append(resampled) + if video_done: + # the video window is complete so the cap is final, but containers + # that interleave audio behind video (fragmented mp4) still owe most + # of it: stop only once the demuxed audio covers the cap + cap = drain_audio() + if pending_audio or samples_written >= cap: + drain_audio(final=True) + audio_done = True + break + + if output is None: + raise ValueError(f"No decodable video frames found in file '{self.__file}'") + if out_audio is not None and not audio_done: + drain_audio(final=True) + window_fill = last_video_end - last_video_pts if video_done and last_video_pts is not None else 0 + for out_packet in out_video.encode(None): + duration = video_frame_durations.pop(out_packet.pts, 0) + if out_packet.pts == last_video_pts: + duration = max(duration, window_fill) + out_packet.duration = duration + output.mux(out_packet) + if out_audio is not None: + output.mux(out_audio.encode(None)) + except BaseException: + if output is not None: + output.close() + if isinstance(path, (str, os.PathLike)) and os.path.exists(path): + os.remove(path) + raise + else: + if output is not None: + output.close() + def _get_first_video_stream(self, container: InputContainer): if len(container.streams.video): return container.streams.video[0] @@ -527,22 +829,12 @@ class VideoFromComponents(VideoInput): bit_depth: int | None = None, ): """Save the video to a file path or BytesIO buffer.""" - if format != VideoContainer.AUTO and format != VideoContainer.MP4: - raise ValueError("Only MP4 format is supported for now") - if codec != VideoCodec.AUTO and codec != VideoCodec.H264: - raise ValueError("Only H264 codec is supported for now") + open_kwargs = mp4_output_open_kwargs(path, format, codec) # None means "use the depth this video was created with" (CreateVideo's choice). if bit_depth is None: bit_depth = self.__bit_depth is_10bit = bit_depth >= 10 - extra_kwargs = {} - if isinstance(format, VideoContainer) and format != VideoContainer.AUTO: - extra_kwargs["format"] = format.value - elif isinstance(path, io.BytesIO): - # BytesIO has no file extension, so av.open can't infer the format. - # Default to mp4 since that's the only supported format anyway. - extra_kwargs["format"] = "mp4" - with av.open(path, mode='w', options={'movflags': 'use_metadata_tags'}, **extra_kwargs) as output: + with av.open(path, **open_kwargs) as output: # Add metadata before writing any streams if metadata is not None: for key, value in metadata.items(): diff --git a/tests-unit/comfy_api_test/video_types_test.py b/tests-unit/comfy_api_test/video_types_test.py index b25fcb1ca..ae758bd40 100644 --- a/tests-unit/comfy_api_test/video_types_test.py +++ b/tests-unit/comfy_api_test/video_types_test.py @@ -2,11 +2,12 @@ import pytest import torch import tempfile import os +import sys import av import io from fractions import Fraction from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents -from comfy_api.util.video_types import VideoComponents +from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec from comfy_api.input.basic_types import AudioInput from av.error import InvalidDataError @@ -237,3 +238,526 @@ def test_duration_consistency(video_components): manual_duration = float(components.images.shape[0] / components.frame_rate) assert duration == pytest.approx(manual_duration) + + +def create_transcode_source( + width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False, + container_format="mov", audio_codec="pcm_s16le", +): + """Create a temp video that save_to must transcode (mpeg4 video, so codec != h264). + + ``undecodable_audio`` trailing PCM streams get their fourcc corrupted so no decoder exists + (``codec_context is None``), like the APAC track in iPhone spatial-audio recordings. + ``rotation`` patches a 90-degree display matrix into the video track header. + """ + buffer = io.BytesIO() + with av.open(buffer, mode="w", format=container_format) as container: + video_stream = container.add_stream("mpeg4", rate=fps) + video_stream.width = width + video_stream.height = height + video_stream.pix_fmt = "yuv420p" + audio = [] + for _ in range(audio_streams + undecodable_audio): + stream = container.add_stream(audio_codec, rate=44100) + stream.sample_rate = 44100 + audio.append(stream) + + for i in range(frames): + frame = av.VideoFrame.from_ndarray( + torch.full((height, width, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), + format="rgb24", + ) + container.mux(video_stream.encode(frame.reformat(format="yuv420p"))) + # write audio in 1024-sample frames, like real decoders produce, so the + # per-frame skip/cap logic in the transcode path actually runs + for stream in audio: + for offset in range(0, 44100 * frames // fps, 1024): + n = min(1024, 44100 * frames // fps - offset) + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(1, n, dtype=torch.int16).numpy(), format="s16", layout="mono" + ) + audio_frame.sample_rate = 44100 + audio_frame.pts = offset + container.mux(stream.encode(audio_frame)) + for stream in [video_stream, *audio]: + container.mux(stream.encode(None)) + + data = bytearray(buffer.getvalue()) + end = len(data) + for _ in range(undecodable_audio): + end = data.rindex(b"sowt", 0, end) + data[end:end + 4] = b"Xpac" + if rotation: + # the 3x3 display matrix sits 40 bytes into the version-0 tkhd payload; first tkhd + # inside moov = video track (search from moov so mdat bytes can't false-match) + matrix_offset = data.index(b"tkhd", data.rindex(b"moov")) + 4 + 40 + values = [0, 1 << 16, 0, -(1 << 16), 0, 0, 0, 0, 1 << 30] + data[matrix_offset:matrix_offset + 36] = b"".join(v.to_bytes(4, "big", signed=True) for v in values) + + tmp = tempfile.NamedTemporaryFile(suffix=f".{container_format}", delete=False) + tmp.write(bytes(data)) + tmp.close() + return tmp.name + + +def transcode_and_probe(video): + buffer = io.BytesIO() + video.save_to(buffer, format=VideoContainer.MP4, codec=VideoCodec.H264) + buffer.seek(0) + with av.open(buffer) as container: + video_stream = container.streams.video[0] + audio_stream = container.streams.audio[0] if container.streams.audio else None + frames = 0 + first_pts = None + for packet in container.demux(video_stream): + for frame in packet.decode(): + if first_pts is None: + first_pts = frame.pts + frames += 1 + return { + "codec": video_stream.codec_context.name, + "width": video_stream.codec_context.width, + "height": video_stream.codec_context.height, + "frames": frames, + "first_pts": first_pts, + "video_seconds": float(video_stream.duration * video_stream.time_base) if video_stream.duration else None, + "audio_seconds": float(audio_stream.duration * audio_stream.time_base) + if audio_stream and audio_stream.duration else None, + "audio_codecs": [s.codec_context.name for s in container.streams.audio], + } + + +def test_save_to_transcode_streams_without_buffering_frames(): + """Transcoding must not decode the whole video into memory first (~2 GiB for this source)""" + resource = pytest.importorskip("resource") # no getrusage on Windows + rss_scale = 1 if sys.platform == "darwin" else 1024 # ru_maxrss: bytes on macOS, KiB elsewhere + # ru_maxrss is a lifetime peak: a heavier test running earlier would shrink the measured + # delta and quietly defang this canary, so keep this source the biggest thing in the suite + file_path = create_transcode_source(width=640, height=480, frames=300) + try: + rss_before = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale + result = transcode_and_probe(VideoFromFile(file_path)) + rss_delta = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale - rss_before + + assert result["codec"] == "h264" + assert result["frames"] == 300 + assert rss_delta < 500 * 2**20, f"transcode buffered frames in RAM (peak grew {rss_delta / 2**20:.0f} MiB)" + finally: + os.unlink(file_path) + + +def test_save_to_transcode_honors_trim_window(): + """start_time/duration trim applies to both video and audio on the streaming path""" + file_path = create_transcode_source(frames=90) # 3s @ 30fps + try: + result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1)) + assert result["frames"] == pytest.approx(30, abs=2) + assert result["first_pts"] == 0 # trimmed output is rebased to start at zero + assert result["video_seconds"] == pytest.approx(1.0, abs=0.1) + assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1) + finally: + os.unlink(file_path) + + +def test_save_to_transcode_keeps_audio_of_sparse_video(): + """Audio that runs ahead of a sparse video track (slideshows, timelapses) must be + kept in full — it is only clamped to the video's end, never to the video cursor.""" + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4") as container: + video_stream = container.add_stream("mpeg4", rate=30) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + audio_stream = container.add_stream("aac", rate=48000, layout="stereo") + for t in (0, 30, 60): # 3 frames spread over 60 seconds + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), t * 4, dtype=torch.uint8).numpy(), format="rgb24" + ).reformat(format="yuv420p") + frame.pts = t * 15360 + frame.time_base = Fraction(1, 15360) + container.mux(video_stream.encode(frame)) + container.mux(video_stream.encode(None)) + for offset in range(0, 48000 * 60, 1024): + n = min(1024, 48000 * 60 - offset) + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo" + ) + audio_frame.sample_rate = 48000 + audio_frame.pts = offset + audio_frame.time_base = Fraction(1, 48000) + container.mux(audio_stream.encode(audio_frame)) + container.mux(audio_stream.encode(None)) + + buffer.seek(0) + result = transcode_and_probe(VideoFromFile(buffer)) + assert result["audio_seconds"] == pytest.approx(60.0, abs=1.0) + + +def test_save_to_transcode_vfr_audio_covers_video_span(): + """A trim window in the sparse region of a VFR file keeps audio for the true pts span + of the kept frames. Deriving the span as frames/average_rate undercuts it badly: the + average is dominated by the dense region (and can be plain wrong on MediaRecorder files).""" + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4") as container: + video_stream = container.add_stream("mpeg4", rate=30) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + audio_stream = container.add_stream("aac", rate=48000, layout="stereo") + # 10 frames inside the first second, then one every 1.25 s + for i, t in enumerate([x / 10 for x in range(10)] + [1.0, 2.25, 3.5, 4.75]): + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), (i * 16) % 256, dtype=torch.uint8).numpy(), format="rgb24" + ).reformat(format="yuv420p") + frame.pts = int(t * 15360) + frame.time_base = Fraction(1, 15360) + container.mux(video_stream.encode(frame)) + container.mux(video_stream.encode(None)) + for offset in range(0, 48000 * 6, 1024): + n = min(1024, 48000 * 6 - offset) + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo" + ) + audio_frame.sample_rate = 48000 + audio_frame.pts = offset + audio_frame.time_base = Fraction(1, 48000) + container.mux(audio_stream.encode(audio_frame)) + container.mux(audio_stream.encode(None)) + + buffer.seek(0) + result = transcode_and_probe(VideoFromFile(buffer, start_time=1, duration=5)) + # kept frames: 1.0/2.25/3.5/4.75 s -> rebased span 3.75 s + one nominal interval + assert result["frames"] == 4 + assert result["audio_seconds"] == pytest.approx(4.0, abs=0.45) + + +def test_save_to_transcode_trims_audio_in_stream_time_base_units(): + """Matroska audio timestamps tick in 1/1000, not 1/sample_rate; trim and audio timing + must convert through the frame's time base instead of assuming sample units. AAC audio, + because it decodes straight to the encoder's format and hits the resampler passthrough + that keeps the source time base on the frames.""" + file_path = create_transcode_source(frames=90, container_format="matroska", audio_codec="aac") + try: + result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1)) + assert result["audio_codecs"] == ["aac"] + assert result["video_seconds"] == pytest.approx(1.0, abs=0.1) + assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1) + finally: + os.unlink(file_path) + + +def test_save_to_transcode_learns_unprobed_audio_params(): + """mpegts is only probed a few seconds deep at open, so an audio stream whose first + packet comes later (live captures where audio kicks in late) still has sample_rate 0 + when the transcode starts; the parameters must be learned from the stream itself.""" + sample_rate, fps, video_seconds, audio_start = 48000, 30, 13, 12 + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mpegts") as container: + video_stream = container.add_stream("mpeg4", rate=fps) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono") + for i in range(video_seconds * fps): + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24" + ) + container.mux(video_stream.encode(frame.reformat(format="yuv420p"))) + for offset in range(0, (video_seconds - audio_start) * sample_rate, 1024): + n = min(1024, (video_seconds - audio_start) * sample_rate - offset) + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono" + ) + audio_frame.sample_rate = sample_rate + audio_frame.pts = audio_start * sample_rate + offset + container.mux(audio_stream.encode(audio_frame)) + for stream in (video_stream, audio_stream): + container.mux(stream.encode(None)) + + buffer.seek(0) + with av.open(buffer) as container: + # the scenario requires unprobed parameters; if a future FFmpeg probes deeper, + # push audio_start/video_seconds further out to restore it + assert container.streams.audio[0].codec_context.sample_rate == 0 + result = transcode_and_probe(VideoFromFile(buffer)) + assert result["frames"] == video_seconds * fps + assert result["audio_codecs"] == ["aac"] + assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1) + + buffer.seek(0) + trimmed_before_audio = transcode_and_probe(VideoFromFile(buffer, duration=1)) + assert trimmed_before_audio["frames"] == fps + assert trimmed_before_audio["audio_codecs"] == [] + assert trimmed_before_audio["audio_seconds"] is None + + buffer.seek(0) + trimmed_crossing_audio = transcode_and_probe(VideoFromFile(buffer, start_time=11.5, duration=1)) + assert trimmed_crossing_audio["frames"] == fps + assert trimmed_crossing_audio["audio_codecs"] == ["aac"] + assert trimmed_crossing_audio["video_seconds"] == pytest.approx(1.0, abs=0.05) + assert trimmed_crossing_audio["audio_seconds"] == pytest.approx(0.5, abs=0.1) + + +def test_save_to_transcode_trimmed_fragmented_mp4_keeps_audio(): + """Fragmented mp4 (MediaRecorder, DASH/HLS-derived files) delivers audio well behind + video, so when the trim window's last video frame arrives the audio demuxed so far + does not cover the window yet; the transcode must keep demuxing audio until it does + instead of finalizing on the first audio frame it sees afterwards.""" + sample_rate, fps, seconds = 48000, 30, 6 + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4", options={"movflags": "frag_keyframe+empty_moov"}) as container: + video_stream = container.add_stream("h264", rate=fps) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono") + next_audio_pts = 0 + for i in range(seconds * fps): + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24" + ) + container.mux(video_stream.encode(frame.reformat(format="yuv420p"))) + while next_audio_pts / sample_rate <= i / fps: # feed audio alongside, like a live pipeline + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(1, 1024, dtype=torch.float32).numpy(), format="fltp", layout="mono" + ) + audio_frame.sample_rate = sample_rate + audio_frame.pts = next_audio_pts + container.mux(audio_stream.encode(audio_frame)) + next_audio_pts += 1024 + for stream in (video_stream, audio_stream): + container.mux(stream.encode(None)) + + result = transcode_and_probe(VideoFromFile(buffer, start_time=0.5, duration=1.0)) + assert result["video_seconds"] == pytest.approx(1.0, abs=0.05) + assert result["audio_seconds"] == pytest.approx(1.0, abs=0.05) + + +def test_save_to_transcode_sparse_video_keeps_true_duration(): + """average_rate is not a frame duration: a 3-frame video spanning 60 s averages + 0.05 fps, and padding the last frame with 1/average_rate used to extend the + output — and the audio kept with it — about 20 s past the source span.""" + sample_rate = 48000 + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4") as container: + video_stream = container.add_stream("mpeg4", rate=30) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono") + for i, second in enumerate((0, 30, 60)): + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), i * 80, dtype=torch.uint8).numpy(), format="rgb24" + ).reformat(format="yuv420p") + frame.pts = second * 30 + frame.time_base = Fraction(1, 30) + container.mux(video_stream.encode(frame)) + for offset in range(0, 90 * sample_rate, 1024): + n = min(1024, 90 * sample_rate - offset) + audio_frame = av.AudioFrame.from_ndarray( + torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono" + ) + audio_frame.sample_rate = sample_rate + audio_frame.pts = offset + container.mux(audio_stream.encode(audio_frame)) + for stream in (video_stream, audio_stream): + container.mux(stream.encode(None)) + + result = transcode_and_probe(VideoFromFile(buffer)) + assert result["frames"] == 3 + # the last frame keeps its true stts duration (1/30 s), not 1/average_rate (~20 s) + assert result["video_seconds"] == pytest.approx(60.03, abs=0.05) + assert result["audio_seconds"] == pytest.approx(60.03, abs=0.1) + + trimmed = transcode_and_probe(VideoFromFile(buffer, duration=45)) + assert trimmed["frames"] == 2 + # a kept frame whose source duration crosses the window end is clamped to it + assert trimmed["video_seconds"] == pytest.approx(45.0, abs=0.05) + assert trimmed["audio_seconds"] == pytest.approx(45.0, abs=0.1) + + +def test_save_to_transcode_clamps_final_pts_to_declared_stream_duration(): + """Some iPhone MOVs report a video stream duration that ends before the final + decoded frame's nominal duration. A transcode must not turn that trailing + timestamp quirk into an extra frame interval compared to the source/remux path.""" + fps = 30 + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4") as container: + video_stream = container.add_stream("mpeg4", rate=fps) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + for i, pts in enumerate([*range(31), 32]): + frame = av.VideoFrame.from_ndarray( + torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24" + ).reformat(format="yuv420p") + frame.pts = pts + frame.time_base = Fraction(1, fps) + container.mux(video_stream.encode(frame)) + container.mux(video_stream.encode(None)) + + class _StreamProxy: + def __init__(self, stream, duration): + self._stream = stream + self.duration = duration + + def __getattr__(self, name): + return getattr(self._stream, name) + + class _StreamsProxy: + def __init__(self, video_stream): + self.video = [video_stream] + self.audio = [] + + class _PacketProxy: + def __init__(self, packet, stream): + self._packet = packet + self.stream = stream + + def __getattr__(self, name): + return getattr(self._packet, name) + + class _ContainerProxy: + def __init__(self, container, stream): + self._container = container + self._stream = stream + self.streams = _StreamsProxy(stream) + + def __getattr__(self, name): + return getattr(self._container, name) + + def demux(self, *streams): + for packet in self._container.demux(self._stream._stream): + yield _PacketProxy(packet, self._stream) + + buffer.seek(0) + output = io.BytesIO() + with av.open(buffer) as container: + real_stream = container.streams.video[0] + declared_duration = 32 * int(round((1 / fps) / real_stream.time_base)) + stream = _StreamProxy(real_stream, declared_duration) + VideoFromFile(buffer)._save_transcoded( + _ContainerProxy(container, stream), output, VideoContainer.MP4, VideoCodec.H264, None, 8 + ) + + output.seek(0) + with av.open(output) as container: + video_stream = container.streams.video[0] + frames = [f for p in container.demux(video_stream) for f in p.decode()] + assert len(frames) == 32 + assert float(video_stream.duration * video_stream.time_base) == pytest.approx(32 / fps, abs=0.01) + assert float(frames[-1].pts * frames[-1].time_base) == pytest.approx(31 / fps, abs=0.01) + + +def test_save_to_transcode_irregular_vfr_keeps_span(): + """B-frames reorder packets, and mp4 sample durations follow decode order: the dts + timeline ends before the pts timeline, so an irregular-VFR source's tail holds fell + out of the container (this 20.23 s span used to come out as 15.27 s, and the 10 s + trim as 6.03 s). The transcode encodes without B-frames so every sample keeps its + true display duration.""" + durations = [1, 1, 60, 1, 1, 120, 1, 180, 1, 1, 150, 90] # 1/30 s ticks, span 20.2333 s + generator = torch.Generator().manual_seed(7) + buffer = io.BytesIO() + with av.open(buffer, mode="w", format="mp4") as container: + video_stream = container.add_stream("mpeg4", rate=30) + video_stream.width = video_stream.height = 64 + video_stream.pix_fmt = "yuv420p" + pts = 0 + for duration in durations: + # textured frames, so an encoder with default settings has B-frames to gain from + frame = av.VideoFrame.from_ndarray( + torch.randint(0, 255, (64, 64, 3), generator=generator, dtype=torch.uint8).numpy(), + format="rgb24", + ).reformat(format="yuv420p") + frame.pts = pts + frame.time_base = Fraction(1, 30) + pts += duration + for packet in video_stream.encode(frame): + packet.duration = duration # exact stts in the source + container.mux(packet) + container.mux(video_stream.encode(None)) + + result = transcode_and_probe(VideoFromFile(buffer)) + assert result["frames"] == len(durations) + assert result["video_seconds"] == pytest.approx(sum(durations) / 30, abs=0.05) + + trimmed = transcode_and_probe(VideoFromFile(buffer, duration=10)) + assert trimmed["frames"] == 8 # frames at 12.167 s+ fall outside the window + assert trimmed["video_seconds"] == pytest.approx(10.0, abs=0.05) + + +def test_save_to_transcode_trim_survives_missing_leading_pts(): + """A trim should survive pts-less kept frames followed by a real-pts frame past the window.""" + nulled_frames = 0 + + class _PacketProxy: + def __init__(self, packet): + self._packet = packet + + def __getattr__(self, name): + return getattr(self._packet, name) + + @property + def stream(self): + return self._packet.stream + + def decode(self): + nonlocal nulled_frames + frames = self._packet.decode() + for frame in frames: + if nulled_frames < 2: + frame.pts = None + nulled_frames += 1 + return frames + + class _ContainerProxy: + def __init__(self, real): + self._real = real + + def __getattr__(self, name): + return getattr(self._real, name) + + def demux(self, *streams): + for packet in self._real.demux(*streams): + yield _PacketProxy(packet) + + file_path = create_transcode_source(frames=10, audio_streams=0) + try: + buffer = io.BytesIO() + with av.open(file_path) as container: + # 0.05 s window: both pts-less frames are kept (synthesized pts 0 and 512), + # and the first real-pts frame (1024 ticks) already lies past end_pts (768) + VideoFromFile(file_path, duration=0.05)._save_transcoded( + _ContainerProxy(container), buffer, VideoContainer.MP4, VideoCodec.H264, None, 8 + ) + assert nulled_frames == 2 + buffer.seek(0) + with av.open(buffer) as container: + video_stream = container.streams.video[0] + frames = [f for p in container.demux(video_stream) for f in p.decode()] + assert len(frames) == 2 + assert float(video_stream.duration * video_stream.time_base) == pytest.approx(2 / 30, abs=0.01) + finally: + os.unlink(file_path) + + +def test_save_to_transcode_bakes_rotation(): + """A 90-degree display-matrix rotation swaps the output dimensions (portrait video)""" + file_path = create_transcode_source(width=64, height=32, rotation=True) + try: + result = transcode_and_probe(VideoFromFile(file_path)) + assert (result["width"], result["height"]) == (32, 64) + assert result["frames"] == 30 + finally: + os.unlink(file_path) + + +def test_save_to_transcode_skips_undecodable_audio(): + """Streaming transcode keeps the decodable audio track and drops undecodable ones; + with no decodable audio at all the output is video-only instead of crashing.""" + mixed = all_bad = None + try: + mixed = create_transcode_source(audio_streams=1, undecodable_audio=1) + all_bad = create_transcode_source(audio_streams=0, undecodable_audio=2) + result = transcode_and_probe(VideoFromFile(mixed)) + assert result["audio_codecs"] == ["aac"] + assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1) + assert transcode_and_probe(VideoFromFile(all_bad))["audio_codecs"] == [] + finally: + for path in (mixed, all_bad): + if path: + os.unlink(path) From 87d23b81765161624889febfb3b81f19f3c8435b Mon Sep 17 00:00:00 2001 From: Alexis Rolland Date: Wed, 15 Jul 2026 16:38:03 +0800 Subject: [PATCH 113/211] [Partner Nodes] feat(client): send ComfyUI Job Id in request headers (#14934) --- comfy_api_nodes/util/_helpers.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/comfy_api_nodes/util/_helpers.py b/comfy_api_nodes/util/_helpers.py index acab10d95..ddfb3b65c 100644 --- a/comfy_api_nodes/util/_helpers.py +++ b/comfy_api_nodes/util/_helpers.py @@ -15,6 +15,7 @@ from comfy.comfy_api_env import normalize_comfy_api_base from comfy.deploy_environment import get_deploy_environment from comfy.model_management import processing_interrupted from comfy_api.latest import IO +from comfy_execution.utils import get_executing_context from comfyui_version import __version__ as comfyui_version from .common_exceptions import ProcessingInterrupted @@ -57,12 +58,16 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]: relative/cloud URLs resolved against ``default_base_url()``; because the result includes auth, callers must not attach it to arbitrary absolute/presigned URLs. """ - return { + headers = { **get_auth_header(node_cls), "Comfy-Env": get_deploy_environment(), "Comfy-Usage-Source": get_usage_source(node_cls), "Comfy-Core-Version": comfyui_version, } + ctx = get_executing_context() + if ctx is not None: + headers["Comfy-Job-Id"] = ctx.prompt_id + return headers def default_base_url() -> str: From 678d42c90e55d05bcb17b3ce19a4e5e765ac53f9 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 15 Jul 2026 20:09:59 -0700 Subject: [PATCH 114/211] Update AGENTS.md (#14955) --- AGENTS.md | 41 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 41 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 05efd834b..20014ce7e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -19,6 +19,9 @@ better to remove a broken feature path than keep a complicated partial fix. - Preserve existing APIs, node names, model-loading behavior, file layout, and workflow compatibility unless the change is explicitly about replacing them. +- When compatibility is explicitly out of scope, remove compatibility-only + aliases, duplicate nodes, legacy entry points, and preset wrappers instead of + retaining parallel ways to perform the same operation. - Code must look hand-written for this repository. Changes that read like generic AI-generated code will be rejected automatically: unnecessary helper layers, vague names, boilerplate comments, defensive branches without a real @@ -96,6 +99,13 @@ unless they are read by current code and change current behavior. Remove pass-through or stored-but-unused values instead of preserving upstream or deprecated API baggage. +- Do not add a model-specific option to a shared helper when only one caller + needs it. Keep one-off behavior at the model integration boundary, or extend + the shared helper only when the option is a coherent reusable capability. +- Implementations of shared model interfaces should accept the standard caller + contract without model-specific rejection branches for optional capabilities + they do not consume. Let supported behavior be determined by implementation + paths that actually use those inputs. - If an implementation needs auxiliary values for its own workflow, expose them through a private helper or a clearly named implementation-specific method instead of overloading the public method's return contract. @@ -154,6 +164,10 @@ `comfy-kitchen` helpers where they already solve the problem. - Use optimized comfy-kitchen ops in places where they improve performance without changing the expected dtype, device, memory, or interface behavior. +- Prefer ComfyUI's shared optimized kernels and backend dispatchers over + handwritten implementations of the same operation. Remove duplicate local + kernels and adapt inputs to the shared operation's documented layout while + preserving the model's original math and output contract. - All models should use the optimized attention function selected by ComfyUI. Treat optimized backend functions, dispatch helpers, and capability-selected callables as opaque. Higher-level code must not inspect function identity, @@ -176,6 +190,12 @@ - Model detection code that inspects linear weight shapes should only use the first dimension. The second dimension may be half the original size for NVFP4 or other 4-bit quantized models. +- A model-detection signature must guard every state-dict key it dereferences. + Do not partially match a format and then raise an incidental `KeyError` while + extracting its configuration. +- Order model-detection checks from established or more-specific signatures to + newer or broader signatures. Put a broad new detector near the generic + fallback when giving it higher precedence could steal another model family. - Avoid adding `einops` usage in core inference code. Use native torch tensor ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`, `unsqueeze`, and `squeeze` instead. @@ -192,11 +212,23 @@ methods for scalar or structural calculations. - Avoid unnecessary casts and transfers. Preserve the intended compute dtype, storage dtype, bias dtype, and original tensor shape metadata. +- Do not cast the result of an optimized backend operation back to its input + dtype unless that backend's documented result contract requires normalization. + In particular, trust the selected optimized-attention implementation to honor + its dtype contract. - Keep model-native latent layout handling inside the model or latent-format owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent dimensions in nodes or other caller-side adapters just to satisfy a model forward; the model path should consume and return the native latent shape for that model family. +- DiT models should accept latent dimensions that are not exact patch-size + multiples. Use `comfy.ldm.common_dit.pad_to_patch_size` on every patchified + target or reference input, then crop only the target output back to its + original dimensions. +- Avoid defensive shape and configuration checks that merely replace the clear + failure from the tensor operation immediately below them. Add explicit + validation only when it provides materially better context at a real boundary + or prevents silent incorrect output. - Assume inputs to the main model forward are already in the compute dtype by default, except integer inputs such as some model timestep tensors. Do not add defensive or convenience casts in model code; it is better for invalid dtype @@ -260,6 +292,15 @@ - Model implementations should add the minimal number of ComfyUI nodes required to run the model. Reuse existing nodes as much as possible; adapting the model to work with existing nodes is strongly preferred over creating new nodes. +- Use `io.Autogrow` for a variable number of repeated inputs instead of a fixed + series of numbered optional sockets. Set its minimum to zero when the model + has a valid no-item path, and cap it only when the model has a real limit. +- Mark inputs optional when execution has a valid path that does not read them. + If one optional input is needed only to process another optional input, do not + force users on the path that supplies neither to connect it. +- Conditioning nodes should normally output conditioning only. Do not expose + input or intermediate images as convenience outputs for downstream sizing or + routing; use the existing image path or a dedicated image operation instead. - Nodes should output only values they own. Do not add pass-through outputs for workflow convenience unless the node is explicitly an output node. Existing models, latents, conditioning, or other inputs should flow directly to the From 03978e1e81475f19eebd7edc065cc55cb4e15e10 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BD=BC=E5=BD=BC?= Date: Thu, 16 Jul 2026 11:48:28 +0800 Subject: [PATCH 115/211] [feat]Add JoyImageEdit native model support (#14428) --- comfy/ldm/joyimage/model.py | 445 ++++++++++++++++++ comfy/model_base.py | 23 + comfy/model_detection.py | 19 + comfy/sd.py | 6 + comfy/supported_models.py | 34 ++ comfy/text_encoders/joyimage.py | 97 ++++ comfy/text_encoders/qwen_vl.py | 4 +- comfy_extras/nodes_joyimage.py | 102 ++++ nodes.py | 5 +- tests-unit/comfy_test/model_detection_test.py | 31 ++ 10 files changed, 762 insertions(+), 4 deletions(-) create mode 100644 comfy/ldm/joyimage/model.py create mode 100644 comfy/text_encoders/joyimage.py create mode 100644 comfy_extras/nodes_joyimage.py diff --git a/comfy/ldm/joyimage/model.py b/comfy/ldm/joyimage/model.py new file mode 100644 index 000000000..bca12c391 --- /dev/null +++ b/comfy/ldm/joyimage/model.py @@ -0,0 +1,445 @@ +# https://github.com/jdopensource/JoyAI-Image-Edit (Apache 2.0) +import math +from typing import Optional, Tuple + +import comfy_kitchen +import torch +import torch.nn as nn + +import comfy.ldm.common_dit +import comfy.ops +import comfy.patcher_extension +from comfy.ldm.lightricks.model import GELU_approx, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps +from comfy.ldm.modules.attention import optimized_attention + + +class JoyImageModulate(nn.Module): + def __init__(self, hidden_size: int, factor: int, dtype=None, device=None): + super().__init__() + self.factor = factor + self.modulate_table = nn.Parameter( + torch.empty(1, factor, hidden_size, dtype=dtype, device=device) + ) + + def forward(self, x: torch.Tensor) -> list: + if x.ndim != 3: + x = x.unsqueeze(1) + table = comfy.ops.cast_to_input(self.modulate_table, x) + return [o.squeeze(1) for o in (table + x).chunk(self.factor, dim=1)] + + +class JoyImageFeedForward(nn.Module): + def __init__( + self, + dim: int, + inner_dim: int, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.net = nn.ModuleList([ + GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations), + nn.Identity(), + operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device), + ]) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + for module in self.net: + x = module(x) + return x + + +class JoyImageAttention(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + eps: float = 1e-6, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.num_attention_heads = num_attention_heads + inner_dim = num_attention_heads * attention_head_dim + + self.img_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device) + self.img_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.img_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.img_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device) + + self.txt_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device) + self.txt_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.txt_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.txt_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device) + + def forward( + self, + img: torch.Tensor, + txt: torch.Tensor, + image_rotary_emb: torch.Tensor, + transformer_options=None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + heads = self.num_attention_heads + + img_q, img_k, img_v = self.img_attn_qkv(img).chunk(3, dim=-1) + txt_q, txt_k, txt_v = self.txt_attn_qkv(txt).chunk(3, dim=-1) + + img_q = img_q.unflatten(-1, (heads, -1)) + img_k = img_k.unflatten(-1, (heads, -1)) + img_v = img_v.unflatten(-1, (heads, -1)) + txt_q = txt_q.unflatten(-1, (heads, -1)) + txt_k = txt_k.unflatten(-1, (heads, -1)) + txt_v = txt_v.unflatten(-1, (heads, -1)) + + img_q = self.img_attn_q_norm(img_q) + img_k = self.img_attn_k_norm(img_k) + txt_q = self.txt_attn_q_norm(txt_q) + txt_k = self.txt_attn_k_norm(txt_k) + + img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb) + + joint_q = torch.cat([img_q, txt_q], dim=1) + joint_k = torch.cat([img_k, txt_k], dim=1) + joint_v = torch.cat([img_v, txt_v], dim=1) + + joint_q = joint_q.flatten(2, 3) + joint_k = joint_k.flatten(2, 3) + joint_v = joint_v.flatten(2, 3) + + joint_out = optimized_attention(joint_q, joint_k, joint_v, heads=heads, transformer_options=transformer_options) + + seq_img = img.shape[1] + img_out = joint_out[:, :seq_img, :] + txt_out = joint_out[:, seq_img:, :] + + img_out = self.img_attn_proj(img_out) + txt_out = self.txt_attn_proj(txt_out) + return img_out, txt_out + + +class JoyImageTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + mlp_width_ratio: float = 4.0, + eps: float = 1e-6, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + mlp_hidden_dim = int(dim * mlp_width_ratio) + + self.img_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device) + self.img_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.img_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.img_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations) + + self.txt_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device) + self.txt_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.txt_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.txt_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations) + + self.attn = JoyImageAttention( + dim=dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + eps=eps, + dtype=dtype, + device=device, + operations=operations, + ) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + image_rotary_emb: torch.Tensor, + transformer_options=None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + ( + img_mod1_shift, + img_mod1_scale, + img_mod1_gate, + img_mod2_shift, + img_mod2_scale, + img_mod2_gate, + ) = self.img_mod(temb) + ( + txt_mod1_shift, + txt_mod1_scale, + txt_mod1_gate, + txt_mod2_shift, + txt_mod2_scale, + txt_mod2_gate, + ) = self.txt_mod(temb) + + img_normed = self.img_norm1(hidden_states) + txt_normed = self.txt_norm1(encoder_hidden_states) + img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1) + txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1) + + img_attn, txt_attn = self.attn(img_modulated, txt_modulated, image_rotary_emb, transformer_options=transformer_options) + + hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1) + encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1) + + img_ffn_normed = self.img_norm2(hidden_states) + txt_ffn_normed = self.txt_norm2(encoder_hidden_states) + img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1) + txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1) + hidden_states = hidden_states + self.img_mlp(img_ffn_input) * img_mod2_gate.unsqueeze(1) + encoder_hidden_states = encoder_hidden_states + self.txt_mlp(txt_ffn_input) * txt_mod2_gate.unsqueeze(1) + + return hidden_states, encoder_hidden_states + + +class JoyImageTimeTextImageEmbedding(nn.Module): + def __init__( + self, + dim: int, + time_freq_dim: int, + time_proj_dim: int, + text_embed_dim: int, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) + self.time_embedder = TimestepEmbedding( + in_channels=time_freq_dim, + time_embed_dim=dim, + dtype=dtype, + device=device, + operations=operations, + ) + self.act_fn = nn.SiLU() + self.time_proj = operations.Linear(dim, time_proj_dim, bias=True, dtype=dtype, device=device) + self.text_embedder = PixArtAlphaTextProjection( + text_embed_dim, dim, act_fn="gelu_tanh", dtype=dtype, device=device, operations=operations, + ) + + def forward(self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor): + timestep = self.timesteps_proj(timestep) + temb = self.time_embedder(timestep.to(dtype=encoder_hidden_states.dtype)).type_as(encoder_hidden_states) + timestep_proj = self.time_proj(self.act_fn(temb)) + encoder_hidden_states = self.text_embedder(encoder_hidden_states) + return temb, timestep_proj, encoder_hidden_states + + +class JoyImageTransformer3DModel(nn.Module): + def __init__( + self, + patch_size: list = [1, 2, 2], + in_channels: int = 16, + out_channels: Optional[int] = None, + hidden_size: int = 3072, + num_attention_heads: int = 24, + text_dim: int = 4096, + mlp_width_ratio: float = 4.0, + num_layers: int = 20, + rope_dim_list: list = [16, 56, 56], + theta: int = 256, + image_model=None, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.dtype = dtype + self.out_channels = out_channels or in_channels + self.patch_size = list(patch_size) + self.rope_dim_list = list(rope_dim_list) + self.theta = theta + + attention_head_dim = hidden_size // num_attention_heads + + self.img_in = operations.Conv3d( + in_channels, + hidden_size, + kernel_size=tuple(self.patch_size), + stride=tuple(self.patch_size), + dtype=dtype, + device=device, + ) + + self.condition_embedder = JoyImageTimeTextImageEmbedding( + dim=hidden_size, + time_freq_dim=256, + time_proj_dim=hidden_size * 6, + text_embed_dim=text_dim, + dtype=dtype, + device=device, + operations=operations, + ) + + self.double_blocks = nn.ModuleList([ + JoyImageTransformerBlock( + dim=hidden_size, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + mlp_width_ratio=mlp_width_ratio, + dtype=dtype, + device=device, + operations=operations, + ) + for _ in range(num_layers) + ]) + + self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.proj_out = operations.Linear( + hidden_size, + self.out_channels * math.prod(self.patch_size), + bias=True, + dtype=dtype, + device=device, + ) + + def _get_rotary_pos_embed_for_range( + self, + start: Tuple[int, int, int], + stop: Tuple[int, int, int], + device=None, + ) -> torch.Tensor: + # 3D RoPE for the patch grid range [start, stop) over (t, h, w). Token order after + # reshape(-1) is (t, h, w), matching the img_in Conv3d flatten. + rope_dim_list = self.rope_dim_list + + grids = [torch.arange(start[i], stop[i], dtype=torch.float32, device=device) for i in range(3)] + mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0) + + angles_parts = [] + for i, dim in enumerate(rope_dim_list): + pos = mesh[i].reshape(-1) + freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device)[: (dim // 2)] / dim)) + angles_parts.append(torch.outer(pos, freqs)) + + angles = torch.cat(angles_parts, dim=1) + cos = angles.cos() + sin = angles.sin() + return torch.stack((cos, -sin, sin, cos), dim=-1).unflatten(-1, (2, 2)) + + def get_rotary_pos_embed_for_components( + self, + component_sizes, + device=None, + ) -> torch.Tensor: + # Per-component 3D RoPE. component_sizes is a list of (t, h, w) patch grid sizes in + # sequence order [target, ref0, ref1, ...]; h/w restart at 0 for each component while t + # continues from the running offset, giving every image its own temporal position band. + freqs_parts = [] + t_offset = 0 + for (t, h, w) in component_sizes: + freqs = self._get_rotary_pos_embed_for_range( + start=(t_offset, 0, 0), + stop=(t_offset + t, h, w), + device=device, + ) + freqs_parts.append(freqs) + t_offset += t + return torch.cat(freqs_parts, dim=0).unsqueeze(0).unsqueeze(2) + + def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor: + c = self.out_channels + pt, ph, pw = self.patch_size + x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c) + x = x.permute(0, 7, 1, 4, 2, 5, 3, 6) + return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw) + + def forward( + self, + hidden_states: torch.Tensor, + timestep: torch.Tensor, + context: torch.Tensor = None, + ref_latents=None, + control=None, + transformer_options=None, + **kwargs, + ) -> torch.Tensor: + transformer_options = {} if transformer_options is None else transformer_options.copy() + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(hidden_states, timestep, context, ref_latents, transformer_options, **kwargs) + + def _forward( + self, + hidden_states: torch.Tensor, + timestep: torch.Tensor, + context: torch.Tensor, + ref_latents=None, + transformer_options=None, + **kwargs, + ) -> torch.Tensor: + pt, ph, pw = self.patch_size + _, _, ot, oh, ow = hidden_states.shape + + components = [hidden_states, *(ref_latents or [])] + component_sizes = [] + img_tokens = [] + for comp in components: + comp = comfy.ldm.common_dit.pad_to_patch_size(comp, self.patch_size) + _, _, ct, ch, cw = comp.shape + component_sizes.append((ct // pt, ch // ph, cw // pw)) + tokens = self.img_in(comp).flatten(2).transpose(1, 2) # (B, n_i, D) + img_tokens.append(tokens) + + img = torch.cat(img_tokens, dim=1) + + _, vec, txt = self.condition_embedder(timestep, context) + vec = vec.unflatten(1, (6, -1)) + + image_rotary_emb = self.get_rotary_pos_embed_for_components( + component_sizes, + device=hidden_states.device, + ) + + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + transformer_options["total_blocks"] = len(self.double_blocks) + transformer_options["block_type"] = "double" + for i, block in enumerate(self.double_blocks): + transformer_options["block_index"] = i + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"], out["txt"] = block( + hidden_states=args["img"], + encoder_hidden_states=args["txt"], + temb=args["vec"], + image_rotary_emb=args["pe"], + transformer_options=args.get("transformer_options"), + ) + return out + + out = blocks_replace[("double_block", i)]({"img": img, + "txt": txt, + "vec": vec, + "pe": image_rotary_emb, + "transformer_options": transformer_options}, + {"original_block": block_wrap}) + txt = out["txt"] + img = out["img"] + else: + img, txt = block( + hidden_states=img, + encoder_hidden_states=txt, + temb=vec, + image_rotary_emb=image_rotary_emb, + transformer_options=transformer_options, + ) + + tt, th, tw = component_sizes[0] + target_tokens = tt * th * tw + img = img[:, :target_tokens, :] + img = self.proj_out(self.norm_out(img)) + img = self.unpatchify(img, tt, th, tw) + return img[:, :, :ot, :oh, :ow] diff --git a/comfy/model_base.py b/comfy/model_base.py index 786a7c127..98f5ba48b 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -58,6 +58,7 @@ import comfy.ldm.omnigen.omnigen2 import comfy.ldm.seedvr.model import comfy.ldm.boogu.model import comfy.ldm.qwen_image.model +import comfy.ldm.joyimage.model import comfy.ldm.ideogram4.model import comfy.ldm.krea2.model import comfy.ldm.kandinsky5.model @@ -2276,6 +2277,28 @@ class QwenImage(BaseModel): out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) return out +class JoyImage(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel) + self.memory_usage_factor_conds = ("ref_latents",) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents]) + return out + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) + return out + class Ideogram4(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 70c8625e3..a1bf047f8 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -1058,6 +1058,25 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["image_model"] = "SAM31" return dit_config + if ( + '{}double_blocks.0.attn.img_attn_qkv.weight'.format(key_prefix) in state_dict_keys + and '{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix) in state_dict_keys + and '{}condition_embedder.time_embedder.linear_1.weight'.format(key_prefix) in state_dict_keys + and '{}img_in.weight'.format(key_prefix) in state_dict_keys + and len(state_dict['{}img_in.weight'.format(key_prefix)].shape) == 5 + ): + img_in = state_dict['{}img_in.weight'.format(key_prefix)] + head_dim = state_dict['{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix)].shape[0] + return { + "image_model": "joyimage", + "in_channels": img_in.shape[1], + "hidden_size": img_in.shape[0], + "patch_size": list(img_in.shape[2:]), + "num_layers": count_blocks(state_dict_keys, '{}double_blocks.'.format(key_prefix) + '{}.'), + "num_attention_heads": img_in.shape[0] // head_dim, + "text_dim": 4096, + } + if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys: return None diff --git a/comfy/sd.py b/comfy/sd.py index 4a0742e7a..9d7fa731f 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -76,6 +76,7 @@ import comfy.text_encoders.gemma4 import comfy.text_encoders.cogvideo import comfy.text_encoders.sa3 import comfy.text_encoders.gpt_oss +import comfy.text_encoders.joyimage import comfy.model_patcher import comfy.lora @@ -1377,6 +1378,7 @@ class CLIPType(Enum): IDEOGRAM4 = 30 BOOGU = 31 KREA2 = 32 + JOYIMAGE = 33 @@ -1706,6 +1708,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer + elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx. + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.joyimage.JoyImageTokenizer elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused. klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b" clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index b82e4178f..e7c8983aa 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -27,6 +27,7 @@ import comfy.text_encoders.z_image import comfy.text_encoders.ideogram4 import comfy.text_encoders.boogu import comfy.text_encoders.krea2 +import comfy.text_encoders.joyimage import comfy.text_encoders.anima import comfy.text_encoders.ace15 import comfy.text_encoders.longcat_image @@ -1911,6 +1912,38 @@ class QwenImage(supported_models_base.BASE): hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.qwen_image.QwenImageTokenizer, comfy.text_encoders.qwen_image.te(**hunyuan_detect)) +class JoyImage(supported_models_base.BASE): + unet_config = { + "image_model": "joyimage", + } + + sampling_settings = { + "multiplier": 1000, + "shift": 1.5, + } + + memory_usage_factor = 1.8 + + unet_extra_config = { + "theta": 10000, + "rope_dim_list": [16, 56, 56], + } + + latent_format = latent_formats.Wan21 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + return model_base.JoyImage(self, device=device) + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + qwen3vl_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.joyimage.JoyImageTokenizer, comfy.text_encoders.joyimage.te(**qwen3vl_detect)) + class HunyuanImage21(HunyuanVideo): unet_config = { "image_model": "hunyuan_video", @@ -2389,6 +2422,7 @@ models = [ Omnigen2, Boogu, QwenImage, + JoyImage, Ideogram4, Krea2, Flux2, diff --git a/comfy/text_encoders/joyimage.py b/comfy/text_encoders/joyimage.py new file mode 100644 index 000000000..143c44250 --- /dev/null +++ b/comfy/text_encoders/joyimage.py @@ -0,0 +1,97 @@ +import torch + +from comfy import sd1_clip +import comfy.text_encoders.qwen_vl +from comfy.text_encoders.qwen3vl import Qwen3VL, Qwen3VLTokenizer + +JOYIMAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>" +JOYIMAGE_TEMPLATE_TEXT = ( + "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" + "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" +) +JOYIMAGE_TEMPLATE_IMAGE = ( + "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" + f"<|im_start|>user\n{JOYIMAGE_VISION_BLOCK}{{}}<|im_end|>\n<|im_start|>assistant\n" +) +# The DiT was trained without the leading system-prompt tokens. +JOYIMAGE_DROP_IDX = 34 +PAD_TOKEN = 151643 + + +class Qwen3VL8B_JoyImage(Qwen3VL): + model_type = "qwen3vl_8b" + + def preprocess_embed(self, embed, device): + if embed["type"] == "image": + image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images( + embed["data"], min_pixels=65536, max_pixels=16777216, patch_size=16, + image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5], + interpolation="bicubic", + ) + merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid) + return merged, {"grid": grid, "deepstack": deepstack} + return None, None + + +class JoyImageTokenizer(Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__( + embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, + model_type="qwen3vl_8b", + ) + self.llama_template = JOYIMAGE_TEMPLATE_TEXT + self.llama_template_images = JOYIMAGE_TEMPLATE_IMAGE + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=None, **kwargs): + kwargs.pop("thinking", None) + return super().tokenize_with_weights( + text, return_word_ids=return_word_ids, llama_template=llama_template, + images=images or [], thinking=True, **kwargs, + ) + + +class _JoyImageClipModel(sd1_clip.SDClipModel): + def __init__(self, device="cpu", layer="hidden", layer_idx=-1, dtype=None, + attention_mask=True, model_options={}): + super().__init__( + device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, + # JoyImage conditions on the pre-final-norm output of the last decoder layer. + dtype=dtype, special_tokens={"pad": PAD_TOKEN}, layer_norm_hidden_state=False, + model_class=Qwen3VL8B_JoyImage, enable_attention_masks=attention_mask, + return_attention_masks=attention_mask, model_options=model_options, + ) + + +class JoyImageTEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + super().__init__( + device=device, dtype=dtype, name="qwen3vl_8b", + clip_model=_JoyImageClipModel, model_options=model_options, + ) + + def encode_token_weights(self, token_weight_pairs): + out, pooled, extra = super().encode_token_weights(token_weight_pairs) + if out.shape[1] <= JOYIMAGE_DROP_IDX: + raise ValueError( + f"JoyImageTEModel: encoded sequence length {out.shape[1]} is shorter " + f"than drop_idx={JOYIMAGE_DROP_IDX}; the prompt did not include the " + f"template prefix." + ) + out = out[:, JOYIMAGE_DROP_IDX:] + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, JOYIMAGE_DROP_IDX:] + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class JoyImageTEModel_(JoyImageTEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + if dtype_llama is not None: + dtype = dtype_llama + super().__init__(device=device, dtype=dtype, model_options=model_options) + return JoyImageTEModel_ diff --git a/comfy/text_encoders/qwen_vl.py b/comfy/text_encoders/qwen_vl.py index 924eb6ad8..f97a88061 100644 --- a/comfy/text_encoders/qwen_vl.py +++ b/comfy/text_encoders/qwen_vl.py @@ -15,6 +15,7 @@ def process_qwen2vl_images( merge_size: int = 2, image_mean: list = None, image_std: list = None, + interpolation: str = "bilinear", ): if image_mean is None: image_mean = [0.48145466, 0.4578275, 0.40821073] @@ -47,10 +48,9 @@ def process_qwen2vl_images( img_resized = F.interpolate( img.unsqueeze(0), size=(h_bar, w_bar), - mode='bilinear', + mode=interpolation, align_corners=False ).squeeze(0) - normalized = img_resized.clone() for c in range(3): normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c] diff --git a/comfy_extras/nodes_joyimage.py b/comfy_extras/nodes_joyimage.py new file mode 100644 index 000000000..539dc44b2 --- /dev/null +++ b/comfy_extras/nodes_joyimage.py @@ -0,0 +1,102 @@ +from typing_extensions import override + +import comfy.utils +import node_helpers +from comfy_api.latest import ComfyExtension, io + + +# fmt: off +BUCKETS_1024 = [ + (512, 1792), (512, 1856), (512, 1920), (512, 1984), (512, 2048), + (576, 1600), (576, 1664), (576, 1728), (576, 1792), + (640, 1472), (640, 1536), (640, 1600), + (704, 1344), (704, 1408), (704, 1472), + (768, 1216), (768, 1280), (768, 1344), + (832, 1152), (832, 1216), + (896, 1088), (896, 1152), + (960, 1024), (960, 1088), + (1024, 960), (1024, 1024), + (1088, 896), (1088, 960), + (1152, 832), (1152, 896), + (1216, 768), (1216, 832), + (1280, 768), + (1344, 704), (1344, 768), + (1408, 704), + (1472, 640), (1472, 704), + (1536, 640), + (1600, 576), (1600, 640), + (1664, 576), + (1728, 576), + (1792, 512), (1792, 576), + (1856, 512), + (1920, 512), + (1984, 512), + (2048, 512), +] +# fmt: on + + +def _find_best_bucket(height: int, width: int) -> tuple[int, int]: + target_ratio = height / width + return min(BUCKETS_1024, key=lambda hw: abs(hw[0] / hw[1] - target_ratio)) + + +def _resize_reference(image): + if image.shape[0] != 1: + raise ValueError("JoyImage reference inputs must contain one image each") + samples = image.movedim(-1, 1) + bucket_h, bucket_w = _find_best_bucket(samples.shape[2], samples.shape[3]) + resized = comfy.utils.common_upscale(samples, bucket_w, bucket_h, "bilinear", "center") + return resized.movedim(1, -1)[:, :, :, :3] + + +def _encode(clip, prompt, vae, images): + resized_images = [_resize_reference(image) for image in images] + conditioning = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=resized_images)) + if vae is not None and resized_images: + ref_latents = [vae.encode(image) for image in resized_images] + conditioning = node_helpers.conditioning_set_values( + conditioning, {"reference_latents": ref_latents}, append=True, + ) + return conditioning + + +class TextEncodeJoyImageEdit(io.ComfyNode): + @classmethod + def define_schema(cls): + image_template = io.Autogrow.TemplatePrefix( + io.Image.Input("image"), + prefix="image", + min=0, + max=6, + ) + return io.Schema( + node_id="TextEncodeJoyImageEdit", + category="model/conditioning/joyimage", + inputs=[ + io.Clip.Input("clip"), + io.String.Input("prompt", multiline=True, dynamic_prompts=True), + io.Vae.Input("vae", optional=True), + io.Autogrow.Input("images", template=image_template, optional=True), + ], + outputs=[ + io.Conditioning.Output(), + ], + ) + + @classmethod + def execute(cls, clip, prompt, vae=None, images: io.Autogrow.Type = None) -> io.NodeOutput: + images = images or {} + return io.NodeOutput(_encode(clip, prompt, vae, list(images.values()))) + + +class JoyImageExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + TextEncodeJoyImageEdit, + ] + + +async def comfy_entrypoint() -> JoyImageExtension: + return JoyImageExtension() diff --git a/nodes.py b/nodes.py index 883258bd1..b03d6c603 100644 --- a/nodes.py +++ b/nodes.py @@ -992,7 +992,7 @@ class CLIPLoader: @classmethod def INPUT_TYPES(s): return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ), - "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2"], ), + "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ), }, "optional": { "device": (["default", "cpu"], {"advanced": True}), @@ -1002,7 +1002,7 @@ class CLIPLoader: CATEGORY = "model/loaders" - DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm" + DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm" def load_clip(self, clip_name, type="stable_diffusion", device="default"): clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION) @@ -2462,6 +2462,7 @@ async def init_builtin_extra_nodes(): "nodes_seedvr.py", "nodes_context_windows.py", "nodes_qwen.py", + "nodes_joyimage.py", "nodes_boogu.py", "nodes_chroma_radiance.py", "nodes_pid.py", diff --git a/tests-unit/comfy_test/model_detection_test.py b/tests-unit/comfy_test/model_detection_test.py index 7c5b271c5..b40ea0d4c 100644 --- a/tests-unit/comfy_test/model_detection_test.py +++ b/tests-unit/comfy_test/model_detection_test.py @@ -112,6 +112,17 @@ def _make_pid_v1_5_sd(latent_proj_channels=16): return sd +def _make_joyimage_edit_plus_sd(): + sd = { + "img_in.weight": torch.empty(4096, 16, 1, 2, 2, device="meta"), + "condition_embedder.time_embedder.linear_1.weight": torch.empty(1, device="meta"), + "double_blocks.0.attn.img_attn_q_norm.weight": torch.empty(128, device="meta"), + } + for i in range(40): + sd[f"double_blocks.{i}.attn.img_attn_qkv.weight"] = torch.empty(1, device="meta") + return sd + + def _add_model_diffusion_prefix(sd): return {f"model.diffusion_model.{k}": v for k, v in sd.items()} @@ -258,6 +269,26 @@ class TestModelDetection: assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,) assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,) + def test_joyimage_edit_plus_detection(self): + sd = _make_joyimage_edit_plus_sd() + unet_config = detect_unet_config(sd, "") + + assert unet_config == { + "image_model": "joyimage", + "in_channels": 16, + "hidden_size": 4096, + "patch_size": [1, 2, 2], + "num_layers": 40, + "num_attention_heads": 32, + "text_dim": 4096, + } + assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "JoyImage" + + def test_incomplete_joyimage_signature_is_not_detected(self): + sd = _make_joyimage_edit_plus_sd() + del sd["double_blocks.0.attn.img_attn_q_norm.weight"] + assert detect_unet_config(sd, "") is None + def test_unet_config_and_required_keys_combination_is_unique(self): """Each model in the registry must have a unique combination of ``unet_config`` and ``required_keys``. If two models share the same From 285a98944c397a4a81f15ac63d69fa3dbc0a27b9 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 16 Jul 2026 15:35:07 +0300 Subject: [PATCH 116/211] [Partner Nodes] feat(OpenAI): add GPT5.6 models (#14957) Signed-off-by: bigcat88 --- comfy_api_nodes/apis/openai.py | 2 +- comfy_api_nodes/nodes_openai.py | 18 ++++++++++++++++++ 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/comfy_api_nodes/apis/openai.py b/comfy_api_nodes/apis/openai.py index bee75d639..827281788 100644 --- a/comfy_api_nodes/apis/openai.py +++ b/comfy_api_nodes/apis/openai.py @@ -128,7 +128,7 @@ class OpenAIResponse(ModelResponseProperties, ResponseProperties): parallel_tool_calls: bool | None = Field(True) status: str | None = Field( None, - description="One of `completed`, `failed`, `in_progress`, or `incomplete`.", + description="One of `completed`, `failed`, `in_progress`, `incomplete`, `queued`, or `cancelled`.", ) usage: ResponseUsage | None = Field(None) diff --git a/comfy_api_nodes/nodes_openai.py b/comfy_api_nodes/nodes_openai.py index ad62f2164..de2c94353 100644 --- a/comfy_api_nodes/nodes_openai.py +++ b/comfy_api_nodes/nodes_openai.py @@ -41,6 +41,9 @@ STARTING_POINT_ID_PATTERN = r"" class SupportedOpenAIModel(str, Enum): + gpt_5_6_sol = "gpt-5.6-sol" + gpt_5_6_terra = "gpt-5.6-terra" + gpt_5_6_luna = "gpt-5.6-luna" gpt_5_5_pro = "gpt-5.5-pro" gpt_5_5 = "gpt-5.5" gpt_5 = "gpt-5" @@ -1063,6 +1066,21 @@ class OpenAIChatNode(IO.ComfyNode): "usd": [0.002, 0.008], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } } + : $contains($m, "gpt-5.6-terra") ? { + "type": "list_usd", + "usd": [0.0025, 0.015], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : $contains($m, "gpt-5.6-luna") ? { + "type": "list_usd", + "usd": [0.001, 0.006], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : $contains($m, "gpt-5.6") ? { + "type": "list_usd", + "usd": [0.005, 0.03], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } : $contains($m, "gpt-5.5-pro") ? { "type": "list_usd", "usd": [0.03, 0.18], From 6a8ff7a929753a4fda2ea60c001a0d42258ef756 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 16 Jul 2026 19:43:12 -0700 Subject: [PATCH 117/211] Various comfy kitchen optimizations and fixes. (#14963) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index e7d301576..13fa237a4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.20 +comfy-kitchen==0.2.21 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 71b73e3b2bbdfb420aca342d61bef980b5a04f63 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 16 Jul 2026 19:44:02 -0700 Subject: [PATCH 118/211] Speed up anima a bit. (#14953) --- comfy/ldm/cosmos/predict2.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/comfy/ldm/cosmos/predict2.py b/comfy/ldm/cosmos/predict2.py index aec874815..371296e21 100644 --- a/comfy/ldm/cosmos/predict2.py +++ b/comfy/ldm/cosmos/predict2.py @@ -14,6 +14,7 @@ from torchvision import transforms import comfy.patcher_extension from comfy.ldm.modules.attention import optimized_attention import comfy.ldm.common_dit +import comfy.ops import comfy.quant_ops @@ -161,11 +162,16 @@ class Attention(nn.Module): def apply_norm_and_rotary_pos_emb( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor] ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - q = self.q_norm(q) - k = self.k_norm(k) v = self.v_norm(v) if self.is_selfattn and rope_emb is not None: # only apply to self-attention! - q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rope_emb) + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, q, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, k, offloadable=True) + q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, rope_emb, q_scale, k_scale, self.q_norm.eps) + comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream) + else: + q = self.q_norm(q) + k = self.k_norm(k) return q, k, v q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb) From 0f42ba51463174fb255f2c4605ae0e0b441fe6d7 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 17 Jul 2026 07:36:21 -0700 Subject: [PATCH 119/211] Support anima lllite control models. (#14954) Put them in the models/model_patches folder. Use the new AnimaLLLiteApply node. --- comfy/ldm/anima/lllite.py | 278 ++++++++++++++++++++++++++++++ comfy/ldm/cosmos/predict2.py | 52 +++++- comfy_extras/nodes_model_patch.py | 53 +++++- 3 files changed, 373 insertions(+), 10 deletions(-) create mode 100644 comfy/ldm/anima/lllite.py diff --git a/comfy/ldm/anima/lllite.py b/comfy/ldm/anima/lllite.py new file mode 100644 index 000000000..5c950ec89 --- /dev/null +++ b/comfy/ldm/anima/lllite.py @@ -0,0 +1,278 @@ +import re + +import torch +from torch import nn +import torch.nn.functional as F + +import comfy.ops +import comfy.utils + + +MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$") + + +def _group_norm(channels, device=None, dtype=None, operations=None): + groups = 8 + while groups > 1 and channels % groups != 0: + groups //= 2 + return operations.GroupNorm(groups, channels, device=device, dtype=dtype) + + +class AnimaLLLiteResBlock(nn.Module): + def __init__(self, channels, device=None, dtype=None, operations=None): + super().__init__() + self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations) + self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype) + self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations) + self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype) + + def forward(self, x): + h = self.conv1(F.silu(self.norm1(x))) + h = self.conv2(F.silu(self.norm2(h))) + return x + h + + +class AnimaLLLiteASPP(nn.Module): + def __init__(self, channels, dilations, device=None, dtype=None, operations=None): + super().__init__() + branches = [] + for dilation in dilations: + if dilation == 1: + conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype) + else: + conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype) + branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU())) + self.branches = nn.ModuleList(branches) + self.global_pool = nn.AdaptiveAvgPool2d(1) + self.global_conv = nn.Sequential( + operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype), + _group_norm(channels, device=device, dtype=dtype, operations=operations), + nn.SiLU(), + ) + self.proj = nn.Sequential( + operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype), + _group_norm(channels, device=device, dtype=dtype, operations=operations), + nn.SiLU(), + ) + + def forward(self, x): + height, width = x.shape[-2:] + outputs = [branch(x) for branch in self.branches] + pooled = self.global_conv(self.global_pool(x)) + outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False)) + return self.proj(torch.cat(outputs, dim=1)) + + +class AnimaLLLiteConditioning(nn.Module): + def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None): + super().__init__() + half_dim = cond_dim // 2 + self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype) + self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations) + self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype) + self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations) + self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype) + self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations) + self.resblocks = nn.ModuleList([ + AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations) + for _ in range(cond_resblocks) + ]) + self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None + self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype) + self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype) + + def forward(self, x): + x = F.silu(self.norm1(self.conv1(x))) + x = F.silu(self.norm2(self.conv2(x))) + x = F.silu(self.norm3(self.conv3(x))) + for block in self.resblocks: + x = block(x) + if self.aspp is not None: + x = self.aspp(x) + x = self.proj(x).flatten(2).transpose(1, 2).contiguous() + return self.out_norm(x) + + +class AnimaLLLiteModule(nn.Module): + def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None): + super().__init__() + self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype) + self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype) + self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype) + self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype) + self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False) + + def forward(self, x, cond_emb, strength): + original_shape = x.shape + if x.ndim == 5: + x = x.flatten(1, 3) + + if x.shape[0] != cond_emb.shape[0]: + if x.shape[0] % cond_emb.shape[0] != 0: + raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}") + cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1) + if x.shape[1] != cond_emb.shape[1]: + raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}") + + cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb) + hidden = F.silu(self.down(x)) + gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1) + hidden = self.mid(torch.cat((cond_local, hidden), dim=-1)) + hidden = F.silu(hidden * (1 + gamma) + beta) + x = x + self.up(hidden) * strength + + if len(original_shape) == 5: + x = x.reshape(original_shape) + return x + + +class AnimaLLLite(nn.Module): + def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None): + super().__init__() + metadata = metadata or {} + version = metadata.get("lllite.version", "2") + if version != "2": + raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported") + + module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")}) + if not module_names: + raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules") + + cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1] + cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0] + cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0] + resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")} + cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0 + use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict) + dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8") + aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else () + + self.cond_in_channels = cond_in_channels + self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true" + self.lllite_conditioning1 = AnimaLLLiteConditioning( + cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, + device=device, dtype=dtype, operations=operations, + ) + + self.module_names = set() + self.block_count = 0 + self.model_dim = None + for name in module_names: + match = MODULE_PATTERN.fullmatch(name) + if match is None: + raise ValueError(f"Unsupported Anima LLLite module name: {name}") + down_shape = state_dict[f"{name}.down.weight"].shape + mlp_dim, in_dim = down_shape + module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1] + if module_cond_dim != cond_emb_dim: + raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}") + if self.model_dim is None: + self.model_dim = in_dim + elif self.model_dim != in_dim: + raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}") + self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations)) + self.module_names.add(name) + self.block_count = max(self.block_count, int(match.group(1)) + 1) + + def encode_conditioning(self, image): + return self.lllite_conditioning1(image) + + def apply(self, x, cond_emb, block_index, target, strength): + name = f"lllite_dit_blocks_{block_index}_{target}" + if name not in self.module_names: + return x + return self.get_submodule(name)(x, cond_emb, strength) + + +class AnimaLLLitePatch: + def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end): + self.model_patch = model_patch + self.image = image + self.mask = mask + self.strength = strength + self.sigma_start = sigma_start + self.sigma_end = sigma_end + + def __call__(self, args): + x = args["x"] + transformer_options = args["transformer_options"] + if self.strength == 0.0: + return args + sigmas = transformer_options.get("sigmas") + if sigmas is not None: + sigma = float(sigmas.max().item()) + if not self.sigma_end <= sigma <= self.sigma_start: + return args + if x.shape[2] != 1: + raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}") + + target_height = x.shape[-2] * 8 + target_width = x.shape[-1] * 8 + image = comfy.utils.common_upscale( + self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center" + ).clamp(0.0, 1.0) + image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0 + + if self.model_patch.model.cond_in_channels == 4: + mask = self.mask + if mask.ndim == 3: + mask = mask.unsqueeze(1) + if mask.ndim != 4 or mask.shape[1] != 1: + raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}") + mask = comfy.utils.common_upscale( + mask.float(), target_width, target_height, "nearest-exact", crop="center" + ) + if mask.shape[0] != image.shape[0]: + if image.shape[0] % mask.shape[0] != 0: + raise ValueError( + f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}" + ) + mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1) + mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype) + if self.model_patch.model.inpaint_masked_input: + image = image * (mask < 0.5).to(image.dtype) + image = torch.cat((image, mask * 2.0 - 1.0), dim=1) + + cond_emb = self.model_patch.model.encode_conditioning(image) + transformer_options["model_patch_data"][self] = cond_emb + return args + + def to(self, device_or_dtype): + return self + + def models(self): + return [self.model_patch] + + +class AnimaLLLiteAttentionPatch: + def __init__(self, patch, targets): + self.patch = patch + self.targets = targets + + def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None): + cond_emb = extra_options["model_patch_data"].get(self.patch) + if cond_emb is None: + return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask} + + block_index = extra_options["block_index"] + values = {"q": q, "k": k, "v": v} + for value_name, target in self.targets.items(): + values[value_name] = self.patch.model_patch.model.apply( + values[value_name], cond_emb, block_index, target, self.patch.strength + ) + + return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask} + + +class AnimaLLLiteMLPPatch: + def __init__(self, patch): + self.patch = patch + + def __call__(self, args): + cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch) + if cond_emb is None: + return args + args["x"] = self.patch.model_patch.model.apply( + args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength + ) + return args diff --git a/comfy/ldm/cosmos/predict2.py b/comfy/ldm/cosmos/predict2.py index 371296e21..d391d50b1 100644 --- a/comfy/ldm/cosmos/predict2.py +++ b/comfy/ldm/cosmos/predict2.py @@ -149,11 +149,29 @@ class Attention(nn.Module): x: torch.Tensor, context: Optional[torch.Tensor] = None, rope_emb: Optional[torch.Tensor] = None, + transformer_options: Optional[dict] = {}, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - q = self.q_proj(x) context = x if context is None else context - k = self.k_proj(context) - v = self.v_proj(context) + q_input = x + k_input = context + v_input = context + + transformer_patches = transformer_options.get("patches", {}) + patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch" + if patch_name in transformer_patches: + extra_options = transformer_options.copy() + extra_options["n_heads"] = self.n_heads + extra_options["dim_head"] = self.head_dim + for patch in transformer_patches[patch_name]: + out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options) + q_input = out.get("q", q_input) + k_input = out.get("k", k_input) + v_input = out.get("v", v_input) + rope_emb = out.get("pe", rope_emb) + + q = self.q_proj(q_input) + k = self.k_proj(k_input) + v = self.v_proj(v_input) q, k, v = map( lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim), (q, k, v), @@ -194,7 +212,7 @@ class Attention(nn.Module): x (Tensor): The query tensor of shape [B, Mq, K] context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None """ - q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb) + q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options) return self.compute_attention(q, k, v, transformer_options=transformer_options) @@ -561,8 +579,14 @@ class Block(nn.Module): self.layer_norm_mlp, scale_mlp_B_T_1_1_D, shift_mlp_B_T_1_1_D, - ) - result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype)) + ).to(compute_dtype) + patches = transformer_options.get("patches", {}) + if "mlp_patch" in patches: + args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options} + for patch in patches["mlp_patch"]: + args = patch(args) + normalized_x_B_T_H_W_D = args["x"] + result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D) x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype)) return x_B_T_H_W_D @@ -869,11 +893,22 @@ class MiniTrainDIT(nn.Module): x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape ), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}" + transformer_options = kwargs.get("transformer_options", {}) + patches = transformer_options.get("patches", {}) + if "post_input" in patches: + transformer_options = transformer_options.copy() + transformer_options["model_patch_data"] = {} + + if "post_input" in patches: + for patch in patches["post_input"]: + out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options}) + x_B_T_H_W_D = out["img"] + block_kwargs = { "rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0), "adaln_lora_B_T_3D": adaln_lora_B_T_3D, "extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D, - "transformer_options": kwargs.get("transformer_options", {}), + "transformer_options": transformer_options, } # The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream @@ -883,7 +918,8 @@ class MiniTrainDIT(nn.Module): if x_B_T_H_W_D.dtype == torch.float16: x_B_T_H_W_D = x_B_T_H_W_D.float() - for block in self.blocks: + for block_index, block in enumerate(self.blocks): + transformer_options["block_index"] = block_index x_B_T_H_W_D = block( x_B_T_H_W_D, t_embedding_B_T_D, diff --git a/comfy_extras/nodes_model_patch.py b/comfy_extras/nodes_model_patch.py index 3f785c8b5..0935af09d 100644 --- a/comfy_extras/nodes_model_patch.py +++ b/comfy_extras/nodes_model_patch.py @@ -8,6 +8,7 @@ import comfy.ldm.common_dit import comfy.latent_formats import comfy.ldm.lumina.controlnet import comfy.ldm.supir.supir_modules +import comfy.ldm.anima.lllite from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel from comfy_api.latest import io from comfy.ldm.supir.supir_patch import SUPIRPatch @@ -236,10 +237,12 @@ class ModelPatchLoader: def load_model_patch(self, name): model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name) - sd = comfy.utils.load_torch_file(model_patch_path, safe_load=True) + sd, metadata = comfy.utils.load_torch_file(model_patch_path, safe_load=True, return_metadata=True) dtype = comfy.utils.weight_dtype(sd) - if 'controlnet_blocks.0.y_rms.weight' in sd: + if 'lllite_conditioning1.conv1.weight' in sd: + model = comfy.ldm.anima.lllite.AnimaLLLite(sd, metadata, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast) + elif 'controlnet_blocks.0.y_rms.weight' in sd: additional_in_dim = sd["img_in.weight"].shape[1] - 64 model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast) elif 'feature_embedder.mid_layer_norm.bias' in sd: @@ -296,6 +299,50 @@ class ModelPatchLoader: return (model_patcher,) +class AnimaLLLiteApply: + @classmethod + def INPUT_TYPES(s): + return {"required": {"model": ("MODEL",), + "model_patch": ("MODEL_PATCH",), + "image": ("IMAGE",), + "strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + }, + "optional": {"mask": ("MASK",), + }} + RETURN_TYPES = ("MODEL",) + FUNCTION = "apply_patch" + EXPERIMENTAL = True + + CATEGORY = "model_patches/anima" + + def apply_patch(self, model, model_patch, image, strength, start_percent, end_percent, mask=None): + image = image[..., :3] + + if model_patch.model.cond_in_channels == 4 and mask is None: + mask = torch.zeros_like(image[..., 0]) + elif model_patch.model.cond_in_channels != 4: + mask = None + + model_sampling = model.get_model_object("model_sampling") + sigma_start = float(model_sampling.percent_to_sigma(start_percent)) + sigma_end = float(model_sampling.percent_to_sigma(end_percent)) + patch = comfy.ldm.anima.lllite.AnimaLLLitePatch(model_patch, image, mask, strength, sigma_start, sigma_end) + model_patched = model.clone() + model_patched.set_model_post_input_patch(patch) + model_patched.set_model_attn1_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch( + patch, + {"q": "self_attn_q_proj", "k": "self_attn_k_proj", "v": "self_attn_v_proj"}, + )) + model_patched.set_model_attn2_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch( + patch, + {"q": "cross_attn_q_proj"}, + )) + model_patched.set_model_patch(comfy.ldm.anima.lllite.AnimaLLLiteMLPPatch(patch), "mlp_patch") + return (model_patched,) + + class DiffSynthCnetPatch: def __init__(self, model_patch, vae, image, strength, mask=None): self.model_patch = model_patch @@ -674,6 +721,7 @@ NODE_CLASS_MAPPINGS = { "ZImageFunControlnet": ZImageFunControlnet, "USOStyleReference": USOStyleReference, "SUPIRApply": SUPIRApply, + "AnimaLLLiteApply": AnimaLLLiteApply, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -682,4 +730,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ZImageFunControlnet": "Apply Z-Image Fun ControlNet", "USOStyleReference": "Apply USO Style Reference", "SUPIRApply": "Apply SUPIR Patch", + "AnimaLLLiteApply": "Apply Anima LLLite", } From 8edea4a65d765b70451e3fcd2604f40a420b7cd7 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:24:34 +0300 Subject: [PATCH 120/211] [Partner Nodes] feat(Google): add Gemini 3.5 Flash LLM model (#14972) Co-authored-by: Alexis Rolland --- comfy_api_nodes/nodes_gemini.py | 21 ++++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index a8eb0a797..283e2233d 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -274,6 +274,10 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N input_tokens_price = 0.25 output_text_tokens_price = 1.50 output_image_tokens_price = 0.0 + elif response.modelVersion == "gemini-3.5-flash": + input_tokens_price = 1.50 + output_text_tokens_price = 9.0 + output_image_tokens_price = 0.0 elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"): input_tokens_price = 2 output_text_tokens_price = 12.0 @@ -619,11 +623,12 @@ class GeminiNode(IO.ComfyNode): GEMINI_V2_MODELS: dict[str, str] = { "Gemini 3.1 Pro": "gemini-3.1-pro-preview", + "Gemini 3.5 Flash": "gemini-3.5-flash", "Gemini 3.1 Flash-Lite": "gemini-3.1-flash-lite-preview", } -def _gemini_text_model_inputs(thinking_default: str) -> list[Input]: +def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str] | None = None) -> list[Input]: """Per-model inputs revealed by the model DynamicCombo (shared media + sampling controls).""" return [ IO.Autogrow.Input( @@ -661,7 +666,7 @@ def _gemini_text_model_inputs(thinking_default: str) -> list[Input]: ), IO.Combo.Input( "thinking_level", - options=["LOW", "HIGH"], + options=thinking_options or ["LOW", "HIGH"], default=thinking_default, tooltip="How hard the model reasons internally before answering. " "HIGH improves quality on difficult tasks but costs more (thinking) tokens and is slower.", @@ -719,6 +724,10 @@ class GeminiNodeV2(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "Gemini 3.5 Flash", + _gemini_text_model_inputs("MEDIUM", ["MINIMAL", "LOW", "MEDIUM", "HIGH"]), + ), IO.DynamicCombo.Option("Gemini 3.1 Pro", _gemini_text_model_inputs("HIGH")), IO.DynamicCombo.Option("Gemini 3.1 Flash-Lite", _gemini_text_model_inputs("LOW")), ], @@ -759,7 +768,13 @@ class GeminiNodeV2(IO.ComfyNode): "type": "list_usd", "usd": [0.00025, 0.0015], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } - } : { + } + : $contains($m, "3.5 flash") ? { + "type": "list_usd", + "usd": [0.0015, 0.009], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : { "type": "list_usd", "usd": [0.002, 0.012], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } From c67f95607f33b5b01c4d3ce062d75a7468775b38 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Sat, 18 Jul 2026 01:29:45 +0900 Subject: [PATCH 121/211] chore(openapi): sync shared API contract from cloud@4acc59a (#14947) --- openapi.yaml | 51 ++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 50 insertions(+), 1 deletion(-) diff --git a/openapi.yaml b/openapi.yaml index e00643bad..a50312226 100644 --- a/openapi.yaml +++ b/openapi.yaml @@ -530,6 +530,10 @@ components: description: Job creation timestamp (Unix timestamp in milliseconds) format: int64 type: integer + execution_end_time: + description: Workflow execution completion timestamp (Unix milliseconds, only present for terminal states) + format: int64 + type: integer execution_error: allOf: - $ref: '#/components/schemas/ExecutionError' @@ -538,6 +542,10 @@ components: additionalProperties: true description: Node-level execution metadata (only for terminal states) type: object + execution_start_time: + description: Workflow execution start timestamp (Unix milliseconds, only present once execution has started) + format: int64 + type: integer execution_status: additionalProperties: true description: ComfyUI execution status and timeline (only for terminal states) @@ -570,6 +578,12 @@ components: description: Last update timestamp (Unix timestamp in milliseconds) format: int64 type: integer + user_id: + description: | + ID of the user that owns this job (see the `workspace_id` + description above for why this is always the caller's own id + on a successful response). + type: string workflow: additionalProperties: true description: | @@ -583,6 +597,18 @@ components: workflow_id: description: UUID identifying the workflow graph definition type: string + workspace_id: + description: | + ID of the workspace that owns this job. A successful (200) + response from this operation is only ever returned for the + caller's own job (see this operation's ownership-scoped + query), so this is always the caller's own workspace — + consumers that also need to correlate this job to its + live-progress broadcast channel (workspace+user scoped; see + the internal common/gateways/broadcast package) can use this + value directly rather than resolving their own identity a + second way. + type: string required: - id - status @@ -1565,7 +1591,13 @@ paths: schema: default: true type: boolean - - description: Filter assets by exact content hash. + - description: | + Filter assets by content hash, in the canonical `blake3:` + form. Matches regardless of which of this asset store's two + internal hash storage formats the matching row was written + under (the canonical form used by from-hash-created references, + or the raw `.`/bare `` storage key used by direct + uploads) — both represent the same content hash. in: query name: hash schema: @@ -2464,6 +2496,23 @@ paths: schema: additionalProperties: true properties: + free_tier_balance: + description: Free-tier job allowance for an authenticated non-paid (FREE-tier) user in the rollout. Absent for paid users and unauthenticated requests. Synthesized from config before a grant row exists so a brand-new user still sees their full allowance. + properties: + allowance: + description: Total free jobs granted for the current period + type: integer + remaining: + description: Free jobs remaining (allowance - used, floored at 0) + type: integer + used: + description: Free jobs consumed so far + type: integer + required: + - allowance + - used + - remaining + type: object max_upload_size: description: Maximum upload size in bytes type: integer From 1d1099bea08efa6904480cf25c54ea92646aea4f Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Sat, 18 Jul 2026 00:53:15 +0800 Subject: [PATCH 122/211] chore: update workflow templates to v0.11.11 (#14973) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 13fa237a4..fbdd6c14a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.21 -comfyui-workflow-templates==0.11.9 +comfyui-workflow-templates==0.11.11 comfyui-embedded-docs==0.5.8 torch torchsde From b08e6cf35fac50d3ca8470dffb3f9a1fbb7187d2 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 17 Jul 2026 23:49:16 +0300 Subject: [PATCH 123/211] [Partner Nodes] feat(HeyGen): add Avatar, Talking Photo, Create Avatar, Video Translate and TTS nodes (#14958) * [Partner Nodes] feat(HeyGen): add Avatar, Talking Photo, Create Avatar, Video Translate and TTS nodes Signed-off-by: bigcat88 * [Partner Nodes] fix(HeyGen): display only Avatars supported by engine Signed-off-by: bigcat88 * [Partner Nodes] remove 4K option --------- Signed-off-by: bigcat88 Co-authored-by: Alexis Rolland --- comfy_api_nodes/apis/heygen.py | 452 ++++++++++++++++++ comfy_api_nodes/nodes_heygen.py | 799 ++++++++++++++++++++++++++++++++ 2 files changed, 1251 insertions(+) create mode 100644 comfy_api_nodes/apis/heygen.py create mode 100644 comfy_api_nodes/nodes_heygen.py diff --git a/comfy_api_nodes/apis/heygen.py b/comfy_api_nodes/apis/heygen.py new file mode 100644 index 000000000..71b724872 --- /dev/null +++ b/comfy_api_nodes/apis/heygen.py @@ -0,0 +1,452 @@ +# (label, avatar_id, avatar_type, supported engines) +HEYGEN_AVATAR_LOOKS: list[tuple[str, str, str, tuple[str, ...]]] = [ + ( + "Annie Lounge Standing Side", + "Annie_Lounge_Standing_Side_public", + "studio_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Yara Modern Lecture Hall", + "fd6814ecc5e143cd899e615a80eaa2dc", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Brandon Business Sitting Front", + "Brandon_Business_Sitting_Front_public", + "studio_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Caroline Business Sitting Side", + "Caroline_Business_Sitting_Side_public", + "studio_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Ursula Lawyer Angle 4", + "f7173d2bb8584c00bfec6905c5e9a492", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Sofia Corporate Presenter 01 Angle 3", + "fe563971fd2d438e957372dac9e2be8c", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Seoyeon Health Nutrition Coach Angle 3", + "fe3c5d5028d941398d064b8fc64a2dea", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Sanne Fitness Coach Angle 4", + "d967f935a8bf4a0c8f0bccfd66c501d2", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ("Sander", "f5cd7b94056f495ca0610602d64a9aa3", "photo_avatar", ("avatar_v", "avatar_iv", "avatar_iii")), + ( + "Rupert Personal Development Coach Angle 4", + "f57b3e626adb4bc997b38f64884adce4", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Olivier Professor Angle 2", + "f6659bbb094b459c87c967edbb9ee481", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Obi Health Nutrition Coach Angle 5", + "f3dc2c38201d414382f506d2d8e8d029", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Matilda Modern Office Setting", + "fda889ac354a440da8dbecc410981273", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Mateo Traditional Law Office", + "ff172d6c499c4e47ba6fcc5de631e9fc", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Marlon Inviting Armchair Setting", + "f5a57db099ab462daa3e7c604a05dacc", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Margaret Professor Angle 1", + "fb472bc29ab04bcca576e3703978fecb", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Marek Therapy Coach Angle 3", + "e197768703f1463a93dc25ada1f421fb", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Maeve Warm, Professional Setting", + "faf66681d8cc48dc82c4283200b3e782", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ( + "Lorenzo Professor Angle 5", + "fc268dc244bb40d7a554663ce723dcf0", + "photo_avatar", + ("avatar_v", "avatar_iv", "avatar_iii"), + ), + ("Luca", "Luca_public", "studio_avatar", ("avatar_iii",)), + ("Bruce", "Bruce_public", "studio_avatar", ("avatar_iii",)), + ("Nico", "Nico_public", "studio_avatar", ("avatar_iii",)), + ("Lisa", "Lisa_public", "studio_avatar", ("avatar_iii",)), + ("Sophie", "Sophie_public", "studio_avatar", ("avatar_iii",)), + ("Aiko", "Aiko_public", "studio_avatar", ("avatar_iii",)), + ("Rebecca (portrait)", "Rebecca_public", "studio_avatar", ("avatar_iii",)), + ("Daphne in Grey blazer (portrait)", "Daphne_public_1", "studio_avatar", ("avatar_iii",)), + ("Bryce in Black t-shirt", "Bryce_public_5", "studio_avatar", ("avatar_iii",)), + ("Diora in White shirt", "Diora_public_3", "studio_avatar", ("avatar_iii",)), + ("Freja in White blazer", "Freja_public_1", "studio_avatar", ("avatar_iii",)), + ("Albert in Blue blazer", "Albert_public_2", "studio_avatar", ("avatar_iii",)), + ("Emery in Red blazer", "Emery_public_1", "studio_avatar", ("avatar_iii",)), + ("Minho in Blue shirt", "Minho_public_6", "studio_avatar", ("avatar_iii",)), + ("Aditya in Brown blazer", "Aditya_public_4", "studio_avatar", ("avatar_iii",)), + ("Nadim in Blue blazer", "Nadim_public_1", "studio_avatar", ("avatar_iii",)), + ("Iker in Black blazer", "Iker_public_1", "studio_avatar", ("avatar_iii",)), + ("Nour in Black blazer", "Nour_public_1", "studio_avatar", ("avatar_iii",)), + ("Saskia in Blue blazer", "Saskia_public_1", "studio_avatar", ("avatar_iii",)), + ("Lucien in Blue blazer", "Lucien_public_1", "studio_avatar", ("avatar_iii",)), + ("Esmond in Blue suit", "Esmond_public_3", "studio_avatar", ("avatar_iii",)), + ("Jinwoo in Blue suit", "Jinwoo_public_5", "studio_avatar", ("avatar_iii",)), + ("Annelore in Red sweater (portrait)", "Annelore_public_3", "studio_avatar", ("avatar_iii",)), + ("Bastien in Blue shirt", "Bastien_public_4", "studio_avatar", ("avatar_iii",)), + ("Zosia in Khaki blazer", "Zosia_public_3", "studio_avatar", ("avatar_iii",)), + ("Tahlia in Dark blue suit", "Tahlia_public_4", "studio_avatar", ("avatar_iii",)), +] +HEYGEN_AVATAR_OPTIONS = [x[0] for x in HEYGEN_AVATAR_LOOKS] +HEYGEN_AVATAR_MAP = {x[0]: (x[1], x[2], x[3]) for x in HEYGEN_AVATAR_LOOKS} + +# (label, voice_id) — Starfish-compatible voices for the TTS endpoint +HEYGEN_VOICE_TTS: list[tuple[str, str]] = [ + ("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"), + ("Zain (English, female)", "0047732240584155b1588455313e78ec"), + ("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"), + ("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"), + ("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"), + ("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"), + ("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"), + ("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"), + ("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"), + ("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"), + ("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"), + ("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"), + ("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"), + ("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"), + ("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"), + ("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"), + ("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"), + ("Rose - UGC -2 (English, female)", "0495e14c2bd74eb3aeeef03583e0bce5"), + ("Derya - Lifelike - Broadcaster 🎙️ (English, female)", "04d0ae1d0af2489ca7d3bb402a39a890"), + ("Dynamic Derek (English, male)", "0516c2d857eb425c94e90b068241914e"), + ("Lotte (English, female)", "052fcfb83d1a4c2f8d0368c226fea4b9"), + ("Thanos - Broadcaster 🎙️ (English, male)", "054af44a167344d0af2722fdfef08d17"), + ("Marcia (English, female)", "05f19352e8f74b0392a8f411eba40de1"), + ("Camden (English, male)", "06468055edd4458aa131a1dfd813c1e9"), + ("Rumi (English, female)", "06672207805f41a9ad0af6797f8aa14b"), + ("Pippa (English, female)", "06b68c4dbb544935b9af984e80efa4fb"), + ("William Prescott - Broadcaster 🎙️ (English, male)", "06c816b952f14fa9b3a6c42aa151f731"), + ("Sammy (English, female)", "06e6facd99654b9dbb9308f67bf3a31c"), + ("Breezy Bagus (Indonesian, male)", "06e81a5d7c8b41818d3f0b38f7cf15a1"), + ("Ben (English, male)", "07ca39b243184dbcb82e7e0f0e524b21"), + ("Smooth Dev (English, male)", "07d2ba65847541feb97abc9b60181555"), + ("Daran inside booth (English, male)", "080d9383c0314056aef392892e009806"), + ("Peppy Stella (English, female)", "084760b4922a44599575c770070ec2d7"), + ("Silas (English, male)", "08f561403ec846dbbd8c691cc448f45a"), + ("Aditya (English, male)", "09c3d65e44e247dd8b78a97a903feb58"), + ("Christy (English, female)", "09d88c036bf449fa905900c08b235a37"), + ("Elio (English, male)", "0a0b38624ac64ec6afcd5842a977ca10"), + ("Luminous Laksh (Hindi, male)", "0adc547b76a5401c856274c379904eb7"), + ("Jeff (English, male)", "0add542e349f4ccaba6ecb3b7ced6034"), + ("Tahlia Brooks - Excited 🤩 (English, female)", "0b440d1ac2454d69a73302fc806522b1"), + ("Riya Mehta (Hindi, female)", "0b464b2f4e2249a4b5a05e60eaf41e7e"), + ("Ben Hart (English, male)", "0b47b5a637e944f9bfd49913999b344b"), + ("Skylar (English, female)", "0bbfbda5aa924a68a9d1da7b8496052a"), + ("Relaxed Reece (English, male)", "0c2151d538844c70a8b096de533f2828"), + ("Daniel (English, male)", "0c23804af39a4946ac6fda42bfff2738"), + ("Melani (English, female)", "0c54c6399ad64551a304e1a346677723"), + ("Clover (English, female)", "0ccb0bea067d4449ad367baeed7ea2e9"), + ("Pedro Lima - Serious 😐 (Portuguese, male)", "0d0e23e8170446e38b18a7380b2d30a8"), + ("Ana Carvalho (Portuguese, female)", "0d23c5b2f6004e909802a2e8bfcd52c2"), + ("Confident Connor - Excited 🤩 (English, male)", "0dd34c3eb79247238219eea35aeb58cd"), + ("Vibrant Victor (Spanish, male)", "1062976ea8bf42f4adc27c7e868b8fde"), + ("Young Olivier (French, male)", "1c5dc9a8f8cf4de0932f91d75f43a15d"), + ("Émile Noir (French, male)", "25a6a67280574d3da78e97b1935ebfc7"), + ("Steadfast Stefan (German, male)", "0eb85e6e8710473b82f7e88609ba3053"), + ("Deep Dieter (German, male)", "118949676b0a46629d1ad52981c3ef84"), + ("Serene Marco (Italian, male)", "72e922488a614041b5ab5f6ee07e3deb"), + ("Murmuring Matteo (Italian, male)", "755902b751654f30a6ef49e8bbcacfec"), + ("Gail in car (Multilingual, female)", "0214ac51f93e420f8711d568dcfbc50e"), + ("Daran outside walking (Multilingual, male)", "0ac81e725f4948dfa9638ceca216bcfa"), + ("BOB - Voice 1 (Chinese, unknown)", "dMkR1XwIkarpNqWUJLnX"), + ("Hakeem Hassan (Arabic, male)", "61a4359785664d01a59664ceb87ce6d4"), + ("Rami Idris (Arabic, male)", "a0bd2e5d41a74643be47ac75ca9171a2"), + ("Bold Kasia - Friendly 😊 (Polish, female)", "331624aec8b24a6c9287b8e16bdf54e8"), + ("Tranquil Tulin (Turkish, female)", "61646c861eb64e2d9036d8db51385356"), + ("Dynamic Derya (Turkish, female)", "664b73058b784aa89ddb2924c141d441"), + ("Quiet Dewa (Indonesian, male)", "1fa1193cf1d74f27ba58531c07ef9862"), + ("Cuong (Vietnamese, male)", "8af68d7ea38f4e7ca05cf46c3f7a590b"), +] +HEYGEN_VOICE_TTS_OPTIONS = [x[0] for x in HEYGEN_VOICE_TTS] +HEYGEN_VOICE_TTS_MAP = dict(HEYGEN_VOICE_TTS) + +# (label, voice_id) — top-ranked voices for video narration (any engine) +HEYGEN_VOICE_GENERAL: list[tuple[str, str]] = [ + ("Cassidy (English, female)", "16a09e4706f74997ba4ed05ea11470f6"), + ("Hope (English, female)", "42d00d4aac5441279d8536cd6b52c53c"), + ("Archer (English, male)", "453c20e1525a429080e2ad9e4b26f2cd"), + ("Brittney (English, female)", "4754e1ec667544b0bd18cdf4bec7d6a7"), + ("Mark (English, male)", "5d8c378ba8c3434586081a52ac368738"), + ("Andrew (English, male)", "6be73833ef9a4eb0aeee399b8fe9d62b"), + ("Spuds Oxley (English, male)", "76940a9adcd0490a9ce2cfe9a64a2664"), + ("Patrick (English, male)", "7e157ec62c9c45f1adca12faae72c86f"), + ("David Castlemore (English, male)", "828b59f834fd4c7188da322b6d9b6c75"), + ("Michael C (English, male)", "8661cd40d6c44c709e2d0031c0186ada"), + ("Adam Stone (English, male)", "88bb9ee1c81b466eb2a08fdde86d3619"), + ("Alex (English, male)", "897d6a9b2c844f56aa077238768fe10a"), + ("Monika Sogam (English, female)", "97dd67ab8ce242b6a9e7689cb00c6414"), + ("Jessica Anne Bogart (English, female)", "b966c31caf124c2a99f19ff1479c964f"), + ("John Doe (English, male)", "c4a8ceb7a2954500bc047fb092bcff3f"), + ("Ivy (English, female)", "cef3bc4e0a84424cafcde6f2cf466c97"), + ("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"), + ("Allison (English, female)", "f8c69e517f424cafaecde32dde57096b"), + ("Mia Starset (Norwegian, female)", "000466f8ac6d47a49f5743d50b3778de"), + ("William Shanks (Spanish, male)", "001248bb63f847888d37b766ee8b3a47"), + ("Zain (English, female)", "0047732240584155b1588455313e78ec"), + ("Jora Slobod (Romanian, male)", "00631519159a402ab5d8f719e51532bb"), + ("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"), + ("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"), + ("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"), + ("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"), + ("Son Tran (Vietnamese, male)", "0132f85950a94d11ba180f885101bf84"), + ("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"), + ("Marc Aurèle (French, male)", "018a94cf15574491a0bab7f6799ac15b"), + ("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"), + ("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"), + ("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"), + ("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"), + ("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"), + ("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"), + ("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"), + ("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"), + ("Tuba (, female)", "034ca0c32b6542028748d6d365d90d6a"), + ("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"), + ("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"), +] +HEYGEN_VOICE_GENERAL_OPTIONS = [x[0] for x in HEYGEN_VOICE_GENERAL] +HEYGEN_VOICE_GENERAL_MAP = dict(HEYGEN_VOICE_GENERAL) + +HEYGEN_TRANSLATE_LANGUAGES = [ + "English", + "Spanish", + "Spanish (Spain)", + "Spanish (Mexico)", + "French", + "French (France)", + "German", + "German (Germany)", + "Portuguese", + "Portuguese (Brazil)", + "Italian", + "Italian (Italy)", + "Japanese", + "Japanese (Japan)", + "Korean", + "Chinese (Mandarin, Simplified)", + "Arabic", + "Hindi", + "Hindi (India)", + "Russian", + "Russian (Russia)", + "Dutch", + "Polish", + "Turkish", + "Indonesian", + "Vietnamese", + "Ukrainian", + "Afrikaans (South Africa)", + "Albanian (Albania)", + "Amharic (Ethiopia)", + "Arabic (Algeria)", + "Arabic (Bahrain)", + "Arabic (Egypt)", + "Arabic (Iraq)", + "Arabic (Jordan)", + "Arabic (Kuwait)", + "Arabic (Lebanon)", + "Arabic (Libya)", + "Arabic (Morocco)", + "Arabic (Oman)", + "Arabic (Qatar)", + "Arabic (Saudi Arabia)", + "Arabic (Syria)", + "Arabic (Tunisia)", + "Arabic (United Arab Emirates)", + "Arabic (World)", + "Arabic (Yemen)", + "Armenian (Armenia)", + "Azerbaijani (Latin, Azerbaijan)", + "Bangla (Bangladesh)", + "Basque", + "Belarusian (Belarus)", + "Bengali (India)", + "Bosnian (Bosnia and Herzegovina)", + "Bulgarian", + "Bulgarian (Bulgaria)", + "Burmese (Myanmar)", + "Catalan", + "Chinese (Cantonese, Traditional)", + "Chinese (Jilu Mandarin, Simplified)", + "Chinese (Northeastern Mandarin, Simplified)", + "Chinese (Southwestern Mandarin, Simplified)", + "Chinese (Taiwanese Mandarin, Traditional)", + "Chinese (Wu, Simplified)", + "Chinese (Zhongyuan Mandarin Henan, Simplified)", + "Chinese (Zhongyuan Mandarin Shaanxi, Simplified)", + "Croatian", + "Croatian (Croatia)", + "Czech", + "Czech (Czechia)", + "Danish", + "Danish (Denmark)", + "Dutch (Belgium)", + "Dutch (Netherlands)", + "English (Australia)", + "English (Canada)", + "English (Hong Kong SAR)", + "English (India)", + "English (Ireland)", + "English (Kenya)", + "English (New Zealand)", + "English (Nigeria)", + "English (Philippines)", + "English (Singapore)", + "English (South Africa)", + "English (Tanzania)", + "English (UK)", + "English (United States)", + "Estonian (Estonia)", + "Filipino", + "Filipino (Cebuano)", + "Filipino (Philippines)", + "Finnish", + "Finnish (Finland)", + "French (Belgium)", + "French (Canada)", + "French (Switzerland)", + "Galician", + "Georgian (Georgia)", + "German (Austria)", + "German (Switzerland)", + "Greek", + "Greek (Greece)", + "Gujarati (India)", + "Haitian Creole (Haiti)", + "Hebrew (Israel)", + "Hungarian (Hungary)", + "Icelandic (Iceland)", + "Indonesian (Indonesia)", + "Irish (Ireland)", + "Javanese (Latin, Indonesia)", + "Kannada (India)", + "Kazakh (Kazakhstan)", + "Khmer (Cambodia)", + "Konkani (India)", + "Korean (Korea)", + "Lao (Laos)", + "Latin (Vatican City)", + "Latvian (Latvia)", + "Lithuanian (Lithuania)", + "Luxembourgish (Luxembourg)", + "Macedonian (North Macedonia)", + "Maithili (India)", + "Malagasy (Madagascar)", + "Malay", + "Malay (Malaysia)", + "Malayalam (India)", + "Maltese (Malta)", + "Mandarin", + "Marathi (India)", + "Mongolian (Mongolia)", + "Nepali (Nepal)", + "Norwegian Bokmål (Norway)", + "Norwegian Nynorsk (Norway)", + "Odia (India)", + "Pashto (Afghanistan)", + "Persian (Iran)", + "Polish (Poland)", + "Portuguese (Portugal)", + "Punjabi (India)", + "Romanian", + "Romanian (Romania)", + "Serbian (Latin, Serbia)", + "Sindhi (India)", + "Sinhala (Sri Lanka)", + "Slovak", + "Slovak (Slovakia)", + "Slovenian (Slovenia)", + "Somali (Somalia)", + "Spanish (Argentina)", + "Spanish (Bolivia)", + "Spanish (Chile)", + "Spanish (Colombia)", + "Spanish (Costa Rica)", + "Spanish (Cuba)", + "Spanish (Dominican Republic)", + "Spanish (Ecuador)", + "Spanish (El Salvador)", + "Spanish (Equatorial Guinea)", + "Spanish (Guatemala)", + "Spanish (Honduras)", + "Spanish (Latin America)", + "Spanish (Nicaragua)", + "Spanish (Panama)", + "Spanish (Paraguay)", + "Spanish (Peru)", + "Spanish (Puerto Rico)", + "Spanish (United States)", + "Spanish (Uruguay)", + "Spanish (Venezuela)", + "Sundanese (Indonesia)", + "Swahili (Kenya)", + "Swahili (Tanzania)", + "Swedish", + "Swedish (Sweden)", + "Tamil", + "Tamil (India)", + "Tamil (Malaysia)", + "Tamil (Singapore)", + "Tamil (Sri Lanka)", + "Telugu (India)", + "Thai (Thailand)", + "Turkish (Türkiye)", + "Ukrainian (Ukraine)", + "Urdu (India)", + "Urdu (Pakistan)", + "Uzbek (Latin, Uzbekistan)", + "Vietnamese (Vietnam)", + "Welsh (United Kingdom)", + "Zulu (South Africa)", +] diff --git a/comfy_api_nodes/nodes_heygen.py b/comfy_api_nodes/nodes_heygen.py new file mode 100644 index 000000000..6c9812c86 --- /dev/null +++ b/comfy_api_nodes/nodes_heygen.py @@ -0,0 +1,799 @@ +import uuid + +import torch +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension, Input +from comfy_api_nodes.apis.heygen import ( + HEYGEN_AVATAR_MAP, + HEYGEN_AVATAR_OPTIONS, + HEYGEN_TRANSLATE_LANGUAGES, + HEYGEN_VOICE_GENERAL_MAP, + HEYGEN_VOICE_GENERAL_OPTIONS, + HEYGEN_VOICE_TTS_MAP, + HEYGEN_VOICE_TTS_OPTIONS, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + audio_bytes_to_audio_input, + download_url_as_bytesio, + download_url_to_image_tensor, + download_url_to_video_output, + downscale_image_tensor_by_max_side, + get_number_of_images, + poll_op_raw, + sync_op_raw, + upload_audio_to_comfyapi, + upload_image_to_comfyapi, + upload_images_to_comfyapi, + upload_video_to_comfyapi, + validate_string, +) +from server import PromptServer + +_VIDEOS_PATH = "/proxy/heygen/v3/videos" +_TRANSLATIONS_PATH = "/proxy/heygen/v3/video-translations" +_SPEECH_PATH = "/proxy/heygen/v3/voices/speech" +_AVATARS_PATH = "/proxy/heygen/v3/avatars" +_LOOKS_PATH = "/proxy/heygen/v3/avatars/looks" + +_DEFAULT_VOICE_OPTION = "(avatar's default voice)" + +_AVATARS_BY_ENGINE = { + e: [label for label, (_aid, _atype, engines) in HEYGEN_AVATAR_MAP.items() if e in engines] + for e in ("avatar_iv", "avatar_iii", "avatar_v") +} + + +async def _apply_speech_source(cls: type[IO.ComfyNode], payload: dict, speech: dict, require_voice: bool) -> None: + """Fill script/audio speech fields of a /v3/videos payload from the DynamicCombo dict.""" + if speech["speech"] == "audio": + payload["audio_url"] = await upload_audio_to_comfyapi( + cls, speech["audio"], container_format="mp3", codec_name="libmp3lame", mime_type="audio/mpeg" + ) + elif speech["speech"] == "script": + validate_string(speech["text"], strip_whitespace=True, min_length=1, max_length=5000) + payload["script"] = speech["text"] + voice_id = speech.get("custom_voice_id", "").strip() + if not voice_id and speech["voice"] != _DEFAULT_VOICE_OPTION: + voice_id = HEYGEN_VOICE_GENERAL_MAP[speech["voice"]] + if voice_id: + payload["voice_id"] = voice_id + elif require_voice: + raise ValueError("A voice is required when driving the video with a text script.") + speed = speech.get("voice_speed", 1.0) + if speed != 1.0: + payload["voice_settings"] = {"speed": round(speed, 2)} + + +async def _create_and_poll_video(cls: type[IO.ComfyNode], payload: dict) -> dict: + """POST a /v3/videos payload, poll until terminal, and return the final video data.""" + created = await sync_op_raw( + cls, + ApiEndpoint(path=_VIDEOS_PATH, method="POST", headers={"Idempotency-Key": uuid.uuid4().hex}), + data=payload, + ) + video_id = (created.get("data") or {}).get("video_id") + if not video_id: + raise ValueError(f"HeyGen did not return a video_id: {created}") + final = await poll_op_raw( + cls, + ApiEndpoint(path=f"{_VIDEOS_PATH}/{video_id}"), + status_extractor=lambda r: (r.get("data") or {}).get("status"), + queued_statuses=["pending", "waiting"], + poll_interval=5.0, + ) + data = final["data"] + if not data.get("video_url"): + raise ValueError(f"HeyGen returned no video_url for video {video_id}.") + return data + + +async def _resolve_avatar( + cls: type[IO.ComfyNode], avatar_label: str, custom_avatar_id: str, engine_choice: str +) -> tuple[str, str | None]: + """Resolve (avatar_id, engine_type) from the combo/override + engine widgets.""" + custom_avatar_id = custom_avatar_id.strip() + if custom_avatar_id: + look = ( + await sync_op_raw( + cls, + ApiEndpoint(path=f"{_LOOKS_PATH}/{custom_avatar_id}"), + final_label_on_success=None, + ) + ).get("data") or {} + avatar_id = custom_avatar_id + avatar_label = look.get("name") or custom_avatar_id + supported = look.get("supported_api_engines") or [] + else: + avatar_id, avatar_type, supported = HEYGEN_AVATAR_MAP[avatar_label] + + if engine_choice == "auto": + engine = next((e for e in ("avatar_iv", "avatar_iii", "avatar_v") if e in supported), None) + else: + engine = engine_choice + if supported and engine not in supported: + raise ValueError( + f"Avatar '{avatar_label}' does not support the {engine} engine " + f"(supported: {', '.join(supported)}). Set engine to 'auto' to pick " + "a compatible engine automatically." + ) + return avatar_id, engine + + +class HeyGenTalkingPhotoNode(IO.ComfyNode): + """Animate a still image of a person into a lip-synced talking video.""" + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="HeyGenTalkingPhotoNode", + display_name="HeyGen Talking Photo", + category="partner/video/HeyGen", + description="Animate any image of a person into a lip-synced talking video " + "(HeyGen Avatar IV). Drive it with a text script or your own audio.", + inputs=[ + IO.Image.Input( + "image", + tooltip="Image of a person to animate. Downscaled automatically if larger than 2K.", + ), + IO.DynamicCombo.Input( + "speech", + display_name="speech source", + options=[ + IO.DynamicCombo.Option( + "script", + [ + IO.String.Input( + "text", + multiline=True, + default="", + tooltip="Text for the avatar to speak (up to 5000 characters). " + "The generated speech must be at least 1 second long.", + ), + IO.Combo.Input( + "voice", + options=HEYGEN_VOICE_GENERAL_OPTIONS, + tooltip="Voice for the script (HeyGen's most popular voices).", + ), + IO.String.Input( + "custom_voice_id", + default="", + optional=True, + tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. " + "Any voice from HeyGen's library (2000+) can be used.", + ), + IO.Float.Input( + "voice_speed", + default=1.0, + min=0.5, + max=1.5, + step=0.05, + optional=True, + tooltip="Speech speed multiplier.", + ), + ], + ), + IO.DynamicCombo.Option( + "audio", + [ + IO.Audio.Input( + "audio", + tooltip="Audio for the avatar to lip-sync, up to 10 minutes.", + ), + ], + ), + ], + tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.", + ), + IO.Combo.Input( + "resolution", + options=["720p", "1080p"], + default="1080p", + optional=True, + tooltip="Output video resolution.", + ), + IO.Combo.Input( + "aspect_ratio", + options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"], + default="auto", + optional=True, + tooltip="Output aspect ratio. 'auto' follows the input image.", + ), + IO.Combo.Input( + "expressiveness", + options=["low", "medium", "high"], + default="low", + optional=True, + tooltip="How expressive the animated face and gestures are.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + optional=True, + tooltip="Not sent to HeyGen; change it to force a re-run.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0715,"format":{"suffix":"/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + speech: dict, + resolution: str = "1080p", + aspect_ratio: str = "auto", + expressiveness: str = "low", + seed: int = 0, + ) -> IO.NodeOutput: + image = downscale_image_tensor_by_max_side(image, max_side=2000) + image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None) + payload = { + "type": "image", + "image": {"type": "url", "url": image_url}, + "resolution": resolution, + "aspect_ratio": aspect_ratio, + "expressiveness": expressiveness, + "title": "ComfyUI Talking Photo", + } + await _apply_speech_source(cls, payload, speech, require_voice=True) + video = await _create_and_poll_video(cls, payload) + return IO.NodeOutput(await download_url_to_video_output(video["video_url"])) + + +class HeyGenAvatarVideoNode(IO.ComfyNode): + """Generate a presenter video from a HeyGen avatar look.""" + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="HeyGenAvatarVideoNode", + display_name="HeyGen Avatar Video", + category="partner/video/HeyGen", + description="Generate a talking-presenter video from a HeyGen avatar. " + "Includes HeyGen's most popular public avatars; any look ID can be supplied as an override.", + inputs=[ + IO.DynamicCombo.Input( + "engine", + options=[ + IO.DynamicCombo.Option( + "auto", + [ + IO.Combo.Input( + "avatar", + options=HEYGEN_AVATAR_OPTIONS, + tooltip="Avatar look to present the video (curated from HeyGen's " + "public library). The best engine the look supports is chosen " + "automatically.", + ), + ], + ), + IO.DynamicCombo.Option( + "avatar_iv", + [ + IO.Combo.Input( + "avatar", + options=_AVATARS_BY_ENGINE["avatar_iv"], + tooltip="Avatar looks that support the Avatar IV engine.", + ), + ], + ), + IO.DynamicCombo.Option( + "avatar_iii", + [ + IO.Combo.Input( + "avatar", + options=_AVATARS_BY_ENGINE["avatar_iii"], + tooltip="Avatar looks that support the Avatar III engine.", + ), + ], + ), + IO.DynamicCombo.Option( + "avatar_v", + [ + IO.Combo.Input( + "avatar", + options=_AVATARS_BY_ENGINE["avatar_v"], + tooltip="Avatar looks that support the Avatar V engine.", + ), + ], + ), + ], + tooltip="Rendering engine; each choice lists only the avatars that support it. " + "'auto' offers every avatar and picks its best engine (Avatar IV preferred). " + "Avatar V is highest fidelity, Avatar III is the most affordable.", + ), + IO.String.Input( + "custom_avatar_id", + default="", + optional=True, + tooltip="Optional HeyGen avatar look ID. When set, overrides the avatar selected above. " + "Any of HeyGen's 3000+ public looks (or your private avatars) can be used.", + ), + IO.DynamicCombo.Input( + "speech", + display_name="speech source", + options=[ + IO.DynamicCombo.Option( + "script", + [ + IO.String.Input( + "text", + multiline=True, + default="", + tooltip="Text for the avatar to speak (up to 5000 characters). " + "The generated speech must be at least 1 second long.", + ), + IO.Combo.Input( + "voice", + options=[_DEFAULT_VOICE_OPTION] + HEYGEN_VOICE_GENERAL_OPTIONS, + tooltip="Voice for the script. The default option uses the voice HeyGen assigned to the avatar.", + ), + IO.String.Input( + "custom_voice_id", + default="", + optional=True, + tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. " + "Any voice from HeyGen's library (2000+) can be used.", + ), + IO.Float.Input( + "voice_speed", + default=1.0, + min=0.5, + max=1.5, + step=0.05, + optional=True, + tooltip="Speech speed multiplier.", + ), + ], + ), + IO.DynamicCombo.Option( + "audio", + [ + IO.Audio.Input( + "audio", + tooltip="Audio for the avatar to lip-sync, up to 10 minutes.", + ), + ], + ), + ], + tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.", + ), + IO.Combo.Input( + "resolution", + options=["720p", "1080p"], + default="1080p", + optional=True, + tooltip="Output video resolution.", + ), + IO.Combo.Input( + "aspect_ratio", + options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"], + default="auto", + optional=True, + tooltip="Output aspect ratio. 'auto' follows the avatar's source footage.", + ), + IO.String.Input( + "background_color", + default="", + optional=True, + tooltip="Optional solid background color as a hex code (e.g. '#00ff00'). " + "Leave empty for the avatar's own background.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + optional=True, + tooltip="Not sent to HeyGen; change it to force a re-run.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["engine"]), + expr=""" + widgets.engine = "avatar_iii" + ? {"type":"range_usd","min_usd":0.023881,"max_usd":0.061919,"format":{"suffix":"/second"}} + : widgets.engine = "avatar_v" + ? {"type":"usd","usd":0.095381,"format":{"suffix":"/second"}} + : widgets.engine = "avatar_iv" + ? {"type":"range_usd","min_usd":0.0715,"max_usd":0.095381,"format":{"suffix":"/second"}} + : {"type":"range_usd","min_usd":0.023881,"max_usd":0.095381,"format":{"suffix":"/second"}} + """, + ), + ) + + @classmethod + async def execute( + cls, + engine: dict, + speech: dict, + custom_avatar_id: str = "", + resolution: str = "1080p", + aspect_ratio: str = "auto", + background_color: str = "", + seed: int = 0, + ) -> IO.NodeOutput: + avatar_id, engine_type = await _resolve_avatar(cls, engine["avatar"], custom_avatar_id, engine["engine"]) + payload = { + "type": "avatar", + "avatar_id": avatar_id, + "resolution": resolution, + "aspect_ratio": aspect_ratio, + "title": "ComfyUI Avatar Video", + } + if engine_type: + payload["engine"] = {"type": engine_type} + background_color = background_color.strip() + if background_color: + if not background_color.startswith("#"): + raise ValueError("background_color must be a hex color code like '#00ff00'.") + payload["background"] = {"type": "color", "value": background_color} + await _apply_speech_source(cls, payload, speech, require_voice=False) + video = await _create_and_poll_video(cls, payload) + return IO.NodeOutput(await download_url_to_video_output(video["video_url"])) + + +class HeyGenCreateAvatarNode(IO.ComfyNode): + """Create a reusable HeyGen avatar from a photo or a text prompt.""" + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="HeyGenCreateAvatarNode", + display_name="HeyGen Create Avatar", + category="partner/video/HeyGen", + description="Create your own reusable HeyGen avatar from a photo of a person or " + "from a text prompt (a generated character). Feed the resulting avatar_id into " + "HeyGen Avatar Video's custom_avatar_id — and save the ID somewhere to reuse the " + "avatar in future workflows.", + inputs=[ + IO.DynamicCombo.Input( + "source", + options=[ + IO.DynamicCombo.Option( + "prompt", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Description of the avatar to generate (up to 1000 characters).", + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("ref_image"), + names=[f"ref_image_{i}" for i in range(1, 4)], + min=0, + ), + tooltip="Up to 3 reference images guiding the generated look.", + ), + ], + ), + IO.DynamicCombo.Option( + "photo", + [ + IO.Image.Input( + "identity_photo", + tooltip="Photo of the person to turn into an avatar. " + "Downscaled automatically if larger than 2K.", + ), + ], + ), + ], + tooltip="Generate a new character from a text prompt, or create the avatar " + "from a connected photo of a person.", + ), + ], + outputs=[ + IO.String.Output( + display_name="avatar_id", + tooltip="Avatar look ID. Pass it to HeyGen Avatar Video's custom_avatar_id; " + "save it to reuse the avatar later.", + ), + IO.Image.Output(display_name="preview"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":1.43}""", + ), + ) + + @classmethod + async def execute( + cls, + source: dict, + ) -> IO.NodeOutput: + payload: dict = {"name": "ComfyUI Avatar"} + if source["source"] == "photo": + image = downscale_image_tensor_by_max_side(source["identity_photo"], max_side=2000) + image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None) + payload["type"] = "photo" + payload["file"] = {"type": "url", "url": image_url} + else: + validate_string(source["prompt"], strip_whitespace=True, min_length=1, max_length=1000) + payload["type"] = "prompt" + payload["prompt"] = source["prompt"] + ref_tensors = [t for t in (source.get("reference_images") or {}).values() if t is not None] + if ref_tensors: + n_images = sum(get_number_of_images(t) for t in ref_tensors) + if n_images > 3: + raise ValueError(f"HeyGen accepts at most 3 reference images; got {n_images}.") + scaled = [downscale_image_tensor_by_max_side(t, max_side=2000) for t in ref_tensors] + ref_urls = await upload_images_to_comfyapi( + cls, scaled, max_images=3, mime_type="image/png", total_pixels=None + ) + payload["reference_images"] = [{"type": "url", "url": u} for u in ref_urls] + created = await sync_op_raw( + cls, + ApiEndpoint(path=_AVATARS_PATH, method="POST"), + data=payload, + ) + look_id = ((created.get("data") or {}).get("avatar_item") or {}).get("id") + if not look_id: + raise ValueError(f"HeyGen did not return an avatar: {created}") + final = await poll_op_raw( + cls, + ApiEndpoint(path=f"{_LOOKS_PATH}/{look_id}"), + # A missing status means the look needed no training and is ready. + status_extractor=lambda r: (r.get("data") or {}).get("status") or "completed", + failed_statuses=["failed", "pending_consent"], + poll_interval=5.0, + ) + data = final["data"] + if data.get("preview_image_url"): + preview = await download_url_to_image_tensor(data["preview_image_url"]) + else: + preview = torch.zeros(1, 64, 64, 3) + PromptServer.instance.send_progress_text( + f"Please save the avatar_id for reuse.\n\navatar_id: {look_id}", + cls.hidden.unique_id, + ) + return IO.NodeOutput(look_id, preview) + + +class HeyGenVideoTranslateNode(IO.ComfyNode): + """Translate a spoken video into another language with voice cloning and lip sync.""" + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="HeyGenVideoTranslateNode", + display_name="HeyGen Video Translate", + category="partner/video/HeyGen", + description="Translate a spoken video into another language. Clones the original " + "speaker's voice and re-animates the mouth to match the translated speech.", + inputs=[ + IO.Video.Input( + "video", + tooltip="Video with speech to translate.", + ), + IO.Combo.Input( + "output_language", + options=HEYGEN_TRANSLATE_LANGUAGES, + tooltip="Target language for the translated video.", + ), + IO.Combo.Input( + "mode", + options=["speed", "precision"], + default="speed", + tooltip="'speed' is faster; 'precision' produces higher-quality lip sync at twice the price.", + ), + IO.Boolean.Input( + "translate_audio_only", + default=False, + optional=True, + tooltip="Only swap the audio track, keeping the original mouth movements (no lip sync).", + ), + IO.Int.Input( + "speaker_count", + default=0, + min=0, + max=10, + optional=True, + tooltip="Number of speakers in the video. 0 = detect automatically.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + optional=True, + tooltip="Not sent to HeyGen; change it to force a re-run.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["mode"]), + expr="""{"type":"usd","usd": widgets.mode = "precision" ? 0.095381 : 0.047619,""" + """"format":{"suffix":"/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + video: Input.Video, + output_language: str, + mode: str, + translate_audio_only: bool = False, + speaker_count: int = 0, + seed: int = 0, + ) -> IO.NodeOutput: + video_url = await upload_video_to_comfyapi(cls, video) + payload = { + "video": {"type": "url", "url": video_url}, + "output_languages": [output_language], + "mode": mode, + "translate_audio_only": translate_audio_only, + "title": "ComfyUI Video Translate", + } + if speaker_count > 0: + payload["speaker_num"] = speaker_count + created = await sync_op_raw( + cls, + ApiEndpoint(path=_TRANSLATIONS_PATH, method="POST"), + data=payload, + ) + translation_ids = (created.get("data") or {}).get("video_translation_ids") or [] + if not translation_ids: + raise ValueError(f"HeyGen did not return a translation ID: {created}") + final = await poll_op_raw( + cls, + ApiEndpoint(path=f"{_TRANSLATIONS_PATH}/{translation_ids[0]}"), + status_extractor=lambda r: (r.get("data") or {}).get("status"), + queued_statuses=["pending"], + poll_interval=5.0, + ) + data = final["data"] + if not data.get("video_url"): + raise ValueError(f"HeyGen returned no video_url for translation {translation_ids[0]}.") + return IO.NodeOutput(await download_url_to_video_output(data["video_url"])) + + +class HeyGenTextToSpeechNode(IO.ComfyNode): + """Synthesize speech audio from text with HeyGen's Starfish TTS engine.""" + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="HeyGenTextToSpeechNode", + display_name="HeyGen Text to Speech", + category="partner/audio/HeyGen", + description="Generate speech audio from text using HeyGen's Starfish TTS engine. " + "Includes HeyGen's most popular voices across 17 languages.", + inputs=[ + IO.String.Input( + "text", + multiline=True, + default="", + tooltip="Text to synthesize (up to 5000 characters). The generated speech " + "must be at least 1 second long.", + ), + IO.Combo.Input( + "voice", + options=HEYGEN_VOICE_TTS_OPTIONS, + tooltip="Voice to use (curated from HeyGen's most popular Starfish-compatible voices).", + ), + IO.String.Input( + "custom_voice_id", + default="", + optional=True, + tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. " + "The voice must support the Starfish engine.", + ), + IO.Float.Input( + "speed", + default=1.0, + min=0.5, + max=2.0, + step=0.05, + optional=True, + tooltip="Speech speed multiplier.", + ), + IO.Boolean.Input( + "ssml", + default=False, + optional=True, + tooltip="Treat the text as SSML markup (for pauses, emphasis, and pronunciation control).", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + control_after_generate=True, + optional=True, + tooltip="Not sent to HeyGen; change it to force a re-run.", + ), + ], + outputs=[IO.Audio.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.00095381,"format":{"approximate":true,"suffix":"/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + text: str, + voice: str, + custom_voice_id: str = "", + speed: float = 1.0, + ssml: bool = False, + seed: int = 0, + ) -> IO.NodeOutput: + validate_string(text, strip_whitespace=True, min_length=1, max_length=5000) + payload = { + "text": text, + "voice_id": custom_voice_id.strip() or HEYGEN_VOICE_TTS_MAP[voice], + "speed": round(speed, 2), + } + if ssml: + payload["input_type"] = "ssml" + response = await sync_op_raw( + cls, + ApiEndpoint(path=_SPEECH_PATH, method="POST"), + data=payload, + ) + audio_url = (response.get("data") or {}).get("audio_url") + if not audio_url: + raise ValueError(f"HeyGen did not return an audio_url: {response}") + audio_bytes = await download_url_as_bytesio(audio_url) + return IO.NodeOutput(audio_bytes_to_audio_input(audio_bytes.getvalue())) + + +class HeyGenExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + HeyGenTalkingPhotoNode, + HeyGenAvatarVideoNode, + HeyGenCreateAvatarNode, + HeyGenVideoTranslateNode, + HeyGenTextToSpeechNode, + ] + + +async def comfy_entrypoint() -> HeyGenExtension: + return HeyGenExtension() From 4800e78518ebb1f2a9443ea5418edbff6c3935f9 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 17 Jul 2026 21:53:15 -0700 Subject: [PATCH 124/211] More comfy-kitchen int8 optimizations. (#14980) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index fbdd6c14a..6e2e895d3 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.21 +comfy-kitchen==0.2.22 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 7774301a7b1d953393a4b34454431d8174fffd6f Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Sat, 18 Jul 2026 23:44:22 +0800 Subject: [PATCH 125/211] chore: update workflow templates to v0.11.12 (#14986) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 6e2e895d3..20dc18dd6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.21 -comfyui-workflow-templates==0.11.11 +comfyui-workflow-templates==0.11.12 comfyui-embedded-docs==0.5.8 torch torchsde From 83082a51c420a364b15ea5f40d61da74e35b2da5 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sat, 18 Jul 2026 19:36:34 +0300 Subject: [PATCH 126/211] [Partner Nodes] fix(Google): switch to Interactions API for Omni model (#14984) --- comfy_api_nodes/apis/gemini.py | 59 ++++++++++++++- comfy_api_nodes/nodes_gemini.py | 130 ++++++++++++++++++++++---------- 2 files changed, 147 insertions(+), 42 deletions(-) diff --git a/comfy_api_nodes/apis/gemini.py b/comfy_api_nodes/apis/gemini.py index 7b2543270..ce89928cb 100644 --- a/comfy_api_nodes/apis/gemini.py +++ b/comfy_api_nodes/apis/gemini.py @@ -1,6 +1,6 @@ from datetime import date from enum import Enum -from typing import Any +from typing import Any, Literal from pydantic import BaseModel, Field @@ -242,3 +242,60 @@ class GeminiGenerateContentResponse(BaseModel): promptFeedback: GeminiPromptFeedback | None = Field(None) usageMetadata: GeminiUsageMetadata | None = Field(None) modelVersion: str | None = Field(None) + + +class GeminiInteractionTextPart(BaseModel): + type: Literal["text"] = "text" + text: str = Field(...) + + +class GeminiInteractionMediaPart(BaseModel): + type: str = Field(..., description="One of: image, video, audio, document.") + data: str | None = Field(None, description="Base64-encoded media bytes.") + uri: str | None = Field(None, description="URI of the media, as an alternative to inline data.") + mime_type: str | None = Field(None) + + +class GeminiInteractionGenerationConfig(BaseModel): + temperature: float | None = Field(None, ge=0.0, le=2.0) + top_p: float | None = Field(None, ge=0.0, le=1.0) + + +class GeminiInteractionRequest(BaseModel): + model: str = Field(...) + input: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = Field(...) + generation_config: GeminiInteractionGenerationConfig | None = Field(None) + + +class GeminiInteractionModalityTokens(BaseModel): + modality: str | None = Field(None, description="One of: text, image, audio, video, document.") + tokens: int | None = Field(None) + + +class GeminiInteractionUsage(BaseModel): + input_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None) + output_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None) + total_thought_tokens: int | None = Field(None) + + +class GeminiInteractionContent(BaseModel): + type: str | None = Field(None) + text: str | None = Field(None) + data: str | None = Field(None) + uri: str | None = Field(None) + mime_type: str | None = Field(None) + + +class GeminiInteractionStep(BaseModel): + type: str | None = Field(None) + content: list[GeminiInteractionContent] | None = Field(None) + + +class GeminiInteraction(BaseModel): + id: str | None = Field(None) + status: str | None = Field( + None, + description="One of: in_progress, requires_action, completed, failed, cancelled, incomplete.", + ) + steps: list[GeminiInteractionStep] | None = Field(None) + usage: GeminiInteractionUsage | None = Field(None) diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 283e2233d..8998b5943 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -24,6 +24,11 @@ from comfy_api_nodes.apis.gemini import ( GeminiImageGenerateContentRequest, GeminiImageGenerationConfig, GeminiInlineData, + GeminiInteraction, + GeminiInteractionGenerationConfig, + GeminiInteractionMediaPart, + GeminiInteractionRequest, + GeminiInteractionTextPart, GeminiMimeType, GeminiPart, GeminiRole, @@ -51,6 +56,7 @@ from comfy_api_nodes.util import ( ) GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini" +GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions" GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB GEMINI_URL_INPUT_BUDGET = 10 GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024 @@ -231,29 +237,10 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug return torch.cat(image_tensors, dim=0) -async def get_video_from_response( - response: GeminiGenerateContentResponse, cls: type[IO.ComfyNode] | None = None -) -> InputImpl.VideoFromFile: - parts = get_parts_by_type(response, "video/*") - for part in parts: - if part.inlineData and part.inlineData.data: - return InputImpl.VideoFromFile(BytesIO(base64.b64decode(part.inlineData.data))) - if part.fileData and part.fileData.fileUri: - return await download_url_to_video_output(part.fileData.fileUri, cls=cls) - model_message = get_text_from_response(response).strip() - if model_message: - raise ValueError(f"Gemini did not generate a video. Model response: {model_message}") - raise ValueError( - "Gemini did not generate a video. Try rephrasing your prompt, " - "shortening the requested duration, or reducing the number of input images/videos." - ) - - def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None: if not response.modelVersion: return None # Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing - output_video_tokens_price = 0.0 if response.modelVersion == "gemini-2.5-pro": input_tokens_price = 1.25 output_text_tokens_price = 10.0 @@ -290,11 +277,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N input_tokens_price = 0.25 output_text_tokens_price = 1.50 output_image_tokens_price = 30.0 - elif response.modelVersion == "gemini-omni-flash-preview": - input_tokens_price = 2.145 - output_text_tokens_price = 12.87 - output_image_tokens_price = 0.0 - output_video_tokens_price = 25.025 else: return None final_price = response.usageMetadata.promptTokenCount * input_tokens_price @@ -302,8 +284,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N for i in response.usageMetadata.candidatesTokensDetails: if i.modality == Modality.IMAGE: final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models - elif i.modality == Modality.VIDEO: - final_price += output_video_tokens_price * i.tokenCount # for Omni Flash else: final_price += output_text_tokens_price * i.tokenCount if response.usageMetadata.thoughtsTokenCount: @@ -311,6 +291,58 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N return final_price / 1_000_000.0 +def get_text_from_interaction(interaction: GeminiInteraction) -> str: + """Extract and concatenate all model output text from an Interactions API response.""" + texts = [] + for step in interaction.steps or []: + if step.type != "model_output": + continue + for content in step.content or []: + if content.type == "text" and content.text: + texts.append(content.text) + return "\n".join(texts) + + +async def get_video_from_interaction( + interaction: GeminiInteraction, cls: type[IO.ComfyNode] | None = None +) -> InputImpl.VideoFromFile: + for step in interaction.steps or []: + if step.type != "model_output": + continue + for content in step.content or []: + if content.type != "video": + continue + if content.data: + return InputImpl.VideoFromFile(BytesIO(base64.b64decode(content.data))) + if content.uri: + return await download_url_to_video_output(content.uri, cls=cls) + model_message = get_text_from_interaction(interaction).strip() + if model_message: + raise ValueError(f"Gemini did not generate a video. Model response: {model_message}") + raise ValueError( + "Gemini did not generate a video. Try rephrasing your prompt, " + "shortening the requested duration, or reducing the number of input images/videos." + ) + + +def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None: + if interaction.usage is None: + return None + input_tokens_price = 1.5 + output_tokens_prices = {"text": 9.0, "video": 17.5} + thoughts_tokens_price = 9.0 + final_price = 0.0 + for i in interaction.usage.input_tokens_by_modality or []: + if i.tokens: + final_price += input_tokens_price * i.tokens + for i in interaction.usage.output_tokens_by_modality or []: + if i.tokens and i.modality in output_tokens_prices: + final_price += output_tokens_prices[i.modality] * i.tokens + if interaction.usage.total_thought_tokens: + final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens + return final_price / 1_000_000.0 + + def create_video_parts(video_input: Input.Video) -> list[GeminiPart]: """Convert a single video input to Gemini API compatible parts (inline MP4/H.264).""" base_64_string = video_to_base64_string( @@ -445,6 +477,15 @@ async def build_gemini_media_parts( return parts +def to_interaction_media_part(part: GeminiPart) -> GeminiInteractionMediaPart: + """Convert a fileData/inlineData GeminiPart into an Interactions API media part.""" + if part.fileData: + mime = part.fileData.mimeType.value + return GeminiInteractionMediaPart(type=mime.split("/")[0], uri=part.fileData.fileUri, mime_type=mime) + mime = part.inlineData.mimeType.value + return GeminiInteractionMediaPart(type=mime.split("/")[0], data=part.inlineData.data, mime_type=mime) + + class GeminiNode(IO.ComfyNode): """ Node to generate text responses from a Gemini model. @@ -1676,7 +1717,7 @@ class GeminiVideoOmni(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr='{"type":"usd","usd":0.146,"format":{"suffix":"/second","approximate":true}}' + expr='{"type":"usd","usd":0.101,"format":{"suffix":"/second","approximate":true}}' ), ) @@ -1695,27 +1736,34 @@ class GeminiVideoOmni(IO.ComfyNode): for video in videos: validate_video_duration(video, max_duration=10) - parts: list[GeminiPart] = [] + parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = [] if images or videos: - parts.extend(await build_gemini_media_parts(cls, images, [], videos)) - parts.append(GeminiPart(text=prompt)) - response = await sync_op( + media_parts = await build_gemini_media_parts(cls, images, [], videos) + parts.extend(to_interaction_media_part(p) for p in media_parts) + parts.append(GeminiInteractionTextPart(text=prompt)) + interaction = await sync_op( cls, - ApiEndpoint(path=f"{GEMINI_BASE_ENDPOINT}/{model_id}", method="POST"), - data=GeminiGenerateContentRequest( - contents=[GeminiContent(role=GeminiRole.user, parts=parts)], - generationConfig=GeminiGenerationConfig( - responseModalities=["TEXT", "VIDEO"], + ApiEndpoint(path=GEMINI_INTERACTIONS_ENDPOINT, method="POST"), + data=GeminiInteractionRequest( + model=model_id, + input=parts, + generation_config=GeminiInteractionGenerationConfig( temperature=model.get("temperature", 1.0), - topP=model.get("top_p", 0.95), + top_p=model.get("top_p", 0.95), ), ), - response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, + response_model=GeminiInteraction, + price_extractor=calculate_interaction_tokens_price, ) + if interaction.status != "completed": + model_message = get_text_from_interaction(interaction).strip() + raise ValueError( + f"Gemini interaction did not complete (status: {interaction.status})." + + (f" Model response: {model_message}" if model_message else "") + ) return IO.NodeOutput( - await get_video_from_response(response, cls=cls), - get_text_from_response(response), + await get_video_from_interaction(interaction, cls=cls), + get_text_from_interaction(interaction), ) From c9602625e445e9ee37d3ac6faf5ea9ec1e0de87e Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 18 Jul 2026 17:12:18 -0700 Subject: [PATCH 127/211] Implement regular and timestep zero reference images to krea 2 for ostris and identity edit ref loras. (#14843) --- comfy/ldm/krea2/model.py | 147 +++++++++++++++++++++++++++++++++------ comfy/model_base.py | 23 ++++-- 2 files changed, 143 insertions(+), 27 deletions(-) diff --git a/comfy/ldm/krea2/model.py b/comfy/ldm/krea2/model.py index ecb16254f..8001812d7 100644 --- a/comfy/ldm/krea2/model.py +++ b/comfy/ldm/krea2/model.py @@ -15,6 +15,7 @@ from einops import rearrange import comfy.model_management import comfy.patcher_extension import comfy.ldm.common_dit +import comfy.utils from comfy.ldm.flux.layers import EmbedND, timestep_embedding from comfy.ldm.flux.math import apply_rope from comfy.ldm.modules.attention import optimized_attention_masked @@ -73,11 +74,20 @@ class Attention(nn.Module): self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) def forward(self, x, freqs=None, mask=None, transformer_options={}): + transformer_patches = transformer_options.get("patches", {}) + extra_options = transformer_options.copy() q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x) q = rearrange(q, "B L (H D) -> B H L D", H=self.heads) k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads) v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads) q, k = self.qknorm(q, k) + + if "block_index" in transformer_options and "attn1_patch" in transformer_patches: + for p in transformer_patches["attn1_patch"]: + out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options) + q, k, v = out.get("q", q), out.get("k", k), out.get("v", v) + freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask) + if freqs is not None: q, k = apply_rope(q, k, freqs) if self.kvheads != self.heads: @@ -86,6 +96,11 @@ class Attention(nn.Module): v = v.repeat_interleave(rep, dim=1) out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True, transformer_options=transformer_options) + + if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches: + for p in transformer_patches["attn1_output_patch"]: + out = p(out, extra_options) + return self.wo(out * F.sigmoid(gate)) @@ -158,8 +173,44 @@ class SingleStreamBlock(nn.Module): self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) - def forward(self, x, vec, freqs, mask=None, transformer_options={}): + def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}): prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec) + if timestep_zero_index is not None: + bs = x.shape[0] + ref_prescale = prescale[bs:] + ref_preshift = preshift[bs:] + ref_pregate = pregate[bs:] + ref_postscale = postscale[bs:] + ref_postshift = postshift[bs:] + ref_postgate = postgate[bs:] + prescale = prescale[:bs] + preshift = preshift[:bs] + pregate = pregate[:bs] + postscale = postscale[:bs] + postshift = postshift[:bs] + postgate = postgate[:bs] + + pre = self.prenorm(x) + pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift) + pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift) + attn = self.attn(pre, freqs, mask, transformer_options=transformer_options) + del pre + attn[:, :timestep_zero_index].mul_(pregate) + attn[:, timestep_zero_index:].mul_(ref_pregate) + x = x + attn + del attn + + post = self.postnorm(x) + post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift) + post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift) + mlp = self.mlp(post) + del post + mlp[:, :timestep_zero_index].mul_(postgate) + mlp[:, timestep_zero_index:].mul_(ref_postgate) + x = x + mlp + del mlp + return x + x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options) x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift) return x @@ -181,7 +232,7 @@ class LastLayer(nn.Module): class SingleStreamDiT(nn.Module): def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4, layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12, - txtheads=20, txtkvheads=20, image_model=None, + txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None, device=None, dtype=None, operations=None, **kwargs): super().__init__() self.dtype = dtype @@ -191,6 +242,7 @@ class SingleStreamDiT(nn.Module): self.heads = heads self.txtdim = txtdim self.txtlayers = txtlayers + self.default_ref_method = default_ref_method headdim = features // heads axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)] @@ -221,61 +273,110 @@ class SingleStreamDiT(nn.Module): operations.Linear(features, features * 6, device=device, dtype=dtype), ) - def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): + def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): return comfy.patcher_extension.WrapperExecutor.new_class_executor( self._forward, self, comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options), - ).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs) + ).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs) - def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs): + def process_img(self, x, index=0): + patch = self.patch + x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch)) + h, w = x.shape[-2] // patch, x.shape[-1] // patch + img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) + + img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32) + img_ids[..., 0] = index + img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None] + img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :] + return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w + + def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): + transformer_options = transformer_options.copy() temporal = x.ndim == 5 if temporal: b5, c5, t5, h5, w5 = x.shape x = x.reshape(b5 * t5, c5, h5, w5) - bs, c, H_orig, W_orig = x.shape + bs, _, h_orig, w_orig = x.shape patch = self.patch - # Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end. - x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch)) - H, W = x.shape[-2], x.shape[-1] - h_, w_ = H // patch, W // patch # context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim). context = self._unpack_context(context) - img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) + img, imgpos, h_, w_ = self.process_img(x) + img_tokens = img.shape[1] + timestep_zero_index = None + ref_method = kwargs.get("ref_latents_method", self.default_ref_method) + if ref_method is not None and ref_latents is not None and len(ref_latents) > 0: + ref_tokens = [] + ref_pos = [] + ref_num_tokens = [] + for index, ref in enumerate(ref_latents, 1): + if ref.ndim == 5: + rb, rc, rt, rh5, rw5 = ref.shape + ref = ref.reshape(rb * rt, rc, rh5, rw5) + ref = comfy.utils.repeat_to_batch_size(ref, bs) + kontext, kontext_ids, _, _ = self.process_img(ref, index=index) + ref_tokens.append(kontext) + ref_pos.append(kontext_ids) + ref_num_tokens.append(kontext.shape[1]) + img = torch.cat([img] + ref_tokens, dim=1) + imgpos = torch.cat([imgpos] + ref_pos, dim=1) + del ref_tokens, ref_pos + if ref_method == "index_timestep_zero": + timestep_zero_index = img_tokens + transformer_options["reference_image_num_tokens"] = ref_num_tokens + img = self.first(img) t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype)) tvec = self.tproj(t) + if timestep_zero_index is not None: + t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype)) + tvec = torch.cat((tvec, self.tproj(t0)), dim=0) context = self.txtfusion(context, mask=None, transformer_options=transformer_options) context = self.txtmlp(context) - txtlen, imglen = context.shape[1], img.shape[1] + txtlen = context.shape[1] + device = context.device + txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32) + + patches = transformer_options.get("patches", {}) + if "post_input" in patches: + for p in patches["post_input"]: + out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options}) + img, context = out["img"], out["txt"] + imgpos, txtpos = out["img_ids"], out["txt_ids"] + combined = torch.cat((context, img), dim=1) + del context, img + if timestep_zero_index is not None: + timestep_zero_index += txtlen # Position ids: text at 0, image at (0, h_idx, w_idx). - device = combined.device - txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32) - imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32) - imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None] - imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :] - imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1) pos = torch.cat((txtpos, imgpos), dim=1) + del txtpos, imgpos freqs = self.pe_embedder(pos) + del pos - for block in self.blocks: - combined = block(combined, tvec, freqs, None, transformer_options=transformer_options) + transformer_options["total_blocks"] = len(self.blocks) + transformer_options["block_type"] = "single" + transformer_options["img_slice"] = [txtlen, combined.shape[1]] + for i, block in enumerate(self.blocks): + transformer_options["block_index"] = i + combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options) final = self.last(combined, t) - out = final[:, txtlen:txtlen + imglen, :] + del combined + out = final[:, txtlen:txtlen + img_tokens, :] out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=h_, w=w_, ph=patch, pw=patch, c=self.channels) - out = out[:, :, :H_orig, :W_orig] # crop padding back off + out = out[:, :, :h_orig, :w_orig] # crop padding back off if temporal: - out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2) + out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2) return out def _unpack_context(self, context): diff --git a/comfy/model_base.py b/comfy/model_base.py index 98f5ba48b..0f705316c 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2227,10 +2227,7 @@ class Omnigen2(BaseModel): out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) ref_latents = kwargs.get("reference_latents", None) if ref_latents is not None: - latents = [] - for lat in ref_latents: - latents.append(self.process_latent_in(lat)) - out['ref_latents'] = comfy.conds.CONDList(latents) + out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents]) return out def extra_conds_shapes(self, **kwargs): @@ -2317,12 +2314,30 @@ class Ideogram4(BaseModel): class Krea2(BaseModel): def __init__(self, model_config, model_type=ModelType.FLUX, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT) + self.memory_usage_factor_conds = ("ref_latents",) def extra_conds(self, **kwargs): out = super().extra_conds(**kwargs) cross_attn = kwargs.get("cross_attn", None) if cross_attn is not None: out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + latents = [] + for lat in ref_latents: + latents.append(self.process_latent_in(lat)) + out['ref_latents'] = comfy.conds.CONDList(latents) + + ref_latents_method = kwargs.get("reference_latents_method", None) + if ref_latents_method is not None: + out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method) + return out + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) return out class HunyuanImage21(BaseModel): From 66655153499f89052aa72d5a869f556b25f0e9c6 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sun, 19 Jul 2026 15:13:49 -0700 Subject: [PATCH 128/211] Fix wan dancer issue with batches. (#14999) --- comfy/model_base.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/comfy/model_base.py b/comfy/model_base.py index 0f705316c..3494925be 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2024,11 +2024,11 @@ class WAN22_WanDancer(WAN21): fps = kwargs.get("fps", None) if fps is not None: - out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps])) + out['fps'] = comfy.conds.CONDConstant(fps) audio_inject_scale = kwargs.get("audio_inject_scale", None) if audio_inject_scale is not None: - out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale])) + out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale) return out class Hunyuan3Dv2(BaseModel): From ecba6f2594755f8d9440d517156771d098b71ba6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Tue, 21 Jul 2026 02:33:26 +0300 Subject: [PATCH 129/211] feat: Support Gemma4 12B (CORE-277) (#14304) --- comfy/sd.py | 9 +- comfy/text_encoders/gemma4.py | 311 +++++++++++++++++++++++++++++----- comfy/text_encoders/llama.py | 4 +- 3 files changed, 276 insertions(+), 48 deletions(-) diff --git a/comfy/sd.py b/comfy/sd.py index 9d7fa731f..e15e0a9fd 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -1434,6 +1434,7 @@ class TEModel(Enum): GPT_OSS_20B = 33 QWEN3VL_4B = 34 QWEN3VL_8B = 35 + GEMMA_4_12B = 36 def detect_te_model(sd): @@ -1463,6 +1464,9 @@ def detect_te_model(sd): if 'model.layers.0.post_feedforward_layernorm.weight' in sd: if 'model.layers.59.self_attn.q_norm.weight' in sd: return TEModel.GEMMA_4_31B + # Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v). + if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd: + return TEModel.GEMMA_4_12B if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd: return TEModel.GEMMA_4_E4B if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd: @@ -1618,10 +1622,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) - elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B): + elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B): variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, - TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model] + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) clip_target.tokenizer = variant.tokenizer tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None) diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 0bba8341b..5163c1676 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1,11 +1,15 @@ import torch import torch.nn as nn +import torchaudio.functional as AF +import torchvision.transforms.functional as TVF import numpy as np +from tokenizers import Tokenizer from dataclasses import dataclass import math from comfy import sd1_clip import comfy.model_management +import comfy.ops from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.rmsnorm import rms_norm from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding @@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3} GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5} +# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space. +GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6} +GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6} + @dataclass class Gemma4Config: vocab_size: int = 262144 @@ -35,6 +43,9 @@ class Gemma4Config: transformer_type: str = "gemma4" head_dim = 256 global_head_dim = 512 + num_global_key_value_heads = None + attention_k_eq_v = False + vision_bidirectional = False rms_norm_add = False mlp_activation = "gelu_pytorch_tanh" qkv_bias = False @@ -51,6 +62,7 @@ class Gemma4Config: num_kv_shared_layers: int = 18 use_double_wide_mlp: bool = False stop_tokens = [1, 50, 106] + suppress_tokens = [] vision_config = GEMMA4_VISION_CONFIG audio_config = GEMMA4_AUDIO_CONFIG mm_tokens_per_image = 280 @@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config): num_hidden_layers: int = 60 num_attention_heads: int = 32 num_key_value_heads: int = 16 + vision_bidirectional = True sliding_attention = [1024, 1024, 1024, 1024, 1024, False] hidden_size_per_layer_input: int = 0 num_kv_shared_layers: int = 0 audio_config = None vision_config = GEMMA4_VISION_31B_CONFIG +@dataclass +class Gemma4_12B_Config(Gemma4Config): + hidden_size: int = 3840 + intermediate_size: int = 15360 + num_hidden_layers: int = 48 + num_attention_heads: int = 16 + num_key_value_heads: int = 8 + num_global_key_value_heads = 1 + attention_k_eq_v = True + vision_bidirectional = True + sliding_attention = [1024, 1024, 1024, 1024, 1024, False] + hidden_size_per_layer_input: int = 0 + num_kv_shared_layers: int = 0 + audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG + vision_config = GEMMA4_UNIFIED_VISION_CONFIG + suppress_tokens = [258883, 258882] + # unfused RoPE as addcmul_ RoPE diverges from reference code def _apply_rotary_pos_emb(x, freqs_cis): @@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis): return out class Gemma4Attention(nn.Module): - def __init__(self, config, head_dim, device=None, dtype=None, ops=None): + def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None): super().__init__() self.num_heads = config.num_attention_heads - self.num_kv_heads = config.num_key_value_heads + self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads self.hidden_size = config.hidden_size self.head_dim = head_dim self.inner_size = self.num_heads * head_dim self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype) self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) - self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) + # k_eq_v: V reuses the K projection (no separate v_proj weight) + self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype) self.q_norm = None @@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module): shareable_kv = None else: xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) - xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) + if self.v_proj is not None: + xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) + else: + xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE) if self.k_norm is not None: xk = self.k_norm(xk) xv = rms_norm(xv) @@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module): head_dim = config.head_dim if self.sliding_attention else config.global_head_dim - self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops) + # k_eq_v only on global layers, which then use num_global_key_value_heads + k_eq_v = config.attention_k_eq_v and not self.sliding_attention + num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads + self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops) num_kv_shared = config.num_kv_shared_layers first_kv_shared = config.num_hidden_layers - num_kv_shared @@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module): self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype) self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype) self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype) - self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype)) - else: - self.layer_scalar = None + + # layer_scalar exists on every gemma4 variant, independent of per-layer input + self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype)) def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None): sliding_window = None @@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module): x = self.post_per_layer_input_norm(x) x = residual + x - if self.layer_scalar is not None: - x = x * self.layer_scalar + x = x * comfy.ops.cast_to_input(self.layer_scalar, x) return x, present_key_value, shareable_kv @@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module): causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val) mask = mask + causal_mask if mask is not None else causal_mask + # Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal). + if self.config.vision_bidirectional and past_len == 0 and embeds_info: + block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device) + group = 0 + for info in embeds_info: + if info.get("type") == "image": + start = info["index"] + block_ids[start:start + info["size"]] = group + group += 1 + if group > 0: + same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0) + mask = mask.masked_fill(same_block, 0.0) + # Per-layer inputs per_layer_inputs = None if self.hidden_size_per_layer_input: @@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module): shared_global_kv = None # KV from last non-shared global layer intermediate = None + all_intermediate = None + only_layers = None + if intermediate_output is not None: + if isinstance(intermediate_output, list): + all_intermediate = [] + only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output} + elif intermediate_output == "all": + all_intermediate = [] + intermediate_output = None + elif intermediate_output < 0: + intermediate_output = len(self.layers) + intermediate_output + next_key_values = [] for i, layer in enumerate(self.layers): + if all_intermediate is not None: + if only_layers is None or (i in only_layers): + all_intermediate.append(x.unsqueeze(1).clone()) + past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None layer_kwargs = {} @@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module): if self.norm is not None: x = self.norm(x) - if len(next_key_values) > 0: + if all_intermediate is not None: + if only_layers is None or (len(self.layers) in only_layers): + all_intermediate.append(x.unsqueeze(1).clone()) + if len(all_intermediate) > 0: + intermediate = torch.cat(all_intermediate, dim=1) + + if intermediate is not None and final_layer_norm_intermediate and self.norm is not None: + intermediate = self.norm(intermediate) + + # Only hand back the KV cache when caching was actually requested; SDClipModel reads + # outputs[2] as the pooled output. + if past_key_values is not None and len(next_key_values) > 0: return x, intermediate, next_key_values return x, intermediate @@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module): cap = self.model.config.final_logit_softcapping if cap: logits = cap * torch.tanh(logits / cap) + if self.model.config.suppress_tokens: + logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min return logits def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): @@ -441,6 +519,28 @@ class Gemma4AudioMixin: return None, None +class Gemma4UnifiedBase(Gemma4Base): + """Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space.""" + def _init_model(self, config, dtype, device, operations): + self.num_layers = config.num_hidden_layers + self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations) + self.dtype = dtype + self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations) + self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations) + self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations) + + def preprocess_embed(self, embed, device): + if embed["type"] == "image": + pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1] + patches, positions = self.vision_model.patchify(pixels) + vision_out = self.vision_model(patches, positions) + return self.multi_modal_projector(vision_out), None + if embed["type"] == "audio": + audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token] + return self.audio_projector(audio), None + return None, None + + # Vision Encoder def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None): @@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector): super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops) +# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space. + +def _patches_merge(patches, positions_xy, length): + patch_size = math.isqrt(patches.shape[-1] // 3) + k = math.isqrt(patches.shape[-2] // length) + batch = patches.shape[:-2] + + max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1 + kidx = torch.div(positions_xy, k, rounding_mode="floor") + rem = torch.remainder(positions_xy, k) + order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1] + perm = order.long().argsort(dim=-1) + + merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches)) + merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3) + merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3) + + pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy)) + pad = (positions_xy == -1).all(dim=-1, keepdim=True) + pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2) + pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0] + return merged, pos + + +class Gemma4UnifiedVisionEmbedder(nn.Module): + """Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector.""" + def __init__(self, config, device=None, dtype=None, ops=None): + super().__init__() + self.patch_size = config["patch_size"] + self.pooling_kernel_size = config["pooling_kernel_size"] + patch_dim = config["model_patch_size"] ** 2 * 3 + mm_embed_dim = config["mm_embed_dim"] + self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype) + self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype) + self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype) + self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype)) + self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype) + + def patchify(self, pixels): + """pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2].""" + ps, k = self.patch_size, self.pooling_kernel_size + out_patches, out_positions = [], [] + for img in pixels: + ph, pw = img.shape[-2] // ps, img.shape[-1] // ps + teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1) + grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy") + tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2) + n_model = teacher.shape[0] // (k * k) + mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model) + out_patches.append(mp.squeeze(0)) + out_positions.append(mpos.squeeze(0)) + return torch.stack(out_patches), torch.stack(out_positions) + + def forward(self, pixel_values, image_position_ids): + x = self.patch_ln1(pixel_values) + x = self.patch_dense(x) + x = self.patch_ln2(x) + + clamped = image_position_ids.clamp(min=0).long() + valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1) + axes = torch.arange(2, device=image_position_ids.device) + pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype) + pos_embs = (pos[clamped, axes] * valid).sum(-2) + x = x + pos_embs + return self.pos_norm(x) + + # Audio Encoder class Gemma4AudioConvSubsampler(nn.Module): @@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector): # Tokenizer and Wrappers +def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size): + target_px = max_patches * patch_size ** 2 + factor = math.sqrt(target_px / (height * width)) + side_mult = pooling_kernel_size * patch_size + target_height = math.floor(factor * height / side_mult) * side_mult + target_width = math.floor(factor * width / side_mult) * side_mult + + if target_height == 0 and target_width == 0: + raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.") + + max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult + if target_height == 0: + target_height = side_mult + target_width = min(math.floor(width / height) * side_mult, max_side_length) + elif target_width == 0: + target_width = side_mult + target_height = min(math.floor(height / width) * side_mult, max_side_length) + + if target_height * target_width > target_px: + raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.") + + return target_height, target_width + + class Gemma4_Tokenizer(): tokenizer_json_data = None @@ -998,25 +1189,35 @@ class Gemma4_Tokenizer(): return {"tokenizer_json": self.tokenizer_json_data} return {} - def _extract_mel_spectrogram(self, waveform, sample_rate): - """Extract 128-bin log mel spectrogram. - Uses numpy for FFT/matmul/log to produce bit-identical results with reference code. - """ - # Mix to mono first, then resample to 16kHz + def _audio_token_count(self, num_samples): + # Default (E2B/E4B): mel frames after two stride-2 conv subsamples. + _fl = 320 # int(round(16000 * 20.0 / 1000.0)) + _hl = 160 # int(round(16000 * 10.0 / 1000.0)) + _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1 + _t = _nmel + for _ in range(2): + _t = (_t + 2 - 3) // 2 + 1 + return min(_t, 750) + + @staticmethod + def _resample_16k(waveform, sample_rate): + """Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers + load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio.""" if waveform.dim() > 1 and waveform.shape[0] > 1: waveform = waveform.mean(dim=0, keepdim=True) if waveform.dim() == 1: waveform = waveform.unsqueeze(0) - audio = waveform.squeeze(0).float().numpy() + audio = waveform.float() if sample_rate != 16000: - # Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match) - from scipy.signal import resample_poly, firwin - from math import gcd - g = gcd(sample_rate, 16000) - up, down = 16000 // g, sample_rate // g - L = max(up, down) - h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5)) - audio = resample_poly(audio, up, down, window=h).astype(np.float32) + audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser", + lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614) + return audio.squeeze(0).contiguous() + + def _extract_audio_features(self, waveform, sample_rate): + """Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder. + Uses numpy for FFT/matmul/log to produce bit-identical results with reference code. + """ + audio = self._resample_16k(waveform, sample_rate).numpy() n = len(audio) # Pad to multiple of 128, build sample-level mask @@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer(): if audio is not None: waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000 - mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate) - audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T]) + feat, feat_mask = self._extract_audio_features(waveform, sample_rate) + audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T]) # Process image/video frames is_video = video is not None @@ -1090,13 +1291,8 @@ class Gemma4_Tokenizer(): pooling_k = 3 max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280) max_patches = max_soft_tokens * pooling_k * pooling_k - target_px = max_patches * patch_size * patch_size - factor = (target_px / (h * w)) ** 0.5 - side_mult = pooling_k * patch_size - target_h = max(int(factor * h // side_mult) * side_mult, side_mult) - target_w = max(int(factor * w // side_mult) * side_mult, side_mult) + target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k) - import torchvision.transforms.functional as TVF for i in range(num_frames): # rescaling to match reference code s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8 @@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer(): llama_text = llama_template.format(text) else: # Build template from modalities present - system = "<|turn>system\n<|think|>\n" if thinking else "" + system = "<|turn>system\n<|think|>\n\n" if thinking else "" media = "" if len(images) > 0: if is_video: @@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer(): if len(audio_features) > 0: # Compute audio token count (always at 16kHz) num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1] - _fl = 320 # int(round(16000 * 20.0 / 1000.0)) - _hl = 160 # int(round(16000 * 10.0 / 1000.0)) - _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1 - _t = _nmel - for _ in range(2): - _t = (_t + 2 - 3) // 2 + 1 - n_audio_tokens = min(_t, 750) + n_audio_tokens = self._audio_token_count(num_samples) media += "<|audio>" + "<|audio|>" * n_audio_tokens + "" - llama_text = f"{system}<|turn>user\n{media}{text}\n<|turn>model\n" + # Non-thinking mode primes an empty thought channel so the model answers directly. + model_open = "" if thinking else "<|channel>thought\n" + llama_text = f"{system}<|turn>user\n{text}{media}\n<|turn>model\n{model_open}" text_tokens = super().tokenize_with_weights(llama_text, return_word_ids) @@ -1178,7 +1370,6 @@ class Gemma4_Tokenizer(): class _Gemma4Tokenizer: """Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)""" def __init__(self, tokenizer_json_bytes=None, **kwargs): - from tokenizers import Tokenizer if isinstance(tokenizer_json_bytes, torch.Tensor): tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist()) self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8")) @@ -1224,6 +1415,30 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class) +class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer): + """Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram.""" + embedding_size = 3840 + + def _extract_audio_features(self, waveform, sample_rate): + audio = self._resample_16k(waveform, sample_rate) + spt = 640 # audio_samples_per_token (40ms at 16kHz) + pad = (-audio.shape[0]) % spt + if pad: + audio = torch.nn.functional.pad(audio, (0, pad)) + num_tokens = audio.shape[0] // spt + feats = audio[:num_tokens * spt].reshape(num_tokens, spt) + feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s) + mask = torch.ones(feats.shape[0], dtype=torch.bool) + return feats, mask + + def _audio_token_count(self, num_samples): + return min((num_samples + 639) // 640, 750) + + +class Gemma4UnifiedTokenizer(Gemma4Tokenizer): + tokenizer_class = Gemma4UnifiedSDTokenizer + + # Model wrappers class Gemma4Model(sd1_clip.SDClipModel): model_class = None @@ -1256,7 +1471,7 @@ class Gemma4Model(sd1_clip.SDClipModel): expanded_idx += 1 initial_token_ids = [ids] input_ids = torch.tensor(initial_token_ids, device=self.execution_device) - return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids) + return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info) def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): @@ -1296,3 +1511,11 @@ def _make_variant(config_cls): Gemma4_E4B = _make_variant(Gemma4Config) Gemma4_E2B = _make_variant(Gemma4_E2B_Config) Gemma4_31B = _make_variant(Gemma4_31B_Config) + + +# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant). +class Gemma4_12B(Gemma4UnifiedBase): + def __init__(self, config_dict, dtype, device, operations): + super().__init__() + self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations) +Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index 3f98fb0a5..40d04007e 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -876,7 +876,7 @@ class BaseGenerate: torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0)) return past_key_values - def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None): + def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None): device = embeds.device if stop_tokens is None: @@ -911,7 +911,7 @@ class BaseGenerate: if step == 0 and deepstack_embeds is not None: extra["deepstack_embeds"] = deepstack_embeds extra["visual_pos_masks"] = visual_pos_masks - x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra) + x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None)) logits = self.logits(x)[:, -1] next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty) token_id = next_token[0].item() From 35c94d6023cab38a557f707b3a2ddd8ed72226c8 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 20 Jul 2026 20:36:03 -0700 Subject: [PATCH 130/211] Fix gfx1035 not being treated like RDNA2 (#15009) --- comfy/model_management.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index 222005b6f..766e9ea89 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -473,7 +473,7 @@ except: SUPPORT_FP8_OPS = args.supports_fp8_compute -AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"] +AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"] AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN' try: From 0384bb25f47ec7f6a3aa724e7141ba0699d71bbf Mon Sep 17 00:00:00 2001 From: Matt Miller Date: Mon, 20 Jul 2026 20:46:05 -0700 Subject: [PATCH 131/211] chore: add /AGENTS.md to CODEOWNERS (#14962) Scope AGENTS.md review to @comfyanonymous, matching the existing /CODEOWNERS, /.ci/, and /.github/ meta-file entries. --- CODEOWNERS | 1 + 1 file changed, 1 insertion(+) diff --git a/CODEOWNERS b/CODEOWNERS index 043c0ec75..634927dd6 100644 --- a/CODEOWNERS +++ b/CODEOWNERS @@ -1,5 +1,6 @@ * @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai /CODEOWNERS @comfyanonymous +/AGENTS.md @comfyanonymous /.ci/ @comfyanonymous /.github/ @comfyanonymous From d0fec2ef7e7086533fde261de3fdb88289bdca9e Mon Sep 17 00:00:00 2001 From: Kohaku-Blueleaf <59680068+KohakuBlueleaf@users.noreply.github.com> Date: Tue, 21 Jul 2026 12:02:54 +0800 Subject: [PATCH 132/211] [Trainer,Dataset/Feature] Video processing nodes, Image Processing Node video support, trainer video support (CORE-81) (#13588) --- comfy_extras/nodes_dataset.py | 435 +++++++++++++++++++++++++++++++++- comfy_extras/nodes_train.py | 5 +- 2 files changed, 434 insertions(+), 6 deletions(-) diff --git a/comfy_extras/nodes_dataset.py b/comfy_extras/nodes_dataset.py index 73fe75b7f..d7e4652cf 100644 --- a/comfy_extras/nodes_dataset.py +++ b/comfy_extras/nodes_dataset.py @@ -2,6 +2,7 @@ import logging import os import json +import av import numpy as np import torch from PIL import Image @@ -9,7 +10,7 @@ from typing_extensions import override import folder_paths import node_helpers -from comfy_api.latest import ComfyExtension, io +from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types def load_and_process_images(image_files, input_dir): @@ -42,6 +43,38 @@ def load_and_process_images(image_files, input_dir): return output_images +VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"] + + +def _decode_selected_frames(video: Input.Video, indices: list[int]) -> Input.Video: + """Decode only the requested frame indices from a video. + + Opens the underlying container once, decodes frames in presentation order, + keeps only the ones whose index is in ``indices``, and returns the result + wrapped in a VideoFromComponents so it still satisfies the VideoInput + contract for downstream nodes. + """ + indices_sorted = sorted(set(indices)) + max_idx = indices_sorted[-1] + source = video.get_stream_source() + + frames_by_idx: dict[int, torch.Tensor] = {} + with av.open(source, mode="r") as container: + stream = container.streams.video[0] + wanted = set(indices_sorted) + for frame_idx, frame in enumerate(container.decode(stream)): + if frame_idx in wanted: + img = frame.to_ndarray(format="rgb24") + frames_by_idx[frame_idx] = torch.from_numpy(img.copy()).float() / 255.0 + if frame_idx >= max_idx: + break + + stacked = torch.stack([frames_by_idx[i] for i in indices]) + return InputImpl.VideoFromComponents( + Types.VideoComponents(images=stacked, frame_rate=video.get_frame_rate()) + ) + + class LoadImageDataSetFromFolderNode(io.ComfyNode): @classmethod def define_schema(cls): @@ -157,6 +190,116 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode): return io.NodeOutput(output_tensor, captions) +class LoadVideoDataSetFromFolderNode(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LoadVideoDataSetFromFolder", + search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"], + display_name="Load Video (from Folder)", + category="video", + description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.", + is_experimental=True, + inputs=[ + io.Combo.Input( + "folder", + options=folder_paths.get_input_subfolders(), + tooltip="The folder containing video files.", + ), + ], + outputs=[ + io.Video.Output( + display_name="videos", + is_output_list=True, + tooltip="Lazy video references; frames are decoded only when needed downstream.", + ), + ], + ) + + @classmethod + def execute(cls, folder): + sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder) + video_files = sorted([ + f for f in os.listdir(sub_input_dir) + if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS) + ]) + + if not video_files: + raise ValueError(f"No video files found in {sub_input_dir}") + + videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files] + logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}") + return io.NodeOutput(videos) + + +class LoadVideoTextDataSetFromFolderNode(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LoadVideoTextDataSetFromFolder", + search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"], + display_name="Load Video-Text (from Folder)", + category="video", + description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.", + is_experimental=True, + inputs=[ + io.Combo.Input( + "folder", + options=folder_paths.get_input_subfolders(), + tooltip="The folder containing video files and .txt captions.", + ), + ], + outputs=[ + io.Video.Output( + display_name="videos", + is_output_list=True, + tooltip="Lazy video references; frames are decoded only when needed downstream.", + ), + io.String.Output( + display_name="texts", + is_output_list=True, + tooltip="List of text captions.", + ), + ], + ) + + @classmethod + def execute(cls, folder): + sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder) + + video_files = [] + for item in sorted(os.listdir(sub_input_dir)): + path = os.path.join(sub_input_dir, item) + if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS): + video_files.append(path) + elif os.path.isdir(path): + # Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/ + repeat = 1 + if item.split("_")[0].isdigit(): + repeat = int(item.split("_")[0]) + video_files.extend([ + os.path.join(path, f) + for f in sorted(os.listdir(path)) + if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS) + ] * repeat) + + if not video_files: + raise ValueError(f"No video files found in {sub_input_dir}") + + captions = [] + for vf in video_files: + caption_path = os.path.splitext(vf)[0] + ".txt" + if os.path.exists(caption_path): + with open(caption_path, "r", encoding="utf-8") as f: + captions.append(f.read().strip()) + else: + captions.append("") + + videos = [InputImpl.VideoFromFile(vf) for vf in video_files] + logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}") + return io.NodeOutput(videos, captions) + + def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True): """Utility function to save a list of image tensors to disk. @@ -470,7 +613,15 @@ class ImageProcessingNode(io.ComfyNode): @classmethod def execute(cls, images, **kwargs): - """Execute the node. Routes to _process or _group_process based on mode.""" + """Execute the node. Routes to _process or _group_process based on mode. + + For individual processing (_process), automatically handles multi-frame + inputs (video tensors [T, H, W, C]) by applying _process per-frame and + concatenating the results. This allows all spatial transform nodes to + work with video without modification. Nodes that natively handle batched + tensors (e.g. pure tensor math) can set per_frame_process = False to + skip the per-frame loop. + """ is_group = cls._detect_processing_mode() if is_group: @@ -489,7 +640,16 @@ class ImageProcessingNode(io.ComfyNode): result = cls._group_process(images, **params) else: # Individual processing: images is single item, call _process - result = cls._process(images, **params) + # Auto-loop over frames for multi-frame inputs (video [T, H, W, C]) + # so that PIL-based spatial transforms work per-frame automatically. + if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True): + results = [] + for i in range(images.shape[0]): + frame_result = cls._process(images[i:i + 1], **params) + results.append(frame_result) + result = torch.cat(results, dim=0) + else: + result = cls._process(images, **params) return io.NodeOutput(result) @@ -803,6 +963,7 @@ class NormalizeImagesNode(ImageProcessingNode): display_name = "Normalize Image Colors" category = "image/color" description = "Normalize images using mean and standard deviation." + per_frame_process = False # Pure tensor math, handles any batch size extra_inputs = [ io.Float.Input( "mean", @@ -833,6 +994,7 @@ class AdjustBrightnessNode(ImageProcessingNode): display_name = "Adjust Brightness" category="image/adjustments" description = "Adjust the brightness of an image." + per_frame_process = False # Pure tensor math, handles any batch size extra_inputs = [ io.Float.Input( "factor", @@ -854,6 +1016,7 @@ class AdjustContrastNode(ImageProcessingNode): display_name = "Adjust Contrast" category="image/adjustments" description = "Adjust the contrast of an image." + per_frame_process = False # Pure tensor math, handles any batch size extra_inputs = [ io.Float.Input( "factor", @@ -935,6 +1098,261 @@ class ShuffleImageTextDatasetNode(io.ComfyNode): return io.NodeOutput(shuffled_images, shuffled_texts) +# ========== Video Processing Nodes ========== + + +class VideoFrameSampleNode(io.ComfyNode): + """Sample a fixed number of frames from a video using various strategies. + + For contiguous strategies ("head"/"tail") the result is a fully lazy + VideoInput (no frames decoded). For non-contiguous strategies + ("uniform"/"random") only the selected indices are decoded. + """ + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="VideoFrameSample", + search_aliases=["sample frames", "extract frames"], + display_name="Sample Video Frame", + category="video", + description="Sample a fixed number of frames from a video using various strategies.", + is_experimental=True, + inputs=[ + io.Video.Input("video", tooltip="Input video."), + io.Int.Input( + "num_frames", + default=16, + min=1, + max=9999, + tooltip="Number of frames to sample.", + ), + io.Combo.Input( + "strategy", + options=["uniform", "head", "tail", "random"], + default="uniform", + tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.", + ), + io.Int.Input( + "seed", + default=0, + min=0, + max=0xFFFFFFFFFFFFFFFF, + tooltip="Random seed (only used with 'random' strategy).", + ), + ], + outputs=[ + io.Video.Output(display_name="video", tooltip="Sampled video."), + ], + ) + + @classmethod + def execute(cls, video, num_frames, strategy, seed): + total_frames = video.get_frame_count() + num_frames = min(num_frames, total_frames) + fps = float(video.get_frame_rate()) + + if strategy == "head": + return io.NodeOutput( + video.as_trimmed(0.0, num_frames / fps, strict_duration=False) + ) + if strategy == "tail": + start_t = (total_frames - num_frames) / fps + return io.NodeOutput( + video.as_trimmed(start_t, num_frames / fps, strict_duration=False) + ) + + if strategy == "uniform": + if num_frames == 1: + indices = [total_frames // 2] + else: + indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)] + elif strategy == "random": + rng = np.random.RandomState(seed % (2**32 - 1)) + indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist()) + else: + raise ValueError(f"Unknown strategy: {strategy}") + + return io.NodeOutput(_decode_selected_frames(video, indices)) + + +class VideoTemporalCropNode(io.ComfyNode): + """Crop a continuous range of frames from a video (fully lazy).""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="VideoTemporalCrop", + search_aliases=["crop", "crop video", "temporal crop", "truncate video"], + display_name="Crop Video (Temporal)", + category="video/transform", + description="Crop a continuous range of frames from a video.", + is_experimental=True, + inputs=[ + io.Video.Input("video", tooltip="Input video."), + io.Int.Input( + "start_frame", + default=0, + min=0, + max=99999, + tooltip="Starting frame index.", + ), + io.Int.Input( + "length", + default=16, + min=1, + max=99999, + tooltip="Number of frames to keep.", + ), + ], + outputs=[ + io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."), + ], + ) + + @classmethod + def execute(cls, video, start_frame, length): + total_frames = video.get_frame_count() + fps = float(video.get_frame_rate()) + start_frame = min(start_frame, max(total_frames - 1, 0)) + length = min(length, total_frames - start_frame) + return io.NodeOutput( + video.as_trimmed(start_frame / fps, length / fps, strict_duration=False) + ) + + +class VideoRandomTemporalCropNode(io.ComfyNode): + """Randomly crop a continuous range of frames from a video (fully lazy).""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="VideoRandomTemporalCrop", + search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"], + display_name="Crop Video (Temporal Random)", + category="video/transform", + description="Randomly crop a continuous range of frames from a video.", + is_experimental=True, + inputs=[ + io.Video.Input("video", tooltip="Input video."), + io.Int.Input( + "length", + default=16, + min=1, + max=99999, + tooltip="Number of frames to keep.", + ), + io.Int.Input( + "seed", + default=0, + min=0, + max=0xFFFFFFFFFFFFFFFF, + tooltip="Random seed.", + ), + ], + outputs=[ + io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."), + ], + ) + + @classmethod + def execute(cls, video, length, seed): + total_frames = video.get_frame_count() + fps = float(video.get_frame_rate()) + length = min(length, total_frames) + max_start = total_frames - length + rng = np.random.RandomState(seed % (2**32 - 1)) + start = rng.randint(0, max_start + 1) if max_start > 0 else 0 + return io.NodeOutput( + video.as_trimmed(start / fps, length / fps, strict_duration=False) + ) + + +class ShuffleVideoDatasetNode(io.ComfyNode): + """Randomly shuffle the order of videos in the dataset.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ShuffleVideoDataset", + search_aliases=["shuffle", "randomize", "mix"], + display_name="Shuffle Videos List", + category="video/batch", + description="Randomly shuffle the order of videos in a list.", + is_experimental=True, + is_input_list=True, + inputs=[ + io.Video.Input("videos", tooltip="List of videos to shuffle."), + io.Int.Input( + "seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed." + ), + ], + outputs=[ + io.Video.Output( + display_name="videos", + is_output_list=True, + tooltip="Shuffled videos", + ), + ], + ) + + @classmethod + def execute(cls, videos, seed): + seed = seed[0] if isinstance(seed, list) else seed + np.random.seed(seed % (2**32 - 1)) + indices = np.random.permutation(len(videos)) + return io.NodeOutput([videos[i] for i in indices]) + + +class ShuffleVideoTextDatasetNode(io.ComfyNode): + """Shuffle videos and their captions together, preserving pairs.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ShuffleVideoTextDataset", + search_aliases=["shuffle", "randomize", "mix"], + display_name="Shuffle Pairs of Video-Text", + category="dataset/video", + description="Randomly shuffle the order of pairs of video-text in a list.", + is_experimental=True, + is_input_list=True, + inputs=[ + io.Video.Input("videos", tooltip="List of videos to shuffle."), + io.String.Input("texts", tooltip="List of texts to shuffle."), + io.Int.Input( + "seed", + default=0, + min=0, + max=0xFFFFFFFFFFFFFFFF, + tooltip="Random seed.", + ), + ], + outputs=[ + io.Video.Output( + display_name="videos", + is_output_list=True, + tooltip="Shuffled videos", + ), + io.String.Output( + display_name="texts", + is_output_list=True, + tooltip="Shuffled texts", + ), + ], + ) + + @classmethod + def execute(cls, videos, texts, seed): + seed = seed[0] if isinstance(seed, list) else seed + np.random.seed(seed % (2**32 - 1)) + indices = np.random.permutation(len(videos)) + return io.NodeOutput( + [videos[i] for i in indices], + [texts[i] for i in indices], + ) + + # ========== Text Transform Nodes ========== @@ -1608,7 +2026,10 @@ class DatasetExtension(ComfyExtension): LoadImageTextDataSetFromFolderNode, SaveImageDataSetToFolderNode, SaveImageTextDataSetToFolderNode, - # Image transform nodes + # Video data loading nodes + LoadVideoDataSetFromFolderNode, + LoadVideoTextDataSetFromFolderNode, + # Image transform nodes (auto-handle video via per-frame processing) ResizeImagesByShorterEdgeNode, ResizeImagesByLongerEdgeNode, CenterCropImagesNode, @@ -1618,6 +2039,12 @@ class DatasetExtension(ComfyExtension): AdjustContrastNode, ShuffleDatasetNode, ShuffleImageTextDatasetNode, + # Video processing nodes (lazy VideoInput in/out) + VideoFrameSampleNode, + VideoTemporalCropNode, + VideoRandomTemporalCropNode, + ShuffleVideoDatasetNode, + ShuffleVideoTextDatasetNode, # Text transform nodes TextToLowercaseNode, TextToUppercaseNode, diff --git a/comfy_extras/nodes_train.py b/comfy_extras/nodes_train.py index a27217b80..0dde97fc9 100644 --- a/comfy_extras/nodes_train.py +++ b/comfy_extras/nodes_train.py @@ -920,10 +920,11 @@ def _run_training_loop( """ sigmas = torch.tensor(range(num_images)) noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed) + ndim = latents[0].ndim if bucket_mode: # Use first bucket's first latent as dummy for guider - dummy_latent = latents[0][:1].repeat(num_images, 1, 1, 1) + dummy_latent = latents[0][:1].repeat(num_images, *[1]*(ndim-1)) guider.sample( noise.generate_noise({"samples": dummy_latent}), dummy_latent, @@ -933,7 +934,7 @@ def _run_training_loop( ) elif multi_res: # use first latent as dummy latent if multi_res - latents = latents[0].repeat(num_images, 1, 1, 1) + latents = latents[0].repeat(num_images, *[1]*(ndim-1)) guider.sample( noise.generate_noise({"samples": latents}), latents, From 593786e4898780e61c5928bc014b5a9a539e75b5 Mon Sep 17 00:00:00 2001 From: TheToxin-git <79914682+TheToxin-git@users.noreply.github.com> Date: Tue, 21 Jul 2026 11:43:34 +0000 Subject: [PATCH 133/211] FreSca: 5D+ (ex. Anima) fix, model-agnostic iteration (#15007) * FreSca: Make fresca work on multi dim --- comfy_extras/nodes_fresca.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/comfy_extras/nodes_fresca.py b/comfy_extras/nodes_fresca.py index 173f42154..a7d181bdf 100644 --- a/comfy_extras/nodes_fresca.py +++ b/comfy_extras/nodes_fresca.py @@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20): Apply frequency-dependent scaling to an image tensor using Fourier transforms. Parameters: - x: Input tensor of shape (B, C, H, W) + x: Input tensor of shape (..., H, W) scale_low: Scaling factor for low-frequency components (default: 1.0) scale_high: Scaling factor for high-frequency components (default: 1.5) freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20) @@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20): # Initialize mask with high-frequency scaling factor mask = torch.ones(x_freq.shape, device=device) * scale_high m = mask - for d in range(len(x_freq.shape) - 2): - dim = d + 2 + for d in range(2): + dim = len(x_freq.shape) - 2 + d cc = x_freq.shape[dim] // 2 f_c = min(freq_cutoff, cc) m = m.narrow(dim, cc - f_c, f_c * 2) From ac3a7a654fb3c694920336c33f8b20a2a1e42ac8 Mon Sep 17 00:00:00 2001 From: Barish Ozbay <17261091+drozbay@users.noreply.github.com> Date: Tue, 21 Jul 2026 08:44:14 -0400 Subject: [PATCH 134/211] Add native Uni3C Controlnet support for Wan models (CORE-365) (#14946) * Add native Uni3C controlnet support for Wan models * Dispatch double_block patches in all Wan model variants * Remove unused grid_sizes assignment in CameraWanModel, WanModel_S2V, HumoWanModel, and AnimateWanModel --- comfy/ldm/wan/model.py | 50 +++++++++ comfy/ldm/wan/model_animate.py | 9 ++ comfy/ldm/wan/model_wandancer.py | 10 ++ comfy/ldm/wan/uni3c.py | 149 +++++++++++++++++++++++++ comfy_extras/nodes_model_patch.py | 178 ++++++++++++++++++++++++++++++ 5 files changed, 396 insertions(+) create mode 100644 comfy/ldm/wan/uni3c.py diff --git a/comfy/ldm/wan/model.py b/comfy/ldm/wan/model.py index 1c9782a38..c042e93c4 100644 --- a/comfy/ldm/wan/model.py +++ b/comfy/ldm/wan/model.py @@ -552,6 +552,7 @@ class WanModel(torch.nn.Module): List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] """ # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] transformer_options["grid_sizes"] = grid_sizes @@ -564,11 +565,13 @@ class WanModel(torch.nn.Module): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # In-context reference (Bernini) context_latents = kwargs.get("context_latents", None) @@ -589,6 +592,7 @@ class WanModel(torch.nn.Module): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -604,6 +608,11 @@ class WanModel(torch.nn.Module): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -777,6 +786,7 @@ class VaceWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] transformer_options["grid_sizes"] = grid_sizes @@ -807,6 +817,7 @@ class VaceWanModel(WanModel): x_orig = x patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -822,6 +833,11 @@ class VaceWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + ii = self.vace_layers_mapping.get(i, None) if ii is not None: for iii in range(len(c)): @@ -887,6 +903,7 @@ class CameraWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) if self.control_adapter is not None and camera_conditions is not None: x = x + self.control_adapter(camera_conditions).to(x.dtype) @@ -909,6 +926,7 @@ class CameraWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -924,6 +942,11 @@ class CameraWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel): # embeddings bs, _, time, height, width = x.shape + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) if control_video is not None: x = x + self.cond_encoder(control_video) @@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel): context = self.text_embedding(context) patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel): x = out["img"] else: x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + if audio_emb is not None: x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len) # head @@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel): bs, _, time, height, width = x.shape # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] x = x.flatten(2).transpose(1, 2) @@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel): audio = None patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel): def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs): + x_input = x + + img_offset = 0 if reference_latent is not None: x = torch.cat((reference_latent, x), dim=2) + img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \ + (reference_latent.shape[3] // self.patch_size[1]) * \ + (reference_latent.shape[4] // self.patch_size[2]) # embeddings x = self.patch_embedding(x.float()).to(x.dtype) @@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) diff --git a/comfy/ldm/wan/model_animate.py b/comfy/ldm/wan/model_animate.py index 84d7adec4..9ebe5694b 100644 --- a/comfy/ldm/wan/model_animate.py +++ b/comfy/ldm/wan/model_animate.py @@ -493,6 +493,7 @@ class AnimateWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values) grid_sizes = x.shape[2:] @@ -505,11 +506,13 @@ class AnimateWanModel(WanModel): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # context context = self.text_embedding(context) @@ -522,6 +525,7 @@ class AnimateWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -537,6 +541,11 @@ class AnimateWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + if i % 5 == 0 and motion_vec is not None: x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec) diff --git a/comfy/ldm/wan/model_wandancer.py b/comfy/ldm/wan/model_wandancer.py index 3caef6dc5..aeec1d725 100644 --- a/comfy/ldm/wan/model_wandancer.py +++ b/comfy/ldm/wan/model_wandancer.py @@ -111,6 +111,7 @@ class WanDancerModel(WanModel): def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs): # embeddings + x_input = x if int(fps + 0.5) != 30: x = self.patch_embedding_global(x.float()).to(x.dtype) else: @@ -128,11 +129,13 @@ class WanDancerModel(WanModel): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # context context = self.text_embedding(context) @@ -163,6 +166,7 @@ class WanDancerModel(WanModel): context_img_len += clip_fea_ref.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -177,6 +181,12 @@ class WanDancerModel(WanModel): x = out["img"] else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + if audio_emb is not None: x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale) diff --git a/comfy/ldm/wan/uni3c.py b/comfy/ldm/wan/uni3c.py new file mode 100644 index 000000000..827ad2339 --- /dev/null +++ b/comfy/ldm/wan/uni3c.py @@ -0,0 +1,149 @@ +# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C +# Converted from the original diffusers based implementation. +import torch +import torch.nn as nn + +from comfy.ldm.flux.layers import EmbedND +from .model import WanSelfAttention + + +class Uni3CLayerNormZero(nn.Module): + def __init__( + self, + conditioning_dim, + embedding_dim, + eps=1e-5, + device=None, dtype=None, operations=None + ): + super().__init__() + self.silu = nn.SiLU() + self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype) + self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype) + + def forward(self, x, temb): + shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1) + x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] + return x, gate[:, None, :] + + +class Uni3CAttentionBlock(nn.Module): + def __init__( + self, + dim, + ffn_dim, + num_heads, + time_embed_dim=5120, + eps=1e-6, + device=None, dtype=None, operations=None + ): + super().__init__() + operation_settings = {"operations": operations, "device": device, "dtype": dtype} + self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) + self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings) + self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) + self.ffn = nn.Sequential( + operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'), + operations.Linear(ffn_dim, dim, device=device, dtype=dtype)) + + def forward(self, x, temb, freqs): + norm_x, gate_msa = self.norm1(x, temb) + x = x + gate_msa * self.self_attn(norm_x, freqs) + norm_x, gate_ff = self.norm2(x, temb) + x = x + gate_ff * self.ffn(norm_x) + return x + + +class MaskCamEmbed(nn.Module): + def __init__( + self, + add_channels=7, + mid_channels=256, + conv_out_dim=5120, + device=None, dtype=None, operations=None + ): + super().__init__() + self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning + self.mask_proj = nn.Sequential( + operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype), + operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype), + nn.SiLU()) + self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype) + + def forward(self, add_inputs): + add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0) + add_embeds = self.mask_proj(add_padded) + add_embeds = self.mask_zero_proj(add_embeds) + add_embeds = add_embeds.flatten(2).transpose(1, 2) + return add_embeds + + +class WanUni3CControlnet(nn.Module): + def __init__( + self, + in_channels=36, + conv_out_dim=5120, + dim=1024, + ffn_dim=8192, + num_heads=16, + num_layers=20, + time_embed_dim=5120, + out_proj_dim=5120, + add_channels=7, + mid_channels=256, + device=None, dtype=None, operations=None + ): + super().__init__() + patch_size = (1, 2, 2) + self.num_layers = num_layers + + self.controlnet_patch_embedding = operations.Conv3d( + in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32) + self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations) + + if conv_out_dim != dim: + self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype) + else: + self.proj_in = nn.Identity() + + self.controlnet_blocks = nn.ModuleList([ + Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations) + for _ in range(num_layers)]) + self.proj_out = nn.ModuleList([ + operations.Linear(dim, out_proj_dim, device=device, dtype=dtype) + for _ in range(num_layers)]) + + head_dim = dim // num_heads + self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)]) + + def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None): + img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype) + img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1) + img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1) + img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1) + img_ids = img_ids.reshape(1, -1, img_ids.shape[-1]) + freqs = self.rope_embedder(img_ids).movedim(1, 2) + return freqs + + def process_input(self, control_input, render_mask=None, camera_embedding=None): + # render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet + hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype) + t_len, h_len, w_len = hidden.shape[2:] + freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype) + hidden = hidden.flatten(2).transpose(1, 2) + + add_inputs = None + if camera_embedding is not None and render_mask is not None: + add_inputs = torch.cat([render_mask, camera_embedding], dim=1) + elif render_mask is not None: + add_inputs = render_mask + + if add_inputs is not None: + hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype)) + + hidden = self.proj_in(hidden) + return hidden, freqs + + def forward_block(self, block_index, hidden, temb, freqs): + hidden = self.controlnet_blocks[block_index](hidden, temb, freqs) + residual = self.proj_out[block_index](hidden) + return hidden, residual diff --git a/comfy_extras/nodes_model_patch.py b/comfy_extras/nodes_model_patch.py index 0935af09d..4d7bf7476 100644 --- a/comfy_extras/nodes_model_patch.py +++ b/comfy_extras/nodes_model_patch.py @@ -9,6 +9,7 @@ import comfy.latent_formats import comfy.ldm.lumina.controlnet import comfy.ldm.supir.supir_modules import comfy.ldm.anima.lllite +import comfy.ldm.wan.uni3c from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel from comfy_api.latest import io from comfy.ldm.supir.supir_patch import SUPIRPatch @@ -264,6 +265,37 @@ class ModelPatchLoader: if torch.count_nonzero(ref_weight) == 0: config['broken'] = True model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config) + elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan + attn_key_replace = {".self_attn.to_q.": ".self_attn.q.", + ".self_attn.to_k.": ".self_attn.k.", + ".self_attn.to_v.": ".self_attn.v.", + ".self_attn.to_out.0.": ".self_attn.o."} + converted_sd = {} + for k, w in sd.items(): + for r, rr in attn_key_replace.items(): + k = k.replace(r, rr) + converted_sd[k] = w + sd = converted_sd + + num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight")) + conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0] + if "proj_in.weight" in sd: + dim = sd["proj_in.weight"].shape[0] + else: + dim = conv_out_dim + model = comfy.ldm.wan.uni3c.WanUni3CControlnet( + in_channels=sd["controlnet_patch_embedding.weight"].shape[1], + conv_out_dim=conv_out_dim, + dim=dim, + ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0], + num_layers=num_layers, + time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1], + out_proj_dim=sd["proj_out.0.weight"].shape[0], + add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1], + mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0], + device=comfy.model_management.unet_offload_device(), + dtype=dtype, + operations=comfy.ops.manual_cast) elif "audio_proj.proj1.weight" in sd: model = MultiTalkModelPatch( audio_window=5, context_tokens=32, vae_scale=4, @@ -561,6 +593,150 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet): CATEGORY = "model/patch/z-image" +class WanUni3CCnetPatch: + def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end): + self.model_patch = model_patch + self.render_video = render_video + self.vae = vae + self.latent_format = latent_format + self.strength = strength + self.sigma_start = sigma_start + self.sigma_end = sigma_end + self.prepared_render = None + self.temp_data = None + + def encode_render_video(self, target_latent_shape): + t_len, h_len, w_len = target_latent_shape + temporal_compression = self.vae.temporal_compression_decode() or 1 + spatial_compression = self.vae.spacial_compression_encode() + target_frames = (t_len - 1) * temporal_compression + 1 + target_height = h_len * spatial_compression + target_width = w_len * spatial_compression + + frames = self.render_video + if frames.shape[0] > target_frames: + frames = frames[:target_frames] + elif frames.shape[0] < target_frames: + last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1) + frames = torch.cat([frames, last_frame], dim=0) + + if frames.shape[1] != target_height or frames.shape[2] != target_width: + frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1) + + loaded_models = comfy.model_management.loaded_models(only_currently_used=True) + render_latent = self.vae.encode(frames) + comfy.model_management.load_models_gpu(loaded_models) + return self.latent_format.process_in(render_latent) + + def build_controlnet_input(self, x, dtype, samples_per_cond): + # first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V) + hidden = x[:samples_per_cond, :20].to(dtype) + if hidden.shape[1] < 20: + pad_shape = list(hidden.shape) + pad_shape[1] = 20 - hidden.shape[1] + hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1) + + render = self.prepared_render + if render is None or render.shape[2:] != hidden.shape[2:]: + render = self.encode_render_video(hidden.shape[2:]) + render = render.to(device=hidden.device, dtype=dtype) + self.prepared_render = render + if render.shape[0] != hidden.shape[0]: + render = render.expand(hidden.shape[0], -1, -1, -1, -1) + return torch.cat([hidden, render], dim=1) + + def __call__(self, kwargs): + img = kwargs.get("img") + block_index = kwargs.get("block_index") + transformer_options = kwargs.get("transformer_options", {}) + + if block_index == 0: + self.temp_data = None + active = True + sigmas = transformer_options.get("sigmas", None) + if sigmas is not None: + sigma = sigmas[0].item() + if sigma > self.sigma_start or sigma < self.sigma_end: + active = False + if active: + x = kwargs.get("x") + # cond and uncond chunks share latents, so we can reuse residuals + num_conds = len(transformer_options.get("cond_or_uncond", [0])) + samples_per_cond = x.shape[0] + if num_conds > 0 and x.shape[0] % num_conds == 0: + samples_per_cond = x.shape[0] // num_conds + temb = kwargs.get("vec")[:samples_per_cond] + if temb.ndim == 3: + temb = temb[:, 0] + model = self.model_patch.model + controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond) + hidden, freqs = model.process_input(controlnet_input) + self.temp_data = (hidden, temb.to(img.dtype), freqs) + + num_layers = self.model_patch.model.num_layers + if self.temp_data is not None and block_index < num_layers: + hidden, temb, freqs = self.temp_data + hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs) + residual = residual.to(img.dtype) * self.strength + if residual.shape[0] != img.shape[0]: + residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1) + img_offset = kwargs.get("img_offset", 0) + img[:, img_offset:img_offset + residual.shape[1]] += residual + if block_index >= num_layers - 1: + self.temp_data = None + else: + self.temp_data = (hidden, temb, freqs) + + return kwargs + + def to(self, device_or_dtype): + if isinstance(device_or_dtype, torch.device): + if self.prepared_render is not None: + self.prepared_render = self.prepared_render.to(device_or_dtype) + self.temp_data = None + return self + + def models(self): + return [self.model_patch] + + +class WanUni3CControlnetApply: + @classmethod + def INPUT_TYPES(s): + return {"required": { "model": ("MODEL",), + "model_patch": ("MODEL_PATCH",), + "vae": ("VAE",), + "render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}), + "strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + }} + RETURN_TYPES = ("MODEL",) + FUNCTION = "apply_patch" + EXPERIMENTAL = True + + CATEGORY = "model/patch/wan" + + def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent): + if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet): + raise ValueError("The connected model patch is not a Uni3C ControlNet.") + cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features + model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None) + if model_dim is None: + raise ValueError("The Uni3C ControlNet only works with Wan models.") + if model_dim != cnet_dim: + raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim)) + + model_patched = model.clone() + model_sampling = model.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + latent_format = model.get_model_object("latent_format") + patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end) + model_patched.set_model_double_block_patch(patch) + return (model_patched,) + + class UsoStyleProjectorPatch: def __init__(self, model_patch, encoded_image): self.model_patch = model_patch @@ -719,6 +895,7 @@ NODE_CLASS_MAPPINGS = { "ModelPatchLoader": ModelPatchLoader, "QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet, "ZImageFunControlnet": ZImageFunControlnet, + "WanUni3CControlnetApply": WanUni3CControlnetApply, "USOStyleReference": USOStyleReference, "SUPIRApply": SUPIRApply, "AnimaLLLiteApply": AnimaLLLiteApply, @@ -728,6 +905,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ModelPatchLoader": "Load Model Patch", "QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet", "ZImageFunControlnet": "Apply Z-Image Fun ControlNet", + "WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet", "USOStyleReference": "Apply USO Style Reference", "SUPIRApply": "Apply SUPIR Patch", "AnimaLLLiteApply": "Apply Anima LLLite", From 78b43d25003ecf3c2b2509259d3a3cc72e941ec8 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 21 Jul 2026 18:19:53 +0300 Subject: [PATCH 135/211] [Partner Nodes] fix(Gemini-Omni): pass videos as inline data (#15014) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_gemini.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 8998b5943..47d028c6c 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -60,6 +60,7 @@ GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions" GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB GEMINI_URL_INPUT_BUDGET = 10 GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024 +GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB GEMINI_IMAGE_SYS_PROMPT = ( "You are an expert image-generation engine. You must ALWAYS produce an image.\n" "Interpret all user input—regardless of " @@ -469,9 +470,10 @@ async def build_gemini_media_parts( part, nbytes = _media_inline_part(kind, payload) inline_bytes += nbytes if inline_bytes > max_inline_bytes: + detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else "" raise ValueError( - f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first " - f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media." + f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). " + "Reduce the number or size of attached media." ) parts.append(part) return parts @@ -1738,7 +1740,14 @@ class GeminiVideoOmni(IO.ComfyNode): parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = [] if images or videos: - media_parts = await build_gemini_media_parts(cls, images, [], videos) + # The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL. + media_parts = await build_gemini_media_parts( + cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES + ) + video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts) + media_parts += await build_gemini_media_parts( + cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes + ) parts.extend(to_interaction_media_part(p) for p in media_parts) parts.append(GeminiInteractionTextPart(text=prompt)) interaction = await sync_op( From 7bf8bfcd078c7f4ae50ca5149c9ff7d8613e1fb1 Mon Sep 17 00:00:00 2001 From: "cloud-code-bot[bot]" <234529496+cloud-code-bot[bot]@users.noreply.github.com> Date: Tue, 21 Jul 2026 12:00:10 -0700 Subject: [PATCH 136/211] ci: bump cursor-review to github-workflows@964d5aa (#15017) --- .github/workflows/ci-cursor-review.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci-cursor-review.yml b/.github/workflows/ci-cursor-review.yml index 2312c0ccd..a7a0692c9 100644 --- a/.github/workflows/ci-cursor-review.yml +++ b/.github/workflows/ci-cursor-review.yml @@ -23,9 +23,9 @@ jobs: # SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up # upstream changes; keep `workflows_ref` matching so prompts/scripts load # from the same commit as the workflow definition. - uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9 + uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa) with: - workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a + workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d diff_excludes: >- :!**/.claude/** :!**/dist/** From 947c2749dd04c51ef0e21b069544d8b0b4f9b411 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 21 Jul 2026 20:02:45 -0700 Subject: [PATCH 137/211] Use optimized rms_rope function in joyai image model. (#15018) --- comfy/ldm/joyimage/model.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/comfy/ldm/joyimage/model.py b/comfy/ldm/joyimage/model.py index bca12c391..9d6951e54 100644 --- a/comfy/ldm/joyimage/model.py +++ b/comfy/ldm/joyimage/model.py @@ -94,12 +94,21 @@ class JoyImageAttention(nn.Module): txt_k = txt_k.unflatten(-1, (heads, -1)) txt_v = txt_v.unflatten(-1, (heads, -1)) - img_q = self.img_attn_q_norm(img_q) - img_k = self.img_attn_k_norm(img_k) txt_q = self.txt_attn_q_norm(txt_q) txt_k = self.txt_attn_k_norm(txt_k) - img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb) + img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True) + img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True) + img_q, img_k = comfy_kitchen.rms_rope( + img_q, + img_k, + image_rotary_emb, + img_q_scale, + img_k_scale, + self.img_attn_q_norm.eps, + ) + comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream) + comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream) joint_q = torch.cat([img_q, txt_q], dim=1) joint_k = torch.cat([img_k, txt_k], dim=1) From ba5226db96abe8aa37c78669285fba377c59f7b1 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:20:18 +0300 Subject: [PATCH 138/211] [Partner Nodes] feat(Openrouter): add new models (#15021) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_openrouter.py | 60 +++++++++++++++++------------ 1 file changed, 36 insertions(+), 24 deletions(-) diff --git a/comfy_api_nodes/nodes_openrouter.py b/comfy_api_nodes/nodes_openrouter.py index ba98133f0..439072e22 100644 --- a/comfy_api_nodes/nodes_openrouter.py +++ b/comfy_api_nodes/nodes_openrouter.py @@ -45,27 +45,39 @@ class _ModelSpec: MODELS: list[_ModelSpec] = [ - _ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.000005, 0.000025, max_images=20), - _ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.00003, 0.00018, max_images=20), - _ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.000005, 0.00003, max_images=20), - _ModelSpec("google/gemini-3.5-flash", "reasoning", 0.0000015, 0.000009, max_images=20, max_videos=4), - _ModelSpec("x-ai/grok-4.20", "reasoning", 0.00000125, 0.0000025, max_images=20), - _ModelSpec("x-ai/grok-4.3", "reasoning", 0.00000125, 0.0000025, max_images=20), - _ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.000000435, 0.00000087), - _ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.000000112, 0.000000224), - _ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.000000252, 0.000000378), - _ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.00000104, 0.00000624), - _ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.000000325, 0.00000195, max_images=10, max_videos=4), - _ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.0000001875, 0.000001125, max_images=10, max_videos=4), - _ModelSpec("mistralai/mistral-large-2512", "standard", 0.0000005, 0.0000015, max_images=8), - _ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.0000015, 0.0000075, max_images=8), - _ModelSpec("z-ai/glm-4.6", "reasoning", 0.00000043, 0.00000174), - _ModelSpec("z-ai/glm-5", "reasoning", 0.0000006, 0.00000192), - _ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.00000073, 0.00000349, max_images=10), - _ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.0000006, 0.0000025), - _ModelSpec("perplexity/sonar-pro", "perplexity", 0.000003, 0.000015), - _ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.000002, 0.000008), - _ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.000002, 0.000008), + _ModelSpec("anthropic/claude-opus-4.8", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20), + _ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20), + _ModelSpec("anthropic/claude-fable-5", "frontier_reasoning", 0.0000143, 0.0000715, max_images=20), + _ModelSpec("anthropic/claude-sonnet-5", "frontier_reasoning", 0.00000286, 0.0000143, max_images=20), + _ModelSpec("anthropic/claude-haiku-4.5", "frontier_reasoning", 0.00000143, 0.00000715, max_images=20), + _ModelSpec("openai/gpt-5.6-sol-pro", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20), + _ModelSpec("openai/gpt-5.6-sol", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20), + _ModelSpec("openai/gpt-5.6-terra-pro", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20), + _ModelSpec("openai/gpt-5.6-terra", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20), + _ModelSpec("openai/gpt-5.6-luna-pro", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20), + _ModelSpec("openai/gpt-5.6-luna", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20), + _ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.0000429, 0.0002574, max_images=20), + _ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20), + _ModelSpec("google/gemini-3.5-flash", "reasoning", 0.000002145, 0.00001287, max_images=20, max_videos=4), + _ModelSpec("x-ai/grok-4.5", "reasoning", 0.00000286, 0.00000858, max_images=20), + _ModelSpec("x-ai/grok-4.20", "reasoning", 0.0000017875, 0.000003575, max_images=20), + _ModelSpec("x-ai/grok-4.3", "reasoning", 0.0000017875, 0.000003575, max_images=20), + _ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.00000062205, 0.0000012441), + _ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.00000016016, 0.00000032032), + _ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.00000036036, 0.00000054054), + _ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.0000014872, 0.0000089232), + _ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.00000046475, 0.0000027885, max_images=10, max_videos=4), + _ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.000000268125, 0.00000160875, max_images=10, max_videos=4), + _ModelSpec("mistralai/mistral-large-2512", "standard", 0.000000715, 0.000002145, max_images=8), + _ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.000002145, 0.000010725, max_images=8), + _ModelSpec("z-ai/glm-4.6", "reasoning", 0.0000006149, 0.0000024882), + _ModelSpec("z-ai/glm-5", "reasoning", 0.000000858, 0.0000027456), + _ModelSpec("moonshotai/kimi-k3", "reasoning", 0.00000429, 0.00002145, max_images=10), + _ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.0000010439, 0.0000049907, max_images=10), + _ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.000000858, 0.000003575), + _ModelSpec("perplexity/sonar-pro", "perplexity", 0.00000429, 0.00002145), + _ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.00000286, 0.00001144), + _ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.00000286, 0.00001144), ] _MODELS_BY_SLUG: dict[str, _ModelSpec] = {m.slug: m for m in MODELS} @@ -148,7 +160,7 @@ def _build_model_options() -> list[IO.DynamicCombo.Option]: def _calculate_price(response: OpenRouterChatResponse) -> float | None: if response.usage and response.usage.cost is not None: - return float(response.usage.cost) + return float(response.usage.cost) * 1.43 return None @@ -269,8 +281,8 @@ class OpenRouterLLMNode(IO.ComfyNode): essentials_category="Text Generation", description=( "Generate text responses through OpenRouter. Routes to a curated set of popular " - "models from xAI, DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and " - "Perplexity Sonar." + "models from Anthropic (Claude), OpenAI (GPT), Google (Gemini), xAI (Grok), " + "DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and Perplexity Sonar." ), inputs=[ IO.String.Input( From f6d30bce9a862d56d9184dd65341621a8905ea3e Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:31:26 +0300 Subject: [PATCH 139/211] [Partner Nodes] feat(Anthropic): add new models (#15023) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_anthropic.py | 90 +++++++++++++++++++++++------- 1 file changed, 71 insertions(+), 19 deletions(-) diff --git a/comfy_api_nodes/nodes_anthropic.py b/comfy_api_nodes/nodes_anthropic.py index 87a870553..218f66ccf 100644 --- a/comfy_api_nodes/nodes_anthropic.py +++ b/comfy_api_nodes/nodes_anthropic.py @@ -28,6 +28,9 @@ ANTHROPIC_IMAGE_MAX_PIXELS = 1568 * 1568 CLAUDE_MAX_IMAGES = 20 CLAUDE_MODELS: dict[str, str] = { + "Opus 4.8": "claude-opus-4-8", + "Fable 5": "claude-fable-5", + "Sonnet 5": "claude-sonnet-5", "Opus 4.7": "claude-opus-4-7", "Opus 4.6": "claude-opus-4-6", "Sonnet 4.6": "claude-sonnet-4-6", @@ -36,9 +39,12 @@ CLAUDE_MODELS: dict[str, str] = { } _THINKING_UNSUPPORTED = {"Haiku 4.5"} -# Models that use the newer "adaptive" thinking mode (Opus 4.7 requires it; older models keep the explicit budget API). +# Models that use the newer "adaptive" thinking mode (Opus 4.7+ require it; older models keep the explicit budget API). # Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint. -_ADAPTIVE_THINKING_MODELS = {"Opus 4.7", "Opus 4.6", "Sonnet 4.6"} +_ADAPTIVE_THINKING_MODELS = {"Opus 4.8", "Sonnet 5", "Opus 4.7", "Opus 4.6", "Sonnet 4.6"} +_ALWAYS_THINKING_MODELS = {"Fable 5"} +_EXPLICIT_THINKING_OFF_MODELS = {"Sonnet 5"} +_NO_TEMPERATURE_MODELS = {"Opus 4.8", "Fable 5", "Sonnet 5"} # Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens. # Sized so even the "high" budget fits comfortably under the default max_tokens=32768. @@ -60,20 +66,33 @@ def _claude_model_inputs(model_label: str): tooltip="Maximum number of tokens to generate (includes reasoning tokens when enabled).", advanced=True, ), - IO.Float.Input( - "temperature", - default=1.0, - min=0.0, - max=1.0, - step=0.01, - tooltip=( - "Controls randomness. 0.0 is deterministic, 1.0 is most random. " - "Ignored for Opus 4.7 and any model when reasoning_effort is set." - ), - advanced=True, - ), ] - if model_label not in _THINKING_UNSUPPORTED: + if model_label not in _NO_TEMPERATURE_MODELS: + inputs.append( + IO.Float.Input( + "temperature", + default=1.0, + min=0.0, + max=1.0, + step=0.01, + tooltip=( + "Controls randomness. 0.0 is deterministic, 1.0 is most random. " + "Ignored for Opus 4.7 and any model when reasoning_effort is set." + ), + advanced=True, + ) + ) + if model_label in _ALWAYS_THINKING_MODELS: + inputs.append( + IO.Combo.Input( + "reasoning_effort", + options=[e for e in _REASONING_EFFORTS if e != "off"], + default="high", + tooltip="Extended thinking effort. Reasoning is always enabled for this model.", + advanced=True, + ) + ) + elif model_label not in _THINKING_UNSUPPORTED: inputs.append( IO.Combo.Input( "reasoning_effort", @@ -88,6 +107,12 @@ def _claude_model_inputs(model_label: str): def _model_price_per_million(model: str) -> tuple[float, float] | None: """Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown.""" + if "fable-5" in model: + return 14.30, 71.50 + if "opus-4-8" in model: + return 7.15, 35.75 + if "sonnet-5" in model: + return 2.86, 14.30 if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model: return 5.0, 25.0 if "sonnet-4" in model: @@ -213,7 +238,22 @@ class ClaudeNode(IO.ComfyNode): expr=""" ( $m := widgets.model; - $contains($m, "opus") ? { + $contains($m, "fable") ? { + "type": "list_usd", + "usd": [0.0143, 0.0715], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : $contains($m, "opus 4.8") ? { + "type": "list_usd", + "usd": [0.00715, 0.03575], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : $contains($m, "sonnet 5") ? { + "type": "list_usd", + "usd": [0.00286, 0.0143], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } + : $contains($m, "opus") ? { "type": "list_usd", "usd": [0.005, 0.025], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } @@ -247,18 +287,23 @@ class ClaudeNode(IO.ComfyNode): model_label = model["model"] max_tokens = model.get("max_tokens", 32768) reasoning_effort = model.get("reasoning_effort", "off") - thinking_enabled = reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED + always_thinking = model_label in _ALWAYS_THINKING_MODELS + thinking_enabled = always_thinking or ( + reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED + ) # Anthropic requires temperature to be unset (defaults to 1.0) when thinking is enabled. # Opus 4.7 also rejects user-supplied temperature. - if thinking_enabled or model_label == "Opus 4.7": + if model_label in _NO_TEMPERATURE_MODELS or thinking_enabled or model_label == "Opus 4.7": temperature = None else: temperature = model.get("temperature", 1.0) thinking_cfg: AnthropicThinkingConfig | None = None output_cfg: AnthropicOutputConfig | None = None - if thinking_enabled: + if always_thinking: + output_cfg = AnthropicOutputConfig(effort=reasoning_effort) + elif thinking_enabled: if model_label in _ADAPTIVE_THINKING_MODELS: # Adaptive mode - Anthropic chooses the budget based on effort hint thinking_cfg = AnthropicThinkingConfig(type="adaptive") @@ -268,6 +313,8 @@ class ClaudeNode(IO.ComfyNode): budget = _REASONING_BUDGET[reasoning_effort] budget = min(budget, max(1024, max_tokens - 1024)) thinking_cfg = AnthropicThinkingConfig(type="enabled", budget_tokens=budget) + elif model_label in _EXPLICIT_THINKING_OFF_MODELS: + thinking_cfg = AnthropicThinkingConfig(type="disabled") image_tensors: list[Input.Image] = [t for t in (images or {}).values() if t is not None] if sum(get_number_of_images(t) for t in image_tensors) > CLAUDE_MAX_IMAGES: @@ -293,6 +340,11 @@ class ClaudeNode(IO.ComfyNode): ), price_extractor=calculate_tokens_price, ) + if response.stop_reason == "refusal": + raise ValueError( + "Claude declined to answer this request for safety reasons. " + "Rephrase the prompt or try a different model." + ) return IO.NodeOutput(_get_text_from_response(response) or "Empty response from Claude model.") From 54ca9193a3862cec3f125677811a72e3be2e75e0 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Thu, 23 Jul 2026 00:49:39 +0800 Subject: [PATCH 140/211] chore: update workflow templates to v0.11.15 (#15030) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 20dc18dd6..10e99cec6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.45.21 -comfyui-workflow-templates==0.11.12 +comfyui-workflow-templates==0.11.15 comfyui-embedded-docs==0.5.8 torch torchsde From 2e47082c8ed1d1a0fe54add57f98b63433cfacbb Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 22 Jul 2026 12:34:27 -0700 Subject: [PATCH 141/211] Make z image/lumina 2 models use comfy kitchen rms rope. (#15036) --- comfy/ldm/lumina/model.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/comfy/ldm/lumina/model.py b/comfy/ldm/lumina/model.py index d0ee97d33..cdf03b2b5 100644 --- a/comfy/ldm/lumina/model.py +++ b/comfy/ldm/lumina/model.py @@ -6,6 +6,9 @@ import torch import torch.nn as nn import torch.nn.functional as F import comfy.ldm.common_dit +import comfy.model_management +import comfy.ops +import comfy.quant_ops from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder from comfy.ldm.modules.attention import optimized_attention_masked @@ -97,6 +100,7 @@ class JointAttention(nn.Module): self.n_local_kv_heads = self.n_kv_heads self.n_rep = self.n_local_heads // self.n_local_kv_heads self.head_dim = dim // n_heads + self.qk_norm = qk_norm self.qkv = operation_settings.get("operations").Linear( dim, @@ -151,10 +155,21 @@ class JointAttention(nn.Module): xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) - xq = self.q_norm(xq) - xk = self.k_norm(xk) - - xq, xk = apply_rope(xq, xk, freqs_cis) + if self.qk_norm and not comfy.model_management.in_training: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True) + epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps + if self.n_local_heads == self.n_local_kv_heads: + xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon) + else: + xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon) + xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon) + comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream) + else: + xq = self.q_norm(xq) + xk = self.k_norm(xk) + xq, xk = apply_rope(xq, xk, freqs_cis) n_rep = self.n_local_heads // self.n_local_kv_heads if n_rep >= 1: From a449f5f987d49ecce18245d1402e4ec68513e7c0 Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Thu, 23 Jul 2026 17:59:00 +0900 Subject: [PATCH 142/211] Bump comfyui-frontend-package to 1.47.10 (#15045) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 10e99cec6..f9480a973 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -comfyui-frontend-package==1.45.21 +comfyui-frontend-package==1.47.10 comfyui-workflow-templates==0.11.15 comfyui-embedded-docs==0.5.8 torch From 7cbe0474475c500420da58b111210377e1fa07c7 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 23 Jul 2026 18:32:46 +0300 Subject: [PATCH 143/211] [Partner Nodes] feat(ByteDance): add new "seed-audio-1.0-multilingual" model (#15034) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_bytedance.py | 22 ++++++++++++++++++++-- 1 file changed, 20 insertions(+), 2 deletions(-) diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index a84399ad3..561d6ae80 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -2690,7 +2690,8 @@ class ByteDanceSeedAudioNode(IO.ComfyNode): "with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music " "and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in " "preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), " - "or derive a voice from a character image. Up to 2 minutes of audio per run." + "or derive a voice from a character image. Up to 2 minutes of audio per run. " + "The multilingual model supports 20 languages and timestamp-based timing control." ), inputs=[ IO.String.Input( @@ -2701,7 +2702,9 @@ class ByteDanceSeedAudioNode(IO.ComfyNode): "Describe the voice(s), emotion, pacing, ambience, background music and sound " "effects, and include the lines to speak (name characters inline for dialogue). " "In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, " - "@Audio3. Maximum 3000 characters." + "@Audio3. With the multilingual model, a quoted line can start with a timestamp " + 'range that controls when and how long it is spoken, e.g. "[5.5s:8.0s] Wait for me!". ' + "Write the prompt in the same language as the lines to speak. Maximum 3000 characters." ), ), IO.DynamicCombo.Input( @@ -2796,6 +2799,19 @@ class ByteDanceSeedAudioNode(IO.ComfyNode): tooltip="Seed controls whether the node should re-run; " "results are non-deterministic regardless of seed.", ), + IO.Combo.Input( + "model", + options=["seed-audio-1.0-multilingual", "seed-audio-1.0"], + default="seed-audio-1.0-multilingual", + optional=True, + tooltip=( + "seed-audio-1.0-multilingual: 20 languages (English, Chinese, Japanese, Korean, " + "Mexican & Castilian Spanish, Indonesian, German, Brazilian Portuguese, French, " + "Thai, Vietnamese, Malay, Filipino, Italian, Russian, Dutch, Polish, Turkish, " + 'Swedish) plus per-sentence timing control via "[5.5s:8.0s] ..." timestamps. ' + "seed-audio-1.0: English and Chinese only, no timing control." + ), + ), ], outputs=[IO.Audio.Output()], hidden=[ @@ -2819,6 +2835,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode): loudness_rate: int, pitch_rate: int, seed: int, + model: str = "seed-audio-1.0-multilingual", ) -> IO.NodeOutput: mode = reference_mode["reference_mode"] audio_indices = connected_audio_indices(reference_mode) @@ -2845,6 +2862,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode): ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"), response_model=SeedAudioResponse, data=SeedAudioRequest( + model=model, text_prompt=text_prompt, references=references, audio_config=SeedAudioConfig( From feca51a8544511dd73d43602f387def0cc601a9d Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Thu, 23 Jul 2026 19:27:55 +0300 Subject: [PATCH 144/211] [Partner Nodes] chore(Runway): deprecate Gen3a model (#15050) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_runway.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/comfy_api_nodes/nodes_runway.py b/comfy_api_nodes/nodes_runway.py index 013a193d9..f58fa636f 100644 --- a/comfy_api_nodes/nodes_runway.py +++ b/comfy_api_nodes/nodes_runway.py @@ -194,6 +194,7 @@ class RunwayImageToVideoNodeGen3a(IO.ComfyNode): depends_on=IO.PriceBadgeDepends(widgets=["duration"]), expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""", ), + is_deprecated=True, ) @classmethod @@ -390,6 +391,7 @@ class RunwayFirstLastFrameNode(IO.ComfyNode): depends_on=IO.PriceBadgeDepends(widgets=["duration"]), expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""", ), + is_deprecated=True, ) @classmethod From 0cb84e7e6e0bdce2fa6e352aa07c6ea9c7cc984b Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 23 Jul 2026 19:06:52 -0700 Subject: [PATCH 145/211] Make Ernie use comfy kitchen rms rope (#15055) --- comfy/ldm/ernie/model.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/comfy/ldm/ernie/model.py b/comfy/ldm/ernie/model.py index f158ca1d2..88a3775d0 100644 --- a/comfy/ldm/ernie/model.py +++ b/comfy/ldm/ernie/model.py @@ -5,6 +5,7 @@ import torch.nn.functional as F from comfy.ldm.modules.attention import optimized_attention import comfy.model_management +import comfy.ops import comfy.quant_ops def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: @@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module): query = q_flat.view(B, S, self.heads, self.head_dim) key = k_flat.view(B, S, self.heads, self.head_dim) - query = self.norm_q(query) - key = self.norm_k(key) - - if image_rotary_emb is not None: - query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb) + if image_rotary_emb is not None and not comfy.model_management.in_training: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True) + query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps) + comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream) + else: + query = self.norm_q(query) + key = self.norm_k(key) + if image_rotary_emb is not None: + query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb) q_flat = query.reshape(B, S, -1) k_flat = key.reshape(B, S, -1) From c0ca3a5991986d76fd85dc21829687f547c2c6a5 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Fri, 24 Jul 2026 19:13:09 +0800 Subject: [PATCH 146/211] chore: update workflow templates to v0.11.17 (#15059) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index f9480a973..123b2e88d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.47.10 -comfyui-workflow-templates==0.11.15 +comfyui-workflow-templates==0.11.17 comfyui-embedded-docs==0.5.8 torch torchsde From 7c59a078d60c85baded8789f10c841221bff80a8 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 24 Jul 2026 12:17:31 -0700 Subject: [PATCH 147/211] Use comfy kitchen rope functions in ltx models. (#15056) --- comfy/ldm/lightricks/embeddings_connector.py | 16 ++- comfy/ldm/lightricks/model.py | 128 +++++++++---------- 2 files changed, 73 insertions(+), 71 deletions(-) diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py index 2811080be..1a6ddcc8d 100644 --- a/comfy/ldm/lightricks/embeddings_connector.py +++ b/comfy/ldm/lightricks/embeddings_connector.py @@ -6,9 +6,8 @@ import torch from comfy.ldm.lightricks.model import ( CrossAttention, FeedForward, + freqs_cis_matrix, generate_freq_grid_np, - interleaved_freqs_cis, - split_freqs_cis, ) from torch import nn @@ -244,12 +243,15 @@ class Embeddings1DConnector(nn.Module): expected_freqs = dim // 2 current_freqs = freqs.shape[-1] pad_size = expected_freqs - current_freqs - cos_freq, sin_freq = split_freqs_cis( - freqs, pad_size, self.num_attention_heads - ) else: - cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) - return cos_freq.to(dtype=out_dtype), sin_freq.to(dtype=out_dtype), self.split_rope + pad_size = dim % n_elem + return freqs_cis_matrix( + freqs, + pad_size, + self.split_rope, + self.num_attention_heads, + out_dtype, + ) def forward( self, diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index 9953b6679..92bb8118c 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -12,6 +12,8 @@ from torch import nn import comfy.patcher_extension import comfy.ldm.modules.attention import comfy.ldm.common_dit +import comfy.model_management +import comfy.quant_ops from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords @@ -322,40 +324,42 @@ class FeedForward(nn.Module): return self.net(x) def apply_rotary_emb(input_tensor, freqs_cis): - cos_freqs, sin_freqs = freqs_cis[0], freqs_cis[1] - split_pe = freqs_cis[2] if len(freqs_cis) > 2 else False - return ( - apply_split_rotary_emb(input_tensor, cos_freqs, sin_freqs) - if split_pe else - apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs) + rotation_matrix, split_pe = freqs_cis + original_shape = input_tensor.shape + input_tensor = input_tensor.reshape( + input_tensor.shape[0], input_tensor.shape[1], rotation_matrix.shape[2], -1 ) -def apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs): # TODO: remove duplicate funcs and pick the best/fastest one - t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2) - t1, t2 = t_dup.unbind(dim=-1) - t_dup = torch.stack((-t2, t1), dim=-1) - input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)") + if comfy.model_management.in_training: + if split_pe: + t = input_tensor.reshape(*input_tensor.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2) + else: + t = input_tensor.reshape(*input_tensor.shape[:-1], -1, 1, 2) + t = t.to(rotation_matrix.dtype) + output = rotation_matrix[..., 0] * t[..., 0] + rotation_matrix[..., 1] * t[..., 1] + if split_pe: + output = output.movedim(-1, -2) + output = output.reshape(input_tensor.shape).type_as(input_tensor) + elif split_pe: + output = comfy.quant_ops.ck.apply_rope_split_half1(input_tensor, rotation_matrix) + else: + output = comfy.quant_ops.ck.apply_rope1(input_tensor, rotation_matrix) + return output.reshape(original_shape) - out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs +def apply_rotary_emb_qk(q, k, freqs_cis): + if comfy.model_management.in_training: + return apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis) - return out - -def apply_split_rotary_emb(input_tensor, cos, sin): - needs_reshape = False - if input_tensor.ndim != 4 and cos.ndim == 4: - B, H, T, _ = cos.shape - input_tensor = input_tensor.reshape(B, T, H, -1).swapaxes(1, 2) - needs_reshape = True - split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2) - first_half_input = split_input[..., :1, :] - second_half_input = split_input[..., 1:, :] - output = split_input * cos.unsqueeze(-2) - first_half_output = output[..., :1, :] - second_half_output = output[..., 1:, :] - first_half_output.addcmul_(-sin.unsqueeze(-2), second_half_input) - second_half_output.addcmul_(sin.unsqueeze(-2), first_half_input) - output = rearrange(output, "... d r -> ... (d r)") - return output.swapaxes(1, 2).reshape(B, T, -1) if needs_reshape else output + rotation_matrix, split_pe = freqs_cis + q_shape = q.shape + k_shape = k.shape + q = q.reshape(q.shape[0], q.shape[1], rotation_matrix.shape[2], -1) + k = k.reshape(k.shape[0], k.shape[1], rotation_matrix.shape[2], -1) + if split_pe: + q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rotation_matrix) + else: + q, k = comfy.quant_ops.ck.apply_rope(q, k, rotation_matrix) + return q.reshape(q_shape), k.reshape(k_shape) class GuideAttentionMask: @@ -461,9 +465,13 @@ class CrossAttention(nn.Module): q = self.q_norm(q) k = self.k_norm(k) + # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent. if pe is not None: - q = apply_rotary_emb(q, pe) - k = apply_rotary_emb(k, pe if k_pe is None else k_pe) + if k_pe is None and q.shape == k.shape: + q, k = apply_rotary_emb_qk(q, k, pe) + else: + q = apply_rotary_emb(q, pe) + k = apply_rotary_emb(k, pe if k_pe is None else k_pe) if mask is None: out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) @@ -653,36 +661,23 @@ def generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid): ) return freqs -def interleaved_freqs_cis(freqs, pad_size): - cos_freq = freqs.cos().repeat_interleave(2, dim=-1) - sin_freq = freqs.sin().repeat_interleave(2, dim=-1) - if pad_size != 0: - cos_padding = torch.ones_like(cos_freq[:, :, : pad_size]) - sin_padding = torch.zeros_like(cos_freq[:, :, : pad_size]) - cos_freq = torch.cat([cos_padding, cos_freq], dim=-1) - sin_freq = torch.cat([sin_padding, sin_freq], dim=-1) - return cos_freq, sin_freq +def freqs_cis_matrix(freqs, pad_size, split_mode, num_attention_heads, out_dtype): + cos_freq = freqs.cos().to(out_dtype) + sin_freq = freqs.sin().to(out_dtype) + if pad_size: + matrix_pad_size = pad_size if split_mode else pad_size // 2 + cos_padding = torch.ones_like(cos_freq[:, :, :matrix_pad_size]) + sin_padding = torch.zeros_like(sin_freq[:, :, :matrix_pad_size]) + cos_freq = torch.cat((cos_padding, cos_freq), dim=-1) + sin_freq = torch.cat((sin_padding, sin_freq), dim=-1) -def split_freqs_cis(freqs, pad_size, num_attention_heads): - cos_freq = freqs.cos() - sin_freq = freqs.sin() - - if pad_size != 0: - cos_padding = torch.ones_like(cos_freq[:, :, :pad_size]) - sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size]) - - cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1) - sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1) - - # Reshape freqs to be compatible with multi-head attention - B , T, half_HD = cos_freq.shape - - cos_freq = cos_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) - sin_freq = sin_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) - - cos_freq = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2) - sin_freq = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2) - return cos_freq, sin_freq + B, T, _ = cos_freq.shape + cos_freq = cos_freq.reshape(B, T, num_attention_heads, -1) + sin_freq = sin_freq.reshape(B, T, num_attention_heads, -1) + rotation_matrix = torch.stack( + (cos_freq, -sin_freq, sin_freq, cos_freq), dim=-1 + ) + return rotation_matrix.reshape(*rotation_matrix.shape[:-1], 2, 2), split_mode class LTXBaseModel(torch.nn.Module, ABC): """ @@ -885,12 +880,17 @@ class LTXBaseModel(torch.nn.Module, ABC): expected_freqs = dim // 2 current_freqs = freqs.shape[-1] pad_size = expected_freqs - current_freqs - cos_freq, sin_freq = split_freqs_cis(freqs, pad_size, num_attention_heads) else: # 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only n_elem = 2 * indices_grid.shape[1] - cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) - return cos_freq.to(out_dtype), sin_freq.to(out_dtype), split_mode + pad_size = dim % n_elem + return freqs_cis_matrix( + freqs, + pad_size, + split_mode, + num_attention_heads, + out_dtype, + ) def _prepare_positional_embeddings(self, pixel_coords, frame_rate, x_dtype): """Prepare positional embeddings.""" From f8a3fd9d79837bd377d4e15e634271b488d9ee26 Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Sat, 25 Jul 2026 06:34:40 +1000 Subject: [PATCH 148/211] upscalers: convert latent_upsampler model to DynamicVram (#15063) These were alll non-dynamic (some non-ModelPatcher) code path calling FreeMemory for management requiring up-front memory freeing. Convert it to dynamic to avoid legacy free behaviour mixing into otherwise dynamic workflows. --- comfy/ldm/lightricks/latent_upsampler.py | 35 ++++++++++++++---------- comfy_extras/nodes_hunyuan.py | 9 ++++-- comfy_extras/nodes_lt_upsampler.py | 24 ++++++---------- comfy_extras/nodes_upscale_model.py | 35 +++++++++++------------- 4 files changed, 52 insertions(+), 51 deletions(-) diff --git a/comfy/ldm/lightricks/latent_upsampler.py b/comfy/ldm/lightricks/latent_upsampler.py index 78ed7653f..6a4beb1bf 100644 --- a/comfy/ldm/lightricks/latent_upsampler.py +++ b/comfy/ldm/lightricks/latent_upsampler.py @@ -97,11 +97,11 @@ class SpatialRationalResampler(nn.Module): For dims==3, work per-frame for spatial scaling (temporal axis untouched). """ - def __init__(self, mid_channels: int, scale: float): + def __init__(self, mid_channels: int, scale: float, operations): super().__init__() self.scale = float(scale) self.num, self.den = _rational_for_scale(self.scale) - self.conv = nn.Conv2d( + self.conv = operations.Conv2d( mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1 ) self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num)) @@ -119,18 +119,18 @@ class SpatialRationalResampler(nn.Module): class ResBlock(nn.Module): def __init__( - self, channels: int, mid_channels: Optional[int] = None, dims: int = 3 + self, channels: int, operations, mid_channels: Optional[int] = None, dims: int = 3 ): super().__init__() if mid_channels is None: mid_channels = channels - Conv = nn.Conv2d if dims == 2 else nn.Conv3d + Conv = operations.Conv2d if dims == 2 else operations.Conv3d self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1) - self.norm1 = nn.GroupNorm(32, mid_channels) + self.norm1 = operations.GroupNorm(32, mid_channels) self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1) - self.norm2 = nn.GroupNorm(32, channels) + self.norm2 = operations.GroupNorm(32, channels) self.activation = nn.SiLU() def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -159,6 +159,7 @@ class LatentUpsampler(nn.Module): def __init__( self, + operations, in_channels: int = 128, mid_channels: int = 512, num_blocks_per_stage: int = 4, @@ -179,34 +180,34 @@ class LatentUpsampler(nn.Module): self.spatial_scale = float(spatial_scale) self.rational_resampler = rational_resampler - Conv = nn.Conv2d if dims == 2 else nn.Conv3d + Conv = operations.Conv2d if dims == 2 else operations.Conv3d self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1) - self.initial_norm = nn.GroupNorm(32, mid_channels) + self.initial_norm = operations.GroupNorm(32, mid_channels) self.initial_activation = nn.SiLU() self.res_blocks = nn.ModuleList( - [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] + [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)] ) if spatial_upsample and temporal_upsample: self.upsampler = nn.Sequential( - nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), + operations.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(3), ) elif spatial_upsample: if rational_resampler: self.upsampler = SpatialRationalResampler( - mid_channels=mid_channels, scale=self.spatial_scale + mid_channels=mid_channels, scale=self.spatial_scale, operations=operations ) else: self.upsampler = nn.Sequential( - nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), + operations.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(2), ) elif temporal_upsample: self.upsampler = nn.Sequential( - nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), + operations.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(1), ) else: @@ -215,11 +216,14 @@ class LatentUpsampler(nn.Module): ) self.post_upsample_res_blocks = nn.ModuleList( - [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] + [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)] ) self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1) + def get_dtype(self): + return getattr(self.initial_conv, "weight_comfy_model_dtype", self.initial_conv.weight.dtype) + def forward(self, latent: torch.Tensor) -> torch.Tensor: b, c, f, h, w = latent.shape @@ -266,7 +270,7 @@ class LatentUpsampler(nn.Module): return x @classmethod - def from_config(cls, config): + def from_config(cls, config, operations): return cls( in_channels=config.get("in_channels", 4), mid_channels=config.get("mid_channels", 128), @@ -276,6 +280,7 @@ class LatentUpsampler(nn.Module): temporal_upsample=config.get("temporal_upsample", False), spatial_scale=config.get("spatial_scale", 2.0), rational_resampler=config.get("rational_resampler", False), + operations=operations, ) def config(self): diff --git a/comfy_extras/nodes_hunyuan.py b/comfy_extras/nodes_hunyuan.py index 8df2c8908..ce2997245 100644 --- a/comfy_extras/nodes_hunyuan.py +++ b/comfy_extras/nodes_hunyuan.py @@ -2,6 +2,8 @@ import nodes import node_helpers import torch import comfy.model_management +import comfy.model_patcher +import comfy.ops from typing_extensions import override from comfy_api.latest import ComfyExtension, io from comfy.ldm.hunyuan_video.upsampler import HunyuanVideo15SRModel @@ -217,8 +219,11 @@ class LatentUpscaleModelLoader(io.ComfyNode): model.load_sd(sd) elif "post_upsample_res_blocks.0.conv2.bias" in sd: config = json.loads(metadata["config"]) - model = LatentUpsampler.from_config(config).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32])) - model.load_state_dict(sd) + model = LatentUpsampler.from_config(config, operations=comfy.ops.disable_weight_init).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32])) + comfy.model_management.archive_model_dtypes(model) + model_patcher = comfy.model_patcher.CoreModelPatcher(model, load_device=comfy.model_management.get_torch_device(), offload_device=comfy.model_management.unet_offload_device()) + model.load_state_dict(sd, assign=model_patcher.is_dynamic()) + model = model_patcher return io.NodeOutput(model) diff --git a/comfy_extras/nodes_lt_upsampler.py b/comfy_extras/nodes_lt_upsampler.py index ef36109d1..7e7975495 100644 --- a/comfy_extras/nodes_lt_upsampler.py +++ b/comfy_extras/nodes_lt_upsampler.py @@ -38,26 +38,20 @@ class LTXVLatentUpsampler(IO.ComfyNode): Returns: tuple: Tuple containing the upsampled latent """ - device = model_management.get_torch_device() - memory_required = model_management.module_size(upscale_model) - - model_dtype = next(upscale_model.parameters()).dtype + device = upscale_model.load_device + model = upscale_model.model + model_dtype = upscale_model.model_dtype() latents = samples["samples"] input_dtype = latents.dtype - memory_required += math.prod(latents.shape) * 3000.0 # TODO: more accurate - model_management.free_memory(memory_required, device) + memory_required = math.prod(latents.shape) * 3000.0 # TODO: more accurate + model_management.load_models_gpu([upscale_model], memory_required=memory_required) - try: - upscale_model.to(device) # TODO: use the comfy model management system. + latents = latents.to(dtype=model_dtype, device=device) - latents = latents.to(dtype=model_dtype, device=device) - - """Upsample latents without tiling.""" - latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents) - upsampled_latents = upscale_model(latents) - finally: - upscale_model.cpu() + """Upsample latents without tiling.""" + latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents) + upsampled_latents = model(latents) upsampled_latents = vae.first_stage_model.per_channel_statistics.normalize( upsampled_latents diff --git a/comfy_extras/nodes_upscale_model.py b/comfy_extras/nodes_upscale_model.py index 1cf5a5d01..a4d692955 100644 --- a/comfy_extras/nodes_upscale_model.py +++ b/comfy_extras/nodes_upscale_model.py @@ -7,6 +7,7 @@ import folder_paths from typing_extensions import override from comfy_api.latest import ComfyExtension, io import comfy.model_management +import comfy.model_patcher try: from spandrel_extra_arches import EXTRA_REGISTRY @@ -42,6 +43,7 @@ class UpscaleModelLoader(io.ComfyNode): if not isinstance(out, ImageModelDescriptor): raise Exception("Upscale model must be a single-image model.") + out.patcher = comfy.model_patcher.CoreModelPatcher(out.model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device()) return io.NodeOutput(out) load_model = execute # TODO: remove @@ -66,14 +68,12 @@ class ImageUpscaleWithModel(io.ComfyNode): @classmethod def execute(cls, upscale_model, image) -> io.NodeOutput: - device = model_management.get_torch_device() + device = upscale_model.patcher.load_device - memory_required = model_management.module_size(upscale_model.model) - memory_required += (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate + memory_required = (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate memory_required += image.nelement() * image.element_size() - model_management.free_memory(memory_required, device) + model_management.load_models_gpu([upscale_model.patcher], memory_required=memory_required) - upscale_model.to(device) in_img = image.movedim(-1,-3).to(device) tile = 512 @@ -82,20 +82,17 @@ class ImageUpscaleWithModel(io.ComfyNode): output_device = comfy.model_management.intermediate_device() oom = True - try: - while oom: - try: - steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap) - pbar = comfy.utils.ProgressBar(steps) - s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device) - oom = False - except Exception as e: - model_management.raise_non_oom(e) - tile //= 2 - if tile < 128: - raise e - finally: - upscale_model.to("cpu") + while oom: + try: + steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap) + pbar = comfy.utils.ProgressBar(steps) + s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device) + oom = False + except Exception as e: + model_management.raise_non_oom(e) + tile //= 2 + if tile < 128: + raise e s = torch.clamp(s.movedim(-3,-1), min=0, max=1.0).to(comfy.model_management.intermediate_dtype()) return io.NodeOutput(s) From 36aec0d086f7321d253cde71b4f3b08f63e35d8f Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Sat, 25 Jul 2026 09:48:52 +1000 Subject: [PATCH 149/211] cli_args: bump clamp to 128BGB (#15068) Some long running chaos testing on a 512GB RAM RTX6000 pro showed that this is a little bit too low for common template workflows switching around. The original number was just a guess from me, so go with the scientific result instead. --- comfy/cli_args.py | 2 +- main.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfy/cli_args.py b/comfy/cli_args.py index e2e0d97ec..8e03ed032 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -112,7 +112,7 @@ parser.add_argument("--preview-method", type=LatentPreviewMethod, default=Latent parser.add_argument("--preview-size", type=int, default=512, help="Sets the maximum preview size for sampler nodes.") cache_group = parser.add_mutually_exclusive_group() -cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 96GB).") +cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 128GB).") cache_group.add_argument("--cache-classic", action="store_true", help="Use the old style (aggressive) caching.") cache_group.add_argument("--cache-lru", type=int, default=0, help="Use LRU caching with a maximum of N node results cached. May use more RAM/VRAM.") cache_group.add_argument("--cache-none", action="store_true", help="Reduced RAM/VRAM usage at the expense of executing every node for each run.") diff --git a/main.py b/main.py index 580074b19..1f16a7f89 100644 --- a/main.py +++ b/main.py @@ -319,7 +319,7 @@ def prompt_worker(q, server_instance): cache_ram_inactive = 0 if not args.cache_classic and not args.cache_none and args.cache_lru <= 0: cache_ram = min(10.0, max(2.0, comfy.model_management.total_ram * 0.10 / 1024.0)) - cache_ram_inactive = min(96.0, comfy.model_management.total_ram / 1024.0) + cache_ram_inactive = min(128.0, comfy.model_management.total_ram / 1024.0) if len(args.cache_ram) > 0: cache_ram = args.cache_ram[0] if len(args.cache_ram) > 1: From 45ffd5430beeccf63682b5f8b569faad45fd60e1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Sat, 25 Jul 2026 06:14:01 +0300 Subject: [PATCH 150/211] feat: Support MageFlow (CORE-372) (#15026) --- comfy/ldm/mage_flow/model.py | 187 ++++++++++++ comfy/ldm/mage_flow/vae.py | 477 +++++++++++++++++++++++++++++++ comfy/model_base.py | 16 +- comfy/model_detection.py | 7 + comfy/sd.py | 18 ++ comfy/supported_models.py | 31 ++ comfy/text_encoders/mage_flow.py | 94 ++++++ comfy/text_encoders/qwen3vl.py | 4 +- comfy_extras/nodes_mage.py | 103 +++++++ nodes.py | 3 +- 10 files changed, 935 insertions(+), 5 deletions(-) create mode 100644 comfy/ldm/mage_flow/model.py create mode 100644 comfy/ldm/mage_flow/vae.py create mode 100644 comfy/text_encoders/mage_flow.py create mode 100644 comfy_extras/nodes_mage.py diff --git a/comfy/ldm/mage_flow/model.py b/comfy/ldm/mage_flow/model.py new file mode 100644 index 000000000..92a5faa52 --- /dev/null +++ b/comfy/ldm/mage_flow/model.py @@ -0,0 +1,187 @@ +# Mage-Flow (https://github.com/microsoft/Mage) native-resolution MMDiT (MIT) +# Architecture is a 12-layer variant of the Qwen-Image double-stream block with +# patch_size=1 (no 2x2 packing), unrotated text tokens and a bf16-rounded +# timestep frequency table. +import math +import torch +import torch.nn as nn +from typing import Optional, Tuple + +from comfy.ldm.lightricks.model import TimestepEmbedding +from comfy.ldm.flux.layers import EmbedND +from comfy.ldm.qwen_image.model import QwenImageTransformerBlock, LastLayer +import comfy.patcher_extension + + +class MageTimestepProjEmbeddings(nn.Module): + def __init__(self, embedding_dim, dtype=None, device=None, operations=None): + super().__init__() + self.timestep_embedder = TimestepEmbedding( + in_channels=256, time_embed_dim=embedding_dim, + dtype=dtype, device=device, operations=operations + ) + + def forward(self, timestep, hidden_states): + timestep = timestep.to(hidden_states.dtype) + half_dim = 128 + exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim + emb = torch.exp(exponent).to(timestep.dtype) + emb = timestep[:, None].float() * emb[None, :] + emb = 1000.0 * emb + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) # flip_sin_to_cos + return self.timestep_embedder(emb.to(dtype=hidden_states.dtype)) + + +class MageFlowTransformer2DModel(nn.Module): + def __init__( + self, + in_channels: int = 128, + out_channels: Optional[int] = 128, + num_layers: int = 12, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + joint_attention_dim: int = 2560, + axes_dims_rope: Tuple[int, int, int] = (16, 56, 56), + image_model=None, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.dtype = dtype + self.patch_size = 1 + self.in_channels = in_channels + self.out_channels = out_channels or in_channels + self.inner_dim = num_attention_heads * attention_head_dim + + self.pe_embedder = EmbedND(dim=attention_head_dim, theta=10000, axes_dim=list(axes_dims_rope)) + + self.time_text_embed = MageTimestepProjEmbeddings(embedding_dim=self.inner_dim, dtype=dtype, device=device, operations=operations) + + self.txt_norm = operations.RMSNorm(joint_attention_dim, eps=1e-6, dtype=dtype, device=device) + self.img_in = operations.Linear(in_channels, self.inner_dim, dtype=dtype, device=device) + self.txt_in = operations.Linear(joint_attention_dim, self.inner_dim, dtype=dtype, device=device) + + self.transformer_blocks = nn.ModuleList([ + QwenImageTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + dtype=dtype, + device=device, + operations=operations + ) + for _ in range(num_layers) + ]) + + self.norm_out = LastLayer(self.inner_dim, self.inner_dim, dtype=dtype, device=device, operations=operations) + self.proj_out = operations.Linear(self.inner_dim, self.out_channels, bias=True, dtype=dtype, device=device) + + def process_img(self, x, index=0): + # patch_size=1: tokens are raw latent pixels, no 2x2 packing. + bs, c, h, w = x.shape + hidden_states = x.movedim(1, -1).reshape(bs, h * w, c) + + img_ids = torch.zeros((h, w, 3), device=x.device) + # Frame axis: positive image index (0 = target, 1..N = reference images). + img_ids[:, :, 0] = index + # Mage scale_rope centering: positions [-ceil(n/2), floor(n/2)), i.e. + # offset by (n - n//2). Differs from Qwen-Image's -(n//2) for odd sizes. + img_ids[:, :, 1] = img_ids[:, :, 1] + torch.arange(h, device=x.device)[:, None] - (h - h // 2) + img_ids[:, :, 2] = img_ids[:, :, 2] + torch.arange(w, device=x.device)[None, :] - (w - w // 2) + return hidden_states, img_ids.reshape(h * w, 3).unsqueeze(0).expand(bs, -1, -1), (h, w) + + def forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(x, timestep, context, attention_mask, ref_latents, transformer_options, **kwargs) + + def _forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, control=None, **kwargs): + if attention_mask is not None and not torch.is_floating_point(attention_mask): + attention_mask = (attention_mask - 1).to(x.dtype) * torch.finfo(x.dtype).max + + hidden_states, img_ids, orig_shape = self.process_img(x) + num_embeds = hidden_states.shape[1] + + if ref_latents is not None: + ref_num_tokens = [] + index = 0 + for ref in ref_latents: + index += 1 + kontext, kontext_ids, _ = self.process_img(ref, index=index) + hidden_states = torch.cat([hidden_states, kontext], dim=1) + img_ids = torch.cat([img_ids, kontext_ids], dim=1) + ref_num_tokens.append(kontext.shape[1]) + transformer_options = transformer_options.copy() + transformer_options["reference_image_num_tokens"] = ref_num_tokens + + # Text tokens are not rotated in Mage-Flow: RoPE at position 0 is the + # identity rotation. + txt_ids = torch.zeros((x.shape[0], context.shape[1], 3), device=x.device) + + hidden_states = self.img_in(hidden_states) + context = self.txt_norm(context) + context = self.txt_in(context) + + temb = self.time_text_embed(timestep, hidden_states) + + patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) + blocks_replace = patches_replace.get("dit", {}) + + if "post_input" in patches: + for p in patches["post_input"]: + out = p({"img": hidden_states, "txt": context, "img_ids": img_ids, "txt_ids": txt_ids, "transformer_options": transformer_options}) + hidden_states = out["img"] + context = out["txt"] + img_ids = out["img_ids"] + txt_ids = out["txt_ids"] + + ids = torch.cat((txt_ids, img_ids), dim=1) + image_rotary_emb = self.pe_embedder(ids).contiguous() + del ids, txt_ids, img_ids + + transformer_options["total_blocks"] = len(self.transformer_blocks) + transformer_options["block_type"] = "double" + for i, block in enumerate(self.transformer_blocks): + transformer_options["block_index"] = i + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["txt"], out["img"] = block(hidden_states=args["img"], encoder_hidden_states=args["txt"], encoder_hidden_states_mask=attention_mask, temb=args["vec"], image_rotary_emb=args["pe"], transformer_options=args["transformer_options"]) + return out + out = blocks_replace[("double_block", i)]({"img": hidden_states, "txt": context, "vec": temb, "pe": image_rotary_emb, "transformer_options": transformer_options}, {"original_block": block_wrap}) + hidden_states = out["img"] + context = out["txt"] + else: + context, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=context, + encoder_hidden_states_mask=attention_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + transformer_options=transformer_options, + ) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": hidden_states, "txt": context, "x": x, "block_index": i, "transformer_options": transformer_options}) + hidden_states = out["img"] + context = out["txt"] + + if control is not None: # Controlnet + control_i = control.get("input") + if i < len(control_i): + add = control_i[i] + if add is not None: + hidden_states[:, :add.shape[1]] += add + + hidden_states = self.norm_out(hidden_states, temb) + hidden_states = self.proj_out(hidden_states) + + hidden_states = hidden_states[:, :num_embeds] + h, w = orig_shape + return hidden_states.reshape(x.shape[0], h, w, self.out_channels).movedim(-1, 1) diff --git a/comfy/ldm/mage_flow/vae.py b/comfy/ldm/mage_flow/vae.py new file mode 100644 index 000000000..e6e21b99f --- /dev/null +++ b/comfy/ldm/mage_flow/vae.py @@ -0,0 +1,477 @@ +# Mage-VAE (https://github.com/microsoft/Mage) (MIT) +# Symmetric one-step diffusion codec: DConvEncoder (image -> 128ch latent) and +# DConvDenoiser + CoD Decoder (latent -> image). 16x downsample, latents in the +# Flux.2-VAE-anchored space (no patch packing, no BN normalization). +# Both encode and decode are single forward passes at t=0. +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops +from comfy.ldm.modules.diffusionmodules.model import vae_attention + +ops = comfy.ops.disable_weight_init + + +def nonlinearity(x): + return torch.nn.functional.silu(x) + + +def Normalize(in_channels): + return ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + + +def modulate(x, shift, scale): + if x.dim() == 4: + b, c = x.shape[:2] + return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1) + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + +class LayerNorm2d(ops.LayerNorm): + def __init__(self, num_channels, eps=1e-6, affine=True): + super().__init__(num_channels, eps=eps, elementwise_affine=affine) + + def forward(self, x): + x = x.permute(0, 2, 3, 1).contiguous() + x = super().forward(x) + return x.permute(0, 3, 1, 2).contiguous() + + +class TimestepEmbedder(nn.Module): + """DConv-style timestep MLP (max_period=10000, freq_size=256).""" + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__() + self.mlp = nn.Sequential( + ops.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + ops.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + @staticmethod + def timestep_embedding(t, dim, max_period=10000): + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half + ).to(t.device) + args = t[:, None].float() * freqs[None] + emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1) + return emb + + def forward(self, t, dtype): + emb = self.timestep_embedding(t, self.frequency_embedding_size) + return self.mlp(emb.to(dtype)) + + +class BottleneckPatchEmbed(nn.Module): + """Image patch embed concatenated with a per-patch conditioning vector.""" + + def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True): + super().__init__() + self.proj1 = ops.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False) + self.proj2 = ops.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias) + + def forward(self, x, cond): + return self.proj2(torch.cat([self.proj1(x), cond], dim=1)) + + +class DiCoBlock(nn.Module): + """DConv block with adaLN modulation.""" + + def __init__(self, hidden_size, mlp_ratio=4.0): + super().__init__() + self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True) + self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + + self.ca = nn.Sequential( + nn.AdaptiveAvgPool2d(1), + ops.Conv2d(hidden_size, hidden_size, 1, bias=True), + nn.Sigmoid(), + ) + + ffn = int(mlp_ratio * hidden_size) + self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True) + self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True) + + self.norm1 = LayerNorm2d(hidden_size, affine=False) + self.norm2 = LayerNorm2d(hidden_size, affine=False) + + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + ops.Linear(hidden_size, 6 * hidden_size, bias=True), + ) + + def forward(self, inp, c): + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1) + x = modulate(self.norm1(inp), shift_msa, scale_msa) + x = F.gelu(self.conv2(self.conv1(x))) + x = x * self.ca(x) + x = self.conv3(x) + x = inp + gate_msa[..., None, None] * x + x = x + gate_mlp[..., None, None] * self.conv5( + F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp))) + ) + return x + + +class EncoderDiCoBlock(nn.Module): + """DiCoBlock without adaLN, for the encoder head pathway.""" + + def __init__(self, hidden_size, mlp_ratio=4.0): + super().__init__() + self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True) + self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.ca = nn.Sequential( + nn.AdaptiveAvgPool2d(1), + ops.Conv2d(hidden_size, hidden_size, 1, bias=True), + nn.Sigmoid(), + ) + ffn = int(mlp_ratio * hidden_size) + self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True) + self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True) + self.norm1 = LayerNorm2d(hidden_size) + self.norm2 = LayerNorm2d(hidden_size) + + def forward(self, inp): + x = self.norm1(inp) + x = F.gelu(self.conv2(self.conv1(x))) + x = x * self.ca(x) + x = self.conv3(x) + x = inp + x + return x + self.conv5(F.gelu(self.conv4(self.norm2(x)))) + + +class NerfEmbedder(nn.Module): + """Patch-position embedder used by the DConv decoder x-pathway.""" + + def __init__(self, in_channels, hidden_size_input, max_freqs=8): + super().__init__() + self.max_freqs = max_freqs + self.embedder = nn.Sequential( + ops.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True), + ) + + def fetch_pos(self, patch_size, device, dtype): + pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype) + pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij") + pos_x = pos_x.reshape(-1, 1, 1) + pos_y = pos_y.reshape(-1, 1, 1) + freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device) + fx = freqs[None, :, None] + fy = freqs[None, None, :] + coeffs = (1 + fx * fy) ** -1 + dct_x = torch.cos(pos_x * fx * torch.pi) + dct_y = torch.cos(pos_y * fy * torch.pi) + return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2) + + def forward(self, x): + B, P2, _ = x.shape + ps = int(P2 ** 0.5) + dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1) + return self.embedder(torch.cat([x, dct], dim=-1)) + + +class NerfFinalLayer(nn.Module): + def __init__(self, hidden_size, out_channels): + super().__init__() + self.norm = ops.RMSNorm(hidden_size, eps=1e-6) + self.linear = ops.Linear(hidden_size, out_channels, bias=True) + + def forward(self, x): + return self.linear(self.norm(x)) + + +class MLPResBlock(nn.Module): + def __init__(self, channels): + super().__init__() + self.in_ln = ops.LayerNorm(channels, eps=1e-6) + self.mlp = nn.Sequential( + ops.Linear(channels, channels, bias=True), + nn.SiLU(), + ops.Linear(channels, channels, bias=True), + ) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + ops.Linear(channels, 3 * channels, bias=True), + ) + + def forward(self, x, y): + shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1) + h = self.in_ln(x) * (1 + scale) + shift + return x + gate * self.mlp(h) + + +class SimpleMLPAdaLN(nn.Module): + """Final small MLP that maps NerfEmbedder features to per-patch RGB.""" + + def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size): + super().__init__() + self.in_channels = in_channels + self.model_channels = model_channels + self.out_channels = out_channels + self.num_res_blocks = num_res_blocks + self.patch_size = patch_size + + self.cond_embed = ops.Linear(z_channels, patch_size ** 2 * model_channels) + self.input_proj = ops.Linear(in_channels, model_channels) + + self.res_blocks = nn.ModuleList(MLPResBlock(model_channels) for _ in range(num_res_blocks)) + + def forward(self, x, c): + x = self.input_proj(x) + c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1) + for block in self.res_blocks: + x = block(x, c) + return x + + +class ResnetBlock(nn.Module): + """GroupNorm + Conv ResBlock used by the CoD Decoder.""" + + def __init__(self, *, in_channels, out_channels=None): + super().__init__() + out_channels = out_channels or in_channels + self.in_channels = in_channels + self.out_channels = out_channels + + self.norm1 = Normalize(in_channels) + self.conv1 = ops.Conv2d(in_channels, out_channels, 3, padding=1) + self.norm2 = Normalize(out_channels) + self.conv2 = ops.Conv2d(out_channels, out_channels, 3, padding=1) + if in_channels != out_channels: + self.nin_shortcut = ops.Conv2d(in_channels, out_channels, 1) + + def forward(self, x): + h = self.conv1(nonlinearity(self.norm1(x))) + h = self.conv2(nonlinearity(self.norm2(h))) + if self.in_channels != self.out_channels: + x = self.nin_shortcut(x) + return x + h + + +class AttnBlock(nn.Module): + """Patched (windowed) self-attention used by the CoD Decoder.""" + + def __init__(self, in_channels, patch_size=32): + super().__init__() + self.in_channels = in_channels + self.patch_size = patch_size + self.norm = Normalize(in_channels) + self.q = ops.Conv2d(in_channels, in_channels, 1) + self.k = ops.Conv2d(in_channels, in_channels, 1) + self.v = ops.Conv2d(in_channels, in_channels, 1) + self.proj_out = ops.Conv2d(in_channels, in_channels, 1) + # VAE attention selection: full-precision backends only (no sage/quantized attention) + self.optimized_attention = vae_attention() + + def forward(self, x): + h_ = self.norm(x) + Q = self.q(h_) + K = self.k(h_) + V = self.v(h_) + + d = self.patch_size + b, c, H, W = Q.shape + pad_h = (d - H % d) % d + pad_w = (d - W % d) % d + if pad_h or pad_w: + Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate") + K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate") + V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate") + _, _, H_pad, W_pad = Q.shape + nph, npw = H_pad // d, W_pad // d + np_ = nph * npw + + def to_patches(t): + return (t.reshape(b, c, nph, d, npw, d) + .permute(0, 2, 4, 1, 3, 5) + .reshape(b * np_, c, d * d)) + + # [b*np, c, d*d]: attention over the d*d spatial positions of each window + Q = to_patches(Q) + K = to_patches(K) + V = to_patches(V) + + h_ = self.optimized_attention(Q, K, V) + h_ = h_.reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad) + if pad_h or pad_w: + h_ = h_[:, :, :H, :W] + return x + self.proj_out(h_) + + +class CoDDecoder(nn.Module): + """CoD Decoder: latent -> conditioning features for the denoiser (ds=16, light).""" + + def __init__(self, out_ch=384, z_ch=128): + super().__init__() + self.conv_in = ops.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1) + self.block = nn.Sequential( + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + AttnBlock(out_ch, patch_size=32), + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + AttnBlock(out_ch, patch_size=32), + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + ) + self.norm_out = Normalize(out_ch) + self.conv_out = ops.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1) + self.ada = nn.Identity() + + def forward(self, z): + h = self.block(self.conv_in(z)) + h = self.conv_out(nonlinearity(self.norm_out(h))) + return self.ada(h) + + +class DConvEncoder(nn.Module): + """DConvEncoder: image -> packed (mean, logvar) latent.""" + + def __init__( + self, + z_ch=128, + hidden_size=384, + num_blocks=21, + patch_size=16, + mlp_ratio=4.0, + head_size=768, + num_head_blocks=2, + out_ch_mult=2, + ): + super().__init__() + self.z_ch = z_ch + self.patch_size = patch_size + self.patch_cond_embed = ops.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True) + self.head_blocks = nn.ModuleList([ + EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks) + ]) + self.proj_down = ops.Conv2d(head_size, hidden_size, kernel_size=1, bias=True) + self.z_proj = ops.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True) + self.fuse_proj = ops.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True) + self.t_embedder = TimestepEmbedder(hidden_size) + self.blocks = nn.ModuleList([ + DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks) + ]) + self.norm_out = LayerNorm2d(hidden_size) + self.proj_out = ops.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True) + + def forward_pred(self, z_t, t, y): + cond = self.patch_cond_embed(y) + for block in self.head_blocks: + cond = block(cond) + cond = self.proj_down(cond) + + s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1)) + c = self.t_embedder(t.view(-1), y.dtype) + for block in self.blocks: + s = block(s, c) + return self.proj_out(self.norm_out(s)) + + +class YEmbedder(nn.Module): + """Holds only the CoD decoder (the original Flux2-VAE encoder side is dropped at load).""" + + def __init__(self, ch=384, z_ch=128): + super().__init__() + self.decoder = CoDDecoder(out_ch=ch, z_ch=z_ch) + + +class DConvDenoiser(nn.Module): + """One-step DConv denoiser: latent (via cond) + zero noise -> reconstructed image.""" + + def __init__( + self, + patch_size=16, + in_channels=3, + hidden_size=384, + hidden_size_x=32, + mlp_ratio=4.0, + num_blocks=24, + num_cond_blocks=21, + bottleneck_dim=128, + ): + super().__init__() + self.in_channels = in_channels + self.patch_size = patch_size + self.hidden_size = hidden_size + self.num_cond_blocks = num_cond_blocks + + self.t_embedder = TimestepEmbedder(hidden_size) + self.y_embedder_x = ops.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0) + self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8) + self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True) + self.blocks = nn.ModuleList([ + DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks) + ]) + self.dec_net = SimpleMLPAdaLN( + in_channels=hidden_size_x, + model_channels=hidden_size_x, + out_channels=in_channels, + z_channels=hidden_size, + num_res_blocks=num_blocks - num_cond_blocks, + patch_size=patch_size, + ) + self.final_layer = NerfFinalLayer(hidden_size_x, in_channels) + self.y_embedder = YEmbedder(ch=hidden_size, z_ch=bottleneck_dim) + + def forward(self, x, t, cond): + b, _, h, w = x.shape + c = self.t_embedder(t.view(-1), x.dtype) + + s = self.s_embedder(x, cond) + for block in self.blocks: + s = block(s, c) + + length = s.shape[-2] * s.shape[-1] + s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size) + + x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size) + x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1) + x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1) + x = self.x_embedder(x) + + x = self.dec_net(x, s) + x = self.final_layer(x) + x = x.transpose(1, 2).reshape(b, length, -1) + return torch.nn.functional.fold( + x.transpose(1, 2).contiguous(), (h, w), + kernel_size=self.patch_size, stride=self.patch_size, + ) + + +class MageVAE(nn.Module): + """ + Encode: DConvEncoder (one-step at t=0) -> posterior mean [B, 128, H/16, W/16] + Decode: DConvDenoiser + CoD Decoder -> image [B, 3, H, W] in [-1, 1] + """ + + latent_channels = 128 + downsample_factor = 16 + + def __init__(self): + super().__init__() + self.dconv_encoder = DConvEncoder() + self.decoder_model = DConvDenoiser() + + def encode(self, x): + B, _, H, W = x.shape + ps = self.dconv_encoder.patch_size + z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype) + t = torch.zeros(B, device=x.device, dtype=x.dtype) + out = self.dconv_encoder.forward_pred(z_t, t, x) + return out[:, : self.latent_channels] # posterior mean (sample_posterior=False) + + def decode(self, z): + cond = self.decoder_model.y_embedder.decoder(z) + B = z.shape[0] + H = z.shape[2] * self.downsample_factor + W = z.shape[3] * self.downsample_factor + noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype) + t = torch.zeros(B, device=z.device, dtype=z.dtype) + return self.decoder_model.forward(noise, t, cond) diff --git a/comfy/model_base.py b/comfy/model_base.py index 3494925be..50c73a431 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -58,6 +58,7 @@ import comfy.ldm.omnigen.omnigen2 import comfy.ldm.seedvr.model import comfy.ldm.boogu.model import comfy.ldm.qwen_image.model +import comfy.ldm.mage_flow.model import comfy.ldm.joyimage.model import comfy.ldm.ideogram4.model import comfy.ldm.krea2.model @@ -2243,8 +2244,8 @@ class Boogu(Omnigen2): self.memory_usage_factor_conds = ("ref_latents",) class QwenImage(BaseModel): - def __init__(self, model_config, model_type=ModelType.FLUX, device=None): - super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel) + def __init__(self, model_config, model_type=ModelType.FLUX, device=None, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel): + super().__init__(model_config, model_type, device=device, unet_model=unet_model) self.memory_usage_factor_conds = ("ref_latents",) def extra_conds(self, **kwargs): @@ -2274,6 +2275,17 @@ class QwenImage(BaseModel): out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) return out +class MageFlow(QwenImage): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel) + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 128, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 128]) + return out + class JoyImage(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index a1bf047f8..39e973d36 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -884,6 +884,13 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): "selected_layer_index": selected_layer_index, } + if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys and '{}proj_out.weight'.format(key_prefix) in state_dict_keys and state_dict['{}txt_norm.weight'.format(key_prefix)].shape[0] == 2560 and state_dict['{}proj_out.weight'.format(key_prefix)].shape[0] == 128: # Mage-Flow (Qwen Image txt_norm/proj_out are 3584/64) + dit_config = {} + dit_config["image_model"] = "mage_flow" + dit_config["in_channels"] = 128 + dit_config["num_layers"] = count_blocks(state_dict_keys, '{}transformer_blocks.'.format(key_prefix) + '{}.') + return dit_config + if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys: # Qwen Image dit_config = {} dit_config["image_model"] = "qwen_image" diff --git a/comfy/sd.py b/comfy/sd.py index e15e0a9fd..caf78222d 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -17,6 +17,7 @@ import comfy.ldm.wan.vae import comfy.ldm.wan.vae2_2 import comfy.ldm.hunyuan3d.vae import comfy.ldm.seedvr.vae +import comfy.ldm.mage_flow.vae import comfy.ldm.triposplat.vae import comfy.ldm.ace.vae.music_dcae_pipeline import comfy.ldm.cogvideo.vae @@ -60,6 +61,7 @@ import comfy.text_encoders.qwen_image import comfy.text_encoders.hunyuan_image import comfy.text_encoders.z_image import comfy.text_encoders.krea2 +import comfy.text_encoders.mage_flow import comfy.text_encoders.ideogram4 import comfy.text_encoders.ovis import comfy.text_encoders.kandinsky5 @@ -567,6 +569,17 @@ class VAE: self.upscale_index_formula = (4, 8, 8) self.process_input = lambda image: image * 2.0 - 1.0 self.crop_input = False + elif "student.dconv_encoder.proj_out.weight" in sd: # Mage-VAE (one-step diffusion codec, Flux2-anchored 128ch/16x latents) + sd = comfy.utils.state_dict_prefix_replace(sd, {"student.dconv_encoder.": "dconv_encoder.", "pipeline.": "decoder_model."}) + # Drop the unused Flux2-VAE anchor encoder carried in the checkpoint. + sd = {k: v for k, v in sd.items() if not k.startswith("decoder_model.y_embedder.encoder.") and not k.startswith("decoder_model.y_embedder.bottleneck.")} + self.first_stage_model = comfy.ldm.mage_flow.vae.MageVAE() + self.latent_channels = 128 + self.downscale_ratio = 16 + self.upscale_ratio = 16 + self.working_dtypes = [torch.bfloat16, torch.float32] + self.memory_used_encode = lambda shape, dtype: (400 * shape[2] * shape[3]) * model_management.dtype_size(dtype) + self.memory_used_decode = lambda shape, dtype: (1000 * shape[2] * shape[3] * 16 * 16) * model_management.dtype_size(dtype) elif "decoder.conv_in.weight" in sd: if sd['decoder.conv_in.weight'].shape[1] == 64: ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True} @@ -1379,6 +1392,7 @@ class CLIPType(Enum): BOOGU = 31 KREA2 = 32 JOYIMAGE = 33 + MAGE = 34 @@ -1713,6 +1727,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer + elif clip_type == CLIPType.MAGE and te_model == TEModel.QWEN3VL_4B: # Mage-Flow: full Qwen3-VL-4B, last hidden state, Qwen-Image-style templates. + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.mage_flow.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.mage_flow.MageFlowTokenizer elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx. clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data)) diff --git a/comfy/supported_models.py b/comfy/supported_models.py index e7c8983aa..ca89850a5 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -27,6 +27,7 @@ import comfy.text_encoders.z_image import comfy.text_encoders.ideogram4 import comfy.text_encoders.boogu import comfy.text_encoders.krea2 +import comfy.text_encoders.mage_flow import comfy.text_encoders.joyimage import comfy.text_encoders.anima import comfy.text_encoders.ace15 @@ -1883,6 +1884,35 @@ class Krea2(supported_models_base.BASE): hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect)) +class MageFlow(supported_models_base.BASE): + unet_config = { + "image_model": "mage_flow", + } + + sampling_settings = { + "multiplier": 1.0, + "shift": 6.0, + } + + memory_usage_factor = 6.5 + + unet_extra_config = {} + latent_format = latent_formats.Flux2 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.MageFlow(self, device=device) + return out + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.mage_flow.MageFlowTokenizer, comfy.text_encoders.mage_flow.te(**hunyuan_detect)) + class QwenImage(supported_models_base.BASE): unet_config = { "image_model": "qwen_image", @@ -2421,6 +2451,7 @@ models = [ ACEStep15, Omnigen2, Boogu, + MageFlow, QwenImage, JoyImage, Ideogram4, diff --git a/comfy/text_encoders/mage_flow.py b/comfy/text_encoders/mage_flow.py new file mode 100644 index 000000000..6542ad315 --- /dev/null +++ b/comfy/text_encoders/mage_flow.py @@ -0,0 +1,94 @@ +"""Mage-Flow text encoder: Qwen3-VL-4B, last hidden state (2560-dim). + +Mage-Flow conditions on the final hidden state of Qwen3-VL-4B with the leading +system + user-opening template tokens stripped (reference start_idx 34 for t2i, +64 for edit). The t2i template is identical to Qwen-Image's; the edit template +uses the same system prompt as Qwen-Image-Edit with "Image N: " reference +prefixes and no block. +""" + +import numbers + +import torch + +import comfy.text_encoders.qwen3vl +from comfy import sd1_clip + +MAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>" + +MAGE_T2I_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" +MAGE_EDIT_TEMPLATE = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + + +class MageFlowTokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b") + self.llama_template = MAGE_T2I_TEMPLATE + self.llama_template_images = MAGE_EDIT_TEMPLATE + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs): + image = kwargs.get("image", None) + if image is not None and len(images) == 0: + images = [image[i:i + 1] for i in range(image.shape[0])] + if llama_template is None: + if len(images) > 0: + # Training-time multi-reference body: "Image 1: Image 2: ...{instruction}" + prefix = "".join("Image {}: {}".format(j + 1, MAGE_VISION_BLOCK) for j in range(len(images))) + llama_template = self.llama_template_images.replace("{}", prefix + "{}", 1) + else: + llama_template = self.llama_template + # thinking=True: Mage templates end at "<|im_start|>assistant\n" with no block. + return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs) + + +class MageFlowQwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel): + def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}, model_type="qwen3vl_4b"): + super().__init__(device=device, dtype=dtype, attention_mask=attention_mask, model_options=model_options, model_type=model_type) + # apply the final RMSNorm to the tapped last layer (HF last_hidden_state) + self.layer_norm_hidden_state = True + + +class MageFlowTEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + clip_model = lambda **kw: MageFlowQwen3VLClipModel(**kw, model_type="qwen3vl_4b") # noqa: E731 + super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=clip_model, model_options=model_options) + + def encode_token_weights(self, token_weight_pairs, template_end=-1): + # Strip the system + user-opening prefix (reference drop_idx: 34 t2i / 64 edit). + out, pooled, extra = super().encode_token_weights(token_weight_pairs) + tok_pairs = token_weight_pairs["qwen3vl_4b"][0] + count_im_start = 0 + if template_end == -1: + for i, v in enumerate(tok_pairs): + elem = v[0] + if not torch.is_tensor(elem): + if isinstance(elem, numbers.Integral): + if elem == 151644 and count_im_start < 2: # <|im_start|> + template_end = i + count_im_start += 1 + + if out.shape[1] > (template_end + 3): + if tok_pairs[template_end + 1][0] == 872: # "user" + if tok_pairs[template_end + 2][0] == 198: # "\n" + template_end += 3 + + out = out[:, template_end:] + + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, template_end:] + if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]): + extra.pop("attention_mask") # attention mask is useless if no masked elements + + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class MageFlowTEModel_(MageFlowTEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if dtype_llama is not None: + dtype = dtype_llama + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + super().__init__(device=device, dtype=dtype, model_options=model_options) + return MageFlowTEModel_ diff --git a/comfy/text_encoders/qwen3vl.py b/comfy/text_encoders/qwen3vl.py index 7a329d2d6..2dd60d4e6 100644 --- a/comfy/text_encoders/qwen3vl.py +++ b/comfy/text_encoders/qwen3vl.py @@ -158,12 +158,12 @@ class Qwen3VLTokenizer(sd1_clip.SD1Tokenizer): self.llama_template = "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" self.llama_template_images = "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n" - def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, **kwargs): + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, skip_template=False, **kwargs): image = kwargs.get("image", None) if image is not None and len(images) == 0: images = [image[i:i + 1] for i in range(image.shape[0])] - skip_template = text.startswith('<|im_start|>') + skip_template = skip_template or text.startswith('<|im_start|>') if prevent_empty_text and text == '': text = ' ' diff --git a/comfy_extras/nodes_mage.py b/comfy_extras/nodes_mage.py new file mode 100644 index 000000000..a3b0d394c --- /dev/null +++ b/comfy_extras/nodes_mage.py @@ -0,0 +1,103 @@ +from typing_extensions import override + +import comfy.utils +import node_helpers +import torch +import comfy.model_management +from comfy_api.latest import ComfyExtension, io + + +class TextEncodeMageFlowEdit(io.ComfyNode): + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="TextEncodeMageFlowEdit", + category="model/conditioning/mage", + description="Encode an edit instruction with one or more reference images for Mage-Flow-Edit. Reference latents are resized to the output resolution (width/height, or the first image's size when 0). Use the latent output for sampling so the sizes always match.", + inputs=[ + io.Clip.Input("clip"), + io.String.Input("prompt", multiline=True, dynamic_prompts=True), + io.String.Input("negative_prompt", multiline=True, dynamic_prompts=True, advanced=True), + io.Vae.Input("vae", optional=True), + io.Autogrow.Input( + "images", + template=io.Autogrow.TemplateNames( + io.Image.Input("image"), + names=[f"image_{i}" for i in range(1, 17)], + min=0, + ), + tooltip="Reference image(s) to edit. All references are resized to the output resolution before encoding.", + ), + io.Int.Input("width", default=0, min=0, max=8192, step=16, tooltip="Output width. 0 = use the first reference image's size."), + io.Int.Input("height", default=0, min=0, max=8192, step=16, tooltip="Output height. 0 = use the first reference image's size."), + io.Int.Input("batch_size", default=1, min=1, max=4096), + ], + outputs=[ + io.Conditioning.Output(display_name="positive"), + io.Conditioning.Output(display_name="negative"), + io.Latent.Output(display_name="latent"), + ], + ) + + @classmethod + def execute(cls, clip, prompt, negative_prompt="", vae=None, images: io.Autogrow.Type = None, width=0, height=0, batch_size=1) -> io.NodeOutput: + ref_latents = [] + images = images or {} + images = [images[name] for name in sorted(images, key=lambda n: int(n.rsplit("_", 1)[-1])) if images[name] is not None] + images_vl = [] + + # Output resolution: explicit width/height, else the primary reference's own size, floored to /16. + # Each dimension falls back independently so a 0 on one axis keeps an explicit value on the other. + if width == 0 or height == 0: + if len(images) > 0: + ref_h, ref_w = images[0].shape[1], images[0].shape[2] + else: + ref_h, ref_w = 1024, 1024 + height = height or ref_h + width = width or ref_w + width = max(16, (width // 16) * 16) + height = max(16, (height // 16) * 16) + + for image in images: + samples = image.movedim(-1, 1) + + # VL conditioning copy: cap the long edge at 384 (training preprocessing). + long_edge = max(samples.shape[3], samples.shape[2]) + if long_edge > 384: + scale_by = 384 / long_edge + s = comfy.utils.common_upscale(samples, max(1, round(samples.shape[3] * scale_by)), max(1, round(samples.shape[2] * scale_by)), "bicubic", "disabled") + images_vl.append(s.movedim(1, -1)) + else: + images_vl.append(image) + + if vae is not None: + # All references are resized to the output resolution before encoding, because Mage's RoPE aligns reference and target content by position + if samples.shape[3] != width or samples.shape[2] != height: + s = comfy.utils.common_upscale(samples, width, height, "bicubic", "disabled") + else: + s = samples + ref_latents.append(vae.encode(s.movedim(1, -1)[:, :, :, :3])) + + # Negative branch keeps the same reference images (VL tokens + ref latents), only the instruction differs. + positive = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=images_vl)) + negative = clip.encode_from_tokens_scheduled(clip.tokenize(negative_prompt if negative_prompt else " ", images=images_vl)) + + if len(ref_latents) > 0: + positive = node_helpers.conditioning_set_values(positive, {"reference_latents": ref_latents}, append=True) + negative = node_helpers.conditioning_set_values(negative, {"reference_latents": ref_latents}, append=True) + + latent = torch.zeros([batch_size, 128, height // 16, width // 16], device=comfy.model_management.intermediate_device()) + return io.NodeOutput(positive, negative, {"samples": latent}) + + +class MageExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + TextEncodeMageFlowEdit, + ] + + +async def comfy_entrypoint() -> MageExtension: + return MageExtension() diff --git a/nodes.py b/nodes.py index b03d6c603..243a55bf2 100644 --- a/nodes.py +++ b/nodes.py @@ -992,7 +992,7 @@ class CLIPLoader: @classmethod def INPUT_TYPES(s): return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ), - "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ), + "type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage", "mage"], ), }, "optional": { "device": (["default", "cpu"], {"advanced": True}), @@ -2462,6 +2462,7 @@ async def init_builtin_extra_nodes(): "nodes_seedvr.py", "nodes_context_windows.py", "nodes_qwen.py", + "nodes_mage.py", "nodes_joyimage.py", "nodes_boogu.py", "nodes_chroma_radiance.py", From 6f6c500c1596b452e5b3c391c16dc7613b7ca8bc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Sat, 25 Jul 2026 14:30:37 +0300 Subject: [PATCH 151/211] Improve LTXV IC-lora detection (#15073) --- comfy_extras/nodes_lt.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/comfy_extras/nodes_lt.py b/comfy_extras/nodes_lt.py index 85d76ecef..044d82cc8 100644 --- a/comfy_extras/nodes_lt.py +++ b/comfy_extras/nodes_lt.py @@ -50,8 +50,8 @@ class GetICLoRAParameters(io.ComfyNode): factor = 1 if metadata: try: - factor = max(1, round(float(metadata.get("reference_downscale_factor", 1)))) - except (TypeError, ValueError): + factor = max(1, round(float(next(v for k, v in metadata.items() if k.endswith("reference_downscale_factor"))))) + except (StopIteration, TypeError, ValueError): factor = 1 parameters = {"reference_downscale_factor": factor} return io.NodeOutput(parameters) From fad06e5da4a757414ea286588240243f876f9996 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sat, 25 Jul 2026 20:25:58 +0300 Subject: [PATCH 152/211] [Partner Nodes] feat(Anthropic): add Claude Opus 5 to OpenRouter node (#15075) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_openrouter.py | 1 + 1 file changed, 1 insertion(+) diff --git a/comfy_api_nodes/nodes_openrouter.py b/comfy_api_nodes/nodes_openrouter.py index 439072e22..e9d6290c2 100644 --- a/comfy_api_nodes/nodes_openrouter.py +++ b/comfy_api_nodes/nodes_openrouter.py @@ -45,6 +45,7 @@ class _ModelSpec: MODELS: list[_ModelSpec] = [ + _ModelSpec("anthropic/claude-opus-5", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20), _ModelSpec("anthropic/claude-opus-4.8", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20), _ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20), _ModelSpec("anthropic/claude-fable-5", "frontier_reasoning", 0.0000143, 0.0000715, max_images=20), From f966a2b38c21702c906ab4103261641c322e0a2d Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 25 Jul 2026 12:09:26 -0700 Subject: [PATCH 153/211] Optimize ideogram model using comfy kitchen rms rope. (#15080) --- comfy/ldm/ideogram4/model.py | 37 +++++++++++++++++++++++++++++++----- 1 file changed, 32 insertions(+), 5 deletions(-) diff --git a/comfy/ldm/ideogram4/model.py b/comfy/ldm/ideogram4/model.py index 4ea5b8aaf..12e1a14fb 100644 --- a/comfy/ldm/ideogram4/model.py +++ b/comfy/ldm/ideogram4/model.py @@ -12,10 +12,13 @@ import torch import torch.nn as nn import torch.nn.functional as F +import comfy.model_management +import comfy.ops import comfy.patcher_extension +import comfy.quant_ops from comfy.ldm.lumina.model import FeedForward from comfy.ldm.modules.attention import optimized_attention_masked -from comfy.text_encoders.llama import apply_rope, precompute_freqs_cis +from comfy.text_encoders.llama import precompute_freqs_cis # Per-token role indicators SEQUENCE_PADDING_INDICATOR = -1 @@ -25,6 +28,22 @@ LLM_TOKEN_INDICATOR = 3 IMAGE_POSITION_OFFSET = 65536 +def _split_half_rope_matrix(freqs_cis): + cos, sin, neg_sin = freqs_cis + half_dim = sin.shape[-1] + matrix = torch.stack( + (cos[..., :half_dim], neg_sin, sin, cos[..., half_dim:]), dim=-1 + ) + return matrix.reshape(*matrix.shape[:-1], 2, 2).unsqueeze(2) + + +def _apply_rope_split_half1(x, freqs_cis): + x_dtype = x.dtype + x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(freqs_cis.dtype) + output = freqs_cis[..., 0] * x[..., 0] + freqs_cis[..., 1] * x[..., 1] + return output.movedim(-1, -2).reshape(*x.shape[:-3], -1).to(x_dtype) + + class Ideogram4Attention(nn.Module): def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None): super().__init__() @@ -42,16 +61,23 @@ class Ideogram4Attention(nn.Module): qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim) q, k, v = qkv.unbind(dim=2) - q = self.norm_q(q) - k = self.norm_k(k) + if comfy.model_management.in_training: + q = _apply_rope_split_half1(self.norm_q(q), freqs_cis) + k = _apply_rope_split_half1(self.norm_k(k), freqs_cis) + else: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, q, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, k, offloadable=True) + q, k = comfy.quant_ops.ck.rms_rope_split_half( + q, k, freqs_cis, q_scale, k_scale, self.norm_q.eps + ) + comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream) # (B, heads, L, head_dim) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) - q, k = apply_rope(q, k, freqs_cis) - out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options) return self.o(out) @@ -181,6 +207,7 @@ class Ideogram4Transformer(nn.Module): self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta, rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device, ) + freqs_cis = _split_half_rope_matrix(freqs_cis) if attn_mask is not None and attn_mask.dtype == torch.bool: attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max) From 806e092ed42772e4ce7abf44c97c50021cc4bd10 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Sun, 26 Jul 2026 04:01:51 +0300 Subject: [PATCH 154/211] Fix MageFlow on cards that don't support bf16 (#15081) --- comfy/ldm/mage_flow/model.py | 1 - comfy/model_base.py | 4 ++++ 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/comfy/ldm/mage_flow/model.py b/comfy/ldm/mage_flow/model.py index 92a5faa52..ac29bb610 100644 --- a/comfy/ldm/mage_flow/model.py +++ b/comfy/ldm/mage_flow/model.py @@ -22,7 +22,6 @@ class MageTimestepProjEmbeddings(nn.Module): ) def forward(self, timestep, hidden_states): - timestep = timestep.to(hidden_states.dtype) half_dim = 128 exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim emb = torch.exp(exponent).to(timestep.dtype) diff --git a/comfy/model_base.py b/comfy/model_base.py index 50c73a431..ee6dc57a2 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2279,6 +2279,10 @@ class MageFlow(QwenImage): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel) + def process_timestep(self, timestep, **kwargs): + # Mage runs in bf16 and rounds its timestep frequency table to the timestep dtype, keep that on fp32 devices. + return timestep.to(torch.bfloat16) + def extra_conds_shapes(self, **kwargs): out = {} ref_latents = kwargs.get("reference_latents", None) From 02c688429e40577510fad10c1e113cceb72b5d6d Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sun, 26 Jul 2026 14:21:05 -0700 Subject: [PATCH 155/211] Update AGENTS.md (#15096) --- AGENTS.md | 22 ++++++++++++++++++++-- 1 file changed, 20 insertions(+), 2 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 20014ce7e..bfe0976fd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -162,8 +162,26 @@ adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`, `comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and `comfy-kitchen` helpers where they already solve the problem. -- Use optimized comfy-kitchen ops in places where they improve performance - without changing the expected dtype, device, memory, or interface behavior. +- Model implementations must use an existing optimized Comfy Kitchen or + ComfyUI operation whenever one supports the required math and tensor layout + without changing expected dtype, device, memory, or interface behavior. This + is the default implementation requirement, not an optional follow-up + optimization. +- Before implementing model math, inspect the operations already exposed by + Comfy Kitchen, `comfy.quant_ops`, and existing ComfyUI model helpers. Check + for optimized single, paired, fused, layout-specific, and quantized variants + before writing a local implementation or composing lower-level torch ops. +- Use the compatible optimized operation first and adapt the model's inputs to + its documented layout while preserving the model's exact math. If several + optimized variants apply, benchmark representative model shapes and select + the fastest valid path. +- Add or retain a local implementation only when no existing optimized + operation supports the required math, layout, dtype, device, autograd, or + patch contract. Keep differentiable or patch-compatible fallbacks when the + optimized inference operation does not provide those contracts. +- Use the existing ComfyUI cast, offload, and cleanup helpers for parameters + passed to optimized operations. Preserve model-specific epsilon, scaling, + layout, dtype, device, and output-shape behavior. - Prefer ComfyUI's shared optimized kernels and backend dispatchers over handwritten implementations of the same operation. Remove duplicate local kernels and adapt inputs to the shared operation's documented layout while From 093d571b83e7a79833200e199b46b9f5a62217f9 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Mon, 27 Jul 2026 08:59:14 +0800 Subject: [PATCH 156/211] chore: update embedded docs to v0.5.9 (#15092) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 123b2e88d..038e0d662 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ comfyui-frontend-package==1.47.10 comfyui-workflow-templates==0.11.17 -comfyui-embedded-docs==0.5.8 +comfyui-embedded-docs==0.5.9 torch torchsde torchvision From c06ee57933f2d4a9b644ab76a918c02e857fd126 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:10:44 +0300 Subject: [PATCH 157/211] [Partner Nodes] feat(Anthropic): add Claude Opus 5 model (#15079) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_anthropic.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/comfy_api_nodes/nodes_anthropic.py b/comfy_api_nodes/nodes_anthropic.py index 218f66ccf..011c3f0cf 100644 --- a/comfy_api_nodes/nodes_anthropic.py +++ b/comfy_api_nodes/nodes_anthropic.py @@ -28,6 +28,7 @@ ANTHROPIC_IMAGE_MAX_PIXELS = 1568 * 1568 CLAUDE_MAX_IMAGES = 20 CLAUDE_MODELS: dict[str, str] = { + "Opus 5": "claude-opus-5", "Opus 4.8": "claude-opus-4-8", "Fable 5": "claude-fable-5", "Sonnet 5": "claude-sonnet-5", @@ -42,9 +43,9 @@ _THINKING_UNSUPPORTED = {"Haiku 4.5"} # Models that use the newer "adaptive" thinking mode (Opus 4.7+ require it; older models keep the explicit budget API). # Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint. _ADAPTIVE_THINKING_MODELS = {"Opus 4.8", "Sonnet 5", "Opus 4.7", "Opus 4.6", "Sonnet 4.6"} -_ALWAYS_THINKING_MODELS = {"Fable 5"} +_ALWAYS_THINKING_MODELS = {"Opus 5", "Fable 5"} _EXPLICIT_THINKING_OFF_MODELS = {"Sonnet 5"} -_NO_TEMPERATURE_MODELS = {"Opus 4.8", "Fable 5", "Sonnet 5"} +_NO_TEMPERATURE_MODELS = {"Opus 5", "Opus 4.8", "Fable 5", "Sonnet 5"} # Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens. # Sized so even the "high" budget fits comfortably under the default max_tokens=32768. @@ -109,7 +110,7 @@ def _model_price_per_million(model: str) -> tuple[float, float] | None: """Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown.""" if "fable-5" in model: return 14.30, 71.50 - if "opus-4-8" in model: + if "opus-5" in model or "opus-4-8" in model: return 7.15, 35.75 if "sonnet-5" in model: return 2.86, 14.30 @@ -253,6 +254,11 @@ class ClaudeNode(IO.ComfyNode): "usd": [0.00286, 0.0143], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } } + : $contains($m, "opus 5") ? { + "type": "list_usd", + "usd": [0.00715, 0.03575], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } : $contains($m, "opus") ? { "type": "list_usd", "usd": [0.005, 0.025], From a3572c4832f3a047dc1fdd4017cad2ff7ffeab9f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Mon, 27 Jul 2026 23:04:14 +0300 Subject: [PATCH 158/211] Allow using float fps for LTXVEmptyLatentAudio (#15106) --- comfy/ldm/lightricks/vae/audio_vae.py | 2 +- comfy_extras/nodes_lt_audio.py | 21 ++++++++++++--------- 2 files changed, 13 insertions(+), 10 deletions(-) diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py index dd5320c8f..b4a8c7524 100644 --- a/comfy/ldm/lightricks/vae/audio_vae.py +++ b/comfy/ldm/lightricks/vae/audio_vae.py @@ -185,7 +185,7 @@ class AudioVAE(torch.nn.Module): self.autoencoder.mel_bins, ) - def num_of_latents_from_frames(self, frames_number: int, frame_rate: int) -> int: + def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int: return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second) def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor: diff --git a/comfy_extras/nodes_lt_audio.py b/comfy_extras/nodes_lt_audio.py index 2d774a0a3..3ff18d8d4 100644 --- a/comfy_extras/nodes_lt_audio.py +++ b/comfy_extras/nodes_lt_audio.py @@ -107,14 +107,17 @@ class LTXVEmptyLatentAudio(io.ComfyNode): display_mode=io.NumberDisplay.number, tooltip="Number of frames.", ), - io.Int.Input( - "frame_rate", - default=25, - min=1, - max=1000, - step=1, - display_mode=io.NumberDisplay.number, - tooltip="Number of frames per second.", + io.MultiType.Input( + io.Float.Input( + "frame_rate", + default=25.0, + min=1.0, + max=1000.0, + step=0.01, + display_mode=io.NumberDisplay.number, + tooltip="Number of frames per second.", + ), + [io.Int], ), io.Int.Input( "batch_size", @@ -137,7 +140,7 @@ class LTXVEmptyLatentAudio(io.ComfyNode): def execute( cls, frames_number: int, - frame_rate: int, + frame_rate: float, batch_size: int, audio_vae, ) -> io.NodeOutput: From 6e36e12970952bca210387c99b941c0bc5390b6f Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:31:55 -0700 Subject: [PATCH 159/211] Update stable portable release workflow. (#15113) --- .github/workflows/release-stable-all.yml | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/.github/workflows/release-stable-all.yml b/.github/workflows/release-stable-all.yml index d7cf69fe2..e33e3f68d 100644 --- a/.github/workflows/release-stable-all.yml +++ b/.github/workflows/release-stable-all.yml @@ -20,7 +20,7 @@ jobs: git_tag: ${{ inputs.git_tag }} cache_tag: "cu130" python_minor: "13" - python_patch: "12" + python_patch: "14" rel_name: "nvidia" rel_extra_name: "" test_release: true @@ -48,13 +48,13 @@ jobs: contents: "write" packages: "write" pull-requests: "read" - name: "Release AMD ROCm 7.2" + name: "Release AMD ROCm 7.14" uses: ./.github/workflows/stable-release.yml with: git_tag: ${{ inputs.git_tag }} - cache_tag: "rocm72" - python_minor: "12" - python_patch: "10" + cache_tag: "rocm714" + python_minor: "13" + python_patch: "14" rel_name: "amd" rel_extra_name: "" test_release: false @@ -71,7 +71,7 @@ jobs: git_tag: ${{ inputs.git_tag }} cache_tag: "xpu" python_minor: "13" - python_patch: "12" + python_patch: "14" rel_name: "intel" rel_extra_name: "" test_release: true From cd0eddaf161656a4a38db4ec7f5d8c4eba6168f5 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 28 Jul 2026 15:17:23 +0300 Subject: [PATCH 160/211] [Partner Nodes] feat(credits): respect "X-Comfy-Credits-Used" header from the comfy-api (#15091) Signed-off-by: Alexander Piskun --- comfy_api_nodes/nodes_anthropic.py | 44 -------------- comfy_api_nodes/nodes_bytedance_llm.py | 26 --------- comfy_api_nodes/nodes_gemini.py | 80 -------------------------- comfy_api_nodes/nodes_grok.py | 3 - comfy_api_nodes/nodes_openai.py | 34 +---------- comfy_api_nodes/nodes_openrouter.py | 7 --- comfy_api_nodes/nodes_reve.py | 10 ---- comfy_api_nodes/util/client.py | 48 ++++++++++++++-- 8 files changed, 46 insertions(+), 206 deletions(-) diff --git a/comfy_api_nodes/nodes_anthropic.py b/comfy_api_nodes/nodes_anthropic.py index 011c3f0cf..76c611b93 100644 --- a/comfy_api_nodes/nodes_anthropic.py +++ b/comfy_api_nodes/nodes_anthropic.py @@ -106,49 +106,6 @@ def _claude_model_inputs(model_label: str): return inputs -def _model_price_per_million(model: str) -> tuple[float, float] | None: - """Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown.""" - if "fable-5" in model: - return 14.30, 71.50 - if "opus-5" in model or "opus-4-8" in model: - return 7.15, 35.75 - if "sonnet-5" in model: - return 2.86, 14.30 - if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model: - return 5.0, 25.0 - if "sonnet-4" in model: - return 3.0, 15.0 - if "haiku-4-5" in model: - return 1.0, 5.0 - return None - - -def calculate_tokens_price(response: AnthropicMessagesResponse) -> float | None: - """Compute approximate USD price from response usage. Server-side billing is authoritative.""" - if not response.usage or not response.model: - return None - rates = _model_price_per_million(response.model) - if rates is None: - return None - input_rate, output_rate = rates - input_tokens = response.usage.input_tokens or 0 - output_tokens = response.usage.output_tokens or 0 - cache_read = response.usage.cache_read_input_tokens or 0 - cache_5m = 0 - cache_1h = 0 - if response.usage.cache_creation: - cache_5m = response.usage.cache_creation.ephemeral_5m_input_tokens or 0 - cache_1h = response.usage.cache_creation.ephemeral_1h_input_tokens or 0 - total = ( - input_tokens * input_rate - + output_tokens * output_rate - + cache_read * input_rate * 0.1 - + cache_5m * input_rate * 1.25 - + cache_1h * input_rate * 2.0 - ) - return total / 1_000_000.0 - - def _get_text_from_response(response: AnthropicMessagesResponse) -> str: if not response.content: return "" @@ -344,7 +301,6 @@ class ClaudeNode(IO.ComfyNode): thinking=thinking_cfg, output_config=output_cfg, ), - price_extractor=calculate_tokens_price, ) if response.stop_reason == "refusal": raise ValueError( diff --git a/comfy_api_nodes/nodes_bytedance_llm.py b/comfy_api_nodes/nodes_bytedance_llm.py index cb41defa0..0403e0c1f 100644 --- a/comfy_api_nodes/nodes_bytedance_llm.py +++ b/comfy_api_nodes/nodes_bytedance_llm.py @@ -34,13 +34,6 @@ SEED_MODELS: dict[str, str] = { "Seed 2.0 Mini": "seed-2-0-mini-260215", } -# USD per 1M tokens: (input, cache_hit_input, output) -_SEED_PRICES_PER_MILLION: dict[str, tuple[float, float, float]] = { - "seed-2-0-pro-260328": (0.50, 0.10, 3.00), - "seed-2-0-lite-260228": (0.25, 0.05, 2.00), - "seed-2-0-mini-260215": (0.10, 0.02, 0.40), -} - def _seed_model_inputs(max_images: int = SEED_MAX_IMAGES, max_videos: int = SEED_MAX_VIDEOS): return [ @@ -74,24 +67,6 @@ def _seed_model_inputs(max_images: int = SEED_MAX_IMAGES, max_videos: int = SEED ] -def _calculate_price(model_id: str, response: BytePlusResponseObject) -> float | None: - """Compute approximate USD price from response usage.""" - if not response.usage: - return None - rates = _SEED_PRICES_PER_MILLION.get(model_id) - if rates is None: - return None - input_rate, cache_hit_rate, output_rate = rates - input_tokens = response.usage.input_tokens or 0 - output_tokens = response.usage.output_tokens or 0 - cached = 0 - if response.usage.input_tokens_details: - cached = response.usage.input_tokens_details.cached_tokens or 0 - fresh_input = max(0, input_tokens - cached) - total = fresh_input * input_rate + cached * cache_hit_rate + output_tokens * output_rate - return total / 1_000_000.0 - - def _get_text_from_response(response: BytePlusResponseObject) -> str: """Extract concatenated text from all assistant message output_text blocks.""" if not response.output: @@ -251,7 +226,6 @@ class ByteDanceSeedNode(IO.ComfyNode): store=False, stream=False, ), - price_extractor=lambda r: _calculate_price(model_id, r), ) if response.error: raise ValueError(f"Seed API error ({response.error.code}): {response.error.message}") diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 47d028c6c..fd9ff04a8 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -35,7 +35,6 @@ from comfy_api_nodes.apis.gemini import ( GeminiSystemInstructionContent, GeminiTextPart, GeminiThinkingConfig, - Modality, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -238,60 +237,6 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug return torch.cat(image_tensors, dim=0) -def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None: - if not response.modelVersion: - return None - # Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing - if response.modelVersion == "gemini-2.5-pro": - input_tokens_price = 1.25 - output_text_tokens_price = 10.0 - output_image_tokens_price = 0.0 - elif response.modelVersion == "gemini-2.5-flash": - input_tokens_price = 0.30 - output_text_tokens_price = 2.50 - output_image_tokens_price = 0.0 - elif response.modelVersion == "gemini-2.5-flash-image": - input_tokens_price = 0.30 - output_text_tokens_price = 2.50 - output_image_tokens_price = 30.0 - elif response.modelVersion in ("gemini-3-pro-preview", "gemini-3.1-pro-preview"): - input_tokens_price = 2 - output_text_tokens_price = 12.0 - output_image_tokens_price = 0.0 - elif response.modelVersion in ("gemini-3.1-flash-lite-preview", "gemini-3.1-flash-lite"): - input_tokens_price = 0.25 - output_text_tokens_price = 1.50 - output_image_tokens_price = 0.0 - elif response.modelVersion == "gemini-3.5-flash": - input_tokens_price = 1.50 - output_text_tokens_price = 9.0 - output_image_tokens_price = 0.0 - elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"): - input_tokens_price = 2 - output_text_tokens_price = 12.0 - output_image_tokens_price = 120.0 - elif response.modelVersion in ("gemini-3.1-flash-image-preview", "gemini-3.1-flash-image"): - input_tokens_price = 0.5 - output_text_tokens_price = 3.0 - output_image_tokens_price = 60.0 - elif response.modelVersion == "gemini-3.1-flash-lite-image": - input_tokens_price = 0.25 - output_text_tokens_price = 1.50 - output_image_tokens_price = 30.0 - else: - return None - final_price = response.usageMetadata.promptTokenCount * input_tokens_price - if response.usageMetadata.candidatesTokensDetails: - for i in response.usageMetadata.candidatesTokensDetails: - if i.modality == Modality.IMAGE: - final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models - else: - final_price += output_text_tokens_price * i.tokenCount - if response.usageMetadata.thoughtsTokenCount: - final_price += output_text_tokens_price * response.usageMetadata.thoughtsTokenCount - return final_price / 1_000_000.0 - - def get_text_from_interaction(interaction: GeminiInteraction) -> str: """Extract and concatenate all model output text from an Interactions API response.""" texts = [] @@ -326,24 +271,6 @@ async def get_video_from_interaction( ) -def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None: - if interaction.usage is None: - return None - input_tokens_price = 1.5 - output_tokens_prices = {"text": 9.0, "video": 17.5} - thoughts_tokens_price = 9.0 - final_price = 0.0 - for i in interaction.usage.input_tokens_by_modality or []: - if i.tokens: - final_price += input_tokens_price * i.tokens - for i in interaction.usage.output_tokens_by_modality or []: - if i.tokens and i.modality in output_tokens_prices: - final_price += output_tokens_prices[i.modality] * i.tokens - if interaction.usage.total_thought_tokens: - final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens - return final_price / 1_000_000.0 - - def create_video_parts(video_input: Input.Video) -> list[GeminiPart]: """Convert a single video input to Gemini API compatible parts (inline MP4/H.264).""" base_64_string = video_to_base64_string( @@ -657,7 +584,6 @@ class GeminiNode(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) output_text = get_text_from_response(response) @@ -872,7 +798,6 @@ class GeminiNodeV2(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) output_text = get_text_from_response(response) @@ -1085,7 +1010,6 @@ class GeminiImage(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) return IO.NodeOutput(await get_image_from_response(response), get_text_from_response(response)) @@ -1225,7 +1149,6 @@ class GeminiImage2(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) return IO.NodeOutput(await get_image_from_response(response), get_text_from_response(response)) @@ -1385,7 +1308,6 @@ class GeminiNanoBanana2(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) return IO.NodeOutput( await get_image_from_response(response), @@ -1610,7 +1532,6 @@ class GeminiNanoBanana2V2(IO.ComfyNode): systemInstruction=gemini_system_prompt, ), response_model=GeminiGenerateContentResponse, - price_extractor=calculate_tokens_price, ) return IO.NodeOutput( await get_image_from_response(response), @@ -1762,7 +1683,6 @@ class GeminiVideoOmni(IO.ComfyNode): ), ), response_model=GeminiInteraction, - price_extractor=calculate_interaction_tokens_price, ) if interaction.status != "completed": model_message = get_text_from_interaction(interaction).strip() diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index dc484536e..a95b35917 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -155,7 +155,6 @@ class GrokImageNode(IO.ComfyNode): resolution=resolution.lower(), ), response_model=ImageGenerationResponse, - price_extractor=_extract_grok_price, ) if len(response.data) == 1: return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url)) @@ -351,7 +350,6 @@ class GrokImageEditNode(IO.ComfyNode): aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio, ), response_model=ImageGenerationResponse, - price_extractor=_extract_grok_price, ) if len(response.data) == 1: return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url)) @@ -488,7 +486,6 @@ class GrokImageEditNodeV2(IO.ComfyNode): aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio, ), response_model=ImageGenerationResponse, - price_extractor=_extract_grok_price, ) if len(response.data) == 1: return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url)) diff --git a/comfy_api_nodes/nodes_openai.py b/comfy_api_nodes/nodes_openai.py index de2c94353..e73319e84 100644 --- a/comfy_api_nodes/nodes_openai.py +++ b/comfy_api_nodes/nodes_openai.py @@ -364,19 +364,6 @@ class OpenAIDalle3(IO.ComfyNode): return IO.NodeOutput(await validate_and_cast_response(response)) -def calculate_tokens_price_image_1(response: OpenAIImageGenerationResponse) -> float | None: - # https://platform.openai.com/docs/pricing - return ((response.usage.input_tokens * 10.0) + (response.usage.output_tokens * 40.0)) / 1_000_000.0 - - -def calculate_tokens_price_image_1_5(response: OpenAIImageGenerationResponse) -> float | None: - return ((response.usage.input_tokens * 8.0) + (response.usage.output_tokens * 32.0)) / 1_000_000.0 - - -def calculate_tokens_price_image_2_0(response: OpenAIImageGenerationResponse) -> float | None: - return ((response.usage.input_tokens * 8.0) + (response.usage.output_tokens * 30.0)) / 1_000_000.0 - - class OpenAIGPTImage1(IO.ComfyNode): @classmethod @@ -570,15 +557,10 @@ class OpenAIGPTImage1(IO.ComfyNode): if size not in ("auto", "1024x1024", "1024x1536", "1536x1024"): raise ValueError(f"Resolution {size} is only supported by GPT Image 2 model") - if model == "gpt-image-1": - price_extractor = calculate_tokens_price_image_1 - elif model == "gpt-image-1.5": - price_extractor = calculate_tokens_price_image_1_5 - elif model == "gpt-image-2": - price_extractor = calculate_tokens_price_image_2_0 + if model == "gpt-image-2": if background == "transparent": raise ValueError("Transparent background is not supported for GPT Image 2 model") - else: + elif model not in ("gpt-image-1", "gpt-image-1.5"): raise ValueError(f"Unknown model: {model}") if image is not None: @@ -633,7 +615,6 @@ class OpenAIGPTImage1(IO.ComfyNode): ), content_type="multipart/form-data", files=files, - price_extractor=price_extractor, ) else: response = await sync_op( @@ -650,7 +631,6 @@ class OpenAIGPTImage1(IO.ComfyNode): size=size, moderation="low", ), - price_extractor=price_extractor, ) return IO.NodeOutput(await validate_and_cast_response(response)) @@ -879,13 +859,7 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode): ) size = f"{custom_width}x{custom_height}" - if model_id == "gpt-image-1": - price_extractor = calculate_tokens_price_image_1 - elif model_id == "gpt-image-1.5": - price_extractor = calculate_tokens_price_image_1_5 - elif model_id == "gpt-image-2": - price_extractor = calculate_tokens_price_image_2_0 - else: + if model_id not in ("gpt-image-1", "gpt-image-1.5", "gpt-image-2"): raise ValueError(f"Unknown model: {model_id}") if image_tensors: @@ -944,7 +918,6 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode): ), content_type="multipart/form-data", files=files, - price_extractor=price_extractor, ) else: response = await sync_op( @@ -960,7 +933,6 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode): size=size, moderation="low", ), - price_extractor=price_extractor, ) return IO.NodeOutput(await validate_and_cast_response(response)) diff --git a/comfy_api_nodes/nodes_openrouter.py b/comfy_api_nodes/nodes_openrouter.py index e9d6290c2..ee93a1228 100644 --- a/comfy_api_nodes/nodes_openrouter.py +++ b/comfy_api_nodes/nodes_openrouter.py @@ -159,12 +159,6 @@ def _build_model_options() -> list[IO.DynamicCombo.Option]: return [IO.DynamicCombo.Option(spec.slug, _inputs_for_model(spec)) for spec in MODELS] -def _calculate_price(response: OpenRouterChatResponse) -> float | None: - if response.usage and response.usage.cost is not None: - return float(response.usage.cost) * 1.43 - return None - - def _price_badge_jsonata() -> str: rates_pairs = [] for spec in MODELS: @@ -372,7 +366,6 @@ class OpenRouterLLMNode(IO.ComfyNode): ApiEndpoint(path=OPENROUTER_CHAT_ENDPOINT, method="POST"), response_model=OpenRouterChatResponse, data=request, - price_extractor=_calculate_price, ) return IO.NodeOutput(_extract_text(response)) diff --git a/comfy_api_nodes/nodes_reve.py b/comfy_api_nodes/nodes_reve.py index 177349a8b..9120c7195 100644 --- a/comfy_api_nodes/nodes_reve.py +++ b/comfy_api_nodes/nodes_reve.py @@ -62,13 +62,6 @@ def _postprocessing_inputs(): ] -def _reve_price_extractor(headers: dict) -> float | None: - credits_used = headers.get("x-reve-credits-used") - if credits_used is not None: - return float(credits_used) / 524.48 - return None - - def _reve_response_header_validator(headers: dict) -> None: error_code = headers.get("x-reve-error-code") if error_code: @@ -180,7 +173,6 @@ class ReveImageCreateNode(IO.ComfyNode): headers={"Accept": "image/webp"}, ), as_binary=True, - price_extractor=_reve_price_extractor, response_header_validator=_reve_response_header_validator, data=ReveImageCreateRequest( prompt=prompt, @@ -279,7 +271,6 @@ class ReveImageEditNode(IO.ComfyNode): headers={"Accept": "image/webp"}, ), as_binary=True, - price_extractor=_reve_price_extractor, response_header_validator=_reve_response_header_validator, data=ReveImageEditRequest( edit_instruction=edit_instruction, @@ -396,7 +387,6 @@ class ReveImageRemixNode(IO.ComfyNode): headers={"Accept": "image/webp"}, ), as_binary=True, - price_extractor=_reve_price_extractor, response_header_validator=_reve_response_header_validator, data=ReveImageRemixRequest( prompt=prompt, diff --git a/comfy_api_nodes/util/client.py b/comfy_api_nodes/util/client.py index 66aab17f8..039e97d58 100644 --- a/comfy_api_nodes/util/client.py +++ b/comfy_api_nodes/util/client.py @@ -2,8 +2,10 @@ import asyncio import contextlib import json import logging +import math import time import uuid +import weakref from collections.abc import Callable, Iterable from dataclasses import dataclass from enum import Enum @@ -84,11 +86,37 @@ class _PollUIState: _RETRY_STATUS = {408, 500, 502, 503, 504} # status 429 is handled separately _MAX_RETRY_AFTER_WAIT = 150.0 # Cap a server Retry-After at this many seconds so a large hint can't block execution + +PRICE_CREDITS_HEADER = "X-Comfy-Credits-Used" +"""Proxy response header with the actual cost in Comfy credits. When present on any successful proxied response, +it takes precedence over ``price_extractor``.""" + +_credits_used_by_execution: "weakref.WeakKeyDictionary[type, float]" = weakref.WeakKeyDictionary() +"""Last PRICE_CREDITS_HEADER value per node execution, keyed by the node's per-execution class clone.""" COMPLETED_STATUSES = ["succeeded", "succeed", "success", "completed", "finished", "done", "complete"] FAILED_STATUSES = ["cancelled", "canceled", "canceling", "fail", "failed", "error"] QUEUED_STATUSES = ["created", "queued", "queueing", "submitted", "initializing", "wait", "in_queue"] +def _maybe_remember_credits_used(node_cls: type[IO.ComfyNode], header_value: str | None) -> None: + """Remember a PRICE_CREDITS_HEADER value from a successful proxied response.""" + if not header_value: + return + try: + credits_used = float(header_value) + except (TypeError, ValueError): + logging.debug("Ignoring malformed %s header: %r", PRICE_CREDITS_HEADER, header_value) + return + if not math.isfinite(credits_used) or credits_used < 0: + logging.debug("Ignoring out-of-range %s header: %r", PRICE_CREDITS_HEADER, header_value) + return + _credits_used_by_execution[node_cls] = credits_used + 0.0 # normalize -0.0 + + +def _get_remembered_credits_used(node_cls: type[IO.ComfyNode]) -> float | None: + return _credits_used_by_execution.get(node_cls) + + async def sync_op( cls: type[IO.ComfyNode], endpoint: ApiEndpoint, @@ -450,10 +478,15 @@ def _display_text( display_lines: list[str] = [] if status: display_lines.append(f"Status: {status.capitalize() if isinstance(status, str) else status}") - if price is not None: + server_credits = _get_remembered_credits_used(node_cls) + if server_credits is not None: + p = f"{server_credits:,.2f}".rstrip("0").rstrip(".") + elif price is not None: p = f"{float(price) * 211:,.1f}".rstrip("0").rstrip(".") - if p != "0": - display_lines.append(f"Price: {p} credits") + else: + p = None + if p is not None and p != "0": + display_lines.append(f"Price: {p} credits") if text is not None: display_lines.append(text) if display_lines: @@ -606,7 +639,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool): """Core request with retries, per-second interruption monitoring, true cancellation, and friendly errors.""" url = cfg.endpoint.path parsed_url = urlparse(url) - if not parsed_url.scheme and not parsed_url.netloc: # is URL relative? + is_comfy_api_request = not parsed_url.scheme and not parsed_url.netloc # is URL relative? + if is_comfy_api_request: url = urljoin(default_base_url().rstrip("/") + "/", url.lstrip("/")) method = cfg.endpoint.method @@ -644,7 +678,7 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool): logging.debug("[DEBUG] HTTP %s %s (attempt %d)", method, url, attempt) payload_headers = {"Accept": "*/*"} if expect_binary else {"Accept": "application/json"} - if not parsed_url.scheme and not parsed_url.netloc: # is URL relative? + if is_comfy_api_request: payload_headers.update(get_comfy_api_headers(cfg.node_cls)) if cfg.endpoint.headers: payload_headers.update(cfg.endpoint.headers) @@ -804,6 +838,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool): ) bytes_payload = bytes(buff) resp_headers = {k.lower(): v for k, v in resp.headers.items()} + if is_comfy_api_request: + _maybe_remember_credits_used(cfg.node_cls, resp.headers.get(PRICE_CREDITS_HEADER)) if cfg.price_extractor: with contextlib.suppress(Exception): extracted_price = cfg.price_extractor(resp_headers) @@ -831,6 +867,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool): except json.JSONDecodeError: payload = {"_raw": text} response_content_to_log = payload if isinstance(payload, dict) else text + if is_comfy_api_request: + _maybe_remember_credits_used(cfg.node_cls, resp.headers.get(PRICE_CREDITS_HEADER)) with contextlib.suppress(Exception): extracted_price = cfg.price_extractor(payload) if cfg.price_extractor else None operation_succeeded = True From f4509ff2136ba6bae8dd3d36a51e023c9414f794 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Tue, 28 Jul 2026 18:39:34 +0300 Subject: [PATCH 161/211] [Partner Nodes] feat(Recraft): add V4.1 model (#15105) * [Partner Nodes] feat(Recraft): add V4.1 models and V4 image edit nodes Signed-off-by: bigcat88 * [Partner Nodes] chore(Recraft): change old nodes names to contain "V3" Signed-off-by: bigcat88 * [Partner Nodes] fix(Recraft): remove v4 model from the image edit nodes; fix the default "strength" value Signed-off-by: bigcat88 * [Partner Nodes] chore(Recraft): remove new image edit nodes Signed-off-by: bigcat88 --------- Signed-off-by: bigcat88 --- comfy_api_nodes/apis/recraft.py | 4 +- comfy_api_nodes/nodes_recraft.py | 130 +++++++++++++++++++++++++++---- 2 files changed, 118 insertions(+), 16 deletions(-) diff --git a/comfy_api_nodes/apis/recraft.py b/comfy_api_nodes/apis/recraft.py index 78ededd94..64780d73b 100644 --- a/comfy_api_nodes/apis/recraft.py +++ b/comfy_api_nodes/apis/recraft.py @@ -244,10 +244,10 @@ RECRAFT_V4_PRO_SIZES = [ "2304x1792", "1792x2304", "1664x2688", - "1434x1024", - "1024x1434", "2560x1792", "1792x2560", + "2688x1536", + "1536x2688", ] diff --git a/comfy_api_nodes/nodes_recraft.py b/comfy_api_nodes/nodes_recraft.py index c44942f50..2605b9021 100644 --- a/comfy_api_nodes/nodes_recraft.py +++ b/comfy_api_nodes/nodes_recraft.py @@ -399,7 +399,7 @@ class RecraftTextToImageNode(IO.ComfyNode): def define_schema(cls): return IO.Schema( node_id="RecraftTextToImageNode", - display_name="Recraft Text to Image", + display_name="Recraft V3 Text to Image", category="partner/image/Recraft", description="Generates images synchronously based on prompt and resolution.", inputs=[ @@ -511,7 +511,7 @@ class RecraftImageToImageNode(IO.ComfyNode): def define_schema(cls): return IO.Schema( node_id="RecraftImageToImageNode", - display_name="Recraft Image to Image", + display_name="Recraft V3 Image to Image", category="partner/image/Recraft", description="Modify image based on prompt and strength.", inputs=[ @@ -731,7 +731,7 @@ class RecraftTextToVectorNode(IO.ComfyNode): def define_schema(cls): return IO.Schema( node_id="RecraftTextToVectorNode", - display_name="Recraft Text to Vector", + display_name="Recraft V3 Text to Vector", category="partner/image/Recraft", description="Generates SVG synchronously based on prompt and resolution.", inputs=[ @@ -1087,7 +1087,7 @@ class RecraftV4TextToImageNode(IO.ComfyNode): node_id="RecraftV4TextToImageNode", display_name="Recraft V4 Text to Image", category="partner/image/Recraft", - description="Generates images using Recraft V4 or V4 Pro models.", + description="Generates images using Recraft V4 and V4.1 models.", inputs=[ IO.String.Input( "prompt", @@ -1097,11 +1097,56 @@ class RecraftV4TextToImageNode(IO.ComfyNode): IO.String.Input( "negative_prompt", multiline=True, - tooltip="An optional text description of undesired elements on an image.", + tooltip="This input is ignored: negative prompt is not supported by " + "Recraft V4 and V4.1 models.", ), IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "recraftv4_1", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_utility", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_pro", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_utility_pro", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), IO.DynamicCombo.Option( "recraftv4", [ @@ -1162,7 +1207,14 @@ class RecraftV4TextToImageNode(IO.ComfyNode): depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]), expr=""" ( - $prices := {"recraftv4": 0.04, "recraftv4_pro": 0.25}; + $prices := { + "recraftv4_1": 0.035, + "recraftv4_1_utility": 0.035, + "recraftv4_1_pro": 0.21, + "recraftv4_1_utility_pro": 0.21, + "recraftv4": 0.04, + "recraftv4_pro": 0.25 + }; {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n} ) """, @@ -1179,14 +1231,13 @@ class RecraftV4TextToImageNode(IO.ComfyNode): seed: int, recraft_controls: RecraftControls | None = None, ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False, min_length=1, max_length=10000) + validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000) response = await sync_op( cls, ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"), response_model=RecraftImageGenerationResponse, data=RecraftImageGenerationRequest( prompt=prompt, - negative_prompt=negative_prompt if negative_prompt else None, model=model["model"], size=model["size"], n=n, @@ -1211,7 +1262,7 @@ class RecraftV4TextToVectorNode(IO.ComfyNode): node_id="RecraftV4TextToVectorNode", display_name="Recraft V4 Text to Vector", category="partner/image/Recraft", - description="Generates SVG using Recraft V4 or V4 Pro models.", + description="Generates SVG using Recraft V4 and V4.1 models.", inputs=[ IO.String.Input( "prompt", @@ -1221,11 +1272,56 @@ class RecraftV4TextToVectorNode(IO.ComfyNode): IO.String.Input( "negative_prompt", multiline=True, - tooltip="An optional text description of undesired elements on an image.", + tooltip="This input is ignored: negative prompt is not supported by " + "Recraft V4 and V4.1 models.", ), IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "recraftv4_1_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_utility_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_pro_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_1_utility_pro_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), IO.DynamicCombo.Option( "recraftv4", [ @@ -1286,7 +1382,14 @@ class RecraftV4TextToVectorNode(IO.ComfyNode): depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]), expr=""" ( - $prices := {"recraftv4": 0.08, "recraftv4_pro": 0.30}; + $prices := { + "recraftv4_1_vector": 0.08, + "recraftv4_1_utility_vector": 0.08, + "recraftv4_1_pro_vector": 0.30, + "recraftv4_1_utility_pro_vector": 0.30, + "recraftv4": 0.08, + "recraftv4_pro": 0.30 + }; {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n} ) """, @@ -1303,18 +1406,17 @@ class RecraftV4TextToVectorNode(IO.ComfyNode): seed: int, recraft_controls: RecraftControls | None = None, ) -> IO.NodeOutput: - validate_string(prompt, strip_whitespace=False, min_length=1, max_length=10000) + validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000) response = await sync_op( cls, ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"), response_model=RecraftImageGenerationResponse, data=RecraftImageGenerationRequest( prompt=prompt, - negative_prompt=negative_prompt if negative_prompt else None, model=model["model"], size=model["size"], n=n, - style="vector_illustration", + style=None if model["model"].endswith("_vector") else "vector_illustration", substyle=None, controls=recraft_controls.create_api_model() if recraft_controls else None, ), From e8f8c2ff432276f711604d21d1547686c2e89253 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Wed, 29 Jul 2026 00:45:57 +0800 Subject: [PATCH 162/211] chore: update workflow templates to v0.11.19 (#15123) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 038e0d662..9b248e69c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.47.10 -comfyui-workflow-templates==0.11.17 +comfyui-workflow-templates==0.11.19 comfyui-embedded-docs==0.5.9 torch torchsde From 99f221c7f5504f1fae012b09daa1060fc44c49ba Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 28 Jul 2026 13:46:44 -0700 Subject: [PATCH 163/211] Go back to older rocm for portable. (#15127) --- .github/workflows/release-stable-all.yml | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/.github/workflows/release-stable-all.yml b/.github/workflows/release-stable-all.yml index e33e3f68d..10f1ccf96 100644 --- a/.github/workflows/release-stable-all.yml +++ b/.github/workflows/release-stable-all.yml @@ -48,13 +48,13 @@ jobs: contents: "write" packages: "write" pull-requests: "read" - name: "Release AMD ROCm 7.14" + name: "Release AMD ROCm 7.2" uses: ./.github/workflows/stable-release.yml with: git_tag: ${{ inputs.git_tag }} - cache_tag: "rocm714" - python_minor: "13" - python_patch: "14" + cache_tag: "rocm72" + python_minor: "12" + python_patch: "10" rel_name: "amd" rel_extra_name: "" test_release: false From a8c44f9b2a0678ac4082e3529a3f43db7472acfe Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Tue, 28 Jul 2026 16:58:41 -0400 Subject: [PATCH 164/211] ComfyUI v0.29.0 --- comfyui_version.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/comfyui_version.py b/comfyui_version.py index dcc0fee96..b7c03631b 100644 --- a/comfyui_version.py +++ b/comfyui_version.py @@ -1,3 +1,3 @@ # This file is automatically generated by the build process when version is # updated in pyproject.toml. -__version__ = "0.28.0" +__version__ = "0.29.0" diff --git a/pyproject.toml b/pyproject.toml index 73de2990f..96ecbb9e5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ComfyUI" -version = "0.28.0" +version = "0.29.0" readme = "README.md" license = { file = "LICENSE" } requires-python = ">=3.10" From 628cdec592c736b65b3db260a06ec4d41b6dad15 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Tue, 28 Jul 2026 14:01:53 -0700 Subject: [PATCH 165/211] Update comfy-kitchen package version to 0.2.23 (#15112) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 9b248e69c..3a8203aff 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.22 +comfy-kitchen==0.2.23 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From 3d41e3ea4e0f0154487759810e00af569c5a5c60 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Wed, 29 Jul 2026 00:02:57 +0300 Subject: [PATCH 166/211] Support int8 convrot embedding lookup (#15035) --- comfy/ops.py | 22 ++++++++++++++++++---- 1 file changed, 18 insertions(+), 4 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 13c2604fb..1f7cc9575 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1469,12 +1469,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec if layer_conf is not None: layer_conf = json.loads(layer_conf.numpy().tobytes()) - # Only fp8 makes sense for embeddings (per-row dequant via index select). + # Only fp8 and int8_tensorwise support per-row dequant via index select. # Block-scaled formats (NVFP4, MXFP8) can't do per-row lookup efficiently. quant_format = layer_conf.get("format") if layer_conf is not None else None manually_loaded_keys = [] - if quant_format in ("float8_e4m3fn", "float8_e5m2") and weight_key in state_dict: + if quant_format in ("float8_e4m3fn", "float8_e5m2", "int8_tensorwise") and weight_key in state_dict: self.quant_format = quant_format qconfig = QUANT_ALGOS[quant_format] self.layout_type = qconfig["comfy_tensor_layout"] @@ -1488,10 +1488,16 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec scale = scale.float() manually_loaded_keys.append(scale_key) + extra = {} + if quant_format == "int8_tensorwise" and layer_conf.get("convrot", False): + # rotated embedding table: record it so the forward un-rotates after lookup + extra["convrot"] = True + extra["convrot_groupsize"] = int(layer_conf.get("convrot_groupsize", 256)) params = layout_cls.Params( scale=scale if scale is not None else torch.ones((), dtype=torch.float32), orig_dtype=MixedPrecisionOps._compute_dtype, orig_shape=(self.num_embeddings, self.embedding_dim), + **extra, ) self.weight = torch.nn.Parameter( QuantizedTensor(weight.to(dtype=qconfig["storage_t"]), qconfig["comfy_tensor_layout"], params), @@ -1513,15 +1519,23 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def forward_comfy_cast_weights(self, input, out_dtype=None): weight = self.weight - # Optimized path: lookup in fp8, dequantize only the selected rows. + # Optimized path: lookup in fp8/int8, dequantize only the selected rows. if isinstance(weight, QuantizedTensor) and len(self.weight_function) == 0: qdata, _, offload_stream = cast_bias_weight(self, device=input.device, dtype=weight.dtype, offloadable=True) if isinstance(qdata, QuantizedTensor): - scale = qdata._params.scale + params = qdata._params + scale = params.scale qdata = qdata._qdata else: + params = weight._params scale = None + # int8: per-row scale possible ConvRot, so let the layout do the gather + if self.quant_format == "int8_tensorwise": + x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input) + uncast_bias_weight(self, qdata, None, offload_stream) + return x if out_dtype is None else x.to(dtype=out_dtype) + x = torch.nn.functional.embedding( input, qdata, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse) From c01175530ed36fcb5961c2f2f2598e19b73287b9 Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Wed, 29 Jul 2026 07:05:57 +1000 Subject: [PATCH 167/211] Load weights to process RAM with MRU policy using pinning infrastructure (#15027) --- comfy/model_management.py | 87 +++++++++++++++++++++++-------------- comfy/model_patcher.py | 89 ++++++++++++++++++++++++++++++++------ comfy/ops.py | 25 ++++++++--- comfy/pinned_memory.py | 70 +++++++++++++++++------------- comfy_execution/caching.py | 4 +- comfy_execution/graph.py | 18 +++++--- execution.py | 6 ++- 7 files changed, 207 insertions(+), 92 deletions(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index 766e9ea89..eb768d783 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -632,18 +632,50 @@ def mark_mmap_dirty(storage): if mmap_refs is not None: DIRTY_MMAPS.add(mmap_refs[0]) -def free_pins(size, evict_active=False): +PIN_SUBSETS = [ "weights", "patches" ] +LOADED_PIN_SUBSETS = [ "weights-loaded", "patches-loaded" ] + +def models_for_pin_eviction(active, current_prompt=None): + for loaded_model in current_loaded_models: + model = loaded_model.model + if model is None or not model.is_dynamic(): + continue + pin_state = model.model.dynamic_pins[model.load_device] + if ((active is None or pin_state["active"] == active) and + (current_prompt is None or pin_state["current_prompt"] == current_prompt)): + yield model + +def free_model_pins(size, subsets, current_prompt, active, registrations=False): freed_total = 0 - for loaded_model in reversed(current_loaded_models): + for model in models_for_pin_eviction(active, current_prompt=current_prompt): if size <= 0: return freed_total - model = loaded_model.model - if model is not None and model.is_dynamic() and (evict_active or not model.model.dynamic_pins[model.load_device]["active"]): - freed = model.partially_unload_ram(size) - freed_total += freed - size -= freed + if registrations: + freed = model.unregister_inactive_pins(size, subsets=subsets) + else: + freed = model.partially_unload_ram(size, subsets=subsets) + freed_total += freed + size -= freed return freed_total +def pin_eviction_tiers(loaded, evict_active): + tiers = [ + (PIN_SUBSETS, False, None), + (LOADED_PIN_SUBSETS, False, None), + (LOADED_PIN_SUBSETS, True, None), + ] + if not loaded: + tiers.append((PIN_SUBSETS, True, False)) + if evict_active: + tiers.append((PIN_SUBSETS, True, True)) + return tiers + +def free_pins(size, evict_active=False, loaded=False): + freed = 0 + for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active): + freed += free_model_pins(size - freed, subsets, current_prompt, active) + return freed + def should_free_pins_for_ram_pressure(shortfall): if shortfall <= 0: return False @@ -653,7 +685,7 @@ def should_free_pins_for_ram_pressure(shortfall): return True return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT -def ensure_pin_budget(size, evict_active=False): +def ensure_pin_budget(size, evict_active=False, loaded=False): if args.high_ram: return True if args.fast_disk: @@ -664,32 +696,21 @@ def ensure_pin_budget(size, evict_active=False): return True to_free = shortfall + PIN_PRESSURE_HYSTERESIS - return free_pins(to_free, evict_active=evict_active) >= shortfall + return free_pins(to_free, evict_active=evict_active, loaded=loaded) >= shortfall -def free_registrations(shortfall, evict_active=True): +def free_registrations(shortfall, evict_active=True, loaded=False): if MAX_PINNED_MEMORY <= 0: return False if shortfall <= 0: return True shortfall += REGISTERABLE_PIN_HYSTERESIS - for loaded_model in reversed(current_loaded_models): - model = loaded_model.model - if model is not None and model.is_dynamic() and not model.model.dynamic_pins[model.load_device]["active"]: - shortfall -= model.unregister_inactive_pins(shortfall) - if shortfall <= 0: - return True - if evict_active: - for loaded_model in current_loaded_models: - model = loaded_model.model - if model is not None and model.is_dynamic() and model.model.dynamic_pins[model.load_device]["active"]: - shortfall -= model.unregister_inactive_pins(shortfall) - if shortfall <= 0: - return True + for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active): + shortfall -= free_model_pins(shortfall, subsets, current_prompt, active, registrations=True) return shortfall <= REGISTERABLE_PIN_HYSTERESIS -def ensure_pin_registerable(size, evict_active=True): - return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active) +def ensure_pin_registerable(size, evict_active=True, loaded=False): + return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active, loaded=loaded) class LoadedModel: def __init__(self, model: ModelPatcher): @@ -1379,15 +1400,17 @@ def reset_cast_buffers(): pin_state = model.model.dynamic_pins[model.load_device] if pin_state["active"]: - *_, buckets = pin_state["weights"] - for size, bucket in list(buckets.items()): - bucket[:] = [ entry for entry in bucket if entry[-1] is not None ] - if not bucket: - del buckets[size] + for subset in ("weights", "weights-loaded"): + *_, buckets = pin_state[subset] + for size, bucket in list(buckets.items()): + bucket[:] = [ entry for entry in bucket if entry[-1] is not None ] + if not bucket: + del buckets[size] pin_state["active"] = False - model.partially_unload_ram(1e30, subsets=[ "patches" ]) - model.model.dynamic_pins[model.load_device]["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {}) + model.partially_unload_ram(1e30, subsets=[ "patches", "patches-loaded" ]) + for subset in ("patches", "patches-loaded"): + pin_state[subset] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {}) STREAM_CAST_BUFFERS.clear() STREAM_AIMDO_CAST_BUFFERS.clear() diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index d70b42bf8..39246b95c 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -42,6 +42,52 @@ from comfy.patcher_extension import CallbacksMP, PatcherInjection, WrappersMP import comfy_aimdo.model_vbar +def is_model_patcher_output(output): + return isinstance(output, ModelPatcher) or isinstance(getattr(output, "patcher", None), ModelPatcher) + +class PromptModelTracker: + def __init__(self): + self.models = {} + + def start(self): + self.end() + + def add(self, outputs): + if isinstance(outputs, collections.abc.Mapping): + outputs = outputs.values() + elif not isinstance(outputs, (list, tuple)): + outputs = (outputs,) + + for output in outputs: + if isinstance(output, (collections.abc.Mapping, list, tuple)): + self.add(output) + continue + + models = [] + if isinstance(output, ModelPatcher): + models.append(output) + models.extend(output.model_patches_models()) + models.extend(output.get_nested_additional_models()) + else: + patcher = getattr(output, "patcher", None) + if isinstance(patcher, ModelPatcher): + models.append(patcher) + get_models = getattr(output, "get_models", None) + if callable(get_models): + models.extend(get_models()) + + for model in models: + if not isinstance(model, ModelPatcher) or not model.is_dynamic(): + continue + key = (id(model.model), model.load_device) + self.models[key] = model + model.set_in_use_by_current_prompt(True) + + def end(self): + for model in self.models.values(): + model.set_in_use_by_current_prompt(False) + self.models.clear() + def set_model_options_patch_replace(model_options, patch, name, block_name, number, transformer_index=None): to = model_options["transformer_options"].copy() @@ -1724,14 +1770,20 @@ class ModelPatcherDynamic(ModelPatcher): self.model.dynamic_pins[device] = { "weights": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), "patches": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), + "weights-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), + "patches-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), "hostbufs_initialized": False, "failed": False, "active": False, + "current_prompt": False, } def is_dynamic(self): return True + def set_in_use_by_current_prompt(self, in_use): + self.model.dynamic_pins[self.load_device]["current_prompt"] = in_use + def _vbar_get(self, create=False): if self.load_device == torch.device("cpu"): return None @@ -1802,6 +1854,8 @@ class ModelPatcherDynamic(ModelPatcher): hostbuf_size = comfy.model_management.pinned_hostbuf_size(self.model_size()) pin_state["weights"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) pin_state["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) + pin_state["weights-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) + pin_state["patches-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) pin_state["hostbufs_initialized"] = True pin_state["failed"] = False pin_state["active"] = True @@ -1943,12 +1997,14 @@ class ModelPatcherDynamic(ModelPatcher): return freed def loaded_ram_size(self): - return (self.model.dynamic_pins[self.load_device]["weights"][0].size) + pin_state = self.model.dynamic_pins[self.load_device] + return pin_state["weights"][0].size + pin_state["weights-loaded"][0].size def pinned_memory_size(self): - return (self.model.dynamic_pins[self.load_device]["weights"][3][0]) + pin_state = self.model.dynamic_pins[self.load_device] + return pin_state["weights"][3][0] + pin_state["weights-loaded"][3][0] - def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights", "patches" ]): + def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]): freed = 0 pin_state = self.model.dynamic_pins[self.load_device] for subset in subsets: @@ -1956,15 +2012,17 @@ class ModelPatcherDynamic(ModelPatcher): split = stack_split[0] while split >= 0: module, offset = stack[split] + module_pin = module._pins[subset] split -= 1 stack_split[0] = split - if not module._pin_registered: + if not module_pin["registered"]: continue - size = module._pin.numel() * module._pin.element_size() - if torch.cuda.cudart().cudaHostUnregister(module._pin.data_ptr()) != 0: + pin = module_pin["pin"] + size = pin.numel() * pin.element_size() + if torch.cuda.cudart().cudaHostUnregister(pin.data_ptr()) != 0: comfy.model_management.discard_cuda_async_error() continue - module._pin_registered = False + module_pin["registered"] = False comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size) pinned_size[0] = max(0, pinned_size[0] - size) freed += size @@ -1973,20 +2031,23 @@ class ModelPatcherDynamic(ModelPatcher): return freed return freed - def partially_unload_ram(self, ram_to_unload, subsets=[ "weights", "patches" ]): + def partially_unload_ram(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]): freed = 0 pin_state = self.model.dynamic_pins[self.load_device] for subset in subsets: hostbuf, stack, stack_split, pinned_size, *_ = pin_state[subset] while len(stack) > 0: module, offset = stack.pop() - size = module._pin.numel() * module._pin.element_size() - module._pin_balancer_entry[-1] = None - del module._pin_balancer_entry - del module._pin - hostbuf.truncate(offset, do_unregister=module._pin_registered) + module_pin = module._pins[subset] + pin = module_pin["pin"] + size = pin.numel() * pin.element_size() + module_pin["balancer_entry"][-1] = None + del module_pin["balancer_entry"] + del module_pin["pin"] + registered = module_pin["registered"] + hostbuf.truncate(offset, do_unregister=registered) stack_split[0] = min(stack_split[0], len(stack) - 1) - if module._pin_registered: + if registered: comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size) pinned_size[0] = max(0, pinned_size[0] - size) freed += size diff --git a/comfy/ops.py b/comfy/ops.py index 1f7cc9575..5e1cce333 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -144,8 +144,13 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin needs_cast = False xfer_source = [ s.weight, s.bias ] - - pin = comfy.pinned_memory.get_pin(s) + subset = "weights" + pin = comfy.pinned_memory.get_pin(s, subset=subset) + if pin is None and not args.fast_disk: + loaded_pin = comfy.pinned_memory.get_pin(s, subset="weights-loaded") + if loaded_pin is not None or signature is not None: + subset = "weights-loaded" + pin = loaded_pin if pin is not None: xfer_source = [ pin ] @@ -182,12 +187,12 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin if pin is not None: cast_maybe_lowvram_patch([pin], dest, offload_stream) return - if signature is None or args.high_ram: + if signature is None or not args.fast_disk or args.high_ram: comfy.pinned_memory.pin_memory(m, subset=subset, size=size) pin = comfy.pinned_memory.get_pin(m, subset=subset) cast_maybe_lowvram_patch(source, pin, offload_stream, xfer_dest2=dest) - handle_pin(s, pin, xfer_source, xfer_dest, size=dest_size) + handle_pin(s, pin, xfer_source, xfer_dest, subset=subset, size=dest_size) for param_key in ("weight", "bias"): lowvram_source = getattr(s, param_key + "_lowvram_function", None) @@ -197,8 +202,16 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin lowvram_dest = get_cast_buffer(lowvram_size) lowvram_source.prepare(lowvram_dest, None, copy=False, commit=True) - pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches") - handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset="patches", size=lowvram_size) + subset = "patches" + pin = comfy.pinned_memory.get_pin(lowvram_source, subset=subset) + if pin is None: + loaded_pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches-loaded") + if loaded_pin is not None: + subset = "patches-loaded" + pin = loaded_pin + elif signature is not None and not args.fast_disk: + subset = "patches-loaded" + handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset=subset, size=lowvram_size) prefetch["xfer_dest"] = xfer_dest diff --git a/comfy/pinned_memory.py b/comfy/pinned_memory.py index cb77c517a..d78ab3c76 100644 --- a/comfy/pinned_memory.py +++ b/comfy/pinned_memory.py @@ -9,14 +9,14 @@ import torch from comfy.cli_args import args -def _add_to_bucket(module, buckets, size, priority): +def _add_to_bucket(module, module_pin, buckets, size, priority): bucket = buckets.setdefault(size, []) entry = [-priority, 0, module] entry[1] = id(entry) bisect.insort(bucket, entry) - module._pin_balancer_entry = entry + module_pin["balancer_entry"] = entry -def _steal_pin(module, stack, buckets, size, priority): +def _steal_pin(module, stack, buckets, size, priority, subset): bucket = buckets.get(size) if bucket is None: return False @@ -31,34 +31,39 @@ def _steal_pin(module, stack, buckets, size, priority): return False *_, victim = bucket.pop() - module._pin = victim._pin - module._pin_registered = victim._pin_registered - module._pin_stack_index = victim._pin_stack_index - stack[module._pin_stack_index] = (module, stack[module._pin_stack_index][1]) + module_pin = module._pins[subset] + victim_pin = victim._pins[subset] + module_pin["pin"] = victim_pin["pin"] + module_pin["registered"] = victim_pin["registered"] + module_pin["stack_index"] = victim_pin["stack_index"] + stack_index = module_pin["stack_index"] + stack[stack_index] = (module, stack[stack_index][1]) - victim._pin_registered = False - del victim._pin - del victim._pin_stack_index - del victim._pin_balancer_entry + victim_pin["registered"] = False + del victim_pin["pin"] + del victim_pin["stack_index"] + del victim_pin["balancer_entry"] - _add_to_bucket(module, buckets, size, priority) + _add_to_bucket(module, module_pin, buckets, size, priority) return True def get_pin(module, subset="weights"): - pin = getattr(module, "_pin", None) - if pin is None or module._pin_registered or args.disable_pinned_memory: + pins = module.__dict__.get("_pins") + module_pin = None if pins is None else pins.get(subset) + pin = None if module_pin is None else module_pin.get("pin") + if pin is None or module_pin["registered"] or args.disable_pinned_memory: return pin _, _, stack_split, pinned_size, *_ = module._pin_state[subset] size = pin.nbytes - comfy.model_management.ensure_pin_registerable(size) + comfy.model_management.ensure_pin_registerable(size, loaded=subset.endswith("-loaded")) if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() return pin - module._pin_registered = True - stack_split[0] = max(stack_split[0], module._pin_stack_index) + module_pin["registered"] = True + stack_split[0] = max(stack_split[0], module_pin["stack_index"]) comfy.model_management.TOTAL_PINNED_MEMORY += size pinned_size[0] += size return pin @@ -72,23 +77,26 @@ def pin_memory(module, subset="weights", size=None): if pin is not None: return + pins = module.__dict__.setdefault("_pins", {}) + module_pin = pins.setdefault(subset, {}) hostbuf, stack, stack_split, pinned_size, counter, buckets = pin_state[subset] if size is None: size = comfy.memory_management.vram_aligned_size([ module.weight, module.bias ]) - offset = hostbuf.size registerable_size = size - priority = getattr(module, "_pin_balancer_priority", None) + loaded = subset.endswith("-loaded") + priority = module_pin.get("balancer_priority") if priority is None: priority = comfy.utils.bit_reverse_range(counter[0], 16) counter[0] += 1 - module._pin_balancer_priority = priority + module_pin["balancer_priority"] = priority comfy.memory_management.extra_ram_release(comfy.memory_management.RAM_CACHE_HEADROOM) - if (not comfy.model_management.ensure_pin_budget(size) or - not comfy.model_management.ensure_pin_registerable(registerable_size)): - return _steal_pin(module, stack, buckets, size, priority) + if (not comfy.model_management.ensure_pin_budget(size, loaded=loaded) or + not comfy.model_management.ensure_pin_registerable(registerable_size, loaded=loaded)): + return _steal_pin(module, stack, buckets, size, priority, subset) + offset = hostbuf.size extended = False try: hostbuf.extend(size=size, register=False) @@ -97,23 +105,23 @@ def pin_memory(module, subset="weights", size=None): pin.untyped_storage()._comfy_hostbuf = hostbuf if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() - comfy.model_management.free_registrations(size) + comfy.model_management.free_registrations(size, loaded=loaded) if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() del pin hostbuf.truncate(offset, do_unregister=False) - return _steal_pin(module, stack, buckets, size, priority) + return _steal_pin(module, stack, buckets, size, priority, subset) except RuntimeError: if extended: hostbuf.truncate(offset, do_unregister=False) - return _steal_pin(module, stack, buckets, size, priority) + return _steal_pin(module, stack, buckets, size, priority, subset) - module._pin = pin + module_pin["pin"] = pin stack.append((module, offset)) - module._pin_registered = True - module._pin_stack_index = len(stack) - 1 - stack_split[0] = max(stack_split[0], module._pin_stack_index) + module_pin["registered"] = True + module_pin["stack_index"] = len(stack) - 1 + stack_split[0] = max(stack_split[0], module_pin["stack_index"]) comfy.model_management.TOTAL_PINNED_MEMORY += size pinned_size[0] += size - _add_to_bucket(module, buckets, size, priority) + _add_to_bucket(module, module_pin, buckets, size, priority) return True diff --git a/comfy_execution/caching.py b/comfy_execution/caching.py index 6bd99b68f..d60aa1e50 100644 --- a/comfy_execution/caching.py +++ b/comfy_execution/caching.py @@ -5,7 +5,7 @@ import psutil import time import torch from typing import Sequence, Mapping, Dict -from comfy.model_patcher import ModelPatcher +from comfy.model_patcher import is_model_patcher_output from comfy_execution.graph import DynamicPrompt from abc import ABC, abstractmethod @@ -567,7 +567,7 @@ class RAMPressureCache(LRUCache): elif isinstance(output, torch.Tensor) and output.device.type == 'cpu': ram_usage += output.numel() * output.element_size() oom_ram_usage += output.numel() * output.element_size() - elif isinstance(output, ModelPatcher) and self.used_generation[key] != self.generation: + elif is_model_patcher_output(output) and self.used_generation[key] != self.generation: #old ModelPatchers are the first to go oom_ram_usage = 1e30 scan_list_for_ram_usage(cache_entry.outputs) diff --git a/comfy_execution/graph.py b/comfy_execution/graph.py index 479ee8a53..64dec2045 100644 --- a/comfy_execution/graph.py +++ b/comfy_execution/graph.py @@ -195,9 +195,10 @@ class ExecutionList(TopologicalSort): ExecutionList implements a topological dissolve of the graph. After a node is staged for execution, it can still be returned to the graph after having further dependencies added. """ - def __init__(self, dynprompt, output_cache): + def __init__(self, dynprompt, output_cache, output_link_callback=None): super().__init__(dynprompt) self.output_cache = output_cache + self.output_link_callback = output_link_callback self.staged_node_id = None self.execution_cache = {} self.execution_cache_listeners = {} @@ -205,13 +206,16 @@ class ExecutionList(TopologicalSort): def is_cached(self, node_id): return self.output_cache.get_local(node_id) is not None - def cache_link(self, from_node_id, to_node_id): + def cache_link(self, from_node_id, to_node_id, from_socket=None): if to_node_id not in self.execution_cache: self.execution_cache[to_node_id] = {} - self.execution_cache[to_node_id][from_node_id] = self.output_cache.get_local(from_node_id) + value = self.output_cache.get_local(from_node_id) + self.execution_cache[to_node_id][from_node_id] = value if from_node_id not in self.execution_cache_listeners: self.execution_cache_listeners[from_node_id] = set() - self.execution_cache_listeners[from_node_id].add(to_node_id) + self.execution_cache_listeners[from_node_id].add((to_node_id, from_socket)) + if value is not None and from_socket is not None and self.output_link_callback is not None: + self.output_link_callback(value.outputs[from_socket]) def get_cache(self, from_node_id, to_node_id): if to_node_id not in self.execution_cache: @@ -225,13 +229,15 @@ class ExecutionList(TopologicalSort): def cache_update(self, node_id, value): if node_id in self.execution_cache_listeners: - for to_node_id in self.execution_cache_listeners[node_id]: + for to_node_id, from_socket in self.execution_cache_listeners[node_id]: if to_node_id in self.execution_cache: self.execution_cache[to_node_id][node_id] = value + if from_socket is not None and self.output_link_callback is not None: + self.output_link_callback(value.outputs[from_socket]) def add_strong_link(self, from_node_id, from_socket, to_node_id): super().add_strong_link(from_node_id, from_socket, to_node_id) - self.cache_link(from_node_id, to_node_id) + self.cache_link(from_node_id, to_node_id, from_socket) async def stage_node_execution(self): assert self.staged_node_id is None diff --git a/execution.py b/execution.py index 387772629..b17ace65a 100644 --- a/execution.py +++ b/execution.py @@ -16,6 +16,7 @@ import torch from comfy.cli_args import args import comfy.memory_management import comfy.model_management +import comfy.model_patcher import comfy.model_prefetch import comfy_aimdo.model_vbar @@ -664,6 +665,7 @@ class PromptExecutor: self.cache_args = cache_args self.cache_type = cache_type self.server = server + self.prompt_model_tracker = comfy.model_patcher.PromptModelTracker() self.reset() def reset(self): @@ -728,6 +730,7 @@ class PromptExecutor: set_preview_method(extra_data.get("preview_method")) nodes.interrupt_processing(False) + self.prompt_model_tracker.start() if "client_id" in extra_data: self.server.client_id = extra_data["client_id"] @@ -770,7 +773,7 @@ class PromptExecutor: pending_async_nodes = {} # TODO - Unify this with pending_subgraph_results ui_node_outputs = {} executed = set() - execution_list = ExecutionList(dynamic_prompt, self.caches.outputs) + execution_list = ExecutionList(dynamic_prompt, self.caches.outputs, self.prompt_model_tracker.add) current_outputs = self.caches.outputs.all_node_ids() for node_id in list(execute_outputs): execution_list.add_node(node_id) @@ -833,6 +836,7 @@ class PromptExecutor: comfy.model_management.unload_all_models() finally: comfy.memory_management.set_ram_cache_release_state(None, 0) + self.prompt_model_tracker.end() self._notify_prompt_lifecycle("end", prompt_id) From fbe6d3ca8fc19ab5bd47690c64bad5e844dc971c Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Wed, 29 Jul 2026 07:31:45 +1000 Subject: [PATCH 168/211] Add configurable DETAIL logging side channel (#15064) --- app/logger.py | 31 +++++++++++++++++++++++++++++-- comfy/cli_args.py | 27 ++++++++++++++++++++++++++- comfy/logging.py | 10 ++++++++++ comfy/model_management.py | 6 ++++++ comfy/model_patcher.py | 21 +++++++++++++++++++-- comfy/samplers.py | 15 ++++++++++++--- comfy_execution/caching.py | 11 +++++++++++ execution.py | 7 +++++-- main.py | 18 +++++++++++++----- 9 files changed, 131 insertions(+), 15 deletions(-) create mode 100644 comfy/logging.py diff --git a/app/logger.py b/app/logger.py index bde815822..1aed54e37 100644 --- a/app/logger.py +++ b/app/logger.py @@ -2,9 +2,12 @@ from collections import deque from datetime import datetime import io import logging +import os import sys import threading +import comfy.logging + ANSI_NAMED_COLORS = { 'black': '\033[30m', 'red': '\033[31m', @@ -18,6 +21,7 @@ ANSI_NAMED_COLORS = { ANSI_LEVEL_COLORS = { 'DEBUG': ANSI_NAMED_COLORS['cyan'], + 'DETAIL': ANSI_NAMED_COLORS['blue'], 'INFO': ANSI_NAMED_COLORS['green'], 'WARNING': ANSI_NAMED_COLORS['yellow'], 'ERROR': ANSI_NAMED_COLORS['red'], @@ -85,7 +89,12 @@ def on_flush(callback): if stderr_interceptor is not None: stderr_interceptor.on_flush(callback) -def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool = False): + +def get_log_level(level): + return comfy.logging.DETAIL if level == "DETAIL" else logging.getLevelName(level) + + +def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300, use_stdout: bool = False): global logs if logs: return @@ -99,13 +108,18 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr) # Setup default global logger + if file_outputs is None: + file_outputs = [('DETAIL', 'comfyui_detail.log')] logger = logging.getLogger() - logger.setLevel(log_level) + console_level = get_log_level(log_level) + file_levels = [get_log_level(level) for level, _ in file_outputs] + logger.setLevel(min(console_level, *file_levels)) formatter = ColoredFormatter("%(message)s") stream_handler = logging.StreamHandler() stream_handler.setFormatter(formatter) + stream_handler.setLevel(console_level) if use_stdout: # Only errors and critical to stderr @@ -114,11 +128,24 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool # Lesser to stdout stdout_handler = logging.StreamHandler(sys.stdout) stdout_handler.setFormatter(formatter) + stdout_handler.setLevel(console_level) stdout_handler.addFilter(lambda record: record.levelno < logging.ERROR) logger.addHandler(stdout_handler) logger.addHandler(stream_handler) + for output_level, output_path in file_outputs: + output_path = os.path.abspath(output_path) + try: + output_handler = logging.FileHandler(output_path, encoding="utf-8") + except OSError as e: + logging.warning("Could not open %s log %s: %s", output_level, output_path, e) + continue + output_handler.setLevel(get_log_level(output_level)) + output_handler.setFormatter(logging.Formatter("[%(asctime)s] [%(levelname)s] %(message)s")) + logger.addHandler(output_handler) + logging.info("%s log: %s", output_level.title(), output_path) + STARTUP_WARNINGS = [] diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 8e03ed032..792148f0a 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -33,6 +33,31 @@ class EnumAction(argparse.Action): setattr(namespace, self.dest, value) +LOG_LEVELS = ('DEBUG', 'DETAIL', 'INFO', 'WARNING', 'ERROR', 'CRITICAL') + + +class VerboseAction(argparse.Action): + def __call__(self, parser, namespace, values, option_string=None): + if len(values) == 0: + output = ('DEBUG', None) + elif len(values) == 1 and values[0] in LOG_LEVELS: + output = (values[0], None) + elif len(values) == 2 and values[0] in LOG_LEVELS: + output = tuple(values) + else: + parser.error(f"{option_string} expects no values, a console LEVEL, or LEVEL FILE") + setattr(namespace, self.dest, [*getattr(namespace, self.dest, []), output]) + + +def get_console_log_level(outputs): + console_levels = [level for level, path in outputs if path is None] + return min(console_levels, key=LOG_LEVELS.index, default='INFO') + + +def get_file_log_outputs(outputs): + return [(level, path) for level, path in outputs if path is not None] + + parser = argparse.ArgumentParser() parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)") @@ -187,7 +212,7 @@ parser.add_argument("--disable-api-nodes", action="store_true", help="Disable lo parser.add_argument("--multi-user", action="store_true", help="Enables per-user storage.") -parser.add_argument("--verbose", default='INFO', const='DEBUG', nargs="?", choices=['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'], help='Set the logging level') +parser.add_argument("--verbose", action=VerboseAction, nargs='*', default=[], metavar='LEVEL FILE', help='Set console logging with no values or LEVEL, or add a LEVEL FILE log output. May be repeated.') parser.add_argument("--log-stdout", action="store_true", help="Send normal process output to stdout instead of stderr (default).") diff --git a/comfy/logging.py b/comfy/logging.py new file mode 100644 index 000000000..cc785296d --- /dev/null +++ b/comfy/logging.py @@ -0,0 +1,10 @@ +import logging + + +DETAIL = 15 +logging.addLevelName(DETAIL, "DETAIL") + + +def detail(message, *args, **kwargs): + kwargs.setdefault("stacklevel", 2) + logging.log(DETAIL, message, *args, **kwargs) diff --git a/comfy/model_management.py b/comfy/model_management.py index eb768d783..f7351224d 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -34,6 +34,7 @@ import comfy.utils import comfy.quant_ops import comfy_aimdo.host_buffer import comfy_aimdo.vram_buffer +from comfy.logging import detail from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -836,6 +837,8 @@ def minimum_inference_memory(): def free_memory(memory_required, device, keep_loaded=[], for_dynamic=False, pins_required=0, ram_required=0): cleanup_models_gc() + if not for_dynamic: + detail("Non dynamic memory free called! memory_required=%s pins_required=%s ram_required=%s", memory_required, pins_required, ram_required) unloaded_model = [] can_unload = [] unloaded_models = [] @@ -974,6 +977,9 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu lowvram_model_memory = 0.1 loaded_model.model_load(lowvram_model_memory, force_patch_weights=force_patch_weights) + vram_used = 0 if is_device_cpu(torch_dev) else loaded_model.model_loaded_memory() + ram_used = model.loaded_ram_size() if model.is_dynamic() else loaded_model.model_memory() - vram_used + detail("Model loaded: patcher=%s model=%s ram_mb=%.1f vram_mb=%.1f", model.__class__.__name__, model.model.__class__.__name__, ram_used / (1024 ** 2), vram_used / (1024 ** 2)) current_loaded_models.insert(0, loaded_model) return diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index 39246b95c..e44322e72 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -22,6 +22,7 @@ import collections import inspect import logging import math +import time import uuid from typing import Callable, Optional @@ -37,6 +38,7 @@ import comfy.patcher_extension import comfy.utils import comfy_aimdo.host_buffer from comfy.comfy_types import UnetWrapperFunction +from comfy.logging import detail from comfy.quant_ops import QuantizedTensor from comfy.patcher_extension import CallbacksMP, PatcherInjection, WrappersMP @@ -1989,10 +1991,25 @@ class ModelPatcherDynamic(ModelPatcher): assert self.load_device != torch.device("cpu") vbar = self._vbar_get() - freed = 0 if vbar is None else vbar.free_memory(memory_to_free) + vbar_freed = 0 if vbar is None else vbar.free_memory(memory_to_free) + freed = vbar_freed + backup_freed = 0 if freed < memory_to_free: - freed += self.restore_loaded_backups() + backup_freed = self.restore_loaded_backups() + freed += backup_freed + + method = "vbar+backups" if vbar_freed and backup_freed else "vbar" if vbar_freed else "backups" if backup_freed else "none" + free_methods = getattr(self, "_free_methods", {}) + free_methods[method] = free_methods.get(method, 0) + 1 + self._free_methods = free_methods + now = time.monotonic() + if now - getattr(self, "_last_free_log_time", 0) >= 5: + requested = "all" if memory_to_free >= 1e30 else f"{memory_to_free / (1024 ** 2):.1f}MB" + prevailing_method = max(free_methods, key=free_methods.get) + detail("AIMDO free: model=%s device=%s prevailing_method=%s methods=%s requested=%s vbar_mb=%.1f backups_mb=%.1f", self.model.__class__.__name__, self.load_device, prevailing_method, free_methods, requested, vbar_freed / (1024 ** 2), backup_freed / (1024 ** 2)) + self._free_methods = {} + self._last_free_log_time = now return freed diff --git a/comfy/samplers.py b/comfy/samplers.py index 25c5a855f..9f571ece9 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -20,6 +20,7 @@ import comfy.hooks import comfy.context_windows import comfy.multigpu import comfy.utils +from comfy.logging import detail import scipy.stats import numpy @@ -991,10 +992,15 @@ class KSAMPLER(Sampler): noise = model_wrap.inner_model.model_sampling.noise_scaling(sigmas[0], noise, latent_image, self.max_denoise(model_wrap, sigmas)) - k_callback = None total_steps = len(sigmas) - 1 - if callback is not None: - k_callback = lambda x: callback(x["i"], x["denoised"], x["x"], total_steps) + first_step = True + def k_callback(x): + nonlocal first_step + if first_step: + detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, self.sampler_function.__name__, x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape)) + first_step = False + if callback is not None: + callback(x["i"], x["denoised"], x["x"], total_steps) samples = self.sampler_function(model_k, noise, sigmas, extra_args=extra_args, callback=k_callback, disable=disable_pbar, **self.extra_options) samples = model_wrap.inner_model.model_sampling.inverse_noise_scaling(sigmas[-1], samples) @@ -1270,10 +1276,13 @@ class CFGGuider: return latent_image if latent_image.is_nested: + sampler_shapes = [tuple(x.shape) for x in latent_image.unbind()] latent_image, latent_shapes = comfy.utils.pack_latents(latent_image.unbind()) noise, _ = comfy.utils.pack_latents(noise.unbind()) else: latent_shapes = [latent_image.shape] + sampler_shapes = [tuple(latent_image.shape)] + detail("Sampler: model=%s latent_shapes=%s", self.model_patcher.model.__class__.__name__, sampler_shapes) if denoise_mask is not None: if denoise_mask.is_nested: diff --git a/comfy_execution/caching.py b/comfy_execution/caching.py index d60aa1e50..3340e5116 100644 --- a/comfy_execution/caching.py +++ b/comfy_execution/caching.py @@ -524,6 +524,13 @@ class RAMPressureCache(LRUCache): def __init__(self, key_class, enable_providers=False): super().__init__(key_class, 0, enable_providers=enable_providers) self.timestamps = {} + self.active_evictions = False + self.full_evictions = False + + async def set_prompt(self, dynprompt, node_ids, is_changed_cache): + self.active_evictions = False + self.full_evictions = False + await super().set_prompt(dynprompt, node_ids, is_changed_cache) def clean_unused(self): self._clean_subcaches() @@ -588,4 +595,8 @@ class RAMPressureCache(LRUCache): self.timestamps.pop(key, None) self.children.pop(key, None) freed += ram_usage + if freed and free_active: + self.active_evictions = True + if min_entry_size == 0: + self.full_evictions = True return freed diff --git a/execution.py b/execution.py index b17ace65a..7cab4b331 100644 --- a/execution.py +++ b/execution.py @@ -13,12 +13,13 @@ import asyncio import torch -from comfy.cli_args import args +from comfy.cli_args import args, get_console_log_level import comfy.memory_management import comfy.model_management import comfy.model_patcher import comfy.model_prefetch import comfy_aimdo.model_vbar +from comfy.logging import detail from latent_preview import set_preview_method import nodes @@ -544,7 +545,7 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed, output_data, output_ui, has_subgraph, has_pending_tasks = await get_output_data(prompt_id, unique_id, obj, input_data_all, execution_block_cb=execution_block_cb, pre_execute_cb=pre_execute_cb, v3_data=v3_data) finally: if comfy.memory_management.aimdo_enabled: - if args.verbose == "DEBUG": + if get_console_log_level(args.verbose) == "DEBUG": comfy_aimdo.control.analyze() comfy.model_management.reset_cast_buffers() comfy.model_prefetch.cleanup_prefetch_queues() @@ -835,6 +836,8 @@ class PromptExecutor: if comfy.model_management.DISABLE_SMART_MEMORY: comfy.model_management.unload_all_models() finally: + if self.cache_type == CacheType.RAM_PRESSURE: + detail("RAM cache evictions: prompt=%s active=%s full=%s", prompt_id, self.caches.outputs.active_evictions, self.caches.outputs.full_evictions) comfy.memory_management.set_ram_cache_release_state(None, 0) self.prompt_model_tracker.end() self._notify_prompt_lifecycle("end", prompt_id) diff --git a/main.py b/main.py index 1f16a7f89..c33e75f62 100644 --- a/main.py +++ b/main.py @@ -2,6 +2,7 @@ import comfy.options comfy.options.enable_args_parsing() from comfy.cli_args import args +from comfy.cli_args import get_console_log_level, get_file_log_outputs if args.list_feature_flags: import json @@ -17,7 +18,9 @@ import folder_paths import time from comfy.cli_args import enables_dynamic_vram from app.logger import setup_logger -setup_logger(log_level=args.verbose, use_stdout=args.log_stdout) +console_log_level = get_console_log_level(args.verbose) +file_log_outputs = [('DETAIL', 'comfyui_detail.log'), *get_file_log_outputs(args.verbose)] +setup_logger(log_level=console_log_level, file_outputs=file_log_outputs, use_stdout=args.log_stdout) from app.assets.seeder import asset_seeder from app.assets.services import register_output_files @@ -251,13 +254,18 @@ if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_managemen aimdo_initialized = comfy_aimdo.control.init_devices(d.index for d in comfy.model_management.get_all_torch_devices()) if aimdo_initialized: - if args.verbose == 'DEBUG': + if console_log_level == 'DEBUG': comfy_aimdo.control.set_log_debug() - elif args.verbose == 'CRITICAL': + elif console_log_level == 'DETAIL': + try: + comfy_aimdo.control.set_log_detail() + except AttributeError: + comfy_aimdo.control.set_log_info() + elif console_log_level == 'CRITICAL': comfy_aimdo.control.set_log_critical() - elif args.verbose == 'ERROR': + elif console_log_level == 'ERROR': comfy_aimdo.control.set_log_error() - elif args.verbose == 'WARNING': + elif console_log_level == 'WARNING': comfy_aimdo.control.set_log_warning() else: #INFO comfy_aimdo.control.set_log_info() From c38171ddb93368ee6a6bbc677b92e4b50cead865 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Wed, 29 Jul 2026 01:25:55 +0300 Subject: [PATCH 169/211] Support Pruna LTX VAE (#15129) --- comfy/ldm/lightricks/vae/causal_conv3d.py | 10 ++++++++-- comfy/ldm/lightricks/vae/causal_video_autoencoder.py | 6 +++--- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/comfy/ldm/lightricks/vae/causal_conv3d.py b/comfy/ldm/lightricks/vae/causal_conv3d.py index 7515f0d4e..bb1803f12 100644 --- a/comfy/ldm/lightricks/vae/causal_conv3d.py +++ b/comfy/ldm/lightricks/vae/causal_conv3d.py @@ -49,6 +49,12 @@ class CausalConv3d(nn.Module): ) self.temporal_cache_state={} + def _empty_output(self, x): + # empty (0 frame) outputs must still have the conv's output channels and spatial dims + h = (x.shape[3] + 2 * self.conv.padding[1] - self.conv.kernel_size[1]) // self.conv.stride[1] + 1 + w = (x.shape[4] + 2 * self.conv.padding[2] - self.conv.kernel_size[2]) // self.conv.stride[2] + 1 + return x.new_empty((x.shape[0], self.out_channels, 0, h, w)) + def forward(self, x, causal: bool = True): tid = threading.get_ident() @@ -58,7 +64,7 @@ class CausalConv3d(nn.Module): if not causal: padding_length = padding_length // 2 if x.shape[2] == 0: - return x + return self._empty_output(x) cached = x[:, :, :1, :, :].repeat((1, 1, padding_length, 1, 1)) pieces = [ cached, x ] if is_end and not causal: @@ -83,7 +89,7 @@ class CausalConv3d(nn.Module): elif is_end: self.temporal_cache_state[tid] = (None, True) - return self.conv(x) if x.shape[2] >= self.time_kernel_size else x[:, :, :0, :, :] + return self.conv(x) if x.shape[2] >= self.time_kernel_size else self._empty_output(x) @property def weight(self): diff --git a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py index 5975015e2..5d0eec5b8 100644 --- a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py +++ b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py @@ -390,10 +390,10 @@ class Decoder(nn.Module): # Compute output channel to be product of all channel-multiplier blocks output_channel = base_channels - for block_name, block_params in list(reversed(blocks)): + for block_name, block_params in blocks: block_params = block_params if isinstance(block_params, dict) else {} if block_name == "res_x_y": - output_channel = output_channel * block_params.get("multiplier", 2) + output_channel = block_params.get("in_channels", output_channel * block_params.get("multiplier", 2)) if block_name == "compress_all": output_channel = output_channel * block_params.get("multiplier", 1) if block_name == "compress_space": @@ -432,7 +432,7 @@ class Decoder(nn.Module): spatial_padding_mode=spatial_padding_mode, ) elif block_name == "res_x_y": - output_channel = output_channel // block_params.get("multiplier", 2) + output_channel = block_params.get("out_channels", output_channel // block_params.get("multiplier", 2)) block = ResnetBlock3D( dims=dims, in_channels=input_channel, From 42d2aa55432b57371ddc9d4078ae250b54227641 Mon Sep 17 00:00:00 2001 From: Kohaku-Blueleaf <59680068+KohakuBlueleaf@users.noreply.github.com> Date: Wed, 29 Jul 2026 07:03:04 +0800 Subject: [PATCH 170/211] [Dataset/Security,Feature] Add dataset folder to avoid arbitrary folder access for dataset stuff. (#14807) --- comfy_extras/nodes_dataset.py | 115 +++++++++++++++++++++++++++++---- extra_model_paths.yaml.example | 1 + folder_paths.py | 2 + 3 files changed, 105 insertions(+), 13 deletions(-) diff --git a/comfy_extras/nodes_dataset.py b/comfy_extras/nodes_dataset.py index d7e4652cf..5e0454d8b 100644 --- a/comfy_extras/nodes_dataset.py +++ b/comfy_extras/nodes_dataset.py @@ -43,6 +43,98 @@ def load_and_process_images(image_files, input_dir): return output_images +def secure_subfolder_path(base_dir, folder_name): + """Resolve folder_name inside base_dir, rejecting anything that escapes it. + + Blocks '..', absolute paths, drive letters and symlink escapes using the + same realpath containment check as the core file endpoints. + """ + target = os.path.abspath(os.path.join(base_dir, folder_name)) + if not folder_paths.is_within_directory(base_dir, target): + raise ValueError(f"Invalid folder name {folder_name!r}: resolves outside of {base_dir}") + return target + + +def list_dataset_folders(): + """Relative paths of dataset folders found under all dataset roots. + + Any subfolder containing a metadata.json or *.safetensors shard counts as + a dataset; the walk doesn't descend into matched folders. + + Symlinked directories are followed, but symlink loops are avoided. + """ + found = set() + + for root in folder_paths.get_folder_paths("datasets"): + if not os.path.isdir(root): + continue + + root = os.path.abspath(root) + seen_dirs = set() + + for dirpath, subdirs, filenames in os.walk(root, followlinks=True): + try: + st = os.stat(dirpath) # follows symlinks + except OSError: + subdirs[:] = [] + continue + + dir_key = (st.st_dev, st.st_ino) + if dir_key in seen_dirs: + subdirs[:] = [] + continue + + seen_dirs.add(dir_key) + + if dirpath != root and ( + "metadata.json" in filenames + or any(f.endswith(".safetensors") for f in filenames) + ): + found.add(os.path.relpath(dirpath, root).replace(os.sep, "/")) + subdirs[:] = [] + continue + + kept_subdirs = [] + for name in subdirs: + child = os.path.join(dirpath, name) + try: + child_st = os.stat(child) # follows symlinks + except OSError: + continue + + child_key = (child_st.st_dev, child_st.st_ino) + if child_key not in seen_dirs: + kept_subdirs.append(name) + + subdirs[:] = kept_subdirs + + return sorted(found) + + +def get_dataset_save_dir(folder_name): + """Resolve the folder to save a new dataset into, inside the default root. + + The folder is not created here; callers makedirs after validation. + """ + root = folder_paths.get_folder_paths("datasets")[0] + target = secure_subfolder_path(root, folder_name) + if os.path.realpath(target) == os.path.realpath(root): + raise ValueError("folder_name must name a subfolder of the datasets directory, e.g. 'my_dataset'.") + return target + + +def get_dataset_dir(folder_name): + """Find an existing dataset folder by relative name across all dataset roots.""" + roots = folder_paths.get_folder_paths("datasets") + for root in roots: + target = secure_subfolder_path(root, folder_name) + if os.path.realpath(target) == os.path.realpath(root): + raise ValueError("folder_name must name a subfolder of the datasets directory, e.g. 'my_dataset'.") + if os.path.isdir(target): + return target + raise ValueError(f"Dataset folder {folder_name!r} not found in: {', '.join(roots)}") + + VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"] @@ -395,7 +487,7 @@ class SaveImageDataSetToFolderNode(io.ComfyNode): filename_prefix = filename_prefix[0] mode = mode[0] - output_dir = os.path.join(folder_paths.get_output_directory(), folder_name) + output_dir = secure_subfolder_path(folder_paths.get_output_directory(), folder_name) saved_files = save_images_to_folder(images, output_dir, filename_prefix, mode=='overwrite') logging.info(f"Saved {len(saved_files)} images to {output_dir}.") @@ -449,7 +541,7 @@ class SaveImageTextDataSetToFolderNode(io.ComfyNode): filename_prefix = filename_prefix[0] mode = mode[0] - output_dir = os.path.join(folder_paths.get_output_directory(), folder_name) + output_dir = secure_subfolder_path(folder_paths.get_output_directory(), folder_name) saved_files = save_images_to_folder(images, output_dir, filename_prefix, mode=='overwrite') # Save captions @@ -1861,7 +1953,7 @@ class SaveTrainingDataset(io.ComfyNode): io.String.Input( "folder_name", default="training_dataset", - tooltip="Name of folder to save dataset (inside output directory).", + tooltip="Name of folder to save the dataset into, inside the datasets directory. Subfolders like 'project/run1' are allowed.", ), io.Int.Input( "shard_size", @@ -1891,8 +1983,8 @@ class SaveTrainingDataset(io.ComfyNode): f"Something went wrong in dataset preparation." ) - # Create output directory - output_dir = os.path.join(folder_paths.get_output_directory(), folder_name) + # Create output directory (inside the datasets root, traversal-safe) + output_dir = get_dataset_save_dir(folder_name) os.makedirs(output_dir, exist_ok=True) # Prepare data pairs @@ -1951,10 +2043,10 @@ class LoadTrainingDataset(io.ComfyNode): description="Load encoded training dataset (latents + conditioning) from disk for use in training.", is_experimental=True, inputs=[ - io.String.Input( + io.Combo.Input( "folder_name", - default="training_dataset", - tooltip="Name of folder containing the saved dataset (inside output directory).", + options=list_dataset_folders(), + tooltip="Saved dataset to load, from the datasets directory.", ), ], outputs=[ @@ -1973,11 +2065,8 @@ class LoadTrainingDataset(io.ComfyNode): @classmethod def execute(cls, folder_name): - # Get dataset directory - dataset_dir = os.path.join(folder_paths.get_output_directory(), folder_name) - - if not os.path.exists(dataset_dir): - raise ValueError(f"Dataset directory not found: {dataset_dir}") + # Get dataset directory (searched across all dataset roots, traversal-safe) + dataset_dir = get_dataset_dir(folder_name) # Find all shard files shard_files = sorted( diff --git a/extra_model_paths.yaml.example b/extra_model_paths.yaml.example index 6a31d8a63..755b8d124 100644 --- a/extra_model_paths.yaml.example +++ b/extra_model_paths.yaml.example @@ -29,6 +29,7 @@ # upscale_models: models/upscale_models/ # latent_upscale_models: models/latent_upscale_models/ # custom_nodes: custom_nodes/ +# datasets: datasets/ # hypernetworks: models/hypernetworks/ # photomaker: models/photomaker/ # classifiers: models/classifiers/ diff --git a/folder_paths.py b/folder_paths.py index 937428c18..bd3f25095 100644 --- a/folder_paths.py +++ b/folder_paths.py @@ -44,6 +44,8 @@ folder_names_and_paths["latent_upscale_models"] = ([os.path.join(models_dir, "la folder_names_and_paths["custom_nodes"] = ([os.path.join(base_path, "custom_nodes")], set()) +folder_names_and_paths["datasets"] = ([os.path.join(base_path, "datasets")], set()) + folder_names_and_paths["hypernetworks"] = ([os.path.join(models_dir, "hypernetworks")], supported_pt_extensions) folder_names_and_paths["photomaker"] = ([os.path.join(models_dir, "photomaker")], supported_pt_extensions) From e651b7bef55a5376343dcb1c0edb79f0142c985e Mon Sep 17 00:00:00 2001 From: Barish Ozbay <17261091+drozbay@users.noreply.github.com> Date: Tue, 28 Jul 2026 23:24:08 -0400 Subject: [PATCH 171/211] Fix LTXAV crash when sampling without an audio latent (#15132) --- comfy/ldm/lightricks/model.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index 92bb8118c..f9de3a38e 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -671,9 +671,9 @@ def freqs_cis_matrix(freqs, pad_size, split_mode, num_attention_heads, out_dtype cos_freq = torch.cat((cos_padding, cos_freq), dim=-1) sin_freq = torch.cat((sin_padding, sin_freq), dim=-1) - B, T, _ = cos_freq.shape - cos_freq = cos_freq.reshape(B, T, num_attention_heads, -1) - sin_freq = sin_freq.reshape(B, T, num_attention_heads, -1) + B, T, half_HD = cos_freq.shape + cos_freq = cos_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) + sin_freq = sin_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) rotation_matrix = torch.stack( (cos_freq, -sin_freq, sin_freq, cos_freq), dim=-1 ) From 4f874c5e3a2fafb9938273c5f3fc47d2da017667 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 29 Jul 2026 14:10:01 -0700 Subject: [PATCH 172/211] Update comfy-kitchen to fix flux kv issue. (#15144) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 3a8203aff..ef30cd7fa 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.23 +comfy-kitchen==0.2.24 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From f73e8cde88794bd9568474e92f4421aa5622ff1a Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Wed, 29 Jul 2026 18:10:28 -0700 Subject: [PATCH 173/211] Fallback to cudnn attention on linux if flash attention doesn't work. (#15146) --- comfy/ops.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 5e1cce333..9d692dcc7 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -41,7 +41,7 @@ def scaled_dot_product_attention(q, k, v, *args, **kwargs): try: - if torch.cuda.is_available() and comfy.model_management.WINDOWS: + if torch.cuda.is_available(): from torch.nn.attention import SDPBackend, sdpa_kernel import inspect if "set_priority" in inspect.signature(sdpa_kernel).parameters: @@ -51,7 +51,10 @@ try: SDPBackend.MATH, ] - SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION) + if comfy.model_management.WINDOWS: + SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION) + else: + SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION) def scaled_dot_product_attention(q, k, v, *args, **kwargs): if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower From c65f9f169cd04637540026f8e4e506715c3c76f0 Mon Sep 17 00:00:00 2001 From: kaalibro <44464226+kaalibro@users.noreply.github.com> Date: Thu, 30 Jul 2026 11:00:16 +0500 Subject: [PATCH 174/211] Fix user.css loading broken by #14734 (#15000) --- app/user_manager.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/app/user_manager.py b/app/user_manager.py index de261ad39..55e7e81e3 100644 --- a/app/user_manager.py +++ b/app/user_manager.py @@ -343,13 +343,22 @@ class UserManager(): # XSS). Content-Disposition: attachment is the load-bearing guard; # the content-type override and nosniff are defence in depth. content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream' - if folder_paths.is_dangerous_content_type(content_type): - content_type = 'application/octet-stream' + + user_root = self.get_request_user_filepath(request, None, create_dir=False) + is_user_css = path == os.path.abspath(os.path.join(user_root, "user.css")) + + if is_user_css: + content_type = "text/css" + disposition = "inline" + else: + if folder_paths.is_dangerous_content_type(content_type): + content_type = 'application/octet-stream' + disposition = "attachment" return web.FileResponse(path, headers={ "Content-Type": content_type, "X-Content-Type-Options": "nosniff", - "Content-Disposition": "attachment", + "Content-Disposition": disposition, }) @routes.post("/userdata/{file}") From 7374157e95aee86ae8c20cbc6a283702bcbb666f Mon Sep 17 00:00:00 2001 From: Denis Date: Thu, 30 Jul 2026 08:26:14 +0200 Subject: [PATCH 175/211] fix(jobs): prefer media over text for job preview_output (#14681) --- comfy_execution/jobs.py | 35 +++++++++++++--- tests/execution/test_jobs.py | 80 ++++++++++++++++++++++++++++++++++++ 2 files changed, 109 insertions(+), 6 deletions(-) diff --git a/comfy_execution/jobs.py b/comfy_execution/jobs.py index f0ad59f86..34c06363b 100644 --- a/comfy_execution/jobs.py +++ b/comfy_execution/jobs.py @@ -170,6 +170,19 @@ def is_previewable(media_type: str, item: dict) -> bool: return False +def is_text_preview(media_type: str, item: dict) -> bool: + """ + Check if a previewable output item is textual rather than visual media. + + Saved text files (SaveText's .txt/.md/.json) are real outputs but must not + outrank visual media when picking the job preview. + """ + if media_type == 'text': + return True + filename = item.get('filename', '').lower() + return any(filename.endswith(ext) for ext in TEXT_EXTENSIONS) + + def normalize_queue_item(item: tuple, status: str) -> dict: """Convert queue item tuple to unified job dict. @@ -259,8 +272,13 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: Returns (outputs_count, preview_output). Preview priority (matching frontend): - 1. type="output" with previewable media - 2. Any previewable media + 1. type="output" visual media (saved images/video/audio/3d) + 2. any other previewable visual media (e.g. temp/preview images) + 3. saved text file (e.g. SaveText's .txt/.md/.json) + 4. raw text (only when the job produced nothing else previewable) + + Text is kept in its own slots so node/execution order can't let a text + output mask a visual one (e.g. a text node that runs before an image). Text content entries (strings under 'text') are preview-only metadata, matching the frontend's METADATA_KEYS: they can serve as the fallback @@ -269,6 +287,8 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: count = 0 preview_output = None fallback_preview = None + text_file_fallback = None + text_fallback = None for node_id, node_outputs in outputs.items(): if not isinstance(node_outputs, dict): @@ -296,8 +316,8 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: 'nodeId': node_id, 'mediaType': media_type } - if fallback_preview is None: - fallback_preview = enriched + if text_fallback is None: + text_fallback = enriched continue # normalize_output_item returned a dict (e.g. 3D file) item = normalized @@ -314,12 +334,15 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: } if 'mediaType' not in item: enriched['mediaType'] = media_type - if item.get('type') == 'output': + if is_text_preview(media_type, item): + if text_file_fallback is None: + text_file_fallback = enriched + elif item.get('type') == 'output': preview_output = enriched elif fallback_preview is None: fallback_preview = enriched - return count, preview_output or fallback_preview + return count, preview_output or fallback_preview or text_file_fallback or text_fallback def apply_sorting(jobs: list[dict], sort_by: str, sort_order: str) -> list[dict]: diff --git a/tests/execution/test_jobs.py b/tests/execution/test_jobs.py index f7cb612e4..cef2b41cb 100644 --- a/tests/execution/test_jobs.py +++ b/tests/execution/test_jobs.py @@ -280,6 +280,86 @@ class TestGetOutputsSummary: assert preview['filename'] == 'model.glb' assert preview['mediaType'] == '3d' + def test_media_preview_preferred_over_text(self): + """A visual output wins the preview even when a text node is iterated + first (regression: text could mask a later temp/preview image).""" + outputs = { + 'text_node': {'text': ['a caption']}, + 'image_node': {'images': [{'filename': 'preview.png', 'type': 'temp'}]}, + } + count, preview = get_outputs_summary(outputs) + # Text is preview-only metadata and not counted; only the image counts. + assert count == 1 + assert preview['filename'] == 'preview.png' + assert preview['mediaType'] == 'images' + + def test_text_used_as_preview_when_no_media(self): + """Text is the preview only when the job produced no media output.""" + outputs = { + 'text_node': {'text': ['hello world']}, + } + count, preview = get_outputs_summary(outputs) + assert count == 0 # text entries are not counted as outputs + assert preview['mediaType'] == 'text' + assert preview['content'] == 'hello world' + + def test_media_preview_preferred_over_saved_text_file(self): + """A visual output wins the preview over a saved text file (SaveText), + even a temp/preview image iterated after the text node.""" + outputs = { + 'save_text': { + 'text': ['the text'], + 'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}], + }, + 'preview_image': {'images': [{'filename': 'preview.png', 'type': 'temp'}]}, + } + count, preview = get_outputs_summary(outputs) + assert count == 2 # the .txt file and the image; raw text is metadata + assert preview['filename'] == 'preview.png' + assert preview['mediaType'] == 'images' + + def test_saved_media_preferred_over_saved_text_file(self): + outputs = { + 'save_text': { + 'text': ['the text'], + 'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}], + }, + 'save_image': {'images': [{'filename': 'result.png', 'type': 'output'}]}, + } + count, preview = get_outputs_summary(outputs) + assert count == 2 + assert preview['filename'] == 'result.png' + + def test_mime_format_file_preferred_over_saved_text_file(self): + """Custom-node outputs previewable via MIME format (e.g. VHS videos + under arbitrary keys) rank as visual media, above saved text files.""" + outputs = { + 'save_text': { + 'files': [{'filename': 'notes.md', 'subfolder': '', 'type': 'output'}], + }, + 'video_node': { + 'files': [{'filename': 'clip.webm', 'format': 'video/webm', 'type': 'output'}], + }, + } + count, preview = get_outputs_summary(outputs) + assert count == 2 + assert preview['filename'] == 'clip.webm' + + + def test_saved_text_file_preferred_over_raw_text(self): + """With no media in the job, the saved text file (a real, counted + output) is the preview rather than the raw text metadata.""" + outputs = { + 'save_text': { + 'text': ['the text'], + 'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}], + }, + } + count, preview = get_outputs_summary(outputs) + assert count == 1 + assert preview['filename'] == 'ComfyUI_00001.txt' + assert preview['mediaType'] == 'files' + class TestHas3DExtension: """Unit tests for has_3d_extension()""" From 9cf91339b708a245762fa38ffeec9702b381e0db Mon Sep 17 00:00:00 2001 From: Matt Miller Date: Wed, 29 Jul 2026 23:29:36 -0700 Subject: [PATCH 176/211] Fix SVG previews broken by the stored-XSS forced-download (#15149) * Fix SVG previews broken by the stored-XSS forced-download /view and the assets download route force every SVG to application/octet-stream + attachment. That blocks the stored XSS from GHSA-779p-m5rp-r4h4, but it also breaks the SVG node output and Media Assets previews, which request the file with a plain . Exempt only that case. An SVG referenced by an loads in secure static mode with scripting and external references disabled, so the payload cannot fire. The attack needs the SVG to become a document, which arrives with a different Sec-Fetch-Dest. Browsers set that header themselves and page script cannot override it. A missing header, from a non-browser client or a proxy that strips it, fails closed. The blocklist itself is unchanged; this is a call-site gate. * Don't let a cache replay the inline SVG into document context The Sec-Fetch-Dest exemption makes /view and the assets content route vary their Content-Type and Content-Disposition on a request header, but neither response said so. FileResponse emits Last-Modified/ETag and the cache_control middleware skips /view (the filename is in the query string, not the path), so the inline image/svg+xml variant is heuristically cacheable. A cache keyed on the URL alone could hand an entry primed by an load to a later top-level navigation of the same URL, turning the SVG back into a document and re-enabling the stored XSS the forced download blocks. Set Vary: Sec-Fetch-Dest and Cache-Control: no-store on both branches, not just the exempt one: a cached attachment replayed to an would re-break the preview this fix exists to restore. Also strip parameters from content_type before building the assets response. mime_type there is uploader-supplied and unvalidated, and aiohttp rejects a charset in the content_type argument with ValueError, so a stored "image/svg+xml; charset=utf-8" turned a valid inline SVG into a 500. Route-level guards now pin the headers on both branches and the parameterised mime type; all three fail against the previous commit. --- app/assets/api/routes.py | 29 ++- folder_paths.py | 17 ++ server.py | 34 ++-- ...test_ghsa_779p_06_inline_svg_image_dest.py | 191 ++++++++++++++++++ 4 files changed, 252 insertions(+), 19 deletions(-) create mode 100644 tests-unit/security_test/test_ghsa_779p_06_inline_svg_image_dest.py diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index 43e60094c..e25b8a57f 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -315,15 +315,29 @@ async def download_asset_content(request: web.Request) -> web.Response: 404, "FILE_NOT_FOUND", "Underlying file not found on disk." ) - # User-controlled asset content must never render inline in the app origin + # User-controlled asset content must not render inline in the app origin # (stored XSS via SVG/HTML/XML). Force dangerous types to download and - # override any requested inline disposition. Centralised through - # folder_paths.is_dangerous_content_type so this can't drift from /view and - # /userdata (the previous inline set here omitted image/svg+xml and missed - # the charset/casing/+xml-dialect bypasses). + # override any requested inline disposition; SVG loaded into an is + # exempt, see renders_safely_as_image. Centralised through folder_paths so + # this can't drift from /view and /userdata (the previous inline set here + # omitted image/svg+xml and missed the charset/casing/+xml-dialect bypasses). + extra_headers = {} + sec_fetch_dest = request.headers.get("Sec-Fetch-Dest") if folder_paths.is_dangerous_content_type(content_type): - content_type = "application/octet-stream" - disposition = "attachment" + # This response now depends on a request header, so it must not be + # reused across destinations by a browser or intermediary cache: an + # inline SVG primed by an fetch and replayed to a document + # navigation of the same URL would re-enable the stored XSS. + extra_headers["Vary"] = "Sec-Fetch-Dest" + extra_headers["Cache-Control"] = "no-store" + if not folder_paths.renders_safely_as_image(content_type, sec_fetch_dest): + content_type = "application/octet-stream" + disposition = "attachment" + + # mime_type is uploader-supplied and unvalidated, so it can carry + # parameters. aiohttp rejects a charset in the content_type argument with + # ValueError, which would turn a valid inline SVG into a 500. + content_type = content_type.split(";", 1)[0].strip() or "application/octet-stream" safe_name = (filename or "").replace("\r", "").replace("\n", "") encoded = urllib.parse.quote(safe_name) @@ -356,6 +370,7 @@ async def download_asset_content(request: web.Request) -> web.Response: "Content-Disposition": cd, "Content-Length": str(file_size), "X-Content-Type-Options": "nosniff", + **extra_headers, }, ) diff --git a/folder_paths.py b/folder_paths.py index bd3f25095..df53542dc 100644 --- a/folder_paths.py +++ b/folder_paths.py @@ -306,6 +306,23 @@ def is_dangerous_content_type(content_type: str | None) -> bool: return normalized.endswith('+xml') or normalized.endswith('/xml') +def renders_safely_as_image(content_type: str | None, sec_fetch_dest: str | None) -> bool: + """Return True if a dangerous `content_type` is safe to serve inline anyway. + + An SVG referenced by an ```` is loaded in secure static mode: scripts + and external references are disabled, so the stored XSS that + ``is_dangerous_content_type`` guards against cannot fire. The attack needs + the SVG to become a document, which is a separate ``Sec-Fetch-Dest``. + Browsers set that header themselves and script cannot override it (the + ``Sec-`` prefix makes it a forbidden header name), so it is trustworthy for + this decision. Anything else, including a missing header from a non-browser + client or a proxy that strips it, fails closed. + """ + if sec_fetch_dest != 'image': + return False + return (content_type or '').split(';', 1)[0].strip().lower() == 'image/svg+xml' + + def is_within_directory(directory: str, target: str) -> bool: """Return True if `target` resolves to a path inside `directory`. diff --git a/server.py b/server.py index e28fe2d22..c9ffcaa0d 100644 --- a/server.py +++ b/server.py @@ -624,8 +624,9 @@ class PromptServer(): # For security, force renderable/active types (HTML, JS, # CSS, SVG, XML — anything that can carry inline ' +ASSET_ID = "00000000-0000-4000-8000-000000000001" +CONTENT_URL = f"/api/assets/{ASSET_ID}/content" + + +class _StubUserManager: + def get_request_user_id(self, request): + return "test-user" + + +@pytest.fixture +def asset_app(monkeypatch, tmp_path): + """Mount the real /api/assets/{id}/content route over a stored SVG.""" + + def _factory(stored_mime_type): + svg = tmp_path / "thumb.svg" + svg.write_bytes(SVG_PAYLOAD) + + monkeypatch.setattr(asset_routes, "_ASSETS_ENABLED", True) + monkeypatch.setattr(asset_routes, "USER_MANAGER", _StubUserManager()) + monkeypatch.setattr( + asset_routes, + "resolve_asset_for_download", + lambda reference_id, owner_id: DownloadResolutionResult( + abs_path=str(svg), + content_type=stored_mime_type, + download_name="thumb.svg", + ), + ) + + app = web.Application() + app.add_routes(asset_routes.ROUTES) + return app + + return _factory + + +@pytest.mark.asyncio +async def test_inline_svg_response_is_not_cacheable_across_destinations( + aiohttp_client, asset_app +): + client = await aiohttp_client(asset_app("image/svg+xml")) + resp = await client.get( + CONTENT_URL, params={"disposition": "inline"}, headers={"Sec-Fetch-Dest": "image"} + ) + + assert resp.status == 200 + assert "image/svg+xml" in resp.headers.get("Content-Type", "").lower() + # The load-bearing assertion: a cache must not be able to hand this inline + # SVG to a later document navigation of the same URL. + assert "sec-fetch-dest" in resp.headers.get("Vary", "").lower(), ( + "The response varies on Sec-Fetch-Dest but does not say so, so a cache " + "keyed on the URL alone can replay the inline SVG into document context." + ) + assert "no-store" in resp.headers.get("Cache-Control", "").lower() + + +@pytest.mark.asyncio +async def test_forced_download_response_also_declares_the_variance( + aiohttp_client, asset_app +): + # The attachment branch needs the same headers, in both directions: a + # cached octet-stream replayed to an re-breaks the preview this fix + # exists to restore. + client = await aiohttp_client(asset_app("image/svg+xml")) + resp = await client.get( + CONTENT_URL, + params={"disposition": "inline"}, + headers={"Sec-Fetch-Dest": "document"}, + ) + + assert resp.status == 200 + assert "application/octet-stream" in resp.headers.get("Content-Type", "").lower() + assert "attachment" in resp.headers.get("Content-Disposition", "").lower() + assert "sec-fetch-dest" in resp.headers.get("Vary", "").lower() + assert "no-store" in resp.headers.get("Cache-Control", "").lower() + + +@pytest.mark.asyncio +async def test_parameterised_svg_mime_type_does_not_500(aiohttp_client, asset_app): + # mime_type is uploader-supplied and unvalidated. aiohttp rejects a charset + # in the content_type argument with ValueError, so the exempt branch must + # strip parameters before building the response. + client = await aiohttp_client(asset_app("image/svg+xml; charset=utf-8")) + resp = await client.get( + CONTENT_URL, params={"disposition": "inline"}, headers={"Sec-Fetch-Dest": "image"} + ) + + assert resp.status == 200, ( + "A charset parameter on the stored mime type must not turn a valid " + "inline SVG request into a 500." + ) + assert "image/svg+xml" in resp.headers.get("Content-Type", "").lower() From 9b3aa0896737845c8c42866f85c837523a97e51a Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 30 Jul 2026 13:15:39 -0700 Subject: [PATCH 177/211] Don't enable comfyui_detail.log by default. (#15159) --- app/logger.py | 2 +- main.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/app/logger.py b/app/logger.py index 1aed54e37..2e6116813 100644 --- a/app/logger.py +++ b/app/logger.py @@ -113,7 +113,7 @@ def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300 logger = logging.getLogger() console_level = get_log_level(log_level) file_levels = [get_log_level(level) for level, _ in file_outputs] - logger.setLevel(min(console_level, *file_levels)) + logger.setLevel(min([console_level, *file_levels])) formatter = ColoredFormatter("%(message)s") diff --git a/main.py b/main.py index c33e75f62..9c318fafe 100644 --- a/main.py +++ b/main.py @@ -19,7 +19,7 @@ import time from comfy.cli_args import enables_dynamic_vram from app.logger import setup_logger console_log_level = get_console_log_level(args.verbose) -file_log_outputs = [('DETAIL', 'comfyui_detail.log'), *get_file_log_outputs(args.verbose)] +file_log_outputs = get_file_log_outputs(args.verbose) setup_logger(log_level=console_log_level, file_outputs=file_log_outputs, use_stdout=args.log_stdout) from app.assets.seeder import asset_seeder From 8352644aa92db0a05b4457050ec90f739c2c8c27 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Thu, 30 Jul 2026 16:53:13 -0700 Subject: [PATCH 178/211] comfy-kitchen AMD support. (#15160) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index ef30cd7fa..37fbc2667 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.24 +comfy-kitchen==0.2.25 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From c20269c1cbb99744c974de0771aaae4d24d31d9f Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Fri, 31 Jul 2026 09:05:51 +0900 Subject: [PATCH 179/211] Bump comfyui-frontend-package to 1.47.11 (#15161) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 37fbc2667..179261c4d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -comfyui-frontend-package==1.47.10 +comfyui-frontend-package==1.47.11 comfyui-workflow-templates==0.11.19 comfyui-embedded-docs==0.5.9 torch From 51d040554010eab8994612b5cc0f359ae76698aa Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 31 Jul 2026 05:18:03 +0300 Subject: [PATCH 180/211] [Partner Nodes] chore(Bria): adjust pricing for the Bria Video Remove Background nodes (#15155) --- comfy_api_nodes/nodes_bria.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/comfy_api_nodes/nodes_bria.py b/comfy_api_nodes/nodes_bria.py index 090154afb..8077f1398 100644 --- a/comfy_api_nodes/nodes_bria.py +++ b/comfy_api_nodes/nodes_bria.py @@ -289,7 +289,7 @@ class BriaRemoveVideoBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.0042,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", ), ) @@ -533,7 +533,7 @@ class BriaTransparentVideoBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.0042,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", ), ) From f65f45511bc40871eca04c14a7d36a6f94f7edcb Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 31 Jul 2026 05:22:34 +0300 Subject: [PATCH 181/211] [Partner Nodes] feat(Ideogram): add new P-Image model support (#15154) --- comfy_api_nodes/apis/ideogram.py | 19 ++++ comfy_api_nodes/nodes_ideogram.py | 144 ++++++++++++++++++++++++++++++ 2 files changed, 163 insertions(+) diff --git a/comfy_api_nodes/apis/ideogram.py b/comfy_api_nodes/apis/ideogram.py index ee3256e96..12c9b2fc3 100644 --- a/comfy_api_nodes/apis/ideogram.py +++ b/comfy_api_nodes/apis/ideogram.py @@ -231,6 +231,25 @@ class IdeogramV3Request(BaseModel): ) +class IdeogramPImageRequest(BaseModel): + prompt: str = Field( + ..., + description="The text prompt, or an Ideogram 4.0 structured JSON caption " + "(used verbatim when prompt_upsampling is 'OFF').", + ) + quality: str | None = Field( + None, description="Generation tier: 'VERY_LOW', 'LOW', 'MEDIUM' or 'HIGH'." + ) + resolution: str | None = Field(None, description="Output size class: '1K' or '2K'.") + aspect_ratio: str | None = Field( + None, description="Aspect ratio in WxH format", examples=['16x9'] + ) + prompt_upsampling: str | None = Field( + None, description="Prompt expansion: 'AUTO', 'ON' or 'OFF'." + ) + seed: int | None = Field(None, ge=0, le=2147483647) + + class IdeogramV4Request(BaseModel): text_prompt: str | None = Field( None, diff --git a/comfy_api_nodes/nodes_ideogram.py b/comfy_api_nodes/nodes_ideogram.py index cc0467987..252617b2c 100644 --- a/comfy_api_nodes/nodes_ideogram.py +++ b/comfy_api_nodes/nodes_ideogram.py @@ -6,6 +6,7 @@ import numpy as np import torch from comfy_api_nodes.apis.ideogram import ( IdeogramGenerateResponse, + IdeogramPImageRequest, IdeogramV3Request, IdeogramV3EditRequest, IdeogramV4Request, @@ -524,12 +525,155 @@ class IdeogramV4(IO.ComfyNode): return IO.NodeOutput(await download_and_process_images(image_urls)) +class IdeogramPImage(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="IdeogramPImage", + display_name="Ideogram P-Image", + category="partner/image/Ideogram", + description="Generates images using P-Image, Ideogram's fast text-to-image model. " + "Strong typography and photorealism; " + "supports Ideogram 4.0 structured JSON captions for exact text, " + "colors and layout.", + inputs=[ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Text prompt. Also accepts an Ideogram 4.0 structured JSON caption " + "(exact colors as #RRGGBB hexes, exact text strings, bounding-box " + "layout) — set prompt_upsampling to OFF to use it verbatim.", + ), + IO.Combo.Input( + "quality", + options=["VERY_LOW", "LOW", "MEDIUM", "HIGH"], + default="MEDIUM", + tooltip="Speed/price/quality tier. MEDIUM is the everyday default; HIGH for " + "complex prompts, fine detail and difficult text; VERY_LOW/LOW for " + "drafts at scale. Difficult text renders poorly below MEDIUM.", + ), + IO.Combo.Input( + "resolution", + options=["1K", "2K"], + default="1K", + tooltip="Output size class (exact pixels follow the aspect ratio, e.g. " + "16:9 gives 1280x720 at 1K and 2560x1440 at 2K). " + "Prefer HIGH + 2K for crisp typography.", + ), + IO.Combo.Input( + "aspect_ratio", + options=list(V3_RATIO_MAP.keys()), + default="1:1", + tooltip="The aspect ratio for image generation.", + ), + IO.Combo.Input( + "prompt_upsampling", + options=["AUTO", "ON", "OFF"], + default="AUTO", + tooltip="Expands short prompts into a detailed structured caption before " + "generation (the rewritten prompt is returned as final_prompt). " + "Set OFF when supplying your own JSON caption or exact wording.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + control_after_generate=True, + display_mode=IO.NumberDisplay.number, + optional=True, + tooltip="Seed for reproducible generation. With prompt_upsampling OFF, " + "the same seed and settings return the same image; with ON/AUTO " + "the prompt rewrite varies per run — reproduce a result by reusing " + "its final_prompt output with prompt_upsampling OFF and the same " + "seed.", + ), + ], + outputs=[ + IO.Image.Output(), + IO.String.Output( + "final_prompt", + tooltip="The prompt the image was actually generated from (the rewritten " + "structured caption when prompt_upsampling ran, else your prompt). " + "Feed it back with prompt_upsampling OFF and the same seed to " + "reproduce this image.", + ), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["quality", "resolution"]), + expr=""" + ( + $q := widgets.quality; + $is2k := $contains(widgets.resolution, "2k"); + $usd := + $contains($q, "very_low") ? ($is2k ? 0.00858 : 0.00429) : + $contains($q, "high") ? ($is2k ? 0.0429 : 0.02145) : + $contains($q, "medium") ? ($is2k ? 0.0286 : 0.0143) : + ($is2k ? 0.02145 : 0.010725); + {"type": "usd", "usd": $usd} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + prompt: str, + quality: str = "MEDIUM", + resolution: str = "1K", + aspect_ratio: str = "1:1", + prompt_upsampling: str = "AUTO", + seed: int = 42, + ): + validate_string(prompt, strip_whitespace=True, min_length=1) + request = IdeogramPImageRequest( + prompt=prompt, + quality=quality, + resolution=resolution, + aspect_ratio=V3_RATIO_MAP[aspect_ratio], + prompt_upsampling=prompt_upsampling, + seed=seed, + ) + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/ideogram/text-to-image/p-image-ideogram", method="POST"), + response_model=IdeogramGenerateResponse, + data=request, + max_retries=1, + ) + if not response.data: + raise Exception("No images were generated in the response") + image_urls = [image_data.url for image_data in response.data if image_data.url] + if not image_urls: + if any(image_data.is_image_safe is False for image_data in response.data): + raise Exception( + "The generation was blocked by Ideogram's content safety filter. " + "Adjust the prompt and try again." + ) + raise Exception("No image URLs were generated in the response") + return IO.NodeOutput( + await download_and_process_images(image_urls), + response.data[0].prompt or prompt, + ) + + class IdeogramExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: return [ IdeogramV3, IdeogramV4, + IdeogramPImage, ] From b6fe23b4f1a7c5c56bfc3e15e2b9934831ba10a9 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Fri, 31 Jul 2026 10:27:06 +0800 Subject: [PATCH 182/211] chore: update workflow templates to v0.11.20 (#15165) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 179261c4d..35dd5e613 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.47.11 -comfyui-workflow-templates==0.11.19 +comfyui-workflow-templates==0.11.20 comfyui-embedded-docs==0.5.9 torch torchsde From 7dd46274601239644fff19b1b069cff199fcf738 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Fri, 31 Jul 2026 10:30:32 +0800 Subject: [PATCH 183/211] Add minimax h3 support (#15167) --- comfy_api_nodes/apis/minimax.py | 78 +++++ comfy_api_nodes/nodes_minimax.py | 506 ++++++++++++++++++++++++++++++- 2 files changed, 582 insertions(+), 2 deletions(-) diff --git a/comfy_api_nodes/apis/minimax.py b/comfy_api_nodes/apis/minimax.py index d747e177a..bac4572d4 100644 --- a/comfy_api_nodes/apis/minimax.py +++ b/comfy_api_nodes/apis/minimax.py @@ -118,3 +118,81 @@ class MinimaxVideoGenerationResponse(BaseModel): task_id: str = Field( ..., description='The task ID for the asynchronous video generation task.' ) + + +class Hailuo03TextContent(BaseModel): + type: str = Field("text") + text: str = Field(...) + + +class Hailuo03ImageContentUrl(BaseModel): + url: str = Field(...) + + +class Hailuo03ImageContent(BaseModel): + type: str = Field("image_url") + image_url: Hailuo03ImageContentUrl = Field(...) + role: str = Field(...) + + +class Hailuo03VideoContentUrl(BaseModel): + url: str = Field(...) + + +class Hailuo03VideoContent(BaseModel): + type: str = Field("video_url") + video_url: Hailuo03VideoContentUrl = Field(...) + role: str = Field("reference_video") + + +class Hailuo03AudioContentUrl(BaseModel): + url: str = Field(...) + + +class Hailuo03AudioContent(BaseModel): + type: str = Field("audio_url") + audio_url: Hailuo03AudioContentUrl = Field(...) + role: str = Field("reference_audio") + + +class Hailuo03TaskCreationRequest(BaseModel): + model: str = Field(...) + content: list[Hailuo03TextContent | Hailuo03ImageContent | Hailuo03VideoContent | Hailuo03AudioContent] = Field( + ..., min_length=1 + ) + resolution: str = Field(...) + duration: int = Field(..., ge=5, le=15) + ratio: str | None = Field(None) + seed: int | None = Field(None, ge=0, le=4294967295) + aigc_watermark: bool | None = Field(None) + + +class Hailuo03TaskCreationResponse(BaseModel): + task_id: str = Field(...) + + +class Hailuo03TaskError(BaseModel): + code: int | str | None = Field(None) + message: str | None = Field(None) + + +class Hailuo03TaskContent(BaseModel): + url: str | None = Field(None) + + +class Hailuo03TaskUsage(BaseModel): + total_seconds: float = Field(0) + input_seconds: float = Field(0) + output_seconds: float = Field(0) + + +class Hailuo03Task(BaseModel): + id: str = Field(...) + status: str = Field(...) + error: Hailuo03TaskError | None = Field(None) + content: Hailuo03TaskContent | None = Field(None) + usage: Hailuo03TaskUsage | None = Field(None) + + +class Hailuo03TaskQueryResponse(BaseModel): + task: Hailuo03Task = Field(...) diff --git a/comfy_api_nodes/nodes_minimax.py b/comfy_api_nodes/nodes_minimax.py index 6250af146..2d7aef654 100644 --- a/comfy_api_nodes/nodes_minimax.py +++ b/comfy_api_nodes/nodes_minimax.py @@ -5,6 +5,16 @@ from typing_extensions import override from comfy_api.latest import IO, ComfyExtension from comfy_api_nodes.apis.minimax import ( + Hailuo03AudioContent, + Hailuo03AudioContentUrl, + Hailuo03ImageContent, + Hailuo03ImageContentUrl, + Hailuo03TaskCreationRequest, + Hailuo03TaskCreationResponse, + Hailuo03TaskQueryResponse, + Hailuo03TextContent, + Hailuo03VideoContent, + Hailuo03VideoContentUrl, MinimaxFileRetrieveResponse, MiniMaxModel, MinimaxTaskResultResponse, @@ -17,7 +27,11 @@ from comfy_api_nodes.util import ( download_url_to_video_output, poll_op, sync_op, + upload_audio_to_comfyapi, upload_images_to_comfyapi, + upload_video_to_comfyapi, + validate_image_aspect_ratio, + validate_image_dimensions, validate_string, ) @@ -293,9 +307,9 @@ class MinimaxHailuoVideoNode(IO.ComfyNode): def define_schema(cls) -> IO.Schema: return IO.Schema( node_id="MinimaxHailuoVideoNode", - display_name="MiniMax Hailuo Video", + display_name="MiniMax Hailuo 02 Video", category="partner/video/MiniMax", - description="Generates videos from prompt, with optional start frame using the new MiniMax Hailuo-02 model.", + description="Generates videos from prompt, with optional start frame using the MiniMax Hailuo-02 model.", inputs=[ IO.String.Input( "prompt_text", @@ -437,6 +451,491 @@ class MinimaxHailuoVideoNode(IO.ComfyNode): return IO.NodeOutput(await download_url_to_video_output(file_url)) +HAILUO_03_CREATE_ENDPOINT = "/proxy/minimax/v2/video_generation" +HAILUO_03_QUERY_ENDPOINT = "/proxy/minimax/v2/query/video_generation" # + /{task_id} +HAILUO_03_MODELS = {"MiniMax H3": "MiniMax-H3"} +HAILUO_03_FAILED_STATUSES = ["failed", "cancelled", "expired"] + + +def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = True): + inputs = [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Text prompt for video generation.", + ), + IO.Combo.Input( + "resolution", + options=["2K"], + tooltip="Resolution of the output video.", + ), + ] + if include_ratio: + ratio_options = ["16:9", "4:3", "1:1", "3:4", "9:16", "21:9"] + if allow_adaptive: + ratio_options.insert(0, "adaptive") + inputs.append( + IO.Combo.Input( + "ratio", + options=ratio_options, + default=ratio_options[0], + tooltip="Aspect ratio of the output video.", + ) + ) + inputs.append( + IO.Int.Input( + "duration", + default=5, + min=5, + max=15, + step=1, + tooltip="Duration of the output video in seconds (5-15).", + display_mode=IO.NumberDisplay.slider, + ) + ) + return inputs + + +async def _hailuo03_run_task( + cls: type[IO.ComfyNode], + *, + model_id: str, + content: list, + resolution: str, + duration: int, + ratio: str | None, + seed: int, + watermark: bool, +) -> IO.NodeOutput: + response = await sync_op( + cls, + ApiEndpoint(path=HAILUO_03_CREATE_ENDPOINT, method="POST"), + response_model=Hailuo03TaskCreationResponse, + data=Hailuo03TaskCreationRequest( + model=model_id, + content=content, + resolution=resolution, + duration=duration, + ratio=ratio, + seed=seed, + aigc_watermark=watermark, + ), + ) + task_result = await poll_op( + cls, + ApiEndpoint(path=f"{HAILUO_03_QUERY_ENDPOINT}/{response.task_id}"), + response_model=Hailuo03TaskQueryResponse, + status_extractor=lambda r: r.task.status, + failed_statuses=HAILUO_03_FAILED_STATUSES, + poll_interval=15, + ) + video_url = task_result.task.content.url if task_result.task.content else None + if not video_url: + raise Exception(f"No video URL in the response: {task_result.model_dump()}") + return IO.NodeOutput(await download_url_to_video_output(video_url)) + + +class MinimaxHailuo03TextToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03TextToVideoNode", + display_name="MiniMax H3 Text to Video", + category="partner/video/MiniMax", + description="Generate video from a text prompt using the MiniMax H3 model.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[IO.DynamicCombo.Option("MiniMax H3", _hailuo03_model_inputs(allow_adaptive=False))], + tooltip="Model to use for video generation.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=4294967295, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Random seed. The same request with the same seed gives similar, " + "but not guaranteed identical, results.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AIGC watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model.duration"]), + expr=""" + ( + $dur := $lookup(widgets, "model.duration"); + {"type": "usd", "usd": $dur * 0.1859} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + seed: int, + watermark: bool, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + return await _hailuo03_run_task( + cls, + model_id=HAILUO_03_MODELS[model["model"]], + content=[Hailuo03TextContent(text=model["prompt"])], + resolution=model["resolution"], + duration=model["duration"], + ratio=model["ratio"], + seed=seed, + watermark=watermark, + ) + + +class MinimaxHailuo03FirstLastFrameNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03FirstLastFrameNode", + display_name="MiniMax H3 First-Last-Frame to Video", + category="partner/video/MiniMax", + description="Generate video from a first frame image and an optional last frame image " + "using the MiniMax H3 model. The aspect ratio of the video follows the supplied images.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[IO.DynamicCombo.Option("MiniMax H3", _hailuo03_model_inputs(include_ratio=False))], + tooltip="Model to use for video generation.", + ), + IO.Image.Input( + "first_frame", + tooltip="First frame image for the video.", + ), + IO.Image.Input( + "last_frame", + tooltip="Optional last frame image for the video.", + optional=True, + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=4294967295, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Random seed. The same request with the same seed gives similar, " + "but not guaranteed identical, results.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AIGC watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model.duration"]), + expr=""" + ( + $dur := $lookup(widgets, "model.duration"); + {"type": "usd", "usd": $dur * 0.1859} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + first_frame: torch.Tensor, + seed: int, + watermark: bool, + last_frame: torch.Tensor | None = None, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + for frame in (first_frame, last_frame): + if frame is not None: + validate_image_aspect_ratio(frame, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(frame, min_width=256, min_height=256) + + content: list = [ + Hailuo03TextContent(text=model["prompt"]), + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, first_frame, max_images=1, wait_label="Uploading first frame" + ) + )[0], + ), + role="first_frame", + ), + ] + if last_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, last_frame, max_images=1, wait_label="Uploading last frame" + ) + )[0], + ), + role="last_frame", + ) + ) + return await _hailuo03_run_task( + cls, + model_id=HAILUO_03_MODELS[model["model"]], + content=content, + resolution=model["resolution"], + duration=model["duration"], + ratio=None, + seed=seed, + watermark=watermark, + ) + + +class MinimaxHailuo03ReferenceNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03ReferenceNode", + display_name="MiniMax H3 Reference to Video", + category="partner/video/MiniMax", + description="Generate video conditioned on reference images, videos, and audio using the " + "MiniMax H3 model. Refer to the references in the prompt by their order: " + "'Image 1', 'Image 2', 'Video 1', 'Audio 1', and so on.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "MiniMax H3", + [ + *_hailuo03_model_inputs(), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image_1", + "image_2", + "image_3", + "image_4", + "image_5", + "image_6", + "image_7", + "image_8", + "image_9", + ], + min=0, + ), + tooltip="Subject or style reference images, referred to in the prompt " + "as 'Image 1'..'Image 9' in connection order. Up to 9 images.", + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=["video_1", "video_2", "video_3"], + min=0, + ), + tooltip="Motion or scene reference videos, referred to in the prompt " + "as 'Video 1'..'Video 3' in connection order. Up to 3 videos, " + "2-15 seconds each, 15 seconds in total.", + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=["audio_1", "audio_2", "audio_3"], + min=0, + ), + tooltip="Audio references, referred to in the prompt as " + "'Audio 1'..'Audio 3' in connection order. Up to 3 clips, " + "2-15 seconds each, 15 seconds in total. Cannot be used without " + "a reference image or video.", + ), + ], + ) + ], + tooltip="Model to use for video generation.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=4294967295, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Random seed. The same request with the same seed gives similar, " + "but not guaranteed identical, results.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AIGC watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + widgets=["model.duration"], + input_groups=["model.reference_images", "model.reference_videos"], + ), + expr=""" + ( + $dur := $lookup(widgets, "model.duration"); + $imgsRaw := $lookup(inputGroups, "model.reference_images"); + $imgs := $imgsRaw ? $imgsRaw : 0; + $vidsRaw := $lookup(inputGroups, "model.reference_videos"); + $vids := $vidsRaw ? $vidsRaw : 0; + $base := $dur * 0.1859 + ($imgs > 5 ? ($imgs - 5) * 0.0572 : 0); + $vids > 0 + ? {"type": "range_usd", "min_usd": $base + $vids * 2 * 0.1859, + "max_usd": $base + 15 * 0.1859, "format": {"approximate": true}} + : {"type": "usd", "usd": $base} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + seed: int, + watermark: bool, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + + reference_images = model.get("reference_images", {}) + reference_videos = model.get("reference_videos", {}) + reference_audios = model.get("reference_audios", {}) + if not reference_images and not reference_videos: + raise ValueError("At least one reference image or video is required.") + + for image in reference_images.values(): + validate_image_aspect_ratio(image, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(image, min_width=256, min_height=256) + + total_video_duration = 0.0 + for i, video in enumerate(reference_videos.values(), 1): + try: + fps = float(video.get_frame_rate()) + except Exception: + fps = 0.0 + if fps and not (23.9 <= fps <= 60.5): + raise ValueError(f"Reference video {i} is {fps:.2f} FPS. Supported range is 23.976-60 FPS.") + try: + dur = video.get_duration() + except Exception: + continue + if dur < 1.8: + raise ValueError(f"Reference video {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_video_duration += dur + if total_video_duration > 15.1: + raise ValueError( + f"Total reference video duration is {total_video_duration:.1f}s. Maximum is 15 seconds." + ) + + total_audio_duration = 0.0 + for i, audio in enumerate(reference_audios.values(), 1): + dur = int(audio["waveform"].shape[-1]) / int(audio["sample_rate"]) + if dur < 1.8: + raise ValueError(f"Reference audio {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_audio_duration += dur + if total_audio_duration > 15.1: + raise ValueError( + f"Total reference audio duration is {total_audio_duration:.1f}s. Maximum is 15 seconds." + ) + + content: list = [Hailuo03TextContent(text=model["prompt"])] + for i, image in enumerate(reference_images.values(), 1): + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, image, max_images=1, wait_label=f"Uploading image {i}" + ) + )[0], + ), + role="reference_image", + ) + ) + for i, video in enumerate(reference_videos.values(), 1): + content.append( + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, video, wait_label=f"Uploading video {i}"), + ), + ) + ) + for audio in reference_audios.values(): + content.append( + Hailuo03AudioContent( + audio_url=Hailuo03AudioContentUrl( + url=await upload_audio_to_comfyapi( + cls, + audio, + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ), + ) + ) + return await _hailuo03_run_task( + cls, + model_id=HAILUO_03_MODELS[model["model"]], + content=content, + resolution=model["resolution"], + duration=model["duration"], + ratio=model["ratio"], + seed=seed, + watermark=watermark, + ) + + class MinimaxExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: @@ -445,6 +944,9 @@ class MinimaxExtension(ComfyExtension): MinimaxImageToVideoNode, # MinimaxSubjectToVideoNode, MinimaxHailuoVideoNode, + MinimaxHailuo03TextToVideoNode, + MinimaxHailuo03FirstLastFrameNode, + MinimaxHailuo03ReferenceNode, ] From 5cc026f5b81b3f01fe7a1438a0fd4131d2ebda25 Mon Sep 17 00:00:00 2001 From: Terry Jia Date: Fri, 31 Jul 2026 03:07:31 -0400 Subject: [PATCH 184/211] fix: resend cached histogram UI for CurveEditor after page refresh (#15152) --- comfy_extras/nodes_curve.py | 1 + 1 file changed, 1 insertion(+) diff --git a/comfy_extras/nodes_curve.py b/comfy_extras/nodes_curve.py index aa2d94bb6..078d0114f 100644 --- a/comfy_extras/nodes_curve.py +++ b/comfy_extras/nodes_curve.py @@ -12,6 +12,7 @@ class CurveEditor(io.ComfyNode): node_id="CurveEditor", display_name="Curve Editor", category="utilities", + has_intermediate_output=True, inputs=[ io.Curve.Input("curve"), io.Histogram.Input("histogram", optional=True), From 831710d25740257ce84691f1367e5e0eae9c3e33 Mon Sep 17 00:00:00 2001 From: blepping <157360029+blepping@users.noreply.github.com> Date: Fri, 31 Jul 2026 10:12:17 -0600 Subject: [PATCH 185/211] Don't assume sampler_function has a __name__ attribute in detail logging (#15179) --- comfy/samplers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/comfy/samplers.py b/comfy/samplers.py index 9f571ece9..29e6bffb3 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -997,7 +997,7 @@ class KSAMPLER(Sampler): def k_callback(x): nonlocal first_step if first_step: - detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, self.sampler_function.__name__, x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape)) + detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, getattr(self.sampler_function, "__name__", "unknown"), x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape)) first_step = False if callback is not None: callback(x["i"], x["denoised"], x["x"], total_steps) From 9bd9815eabaf02feb6adb07ec3b5e8123fb14cd6 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 31 Jul 2026 19:59:53 +0300 Subject: [PATCH 186/211] [Partner Nodes] chore(Bria): adjust pricing for other video endpoints (#15171) Signed-off-by: bigcat88 --- comfy_api_nodes/nodes_bria.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/comfy_api_nodes/nodes_bria.py b/comfy_api_nodes/nodes_bria.py index 8077f1398..9e5a93330 100644 --- a/comfy_api_nodes/nodes_bria.py +++ b/comfy_api_nodes/nodes_bria.py @@ -357,7 +357,7 @@ class BriaVideoGreenScreen(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.0042,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", ), ) @@ -433,7 +433,7 @@ class BriaVideoReplaceBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.0042,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", ), ) From de5625a64aba29080dc8cbea3a85167da9b8e77f Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Sat, 1 Aug 2026 03:29:47 +1000 Subject: [PATCH 187/211] Delay dynamic pin cleanup until model destruction (#15183) --- comfy/model_patcher.py | 20 ++++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index e44322e72..6b698f767 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -558,12 +558,9 @@ class ModelPatcher: new_multigpu_models = [] for mm in multigpu_models: # clone main model, but bring over relevant props from existing multigpu clone - n = self.clone() + n = self.clone(model_override=mm.get_clone_model_override()) n.load_device = mm.load_device - n.backup = mm.backup - n.object_patches_backup = mm.object_patches_backup n.hook_backup = mm.hook_backup - n.model = mm.model n.is_multigpu_base_clone = mm.is_multigpu_base_clone n.remove_additional_models("multigpu") orig_additional_models: dict[str, list[ModelPatcher]] = comfy.patcher_extension.copy_nested_dicts(n.additional_models) @@ -1758,6 +1755,9 @@ class ModelPatcherDynamic(ModelPatcher): self.register_load_device(self.load_device) self.non_dynamic_delegate_model = None assert load_device is not None + if not hasattr(self.model, "dynamic_patchers"): + self.model.dynamic_patchers = set() + self.model.dynamic_patchers.add(id(self)) def register_load_device(self, device): """Ensure dynamic_pins has an entry for *device*. @@ -1813,6 +1813,18 @@ class ModelPatcherDynamic(ModelPatcher): def unpin_all_weights(self): self.partially_unload_ram(1e32) + def __del__(self): + model = getattr(self, "model", None) + dynamic_patchers = getattr(model, "dynamic_patchers", None) + if dynamic_patchers is None or id(self) not in dynamic_patchers: + return + dynamic_patchers.discard(id(self)) + try: + if not dynamic_patchers: + self.unpin_all_weights() + finally: + self.detach(unpatch_all=False) + def memory_required(self, input_shape): #Pad this significantly. We are trying to get away from precise estimates. This #estimate is only used when using the ModelPatcherDynamic after ModelPatcher. If you From 7c806288d5210b497cedfca0d111bb56d84a20de Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Sat, 1 Aug 2026 03:44:09 +1000 Subject: [PATCH 188/211] ops: apply the custom placeholder logic to Linux too (#15181) Windows has proven this logic works for a long time and there are corner cases where this materialization actual consumes real RAM on linux. Its not as bad as the original windows commit charge surge, but its still a detectable transient leak. So simplify and unify. --- comfy/ops.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/comfy/ops.py b/comfy/ops.py index 9d692dcc7..f4bdd0aef 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -464,8 +464,7 @@ class disable_weight_init: def __init__(self, in_features, out_features, bias=True, device=None, dtype=None): # don't trust subclasses that BYO state dict loader to call us. - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Linear._load_from_state_dict): super().__init__(in_features, out_features, bias, device, dtype) return @@ -487,8 +486,7 @@ class disable_weight_init: def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs): - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Linear._load_from_state_dict): return super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) @@ -716,8 +714,7 @@ class disable_weight_init: norm_type=2.0, scale_grad_by_freq=False, sparse=False, _weight=None, _freeze=False, device=None, dtype=None): # don't trust subclasses that BYO state dict loader to call us. - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Embedding._load_from_state_dict): super().__init__(num_embeddings, embedding_dim, padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse, _weight, @@ -744,8 +741,7 @@ class disable_weight_init: def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs): - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Embedding._load_from_state_dict): return super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) From 0b73ae83a1284b2e288aae656d88dede6dbb623b Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Sat, 1 Aug 2026 02:20:55 +0800 Subject: [PATCH 189/211] chore: update workflow templates to v0.11.23 (#15187) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 35dd5e613..af414d413 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ comfyui-frontend-package==1.47.11 -comfyui-workflow-templates==0.11.20 +comfyui-workflow-templates==0.11.23 comfyui-embedded-docs==0.5.9 torch torchsde From 6cedd34343ba3214ca9591397bd106e02ef2acf6 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Fri, 31 Jul 2026 21:57:45 +0300 Subject: [PATCH 190/211] [Partner Nodes] fix(ByteDance): encode stereo reference audio without doubling its duration (#15177) * [Partner Nodes] fix(ByteDance): encode stereo reference audio without doubling its duration Signed-off-by: bigcat88 --- comfy_api_nodes/util/conversions.py | 23 +++--- .../audio_conversions_test.py | 78 +++++++++++++++++++ 2 files changed, 89 insertions(+), 12 deletions(-) create mode 100644 tests-unit/comfy_api_nodes_test/audio_conversions_test.py diff --git a/comfy_api_nodes/util/conversions.py b/comfy_api_nodes/util/conversions.py index 9cd644fc0..f46cac3f8 100644 --- a/comfy_api_nodes/util/conversions.py +++ b/comfy_api_nodes/util/conversions.py @@ -266,16 +266,14 @@ def audio_tensor_to_contiguous_ndarray(waveform: torch.Tensor) -> np.ndarray: waveform: a tensor of shape (1, channels, samples) derived from a Comfy `AUDIO` type. Returns: - Contiguous numpy array of the audio waveform. If the audio was batched, - the first item is taken. + Contiguous numpy array of the audio waveform. + + Raises: + ValueError: If the waveform is not shaped (1, channels, samples). """ if waveform.ndim != 3 or waveform.shape[0] != 1: raise ValueError("Expected waveform tensor shape (1, channels, samples)") - # If batch is > 1, take first item - if waveform.shape[0] > 1: - waveform = waveform[0] - # Prepare for av: remove batch dim, move to CPU, make contiguous, convert to numpy array audio_data_np = waveform.squeeze(0).cpu().contiguous().numpy() if audio_data_np.dtype != np.float32: @@ -285,20 +283,21 @@ def audio_tensor_to_contiguous_ndarray(waveform: torch.Tensor) -> np.ndarray: def audio_input_to_mp3(audio: Input.Audio) -> BytesIO: - waveform = audio["waveform"].cpu() + audio_data_np = audio_tensor_to_contiguous_ndarray(audio["waveform"]) + sample_rate = int(audio["sample_rate"]) output_buffer = BytesIO() output_container = av.open(output_buffer, mode="w", format="mp3") - out_stream = output_container.add_stream("libmp3lame", rate=audio["sample_rate"]) + out_stream = output_container.add_stream("libmp3lame", rate=sample_rate) out_stream.bit_rate = 320000 frame = av.AudioFrame.from_ndarray( - waveform.movedim(0, 1).reshape(1, -1).float().numpy(), - format="flt", - layout="mono" if waveform.shape[0] == 1 else "stereo", + audio_data_np, + format="fltp", + layout="stereo" if audio_data_np.shape[0] > 1 else "mono", ) - frame.sample_rate = audio["sample_rate"] + frame.sample_rate = sample_rate frame.pts = 0 output_container.mux(out_stream.encode(frame)) output_container.mux(out_stream.encode(None)) diff --git a/tests-unit/comfy_api_nodes_test/audio_conversions_test.py b/tests-unit/comfy_api_nodes_test/audio_conversions_test.py new file mode 100644 index 000000000..d7c9b3899 --- /dev/null +++ b/tests-unit/comfy_api_nodes_test/audio_conversions_test.py @@ -0,0 +1,78 @@ +import math + +import av +import numpy as np +import pytest +import torch + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +from comfy_api_nodes.util.conversions import audio_input_to_mp3 # noqa: E402 + +SAMPLE_RATE = 48000 +DURATION = 2.0 +LEFT_HZ = 440.0 +RIGHT_HZ = 880.0 + + +def tone(freq, duration=DURATION, sample_rate=SAMPLE_RATE): + t = torch.arange(int(sample_rate * duration), dtype=torch.float32) / sample_rate + return 0.5 * torch.sin(2 * math.pi * freq * t) + + +@pytest.fixture +def stereo_audio(): + """Comfy AUDIO with two tones that stay distinguishable through mp3.""" + waveform = torch.stack([tone(LEFT_HZ), tone(RIGHT_HZ)]).unsqueeze(0) + return {"waveform": waveform, "sample_rate": SAMPLE_RATE} + + +@pytest.fixture +def mono_audio(): + return {"waveform": tone(LEFT_HZ).unsqueeze(0).unsqueeze(0), "sample_rate": SAMPLE_RATE} + + +def decode(buffer): + """(planes[C, N], sample_rate, channels) of an encoded mp3 buffer""" + buffer.seek(0) + with av.open(buffer, mode="r") as container: + stream = container.streams.audio[0] + planes = [] + for frame in container.decode(audio=0): + array = frame.to_ndarray() + if frame.format.is_planar: + planes.append(array) + else: + planes.append(array.reshape(-1, len(frame.layout.channels)).T) + return np.concatenate(planes, axis=1), stream.codec_context.sample_rate, len(stream.layout.channels) + + +def dominant_hz(signal, sample_rate): + """Peak frequency, ignoring the encoder's padding at either edge""" + edge = int(0.2 * sample_rate) + window = signal[edge:-edge] + spectrum = np.abs(np.fft.rfft(window * np.hanning(window.size))) + return np.fft.rfftfreq(window.size, 1.0 / sample_rate)[int(np.argmax(spectrum))] + + +def test_stereo_duration_is_preserved(stereo_audio): + planes, sample_rate, channels = decode(audio_input_to_mp3(stereo_audio)) + assert channels == 2 + assert sample_rate == SAMPLE_RATE + assert planes.shape[1] / sample_rate == pytest.approx(DURATION, abs=0.15) + + +def test_stereo_channels_are_not_concatenated(stereo_audio): + """The channels must be interleaved; concatenating them plays the clip twice.""" + planes, sample_rate, _ = decode(audio_input_to_mp3(stereo_audio)) + assert dominant_hz(planes[0], sample_rate) == pytest.approx(LEFT_HZ, abs=15) + assert dominant_hz(planes[1], sample_rate) == pytest.approx(RIGHT_HZ, abs=15) + + +def test_mono_duration_is_preserved(mono_audio): + planes, sample_rate, _ = decode(audio_input_to_mp3(mono_audio)) + assert planes.shape[1] / sample_rate == pytest.approx(DURATION, abs=0.15) + assert dominant_hz(planes[0], sample_rate) == pytest.approx(LEFT_HZ, abs=15) From a1c421994cdcc5044dbce2bb7628e89386311cc5 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 31 Jul 2026 15:17:56 -0700 Subject: [PATCH 191/211] Expand k, v when attention backend would fall back to math because gqa. (#15190) --- comfy/ldm/modules/attention.py | 34 ++++++------------------- comfy/ops.py | 46 ++++++++++++++++++++++++++++------ 2 files changed, 46 insertions(+), 34 deletions(-) diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index e6500cff4..2c549e095 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -90,22 +90,6 @@ def default(val, d): return val return d -def _gqa_repeat_factor(query_heads, key_heads, value_heads): - if key_heads != value_heads: - raise ValueError(f"Key/value head count mismatch for GQA: {key_heads} != {value_heads}") - if query_heads == key_heads: - return 1 - if query_heads % key_heads != 0: - raise ValueError(f"Query heads must be divisible by key/value heads for GQA: {query_heads} vs {key_heads}") - return query_heads // key_heads - -def _repeat_kv_for_gqa(k, v, query_heads, head_dim): - n_rep = _gqa_repeat_factor(query_heads, k.shape[head_dim], v.shape[head_dim]) - if n_rep > 1: - k = k.repeat_interleave(n_rep, dim=head_dim) - v = v.repeat_interleave(n_rep, dim=head_dim) - return k, v - def _heads_from_dim(tensor, dim_head, name): inner_dim = tensor.shape[-1] if inner_dim % dim_head != 0: @@ -122,10 +106,8 @@ def _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa=False, expand_ value_heads = heads k = k.unsqueeze(3).reshape(b, -1, key_heads, dim_head) v = v.unsqueeze(3).reshape(b, -1, value_heads, dim_head) - if enable_gqa: - _gqa_repeat_factor(heads, key_heads, value_heads) - if expand_kv: - k, v = _repeat_kv_for_gqa(k, v, heads, -2) + if enable_gqa and expand_kv: + k, v = comfy.ops.repeat_kv_for_gqa(k, v, heads, -2) return q, k, v @@ -196,7 +178,7 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape h = heads if skip_reshape: if kwargs.get("enable_gqa", False): - k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), @@ -262,7 +244,7 @@ def attention_sub_quad(query, key, value, heads, mask=None, attn_precision=None, if skip_reshape: if kwargs.get("enable_gqa", False): - key, value = _repeat_kv_for_gqa(key, value, query.shape[-3], -3) + key, value = comfy.ops.repeat_kv_for_gqa(key, value, query.shape[-3], -3) query = query.reshape(b * heads, -1, dim_head) value = value.reshape(b * heads, -1, dim_head) key = key.reshape(b * heads, -1, dim_head).movedim(1, 2) @@ -338,7 +320,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape if skip_reshape: if kwargs.get("enable_gqa", False): - k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), @@ -476,7 +458,7 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh (q, k, v), ) if kwargs.get("enable_gqa", False): - k, v = _repeat_kv_for_gqa(k, v, q.shape[-2], -2) + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-2], -2) # actually do the reshaping else: dim_head //= heads @@ -573,7 +555,7 @@ def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape= b, _, _, dim_head = q.shape tensor_layout = "HND" if kwargs.get("enable_gqa", False): - k, v = _repeat_kv_for_gqa(k, v, q.shape[-3], -3) + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) else: b, _, dim_head = q.shape dim_head //= heads @@ -671,7 +653,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape if skip_reshape: q_s = q if kwargs.get("enable_gqa", False): - k_s, v_s = _repeat_kv_for_gqa(k, v, H, -3) + k_s, v_s = comfy.ops.repeat_kv_for_gqa(k, v, H, -3) else: k_s, v_s = k, v else: diff --git a/comfy/ops.py b/comfy/ops.py index f4bdd0aef..6c3845eef 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -19,6 +19,7 @@ import torch import logging import contextlib +import inspect import comfy.model_management from comfy.cli_args import args, PerformanceFeature import comfy.float @@ -36,30 +37,59 @@ def run_every_op(): comfy.model_management.throw_exception_if_processing_interrupted() +def gqa_repeat_factor(query_heads, key_heads, value_heads): + if key_heads != value_heads: + raise ValueError(f"Key/value head count mismatch for GQA: {key_heads} != {value_heads}") + if query_heads == key_heads: + return 1 + if query_heads % key_heads != 0: + raise ValueError(f"Query heads must be divisible by key/value heads for GQA: {query_heads} vs {key_heads}") + return query_heads // key_heads + +def repeat_kv_for_gqa(k, v, query_heads, head_dim): + n_rep = gqa_repeat_factor(query_heads, k.shape[head_dim], v.shape[head_dim]) + if n_rep > 1: + k = k.repeat_interleave(n_rep, dim=head_dim) + v = v.repeat_interleave(n_rep, dim=head_dim) + return k, v + def scaled_dot_product_attention(q, k, v, *args, **kwargs): + attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask") + if kwargs.get("enable_gqa", False) and attn_mask is not None: + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) try: if torch.cuda.is_available(): from torch.nn.attention import SDPBackend, sdpa_kernel - import inspect if "set_priority" in inspect.signature(sdpa_kernel).parameters: SDPA_BACKEND_PRIORITY = [ SDPBackend.FLASH_ATTENTION, + SDPBackend.CUDNN_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH, ] - if comfy.model_management.WINDOWS: - SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION) - else: - SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION) - def scaled_dot_product_attention(q, k, v, *args, **kwargs): - if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower - return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) + attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask") + if kwargs.get("enable_gqa", False) and attn_mask is not None and not comfy.model_management.is_nvidia(): + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False with sdpa_kernel(SDPA_BACKEND_PRIORITY, set_priority=True): + if kwargs.get("enable_gqa", False) and attn_mask is not None and q.shape[-3] != k.shape[-3]: + dropout_p = args[1] if len(args) > 1 else kwargs.get("dropout_p", 0.0) + is_causal = args[2] if len(args) > 2 else kwargs.get("is_causal", False) + params = torch.backends.cuda.SDPAParams(q, k, v, attn_mask, dropout_p, is_causal, True) + supports_native_gqa = ( + torch.backends.cuda.can_use_flash_attention(params) + or torch.backends.cuda.can_use_cudnn_attention(params) + or torch.backends.cuda.can_use_efficient_attention(params) + ) + if not supports_native_gqa: + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) else: logging.warning("Torch version too old to set sdpa backend priority.") From 235b466a0cb26d47c24f2ab66d1a8c5e70b21070 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 31 Jul 2026 21:27:48 -0700 Subject: [PATCH 192/211] Add crf option to save video node. (#15191) --- comfy_api/latest/_input/video_types.py | 2 ++ comfy_api/latest/_input_impl/video_types.py | 11 +++++- comfy_extras/nodes_video.py | 35 ++++++++++++++++--- tests-unit/comfy_api_test/video_types_test.py | 18 ++++++++++ 4 files changed, 61 insertions(+), 5 deletions(-) diff --git a/comfy_api/latest/_input/video_types.py b/comfy_api/latest/_input/video_types.py index e2e99521f..b700d44f5 100644 --- a/comfy_api/latest/_input/video_types.py +++ b/comfy_api/latest/_input/video_types.py @@ -29,11 +29,13 @@ class VideoInput(ABC): codec: VideoCodec = VideoCodec.AUTO, metadata: Optional[dict] = None, bit_depth: int | None = None, + crf: float | None = None, ): """ Abstract method to save the video input to a file. bit_depth selects the encoded bit depth; None keeps the video's native depth. + crf selects the H.264 constant rate factor; None uses the encoder default. """ pass diff --git a/comfy_api/latest/_input_impl/video_types.py b/comfy_api/latest/_input_impl/video_types.py index f5af41973..14d663881 100644 --- a/comfy_api/latest/_input_impl/video_types.py +++ b/comfy_api/latest/_input_impl/video_types.py @@ -460,6 +460,7 @@ class VideoFromFile(VideoInput): codec: VideoCodec = VideoCodec.AUTO, metadata: Optional[dict] = None, bit_depth: int | None = None, + crf: float | None = None, ): if isinstance(self.__file, io.BytesIO): self.__file.seek(0) # Reset the BytesIO object to the beginning @@ -475,13 +476,15 @@ class VideoFromFile(VideoInput): reuse_streams = False if bit_depth is not None and video_encoding is not None and bit_depth != source_bit_depth: reuse_streams = False + if crf is not None: + reuse_streams = False if self.__start_time or self.__duration: reuse_streams = False if not reuse_streams: if bit_depth is None: bit_depth = source_bit_depth - return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth) + return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth, crf=crf) streams = container.streams @@ -514,6 +517,7 @@ class VideoFromFile(VideoInput): codec: VideoCodec, metadata: dict | None, bit_depth: int, + crf: float | None = None, ): """Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length.""" open_kwargs = mp4_output_open_kwargs(path, format, codec) @@ -659,6 +663,8 @@ class VideoFromFile(VideoInput): out_video.width = out_width out_video.height = out_height out_video.pix_fmt = pix_fmt + if crf is not None: + out_video.options = {"crf": str(crf)} # source pts pass through (rebased to 0), so variable frame rate survives out_video.codec_context.time_base = video_stream.time_base if audio_stream is not None: @@ -827,6 +833,7 @@ class VideoFromComponents(VideoInput): codec: VideoCodec = VideoCodec.AUTO, metadata: Optional[dict] = None, bit_depth: int | None = None, + crf: float | None = None, ): """Save the video to a file path or BytesIO buffer.""" open_kwargs = mp4_output_open_kwargs(path, format, codec) @@ -847,6 +854,8 @@ class VideoFromComponents(VideoInput): video_stream.width = self.__components.images.shape[2] video_stream.height = self.__components.images.shape[1] video_stream.pix_fmt = pix_fmt + if crf is not None: + video_stream.options = {"crf": str(crf)} # Create an audio stream audio_sample_rate = 1 diff --git a/comfy_extras/nodes_video.py b/comfy_extras/nodes_video.py index 3bfd00be4..45394ce4d 100644 --- a/comfy_extras/nodes_video.py +++ b/comfy_extras/nodes_video.py @@ -86,7 +86,31 @@ class SaveVideo(io.ComfyNode): io.Video.Input("video", tooltip="The video to save."), io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."), io.Combo.Input("format", options=Types.VideoContainer.as_input(), default="auto", tooltip="The format to save the video as."), - io.Combo.Input("codec", options=Types.VideoCodec.as_input(), default="auto", tooltip="The codec to use for the video."), + io.DynamicCombo.Input( + "codec", + options=[ + io.DynamicCombo.Option("auto", []), + io.DynamicCombo.Option( + "h264", + [ + io.DynamicCombo.Input( + "encoding", + display_name="encoding mode", + options=[ + io.DynamicCombo.Option("auto", []), + io.DynamicCombo.Option( + "re-encode", + [io.Float.Input("crf", default=23.0, min=0.0, max=51.0, step=1.0, tooltip="Lower values produce higher quality and larger files.")], + ), + ], + optional=True, + tooltip="Automatic preserves compatible H.264 streams. Re-encode applies a custom CRF.", + ), + ], + ), + ], + tooltip="The codec to use for the video.", + ), ], hidden=[io.Hidden.prompt, io.Hidden.extra_pnginfo], is_output_node=True, @@ -94,7 +118,9 @@ class SaveVideo(io.ComfyNode): ) @classmethod - def execute(cls, video: Input.Video, filename_prefix, format: str, codec) -> io.NodeOutput: + def execute(cls, video: Input.Video, filename_prefix, format: str, codec: io.DynamicCombo.Type) -> io.NodeOutput: + codec_name = codec["codec"] + encoding = codec.get("encoding") or {} width, height = video.get_dimensions() full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( filename_prefix, @@ -115,8 +141,9 @@ class SaveVideo(io.ComfyNode): video.save_to( os.path.join(full_output_folder, file), format=Types.VideoContainer(format), - codec=codec, - metadata=saved_metadata + codec=codec_name, + metadata=saved_metadata, + crf=encoding.get("crf"), ) return io.NodeOutput(video, ui=ui.PreviewVideo([ui.SavedResult(file, subfolder, io.FolderType.output)])) diff --git a/tests-unit/comfy_api_test/video_types_test.py b/tests-unit/comfy_api_test/video_types_test.py index ae758bd40..dd95dc843 100644 --- a/tests-unit/comfy_api_test/video_types_test.py +++ b/tests-unit/comfy_api_test/video_types_test.py @@ -240,6 +240,24 @@ def test_duration_consistency(video_components): assert duration == pytest.approx(manual_duration) +def test_save_to_h264_crf_controls_quality(tmp_path): + generator = torch.Generator().manual_seed(7) + components = VideoComponents( + images=torch.rand(12, 64, 64, 3, generator=generator), + frame_rate=Fraction(30), + ) + high_quality = str(tmp_path / "high_quality.mp4") + low_quality = str(tmp_path / "low_quality.mp4") + transcoded = str(tmp_path / "transcoded.mp4") + + VideoFromComponents(components).save_to(high_quality, codec=VideoCodec.H264, crf=0) + VideoFromComponents(components).save_to(low_quality, codec=VideoCodec.H264, crf=51) + assert os.path.getsize(high_quality) > os.path.getsize(low_quality) + + VideoFromFile(high_quality).save_to(transcoded, codec=VideoCodec.H264, crf=51) + assert os.path.getsize(transcoded) < os.path.getsize(high_quality) + + def create_transcode_source( width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False, container_format="mov", audio_codec="pcm_s16le", From 2881e6161081439b1c3fb3b6c1f51b3d272da710 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 00:21:28 -0700 Subject: [PATCH 193/211] Store mp4 metadata at the beginning of the file when possible. (#15195) --- comfy_api/latest/_input_impl/video_types.py | 10 +++++++--- tests-unit/comfy_api_test/input_impl_test.py | 5 +++-- tests-unit/comfy_api_test/video_types_test.py | 12 ++++++++++++ 3 files changed, 22 insertions(+), 5 deletions(-) diff --git a/comfy_api/latest/_input_impl/video_types.py b/comfy_api/latest/_input_impl/video_types.py index 14d663881..cf4119250 100644 --- a/comfy_api/latest/_input_impl/video_types.py +++ b/comfy_api/latest/_input_impl/video_types.py @@ -36,13 +36,15 @@ def get_open_write_kwargs( dest: str | io.BytesIO, container_format: str, to_format: str | None ) -> dict: """Get kwargs for writing a `VideoFromFile` to a file/stream with `av.open`""" + is_write_to_buffer = isinstance(dest, io.BytesIO) + is_mp4_file = not is_write_to_buffer and os.path.splitext(dest)[1].lower() == ".mp4" + movflags = "use_metadata_tags+faststart" if is_mp4_file else "use_metadata_tags" open_kwargs = { "mode": "w", # If isobmff, preserve custom metadata tags (workflow, prompt, extra_pnginfo) - "options": {"movflags": "use_metadata_tags"}, + "options": {"movflags": movflags}, } - is_write_to_buffer = isinstance(dest, io.BytesIO) if is_write_to_buffer: # Set output format explicitly, since it cannot be inferred from file extension if to_format == VideoContainer.AUTO: @@ -103,7 +105,9 @@ def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec raise ValueError("Only MP4 format is supported for now") if codec != VideoCodec.AUTO and codec != VideoCodec.H264: raise ValueError("Only H264 codec is supported for now") - open_kwargs = {"mode": "w", "options": {"movflags": "use_metadata_tags"}} + # FFmpeg's faststart pass reopens the output by filename, so it cannot be used with file-like objects. + movflags = "use_metadata_tags+faststart" if isinstance(path, (str, os.PathLike)) else "use_metadata_tags" + open_kwargs = {"mode": "w", "options": {"movflags": movflags}} if isinstance(format, VideoContainer) and format != VideoContainer.AUTO: open_kwargs["format"] = format.value elif isinstance(path, io.BytesIO): diff --git a/tests-unit/comfy_api_test/input_impl_test.py b/tests-unit/comfy_api_test/input_impl_test.py index 5fc21a9a7..f1924f163 100644 --- a/tests-unit/comfy_api_test/input_impl_test.py +++ b/tests-unit/comfy_api_test/input_impl_test.py @@ -36,6 +36,7 @@ def test_get_open_write_kwargs_filepath_no_format(): kwargs_specific = get_open_write_kwargs("output.avi", "mp4", "avi") fail_msg = "Format should not be set for file paths (Specific)" assert "format" not in kwargs_specific, fail_msg + assert kwargs_specific["options"]["movflags"] == "use_metadata_tags" def test_get_open_write_kwargs_base_options_mode(): @@ -43,9 +44,9 @@ def test_get_open_write_kwargs_base_options_mode(): kwargs = get_open_write_kwargs("output.mp4", "mp4", VideoContainer.AUTO) assert kwargs["mode"] == "w", "mode should be set to write" - fail_msg = "movflags should be set to preserve custom metadata tags" + fail_msg = "movflags should preserve custom metadata tags and enable faststart for MP4 files" assert "movflags" in kwargs["options"], fail_msg - assert kwargs["options"]["movflags"] == "use_metadata_tags", fail_msg + assert kwargs["options"]["movflags"] == "use_metadata_tags+faststart", fail_msg def test_get_open_write_kwargs_bytesio_auto_format(): diff --git a/tests-unit/comfy_api_test/video_types_test.py b/tests-unit/comfy_api_test/video_types_test.py index dd95dc843..f688d3eca 100644 --- a/tests-unit/comfy_api_test/video_types_test.py +++ b/tests-unit/comfy_api_test/video_types_test.py @@ -258,6 +258,18 @@ def test_save_to_h264_crf_controls_quality(tmp_path): assert os.path.getsize(transcoded) < os.path.getsize(high_quality) +def test_save_to_mp4_writes_metadata_before_media(video_components, tmp_path): + encoded = tmp_path / "encoded.mp4" + remuxed = tmp_path / "remuxed.mp4" + + VideoFromComponents(video_components).save_to(str(encoded), metadata={"prompt": {"test": "value"}}) + VideoFromFile(str(encoded)).save_to(str(remuxed), metadata={"prompt": {"test": "value"}}) + + for path in (encoded, remuxed): + data = path.read_bytes() + assert data.index(b"moov") < data.index(b"mdat") + + def create_transcode_source( width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False, container_format="mov", audio_codec="pcm_s16le", From d8e6aa55f377d9914e64bf07d7633dfe6abbf643 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sat, 1 Aug 2026 20:30:00 +0300 Subject: [PATCH 194/211] [Partner Nodes] feat(xAI): update nodes for grok-imagine-video-1.5 model (#15197) * [Partner Nodes] feat(xAI): update nodes for grok-imagine-video-1.5 model Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/grok.py | 5 ++ comfy_api_nodes/nodes_grok.py | 158 +++++++++++++++++++++++++++++++--- 2 files changed, 152 insertions(+), 11 deletions(-) diff --git a/comfy_api_nodes/apis/grok.py b/comfy_api_nodes/apis/grok.py index fbedb53e0..526d8c8ab 100644 --- a/comfy_api_nodes/apis/grok.py +++ b/comfy_api_nodes/apis/grok.py @@ -15,6 +15,10 @@ class InputUrlObject(BaseModel): url: str = Field(...) +class VoiceReferenceObject(BaseModel): + voice_id: str = Field(...) + + class ImageEditRequest(BaseModel): model: str = Field(...) images: list[InputUrlObject] = Field(...) @@ -31,6 +35,7 @@ class VideoGenerationRequest(BaseModel): prompt: str = Field(...) image: InputUrlObject | None = Field(None) reference_images: list[InputUrlObject] | None = Field(None) + reference_audios: list[VoiceReferenceObject] | None = Field(None) duration: int = Field(...) aspect_ratio: str | None = Field(...) resolution: str = Field(...) diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index a95b35917..672a3e537 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -1,3 +1,5 @@ +import re + import torch from typing_extensions import override @@ -12,6 +14,7 @@ from comfy_api_nodes.apis.grok import ( VideoGenerationRequest, VideoGenerationResponse, VideoStatusResponse, + VoiceReferenceObject, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -33,6 +36,75 @@ _GROK_VIDEO_MODEL_API_IDS = { "grok-imagine-video-1.5": "grok-imagine-video-1.5", } +_GROK_VOICE_OPTIONS = [ + "none", + "ara", + "eve", + "leo", + "rex", + "sal", + "carina", + "zagan", + "helix", + "orion", + "luna", + "iris", + "altair", + "zenith", + "perseus", + "helios", + "lux", + "kepler", + "rigel", + "cosmo", + "celeste", + "ursa", + "sirius", + "lumen", + "castor", + "naksh", + "atlas", +] + + +_GROK_REF_TAG_RE = re.compile(r"(?\d*)(?!\w)", re.IGNORECASE | re.ASCII) + + +def _normalize_grok_reference_prompt(prompt: str, total_images: int, voices: list[str]) -> str: + """Rewrite @Image1/@Audio1 style references (1-based, shared partner-node syntax) + into Grok's native / tags; an unnumbered @image/@audio means the first one. + Native tags pass through untouched. @ImageN refers to the Nth reference image overall, in + input order — a batched input contributes one number per image. @AudioN refers to the + 'voice_N' widget; the API only accepts compact arrays, so voices are remapped to array + positions and 'none' slots between selected voices are harmless. Substitution repeats until + stable so adjacent tags like '@Image1@Image2' all resolve.""" + audio_indices: dict[int, int] = {} + for slot, voice in enumerate(voices, start=1): + if voice != "none": + audio_indices[slot] = len(audio_indices) + + def repl(match: re.Match) -> str: + kind = match.group(1).lower() + idx = int(match.group("idx") or 1) + if kind == "image": + if not 1 <= idx <= total_images: + raise ValueError( + f"The prompt references @Image{idx}, but only {total_images} " + f"reference images are connected (a batched input counts once per image)." + ) + return f"" + if idx not in audio_indices: + if 1 <= idx <= len(voices): + raise ValueError(f"The prompt references @Audio{idx}, but 'voice_{idx}' is set to 'none'.") + raise ValueError(f"The prompt references @Audio{idx}, but only voices 1..{len(voices)} exist.") + return f"" + + prev = None + while prev != prompt: + prev = prompt + prompt = _GROK_REF_TAG_RE.sub(repl, prompt) + return prompt + def _extract_grok_price(response) -> float | None: if response.usage and response.usage.cost_in_usd_ticks is not None: @@ -509,12 +581,13 @@ class GrokVideoNode(IO.ComfyNode): IO.Combo.Input( "model", options=["grok-imagine-video", "grok-imagine-video-1.5"], - tooltip="grok-imagine-video-1.5 currently always requires an input image.", + tooltip="The model to use for video generation.", ), IO.String.Input( "prompt", multiline=True, - tooltip="Text description of the desired video.", + tooltip="Text description of the desired video. " + "Optional for grok-imagine-video-1.5 when an input image is provided.", ), IO.Combo.Input( "resolution", @@ -549,7 +622,7 @@ class GrokVideoNode(IO.ComfyNode): IO.Image.Input( "image", optional=True, - tooltip="Optional starting image for grok-imagine-video. Required for grok-imagine-video-1.5.", + tooltip="Optional starting image. If omitted, the video is generated from the text prompt alone.", ), ], outputs=[ @@ -589,8 +662,6 @@ class GrokVideoNode(IO.ComfyNode): seed: int, image: Input.Image | None = None, ) -> IO.NodeOutput: - if image is None and model == "grok-imagine-video-1.5": - raise ValueError(f"The '{model}' model requires an input image; connect one to the 'image' input.") if resolution == "1080p" and model != "grok-imagine-video-1.5": raise ValueError(f"1080p resolution is only available for grok-imagine-video-1.5, not '{model}'.") image_url = None @@ -598,7 +669,8 @@ class GrokVideoNode(IO.ComfyNode): if get_number_of_images(image) != 1: raise ValueError("Only one input image is supported.") image_url = InputUrlObject(url=f"data:image/png;base64,{tensor_to_base64_string(image)}") - validate_string(prompt, strip_whitespace=True, min_length=1) + if image is None or model != "grok-imagine-video-1.5": + validate_string(prompt, strip_whitespace=True, min_length=1) initial_response = await sync_op( cls, ApiEndpoint(path="/proxy/xai/v1/videos/generations", method="POST"), @@ -709,7 +781,7 @@ class GrokVideoReferenceNode(IO.ComfyNode): node_id="GrokVideoReferenceNode", display_name="Grok Reference-to-Video", category="partner/video/Grok", - description="Generate video guided by reference images as style and content references.", + description="Generate video guided by reference images, with optional preset voice references.", inputs=[ IO.String.Input( "prompt", @@ -719,6 +791,57 @@ class GrokVideoReferenceNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "grok-imagine-video-1.5", + [ + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("image"), + names=[f"reference_{i}" for i in range(1, 8)], + min=1, + ), + tooltip="Up to 7 reference images to guide the video generation. " + "Refer to them in the prompt as @Image1 ... @Image7, numbered " + "in input order; a batched input counts once per image.", + ), + IO.Combo.Input( + "voice_1", + options=_GROK_VOICE_OPTIONS, + tooltip="Optional preset voice reference; refer to it in the prompt as @Audio1. " + "The API supports only these preset voices, not custom audio.", + ), + IO.Combo.Input( + "voice_2", + options=_GROK_VOICE_OPTIONS, + tooltip="Optional second voice reference; @Audio2 in the prompt.", + ), + IO.Combo.Input( + "voice_3", + options=_GROK_VOICE_OPTIONS, + tooltip="Optional third voice reference; @Audio3 in the prompt.", + ), + IO.Combo.Input( + "resolution", + options=["480p", "720p"], + tooltip="The resolution of the output video.", + ), + IO.Combo.Input( + "aspect_ratio", + options=["16:9", "4:3", "3:2", "1:1", "2:3", "3:4", "9:16"], + tooltip="The aspect ratio of the output video.", + ), + IO.Int.Input( + "duration", + default=6, + min=1, + max=15, + step=1, + tooltip="The duration of the output video in seconds.", + display_mode=IO.NumberDisplay.slider, + ), + ], + ), IO.DynamicCombo.Option( "grok-imagine-video", [ @@ -779,16 +902,20 @@ class GrokVideoReferenceNode(IO.ComfyNode): is_api_node=True, price_badge=IO.PriceBadge( depends_on=IO.PriceBadgeDepends( - widgets=["model.duration", "model.resolution"], + widgets=["model", "model.duration", "model.resolution"], input_groups=["model.reference_images"], ), expr=""" ( + $is15 := $contains(widgets.model, "1.5"); $res := $lookup(widgets, "model.resolution"); $dur := $lookup(widgets, "model.duration"); $refs := $lookup(inputGroups, "model.reference_images"); - $rate := $res = "720p" ? 0.07 : 0.05; - $price := ($rate * $dur + 0.002 * $refs) * 1.43; + $rate := $is15 + ? ($res = "720p" ? 0.14 : 0.08) + : ($res = "720p" ? 0.07 : 0.05); + $imgCost := $is15 ? 0.01 : 0.002; + $price := ($rate * $dur + $imgCost * $refs) * 1.43; {"type":"usd","usd": $price} ) """, @@ -803,6 +930,14 @@ class GrokVideoReferenceNode(IO.ComfyNode): seed: int, ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1) + total_images = sum(get_number_of_images(t) for t in model["reference_images"].values()) + if total_images > 7: + raise ValueError(f"A maximum of 7 reference images is supported; {total_images} are connected.") + reference_audios = None + if model["model"] == "grok-imagine-video-1.5": + voices = [model.get(f"voice_{i}", "none") for i in range(1, 4)] + reference_audios = [VoiceReferenceObject(voice_id=v) for v in voices if v != "none"] or None + prompt = _normalize_grok_reference_prompt(prompt, total_images=total_images, voices=voices) ref_image_urls = await upload_images_to_comfyapi( cls, list(model["reference_images"].values()), @@ -814,8 +949,9 @@ class GrokVideoReferenceNode(IO.ComfyNode): cls, ApiEndpoint(path="/proxy/xai/v1/videos/generations", method="POST"), data=VideoGenerationRequest( - model=model["model"], + model=_GROK_VIDEO_MODEL_API_IDS.get(model["model"], model["model"]), reference_images=[InputUrlObject(url=i) for i in ref_image_urls], + reference_audios=reference_audios, prompt=prompt, resolution=model["resolution"], duration=model["duration"], From 1091c47b3ae9ac423097649a5b7593d369d03bec Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sat, 1 Aug 2026 20:43:42 +0300 Subject: [PATCH 195/211] [Partner Nodes] chore(Bria): increase price for video endpoints (#15186) Signed-off-by: Alexander Piskun --- comfy_api_nodes/nodes_bria.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/comfy_api_nodes/nodes_bria.py b/comfy_api_nodes/nodes_bria.py index 9e5a93330..77f780a3b 100644 --- a/comfy_api_nodes/nodes_bria.py +++ b/comfy_api_nodes/nodes_bria.py @@ -289,7 +289,7 @@ class BriaRemoveVideoBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.05,"format":{"suffix":"/second"}}""", ), ) @@ -357,7 +357,7 @@ class BriaVideoGreenScreen(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.05,"format":{"suffix":"/second"}}""", ), ) @@ -433,7 +433,7 @@ class BriaVideoReplaceBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.05,"format":{"suffix":"/second"}}""", ), ) @@ -533,7 +533,7 @@ class BriaTransparentVideoBackground(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - expr="""{"type":"usd","usd":0.005,"format":{"suffix":"/second"}}""", + expr="""{"type":"usd","usd":0.05,"format":{"suffix":"/second"}}""", ), ) From e8e233cdf61c19a76fa60918efe2a6a6172f88cd Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 12:10:39 -0700 Subject: [PATCH 196/211] Update comfy-kitchen version to 0.2.26 (#15208) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index af414d413..aa3fbba8e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,7 +22,7 @@ alembic SQLAlchemy>=2.0.0 filelock av>=16.0.0 -comfy-kitchen==0.2.25 +comfy-kitchen==0.2.26 comfy-aimdo==0.4.10 requests simpleeval>=1.0.0 From e803f24ea090de7108772d65957fd6388d3f2085 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 13:38:42 -0700 Subject: [PATCH 197/211] Let the VAEDecodeAudio node decode nested audio. (#15211) --- comfy_extras/nodes_audio.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/comfy_extras/nodes_audio.py b/comfy_extras/nodes_audio.py index 4ac5ced53..1c6e91fe4 100644 --- a/comfy_extras/nodes_audio.py +++ b/comfy_extras/nodes_audio.py @@ -96,10 +96,14 @@ class VAEEncodeAudio(IO.ComfyNode): def vae_decode_audio(vae, samples, tile=None, overlap=None): + latent = samples["samples"] + if latent.is_nested: + latent = latent.unbind()[-1] + if tile is not None: - audio = vae.decode_tiled(samples["samples"], tile_x=tile, tile_y=tile, overlap=overlap).movedim(-1, 1) + audio = vae.decode_tiled(latent, tile_x=tile, tile_y=tile, overlap=overlap).movedim(-1, 1) else: - audio = vae.decode(samples["samples"]).movedim(-1, 1) + audio = vae.decode(latent).movedim(-1, 1) std = torch.std(audio, dim=[1, 2], keepdim=True) * 5.0 std[std < 1.0] = 1.0 From 49a74228920714b3756d206d2fec8ff61145dd2d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Sun, 2 Aug 2026 02:57:15 +0300 Subject: [PATCH 198/211] Support latent previews for nested latents (#15196) --- comfy/latent_formats.py | 134 ++++++++++++++++++++++++++- comfy/nested_tensor.py | 3 + comfy/samplers.py | 8 ++ comfy_extras/nodes_custom_sampler.py | 14 +-- latent_preview.py | 2 + 5 files changed, 153 insertions(+), 8 deletions(-) diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index 8a16cfe55..63eefd420 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -434,8 +434,138 @@ class LTXV(LatentFormat): class LTXAV(LTXV): def __init__(self): - self.latent_rgb_factors = None - self.latent_rgb_factors_bias = None + # video-stream preview factors for the packed AV latent (audio stream is not previewed) + self.latent_rgb_factors = [ + [ 0.001135, -0.010555, -0.004925], + [-0.008019, -0.006231, -0.005564], + [ 0.012637, 0.005605, 0.012713], + [ 0.023454, 0.020771, 0.017844], + [-0.011940, -0.000932, 0.009292], + [ 0.018602, 0.011018, 0.013969], + [-0.036369, -0.046631, -0.057898], + [-0.031919, 0.000131, 0.015214], + [ 0.014519, 0.021041, 0.015325], + [ 0.018889, 0.016149, -0.002836], + [-0.003784, -0.006057, -0.008195], + [ 0.013262, 0.030259, 0.029775], + [ 0.050465, 0.050366, 0.025255], + [ 0.018628, 0.007691, 0.002893], + [-0.015698, -0.008451, -0.000676], + [-0.013600, -0.012587, -0.004437], + [ 0.012482, 0.021469, 0.027913], + [-0.018241, -0.013488, -0.010975], + [ 0.013828, 0.012568, 0.021984], + [ 0.017911, 0.006552, 0.005567], + [ 0.026769, 0.006803, -0.009360], + [-0.006794, -0.008447, -0.013921], + [ 0.029708, 0.018671, 0.022811], + [-0.014732, -0.019169, 0.000903], + [ 0.019607, 0.032595, 0.053409], + [-0.003721, 0.003976, 0.010364], + [-0.020193, -0.026076, -0.036068], + [-0.002328, 0.006527, 0.013052], + [ 0.017171, 0.009224, 0.006548], + [ 0.001104, -0.000591, 0.000147], + [-0.000217, 0.011834, 0.017945], + [-0.015329, -0.012463, -0.006178], + [-0.009478, -0.008680, -0.004107], + [-0.005565, -0.006006, -0.001493], + [ 0.009451, 0.008794, 0.013207], + [-0.009989, -0.008027, -0.009568], + [-0.001505, -0.008805, -0.006828], + [ 0.001105, 0.008999, 0.009079], + [ 0.025935, 0.016426, 0.008036], + [ 0.006313, 0.000694, -0.006039], + [-0.001893, -0.006951, -0.009560], + [-0.007082, -0.002566, -0.007152], + [-0.005231, 0.004829, 0.008220], + [-0.004333, 0.001251, -0.004852], + [-0.017024, -0.012730, -0.007457], + [ 0.024988, 0.032963, 0.036556], + [ 0.013697, 0.012278, 0.009979], + [-0.013751, -0.008369, -0.015446], + [-0.009348, -0.001047, 0.007622], + [-0.003135, -0.003350, -0.003766], + [ 0.007436, 0.004957, 0.010480], + [ 0.018315, 0.022066, 0.021104], + [-0.005621, -0.006770, -0.008219], + [-0.007427, 0.001911, -0.001231], + [-0.007413, 0.000486, -0.006039], + [-0.014698, -0.007160, 0.006509], + [ 0.013775, 0.014185, 0.008203], + [ 0.060246, 0.069787, 0.072833], + [ 0.009861, 0.004870, 0.001194], + [-0.003660, 0.003251, 0.008015], + [ 0.003696, -0.003680, -0.008851], + [ 0.014924, 0.006196, 0.005282], + [-0.006740, -0.004319, -0.006729], + [ 0.020635, 0.015163, 0.012385], + [-0.032623, -0.006105, 0.010436], + [-0.058988, -0.030162, -0.037961], + [-0.035614, -0.021929, -0.011062], + [-0.023412, -0.011305, -0.005054], + [-0.002716, -0.005184, -0.004084], + [ 0.014591, 0.015294, 0.014045], + [ 0.008310, 0.002466, -0.003225], + [ 0.005176, 0.001119, 0.000695], + [-0.021569, -0.030886, -0.044732], + [ 0.007517, 0.003891, 0.000551], + [-0.006793, 0.004059, 0.010184], + [-0.086481, -0.082033, -0.083414], + [ 0.004192, 0.000762, -0.008658], + [ 0.010970, 0.009002, 0.007384], + [ 0.004042, -0.006732, -0.011031], + [ 0.012164, 0.006401, 0.007483], + [ 0.029252, 0.013990, 0.011128], + [ 0.048452, 0.034648, 0.016269], + [ 0.024104, 0.012647, 0.011754], + [-0.013216, -0.020192, -0.019752], + [-0.010799, -0.008535, -0.005467], + [ 0.005823, 0.001403, 0.001890], + [ 0.052393, 0.044771, 0.032777], + [ 0.007576, -0.008080, -0.012453], + [ 0.009830, 0.004244, 0.001213], + [-0.025867, -0.013169, -0.010636], + [ 0.008494, 0.003135, 0.000790], + [ 0.003969, -0.002625, -0.010204], + [ 0.006509, 0.008272, 0.020819], + [-0.004943, -0.013424, -0.015351], + [ 0.005541, 0.009136, -0.003666], + [-0.014300, -0.015864, -0.016853], + [ 0.002650, 0.028393, 0.014125], + [-0.027661, -0.045422, -0.064995], + [ 0.009220, 0.015522, 0.010574], + [-0.002236, 0.002915, 0.004557], + [-0.020269, -0.008212, -0.000532], + [ 0.019294, 0.003655, -0.002809], + [ 0.007116, -0.002784, 0.000017], + [ 0.057277, 0.073270, 0.074401], + [-0.002616, -0.001696, -0.000498], + [ 0.007248, 0.009793, 0.022829], + [-0.002590, -0.005601, -0.000436], + [-0.007681, 0.003893, -0.004119], + [-0.057392, -0.045545, -0.025290], + [ 0.045188, 0.047985, 0.054059], + [ 0.000937, -0.008861, -0.038406], + [-0.010192, -0.008036, -0.005385], + [-0.030222, -0.027498, -0.030765], + [-0.008359, 0.013247, 0.010918], + [ 0.004102, 0.002093, 0.006934], + [ 0.039461, 0.027339, 0.008284], + [-0.075747, -0.076340, -0.071625], + [ 0.002692, 0.005096, -0.002247], + [-0.002453, -0.002785, -0.010483], + [ 0.012265, 0.005481, 0.001729], + [ 0.017755, 0.008655, 0.003532], + [ 0.055560, 0.049128, 0.044137], + [-0.025861, -0.023798, -0.018815], + [-0.014876, -0.010770, -0.010713], + [-0.017315, -0.012599, -0.008661], + [-0.008461, -0.006210, -0.007744], + [-0.040175, -0.042255, -0.048119], + [-0.019355, -0.021055, -0.021919], + ] + self.latent_rgb_factors_bias = [-0.347892, -0.363814, -0.370287] class HunyuanVideo(LatentFormat): latent_channels = 16 diff --git a/comfy/nested_tensor.py b/comfy/nested_tensor.py index b700816fa..08c7133f8 100644 --- a/comfy/nested_tensor.py +++ b/comfy/nested_tensor.py @@ -51,6 +51,9 @@ class NestedTensor: def float(self): return self.to(dtype=torch.float) + def cpu(self): + return self.to(device="cpu") + def chunk(self, *args, **kwargs): return self.apply_operation(None, lambda x, y: x.chunk(*args, **kwargs)) diff --git a/comfy/samplers.py b/comfy/samplers.py index 29e6bffb3..6fea0913e 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -1284,6 +1284,14 @@ class CFGGuider: sampler_shapes = [tuple(latent_image.shape)] detail("Sampler: model=%s latent_shapes=%s", self.model_patcher.model.__class__.__name__, sampler_shapes) + if len(latent_shapes) > 1 and callback is not None: + # samplers run on the flat pack, hand callbacks (previews, x0 output) the nested view + packed_callback = callback + def callback(step, x0, x, total_steps): + x0 = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0, latent_shapes)) + x = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x, latent_shapes)) + return packed_callback(step, x0, x, total_steps) + if denoise_mask is not None: if denoise_mask.is_nested: denoise_masks = denoise_mask.unbind() diff --git a/comfy_extras/nodes_custom_sampler.py b/comfy_extras/nodes_custom_sampler.py index 56ef5f526..e81b6328b 100644 --- a/comfy_extras/nodes_custom_sampler.py +++ b/comfy_extras/nodes_custom_sampler.py @@ -775,10 +775,11 @@ class SamplerCustom(io.ComfyNode): out.pop("downscale_ratio_temporal", None) out["samples"] = samples if "x0" in x0_output: - x0_out = model.model.process_latent_out(x0_output["x0"].cpu()) - if samples.is_nested: + x0 = x0_output["x0"] + if samples.is_nested and not x0.is_nested: latent_shapes = [x.shape for x in samples.unbind()] - x0_out = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0_out, latent_shapes)) + x0 = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0, latent_shapes)) + x0_out = model.model.process_latent_out(x0.cpu()) out_denoised = latent.copy() out_denoised["samples"] = x0_out else: @@ -1053,10 +1054,11 @@ class SamplerCustomAdvanced(io.ComfyNode): out.pop("downscale_ratio_temporal", None) out["samples"] = samples if "x0" in x0_output: - x0_out = guider.model_patcher.model.process_latent_out(x0_output["x0"].cpu()) - if samples.is_nested: + x0 = x0_output["x0"] + if samples.is_nested and not x0.is_nested: latent_shapes = [x.shape for x in samples.unbind()] - x0_out = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0_out, latent_shapes)) + x0 = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0, latent_shapes)) + x0_out = guider.model_patcher.model.process_latent_out(x0.cpu()) out_denoised = latent.copy() out_denoised["samples"] = x0_out else: diff --git a/latent_preview.py b/latent_preview.py index a9d777661..6bf2c1869 100644 --- a/latent_preview.py +++ b/latent_preview.py @@ -123,6 +123,8 @@ def prepare_callback(model, steps, x0_output_dict=None): preview_bytes = None if previewer: + if x0.is_nested: + x0 = x0.tensors[0] preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0) pbar.update_absolute(step + 1, total_steps, preview_bytes) return callback From 41a3e160d714cd99a42fa34bbfa3058839c2ceb1 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 19:15:17 -0700 Subject: [PATCH 199/211] Use the actual function to check if the path is within the dir. (#15216) --- folder_paths.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/folder_paths.py b/folder_paths.py index df53542dc..6723e52c8 100644 --- a/folder_paths.py +++ b/folder_paths.py @@ -549,11 +549,10 @@ def get_save_image_path(filename_prefix: str, output_dir: str, image_width=0, im full_output_folder = os.path.join(output_dir, subfolder) - if os.path.commonpath((output_dir, os.path.abspath(full_output_folder))) != output_dir: + if not is_within_directory(output_dir, full_output_folder): err = "**** ERROR: Saving image outside the output folder is not allowed." + \ "\n full_output_folder: " + os.path.abspath(full_output_folder) + \ - "\n output_dir: " + output_dir + \ - "\n commonpath: " + os.path.commonpath((output_dir, os.path.abspath(full_output_folder))) + "\n output_dir: " + output_dir logging.error(err) raise Exception(err) From e0951a9c7e47f534d47ec1ab089c94faf484cbc0 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 19:27:37 -0700 Subject: [PATCH 200/211] Let the save text node save in csv format. (#15217) --- comfy_extras/nodes_text.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/comfy_extras/nodes_text.py b/comfy_extras/nodes_text.py index a485f5df8..82efc247f 100644 --- a/comfy_extras/nodes_text.py +++ b/comfy_extras/nodes_text.py @@ -8,6 +8,13 @@ import folder_paths class SaveTextNode(io.ComfyNode): """Save text content to .txt, .md, or .json.""" + FORMAT_EXTENSIONS = { + "txt": "txt", + "csv": "csv", + "md": "md", + "json": "json", + } + @classmethod def define_schema(cls): return io.Schema( @@ -19,7 +26,7 @@ class SaveTextNode(io.ComfyNode): inputs=[ io.String.Input("text", force_input=True), io.String.Input("filename_prefix", default="ComfyUI"), - io.Combo.Input("format", options=["txt", "md", "json"], default="txt"), + io.Combo.Input("format", options=list(cls.FORMAT_EXTENSIONS), default="txt"), ], outputs=[io.String.Output(display_name="text")], is_output_node=True, @@ -27,6 +34,10 @@ class SaveTextNode(io.ComfyNode): @classmethod def execute(cls, text, filename_prefix, format): + extension = cls.FORMAT_EXTENSIONS.get(format) + if extension is None: + raise ValueError(f"Unsupported text format: {format!r}") + full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( filename_prefix, folder_paths.get_output_directory(), @@ -34,10 +45,10 @@ class SaveTextNode(io.ComfyNode): 1, ) - file = f"{filename}_{counter:05}.{format}" + file = f"{filename}_{counter:05}.{extension}" filepath = os.path.join(full_output_folder, file) - if format == "json": + if extension == "json": # tries to pretty print otherwise saves normally try: data = json.loads(text) From 532a16f3b9557f49ef8e3d5dc7978d2040a0f4a1 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 20:15:51 -0700 Subject: [PATCH 201/211] Disable gradients on diffusion models. (#15218) --- comfy/model_base.py | 1 + 1 file changed, 1 insertion(+) diff --git a/comfy/model_base.py b/comfy/model_base.py index ee6dc57a2..7c68c6eba 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -169,6 +169,7 @@ class BaseModel(torch.nn.Module): else: operations = model_config.custom_operations self.diffusion_model = unet_model(**unet_config, device=device, operations=operations) + self.diffusion_model.requires_grad_(False) self.diffusion_model.eval() if comfy.model_management.force_channels_last(): self.diffusion_model.to(memory_format=torch.channels_last) From f06a187f50f896e4a0ba5be1ce1f2d2dcd13b77b Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Sat, 1 Aug 2026 22:12:12 -0700 Subject: [PATCH 202/211] Handle case where swap memory query fails on windows. (#15219) --- comfy/model_management.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index f7351224d..e3c94c15a 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -684,7 +684,11 @@ def should_free_pins_for_ram_pressure(shortfall): return True if psutil.virtual_memory().available < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE: return True - return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT + try: + return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT + except RuntimeError as err: + logging.warning("Could not read Windows swap usage; falling back to RAM-pressure pin eviction: %s", err) + return True def ensure_pin_budget(size, evict_active=False, loaded=False): if args.high_ram: From 8084083d4b085e0ba4bf09b31a050e48ea2a970a Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Sun, 2 Aug 2026 20:55:37 +1000 Subject: [PATCH 203/211] comfy-aimdo 0.4.11 (#15215) Changes: Remove sequential scan hint Prefer NVML pressure on windows Add async malloc clamp option (unused by comfy so far) Workaround AMD windows GPU virtual address space leak The largest change is the NVML pressure, which works around a cuMemGetInfo drift from actual VRAM in some circumstances. --- comfy/cli_args.py | 1 + main.py | 11 ++++++++--- requirements.txt | 2 +- 3 files changed, 10 insertions(+), 4 deletions(-) diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 792148f0a..ee9e1ce9f 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -172,6 +172,7 @@ vram_group.add_argument("--cpu", action="store_true", help="To use the CPU for e parser.add_argument("--reserve-vram", type=float, default=None, help="Set the amount of vram in GB you want to reserve for use by your OS/other software. By default some amount is reserved depending on your OS.") parser.add_argument("--vram-headroom", type=float, default=0, help="Set the amount of vram in GB for DynamicVRAM to maintain as extra headroom above default. ComfyUI will try and keep this much VRAM completely free and unused, even counting VRAM from other apps.") +parser.add_argument("--disable-nvml-pressure", action="store_true", help="Use CUDA instead of NVML for DynamicVRAM memory pressure.") parser.add_argument("--async-offload", nargs='?', const=2, type=int, default=None, metavar="NUM_STREAMS", help="Use async weight offloading. An optional argument controls the amount of offload streams. Default is 2. Enabled by default on Nvidia.") parser.add_argument("--disable-async-offload", action="store_true", help="Disable async weight offloading.") diff --git a/main.py b/main.py index 9c318fafe..361b1fc89 100644 --- a/main.py +++ b/main.py @@ -58,11 +58,16 @@ if __name__ == "__main__" and args.debug_hang: import comfy_aimdo.control if enables_dynamic_vram(): + simple_vram_headroom = None if args.reserve_vram is None else int(args.reserve_vram * 1024 ** 3) try: - comfy_aimdo.control.init(simple_vram_headroom=None if args.reserve_vram is None else int(args.reserve_vram * 1024 ** 3)) + comfy_aimdo.control.init(simple_vram_headroom=simple_vram_headroom, nvml_pressure=not args.disable_nvml_pressure) except TypeError: - # comfy-aimdo 0.4.9 protocol. - comfy_aimdo.control.init() + # comfy-aimdo 0.4.10 protocol. + try: + comfy_aimdo.control.init(simple_vram_headroom=simple_vram_headroom) + except TypeError: + # comfy-aimdo 0.4.9 protocol. + comfy_aimdo.control.init() if os.name == "nt": os.environ['MIMALLOC_PURGE_DELAY'] = '0' diff --git a/requirements.txt b/requirements.txt index aa3fbba8e..c1eeb9752 100644 --- a/requirements.txt +++ b/requirements.txt @@ -23,7 +23,7 @@ SQLAlchemy>=2.0.0 filelock av>=16.0.0 comfy-kitchen==0.2.26 -comfy-aimdo==0.4.10 +comfy-aimdo==0.4.11 requests simpleeval>=1.0.0 blake3 From 364081170f1c15853608db1f24ed9897763bdd09 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sun, 2 Aug 2026 17:10:47 +0300 Subject: [PATCH 204/211] [Partner Nodes] feat(Minimax): add 768P resolution for H3 model (#15227) Signed-off-by: Alexander Piskun --- comfy_api_nodes/nodes_minimax.py | 21 ++++++++++++--------- 1 file changed, 12 insertions(+), 9 deletions(-) diff --git a/comfy_api_nodes/nodes_minimax.py b/comfy_api_nodes/nodes_minimax.py index 2d7aef654..3c1d29257 100644 --- a/comfy_api_nodes/nodes_minimax.py +++ b/comfy_api_nodes/nodes_minimax.py @@ -467,7 +467,7 @@ def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = Tr ), IO.Combo.Input( "resolution", - options=["2K"], + options=["768P", "2K"], tooltip="Resolution of the output video.", ), ] @@ -578,11 +578,12 @@ class MinimaxHailuo03TextToVideoNode(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model.duration"]), + depends_on=IO.PriceBadgeDepends(widgets=["model.resolution", "model.duration"]), expr=""" ( $dur := $lookup(widgets, "model.duration"); - {"type": "usd", "usd": $dur * 0.1859} + $rate := $lookup(widgets, "model.resolution") = "768p" ? 0.1287 : 0.1859; + {"type": "usd", "usd": $dur * $rate} ) """, ), @@ -660,11 +661,12 @@ class MinimaxHailuo03FirstLastFrameNode(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model.duration"]), + depends_on=IO.PriceBadgeDepends(widgets=["model.resolution", "model.duration"]), expr=""" ( $dur := $lookup(widgets, "model.duration"); - {"type": "usd", "usd": $dur * 0.1859} + $rate := $lookup(widgets, "model.resolution") = "768p" ? 0.1287 : 0.1859; + {"type": "usd", "usd": $dur * $rate} ) """, ), @@ -818,20 +820,21 @@ class MinimaxHailuo03ReferenceNode(IO.ComfyNode): is_api_node=True, price_badge=IO.PriceBadge( depends_on=IO.PriceBadgeDepends( - widgets=["model.duration"], + widgets=["model.resolution", "model.duration"], input_groups=["model.reference_images", "model.reference_videos"], ), expr=""" ( $dur := $lookup(widgets, "model.duration"); + $rate := $lookup(widgets, "model.resolution") = "768p" ? 0.1287 : 0.1859; $imgsRaw := $lookup(inputGroups, "model.reference_images"); $imgs := $imgsRaw ? $imgsRaw : 0; $vidsRaw := $lookup(inputGroups, "model.reference_videos"); $vids := $vidsRaw ? $vidsRaw : 0; - $base := $dur * 0.1859 + ($imgs > 5 ? ($imgs - 5) * 0.0572 : 0); + $base := $dur * $rate + ($imgs > 5 ? ($imgs - 5) * 0.0572 : 0); $vids > 0 - ? {"type": "range_usd", "min_usd": $base + $vids * 2 * 0.1859, - "max_usd": $base + 15 * 0.1859, "format": {"approximate": true}} + ? {"type": "range_usd", "min_usd": $base + $vids * 2 * $rate, + "max_usd": $base + 15 * $rate, "format": {"approximate": true}} : {"type": "usd", "usd": $base} ) """, From 611f2a4e0f30ea7f50451fd34ab25b8f9365ff2f Mon Sep 17 00:00:00 2001 From: rattus <46076784+rattus128@users.noreply.github.com> Date: Mon, 3 Aug 2026 01:16:06 +1000 Subject: [PATCH 205/211] fix pin registration priority (#15226) This priority scheme was broken in the case where you have pin registration exhaustion while loading a VBAR that gets a big evicition. The weight would stay in the loaded set but inherit the MRU priority against other workflow models WRT pin registration which leads to async offload without pinning. Fix by universally promiting active pin registration above workflow pins without concern for the weights/weights-loaded split. This diverges from the actual budgeting where the split still makes sense. --- comfy/model_management.py | 21 +++++++++++++++++---- comfy/pinned_memory.py | 6 +++--- 2 files changed, 20 insertions(+), 7 deletions(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index e3c94c15a..1000f69e1 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -671,6 +671,19 @@ def pin_eviction_tiers(loaded, evict_active): tiers.append((PIN_SUBSETS, True, True)) return tiers +def registration_eviction_tiers(evict_active): + subsets = PIN_SUBSETS + LOADED_PIN_SUBSETS + tiers = [ + (subsets, False, False), + (subsets, True, False), + ] + if evict_active: + tiers.extend([ + (subsets, False, True), + (subsets, True, True), + ]) + return tiers + def free_pins(size, evict_active=False, loaded=False): freed = 0 for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active): @@ -703,19 +716,19 @@ def ensure_pin_budget(size, evict_active=False, loaded=False): to_free = shortfall + PIN_PRESSURE_HYSTERESIS return free_pins(to_free, evict_active=evict_active, loaded=loaded) >= shortfall -def free_registrations(shortfall, evict_active=True, loaded=False): +def free_registrations(shortfall, evict_active=True): if MAX_PINNED_MEMORY <= 0: return False if shortfall <= 0: return True shortfall += REGISTERABLE_PIN_HYSTERESIS - for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active): + for subsets, current_prompt, active in registration_eviction_tiers(evict_active): shortfall -= free_model_pins(shortfall, subsets, current_prompt, active, registrations=True) return shortfall <= REGISTERABLE_PIN_HYSTERESIS -def ensure_pin_registerable(size, evict_active=True, loaded=False): - return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active, loaded=loaded) +def ensure_pin_registerable(size, evict_active=True): + return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active) class LoadedModel: def __init__(self, model: ModelPatcher): diff --git a/comfy/pinned_memory.py b/comfy/pinned_memory.py index d78ab3c76..e9a9a70e2 100644 --- a/comfy/pinned_memory.py +++ b/comfy/pinned_memory.py @@ -56,7 +56,7 @@ def get_pin(module, subset="weights"): _, _, stack_split, pinned_size, *_ = module._pin_state[subset] size = pin.nbytes - comfy.model_management.ensure_pin_registerable(size, loaded=subset.endswith("-loaded")) + comfy.model_management.ensure_pin_registerable(size) if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() @@ -93,7 +93,7 @@ def pin_memory(module, subset="weights", size=None): comfy.memory_management.extra_ram_release(comfy.memory_management.RAM_CACHE_HEADROOM) if (not comfy.model_management.ensure_pin_budget(size, loaded=loaded) or - not comfy.model_management.ensure_pin_registerable(registerable_size, loaded=loaded)): + not comfy.model_management.ensure_pin_registerable(registerable_size)): return _steal_pin(module, stack, buckets, size, priority, subset) offset = hostbuf.size @@ -105,7 +105,7 @@ def pin_memory(module, subset="weights", size=None): pin.untyped_storage()._comfy_hostbuf = hostbuf if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() - comfy.model_management.free_registrations(size, loaded=loaded) + comfy.model_management.free_registrations(size) if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0: comfy.model_management.discard_cuda_async_error() del pin From b53e247c94f9225dc206bcfef5d64a2f7bc85232 Mon Sep 17 00:00:00 2001 From: Oliver Freyermuth Date: Sun, 2 Aug 2026 22:20:27 +0200 Subject: [PATCH 206/211] rename comfy/logging.py to comfy/internal_logging.py (#15231) This avoids name collision (circular imports) for external custom nodes, for which the comfy path is pushed into sys.path so Python's own logging module is shadowed otherwise. fixes: #15229 --- app/logger.py | 4 ++-- comfy/{logging.py => internal_logging.py} | 0 comfy/model_management.py | 2 +- comfy/model_patcher.py | 2 +- comfy/samplers.py | 2 +- execution.py | 2 +- 6 files changed, 6 insertions(+), 6 deletions(-) rename comfy/{logging.py => internal_logging.py} (100%) diff --git a/app/logger.py b/app/logger.py index 2e6116813..fe82c40c9 100644 --- a/app/logger.py +++ b/app/logger.py @@ -6,7 +6,7 @@ import os import sys import threading -import comfy.logging +import comfy.internal_logging ANSI_NAMED_COLORS = { 'black': '\033[30m', @@ -91,7 +91,7 @@ def on_flush(callback): def get_log_level(level): - return comfy.logging.DETAIL if level == "DETAIL" else logging.getLevelName(level) + return comfy.internal_logging.DETAIL if level == "DETAIL" else logging.getLevelName(level) def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300, use_stdout: bool = False): diff --git a/comfy/logging.py b/comfy/internal_logging.py similarity index 100% rename from comfy/logging.py rename to comfy/internal_logging.py diff --git a/comfy/model_management.py b/comfy/model_management.py index 1000f69e1..ec91fb929 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -34,7 +34,7 @@ import comfy.utils import comfy.quant_ops import comfy_aimdo.host_buffer import comfy_aimdo.vram_buffer -from comfy.logging import detail +from comfy.internal_logging import detail from typing import TYPE_CHECKING if TYPE_CHECKING: diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index 6b698f767..ae3f0191d 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -38,7 +38,7 @@ import comfy.patcher_extension import comfy.utils import comfy_aimdo.host_buffer from comfy.comfy_types import UnetWrapperFunction -from comfy.logging import detail +from comfy.internal_logging import detail from comfy.quant_ops import QuantizedTensor from comfy.patcher_extension import CallbacksMP, PatcherInjection, WrappersMP diff --git a/comfy/samplers.py b/comfy/samplers.py index 6fea0913e..a280f3bb6 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -20,7 +20,7 @@ import comfy.hooks import comfy.context_windows import comfy.multigpu import comfy.utils -from comfy.logging import detail +from comfy.internal_logging import detail import scipy.stats import numpy diff --git a/execution.py b/execution.py index 7cab4b331..0b858969d 100644 --- a/execution.py +++ b/execution.py @@ -19,7 +19,7 @@ import comfy.model_management import comfy.model_patcher import comfy.model_prefetch import comfy_aimdo.model_vbar -from comfy.logging import detail +from comfy.internal_logging import detail from latent_preview import set_preview_method import nodes From 57500fc5bc92566a63f2046824f522cd55c335ca Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Mon, 3 Aug 2026 05:28:29 +0300 Subject: [PATCH 207/211] feat: Support MiniMax-H3 (CORE-375) (#15224) --- comfy/latent_formats.py | 39 ++ comfy/ldm/minimax/audio_vae.py | 442 ++++++++++++++++++++ comfy/ldm/minimax/model.py | 646 ++++++++++++++++++++++++++++ comfy/ldm/minimax/vae.py | 694 +++++++++++++++++++++++++++++++ comfy/model_base.py | 52 +++ comfy/model_detection.py | 29 ++ comfy/ops.py | 55 ++- comfy/quant_ops.py | 2 +- comfy/sd.py | 55 +++ comfy/supported_models.py | 29 ++ comfy/text_encoders/llama.py | 11 + comfy/text_encoders/minimax.py | 201 +++++++++ comfy/text_encoders/qwen3vl.py | 5 +- comfy_extras/nodes_lt.py | 5 +- comfy_extras/nodes_minimax_h3.py | 337 +++++++++++++++ nodes.py | 3 +- 16 files changed, 2599 insertions(+), 6 deletions(-) create mode 100644 comfy/ldm/minimax/audio_vae.py create mode 100644 comfy/ldm/minimax/model.py create mode 100644 comfy/ldm/minimax/vae.py create mode 100644 comfy/text_encoders/minimax.py create mode 100644 comfy_extras/nodes_minimax_h3.py diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index 63eefd420..c4270022b 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -567,6 +567,45 @@ class LTXAV(LTXV): ] self.latent_rgb_factors_bias = [-0.347892, -0.363814, -0.370287] +class MiniMaxH3Video(LatentFormat): + latent_channels = 24 + latent_dimensions = 3 + spacial_downscale_ratio = 16 + temporal_downscale_ratio = 4 + scale_factor = 1.0 + + latent_rgb_factors = [ + [-0.018555, 0.024344, -0.017536], + [ 0.150164, 0.137244, 0.129221], + [ 0.027367, -0.050369, -0.208606], + [-0.000793, -0.164622, -0.323161], + [-0.048556, 0.013970, -0.074286], + [ 0.011740, 0.014172, -0.006906], + [ 0.061517, 0.061212, 0.110025], + [ 0.035321, 0.086879, 0.110059], + [-0.017426, 0.002997, 0.035356], + [ 0.531539, 0.548819, 0.624404], + [-0.024968, -0.040234, -0.034302], + [-0.032549, -0.029096, -0.017221], + [ 0.022609, 0.020286, 0.050661], + [-0.084001, -0.038131, -0.020805], + [-0.018830, 0.010412, 0.061120], + [ 0.020777, 0.011196, -0.030994], + [-0.008390, -0.012201, -0.025687], + [-0.013281, -0.002924, 0.006331], + [ 0.000260, 0.001833, -0.011038], + [ 0.105471, 0.100482, 0.132106], + [ 0.016529, 0.015213, 0.009999], + [-0.014015, -0.017438, -0.019134], + [-0.033787, -0.009984, -0.019725], + [ 0.004224, 0.017284, 0.027196], + ] + latent_rgb_factors_bias = [ 0.057426, -0.022078, -0.071449] + +class MiniMaxH3AV(MiniMaxH3Video): + # max channels across the two streams (video 24, audio 32) so per-stream slices keep both streams whole + latent_channels = 32 + class HunyuanVideo(LatentFormat): latent_channels = 16 latent_dimensions = 3 diff --git a/comfy/ldm/minimax/audio_vae.py b/comfy/ldm/minimax/audio_vae.py new file mode 100644 index 000000000..033ae6966 --- /dev/null +++ b/comfy/ldm/minimax/audio_vae.py @@ -0,0 +1,442 @@ +# MiniMax H3 audio VAE: DAC-lineage waveform encoder + BigVGAN decoder. +# Weight-norm parametrizations are folded into plain conv weights, so this +# module uses ordinary ops.Conv1d / ops.ConvTranspose1d and loads the converted +# checkpoint (plain "*.weight" tensors) with strict=True. +# +# Lineage / licenses of the reference implementation: +# DAC encoder: descript-audio-codec (MIT) +# BigVGAN decoder: NVIDIA BigVGAN (MIT), adapted from hifi-gan (MIT) +# Alias-free ops: junjun3518/alias-free-torch (Apache-2.0), julius (MIT) + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops + +ops = comfy.ops.disable_weight_init + + +# Snake activations + +def snake(x, alpha, beta): + # x + 1/beta * sin^2(alpha * x) + t = torch.sin(alpha * x) + return t.mul_(t).mul_((beta + 1e-9).reciprocal()).add_(x) + + +class Snake1d(nn.Module): + """Snake activation with per-channel alpha (encoder side).""" + + def __init__(self, channels): + super().__init__() + self.alpha = nn.Parameter(torch.empty(1, channels, 1)) + + def forward(self, x): + return snake(x, self.alpha, self.alpha) + + +class SnakeBeta(nn.Module): + """SnakeBeta := x + 1/beta * sin^2(alpha * x); alpha/beta stored in log scale.""" + + def __init__(self, in_features): + super().__init__() + self.alpha = nn.Parameter(torch.empty(in_features)) + self.beta = nn.Parameter(torch.empty(in_features)) + + def forward(self, x): + alpha = torch.exp(self.alpha).view(1, -1, 1) + beta = torch.exp(self.beta).view(1, -1, 1) + return snake(x, alpha, beta) + + +# Alias-free (anti-aliased) activation: kaiser-windowed sinc resampling + +def kaiser_sinc_filter1d(cutoff, half_width, kernel_size): + # returns filter [1, 1, kernel_size] + even = kernel_size % 2 == 0 + half_size = kernel_size // 2 + + # kaiser window design + delta_f = 4 * half_width + A = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95 + if A > 50.0: + beta = 0.1102 * (A - 8.7) + elif A >= 21.0: + beta = 0.5842 * (A - 21) ** 0.4 + 0.07886 * (A - 21.0) + else: + beta = 0.0 + window = torch.kaiser_window(kernel_size, beta=beta, periodic=False) + + if even: + time = torch.arange(-half_size, half_size) + 0.5 + else: + time = torch.arange(kernel_size) - half_size + + filter_ = 2 * cutoff * window * torch.sinc(2 * cutoff * time) + # Normalize filter to have sum = 1, otherwise there is a small leakage of + # the constant component in the input signal. + filter_ /= filter_.sum() + return filter_.view(1, 1, kernel_size) + + +class UpSample1d(nn.Module): + def __init__(self, ratio=2, kernel_size=12): + super().__init__() + self.ratio = ratio + self.stride = ratio + self.pad = kernel_size // ratio - 1 + self.pad_left = self.pad * ratio + (kernel_size - ratio) // 2 + self.pad_right = self.pad * ratio + (kernel_size - ratio + 1) // 2 + self.register_buffer( + "filter", + kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=kernel_size), + ) + + def forward(self, x): + _, C, _ = x.shape + x = F.pad(x, (self.pad, self.pad), mode="replicate") + x = F.conv_transpose1d(x, self.filter.expand(C, -1, -1).to(x.dtype), stride=self.stride, groups=C).mul_(self.ratio) + x = x[..., self.pad_left:-self.pad_right] + return x + + +class LowPassFilter1d(nn.Module): + def __init__(self, cutoff=0.5, half_width=0.6, stride=1, kernel_size=12): + super().__init__() + self.pad_left = kernel_size // 2 - int(kernel_size % 2 == 0) + self.pad_right = kernel_size // 2 + self.stride = stride + self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size)) + + def forward(self, x): + _, C, _ = x.shape + x = F.pad(x, (self.pad_left, self.pad_right), mode="replicate") + return F.conv1d(x, self.filter.expand(C, -1, -1).to(x.dtype), stride=self.stride, groups=C) + + +class DownSample1d(nn.Module): + def __init__(self, ratio=2, kernel_size=12): + super().__init__() + self.ratio = ratio + self.kernel_size = kernel_size + self.lowpass = LowPassFilter1d( + cutoff=0.5 / ratio, + half_width=0.6 / ratio, + stride=ratio, + kernel_size=self.kernel_size, + ) + + def forward(self, x): + return self.lowpass(x) + + +class Activation1d(nn.Module): + """upsample x2 -> pointwise activation -> downsample x2 (anti-aliased).""" + + def __init__(self, activation, up_ratio=2, down_ratio=2, up_kernel_size=12, down_kernel_size=12): + super().__init__() + self.act = activation + self.upsample = UpSample1d(up_ratio, up_kernel_size) + self.downsample = DownSample1d(down_ratio, down_kernel_size) + + def forward(self, x): + x = self.upsample(x) + x = self.act(x) + x = self.downsample(x) + return x + + +# DAC encoder + +class ResidualUnit(nn.Module): + def __init__(self, dim=16, dilation=1): + super().__init__() + pad = ((7 - 1) * dilation) // 2 + self.block = nn.Sequential( + Snake1d(dim), + ops.Conv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad), + Snake1d(dim), + ops.Conv1d(dim, dim, kernel_size=1), + ) + + def forward(self, x): + y = self.block(x) + pad = (x.shape[-1] - y.shape[-1]) // 2 + if pad > 0: + x = x[..., pad:-pad] + return y.add_(x) + + +class EncoderBlock(nn.Module): + def __init__(self, dim=16, stride=1): + super().__init__() + self.block = nn.Sequential( + ResidualUnit(dim // 2, dilation=1), + ResidualUnit(dim // 2, dilation=3), + ResidualUnit(dim // 2, dilation=9), + Snake1d(dim // 2), + ops.Conv1d( + dim // 2, + dim, + kernel_size=2 * stride, + stride=stride, + padding=math.ceil(stride / 2), + ), + ) + + def forward(self, x): + return self.block(x) + + +class Encoder(nn.Module): + def __init__(self, d_model=64, strides=(2, 4, 4, 5, 5), d_latent=2048): + super().__init__() + block = [ops.Conv1d(1, d_model, kernel_size=7, padding=3)] + for stride in strides: + d_model *= 2 + block += [EncoderBlock(d_model, stride=stride)] + block += [ + Snake1d(d_model), + ops.Conv1d(d_model, d_latent, kernel_size=3, padding=1), + ] + self.block = nn.Sequential(*block) + + def forward(self, x): + return self.block(x) + + +# Attention projection (encoder posterior head) + +class GeGluMlp(nn.Module): + def __init__(self, in_features, hidden_features): + super().__init__() + self.norm = ops.LayerNorm(in_features) + self.act = nn.GELU(approximate="tanh") + self.w0 = ops.Linear(in_features, hidden_features) + self.w1 = ops.Linear(in_features, hidden_features) + self.w2 = ops.Linear(hidden_features, in_features) + + def forward(self, x): + x = self.norm(x) + return self.w2(self.act(self.w0(x)).mul_(self.w1(x))) + + +class CausalAttention(nn.Module): + def __init__(self, in_dim, out_dim, num_heads): + super().__init__() + self.head_dim = in_dim // num_heads + self.num_heads = num_heads + self.out_dim = out_dim + self.qkv = ops.Linear(in_dim, in_dim * 3, bias=False) + self.q_bias = nn.Parameter(torch.empty(in_dim)) + self.v_bias = nn.Parameter(torch.empty(in_dim)) + self.register_buffer("zero_k_bias", torch.empty(in_dim)) + self.proj = ops.Linear(out_dim, out_dim) + + def forward(self, x): + B, N, C = x.shape + weight, _, offload_stream = comfy.ops.cast_bias_weight(self.qkv, x, offloadable=True) + qkv = F.linear(x, weight=weight, bias=torch.cat((self.q_bias, self.zero_k_bias, self.v_bias))) + comfy.ops.uncast_bias_weight(self.qkv, weight, None, offload_stream) + q, k, v = qkv.reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4).unbind(0) + + # mean over heads then pool down to the latent width (in_dim >> out_dim) + x = comfy.ops.scaled_dot_product_attention(q, k, v, is_causal=True) + x = F.adaptive_avg_pool1d(torch.mean(x, dim=1), self.out_dim) + return self.proj(x) + + +class AttnProjection(nn.Module): + def __init__(self, in_dim, out_dim, num_heads, mlp_ratio=2): + super().__init__() + self.norm1 = ops.LayerNorm(in_dim) + self.attn = CausalAttention(in_dim, out_dim, num_heads) + self.proj = ops.Linear(in_dim, out_dim) + self.norm3 = ops.LayerNorm(in_dim) + + self.norm2 = ops.LayerNorm(out_dim) + hidden_dim = int(out_dim * mlp_ratio) + self.mlp = GeGluMlp(in_features=out_dim, hidden_features=hidden_dim) + + def forward(self, x): + # x: [B, T, in_dim] + x = self.proj(self.norm3(x)).add_(self.attn(self.norm1(x))) + return x.add_(self.mlp(self.norm2(x))) + + +# BigVGAN decoder + +def get_padding(kernel_size, dilation=1): + return int((kernel_size * dilation - dilation) / 2) + + +class AMPBlock1(nn.Module): + def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)): + super().__init__() + self.convs1 = nn.ModuleList( + [ + ops.Conv1d(channels, channels, kernel_size, stride=1, dilation=d, padding=get_padding(kernel_size, d)) + for d in dilation + ] + ) + self.convs2 = nn.ModuleList( + [ + ops.Conv1d(channels, channels, kernel_size, stride=1, dilation=1, padding=get_padding(kernel_size, 1)) + for _ in range(len(dilation)) + ] + ) + self.num_layers = len(self.convs1) + len(self.convs2) + self.activations = nn.ModuleList( + [Activation1d(activation=SnakeBeta(channels)) for _ in range(self.num_layers)] + ) + + def forward(self, x): + acts1, acts2 = self.activations[::2], self.activations[1::2] + for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2): + xt = a1(x) + xt = c1(xt) + xt = a2(xt) + xt = c2(xt) + x = xt.add_(x) + return x + + +class BigVGAN(nn.Module): + """BigVGAN vocoder (MiniMax H3 32 kHz configuration). + + use_bias_at_final=False, use_tanh_at_final=False (output clamped to [-1, 1]). + """ + + def __init__( + self, + num_mels=2048, + upsample_initial_channel=1024, + upsample_rates=(5, 5, 2, 2, 2, 2, 2), + upsample_kernel_sizes=(9, 9, 4, 4, 4, 4, 4), + resblock_kernel_sizes=(3, 7, 11), + resblock_dilation_sizes=((1, 3, 5), (1, 3, 5), (1, 3, 5)), + ): + super().__init__() + self.num_kernels = len(resblock_kernel_sizes) + self.num_upsamples = len(upsample_rates) + + self.conv_pre = ops.Conv1d(num_mels, upsample_initial_channel, 7, 1, padding=3) + + self.ups = nn.ModuleList() + for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): + self.ups.append( + nn.ModuleList( + [ + ops.ConvTranspose1d( + upsample_initial_channel // (2 ** i), + upsample_initial_channel // (2 ** (i + 1)), + k, + u, + padding=(k - u) // 2, + ) + ] + ) + ) + + self.resblocks = nn.ModuleList() + for i in range(len(self.ups)): + ch = upsample_initial_channel // (2 ** (i + 1)) + for k, d in zip(resblock_kernel_sizes, resblock_dilation_sizes): + self.resblocks.append(AMPBlock1(ch, k, d)) + + self.activation_post = Activation1d(activation=SnakeBeta(ch)) + self.conv_post = ops.Conv1d(ch, 1, 7, 1, padding=3, bias=False) + + def forward(self, x): + x = self.conv_pre(x) + + for i in range(self.num_upsamples): + for i_up in range(len(self.ups[i])): + x = self.ups[i][i_up](x) + xs = None + for j in range(self.num_kernels): + if xs is None: + xs = self.resblocks[i * self.num_kernels + j](x) + else: + xs += self.resblocks[i * self.num_kernels + j](x) + x = xs.div_(self.num_kernels) + + x = self.activation_post(x) + return self.conv_post(x).clamp_(-1.0, 1.0) + + +# Top-level VAE + +class MiniMaxH3AudioVAE(nn.Module): + """MiniMax H3 stereo audio VAE at 32 kHz. + + Latents are [B, 32, 2, T]: 32 channels, 2 stereo channels, T frames at + 40 latent frames per second (800 audio samples per latent frame). The + stereo channels are processed independently by the mono encoder/decoder. + Latents are normalized with the stored per-channel latents_mean/std. + """ + + def __init__( + self, + encoder_dim=64, + encoder_rates=(2, 4, 4, 5, 5), + latent_dim=2048, + decoder_dim=1024, + vae_latent_channels=32, + ): + super().__init__() + self.sample_rate = 32000 + + self.hop_length = 1 + for r in encoder_rates: + self.hop_length *= r + self.samples_per_latent = self.hop_length # 800 + self.latents_per_second = self.sample_rate // self.hop_length # 40 + self.output_sample_rate = self.sample_rate # read by LTXVAudioVAEDecode + + self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) + + self.pre_block = AttnProjection(latent_dim, vae_latent_channels, num_heads=8) + + self.mean_proj = ops.Conv1d(vae_latent_channels, vae_latent_channels, 1) + # logs_proj exists in the checkpoint but is unused at inference + # (encode returns the posterior mean, no sampling). + self.logs_proj = ops.Conv1d(vae_latent_channels, vae_latent_channels, 1) + + self.dec_in_proj = ops.Conv1d(vae_latent_channels, latent_dim, 1) + self.decoder = BigVGAN(num_mels=latent_dim, upsample_initial_channel=decoder_dim) + + self.register_buffer("latents_mean", torch.empty(vae_latent_channels)) + self.register_buffer("latents_std", torch.empty(vae_latent_channels)) + + def decode(self, z): + """Decode normalized latents [B, 32, 2, T] to stereo waveforms [B, 2, L] at 32 kHz.""" + b, c, s, t = z.shape + z = z.permute(0, 2, 1, 3).reshape(b * s, c, t) + mean = self.latents_mean.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + std = self.latents_std.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + z = z * std + mean + x = self.dec_in_proj(z) + x = self.decoder(x) # [b * s, 1, L], already clamped to [-1, 1] + return x.reshape(b, s, -1) + + def encode(self, waveform): + """Encode stereo waveforms [B, 2, L] at 32 kHz (in [-1, 1]) to normalized latents [B, 32, 2, T]. + + L is right-padded with zeros to a multiple of 800 samples; the returned + posterior mean is used directly (no sampling). + """ + b, s, length = waveform.shape + right_pad = math.ceil(length / self.hop_length) * self.hop_length - length + waveform = F.pad(waveform, (0, right_pad)) + x = waveform.reshape(b * s, 1, -1) + x = self.encoder(x) # [b * s, latent_dim, T] + x = self.pre_block(x.transpose(1, 2)).transpose(1, 2) # [b * s, 32, T] + z = self.mean_proj(x) + mean = self.latents_mean.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + std = self.latents_std.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + z = (z - mean) / std + return z.reshape(b, s, z.shape[1], z.shape[2]).permute(0, 2, 1, 3) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py new file mode 100644 index 000000000..494350d40 --- /dev/null +++ b/comfy/ldm/minimax/model.py @@ -0,0 +1,646 @@ +"""MiniMax H3 audio-video DiT. + +Single-stream packed-token transformer denoising video (24ch, patch 1x2x2) and +stereo audio (32ch, 40 Hz) latents jointly, conditioned on Qwen3-VL layer-50 hidden states. +The packed sequence is: +[text | cond rows | audio | video] for t2va/fl2va +[text | reference blocks | audio | video] for ref2va + +Timestep domain: the model receives the *video* sigma from the sampler and +derives per-token timesteps t = 1 - sigma internally; the audio stream runs on +its own shifted schedule (sigma_shift video 12.0 / audio 3.0), mapped from the +video sigma in closed form. The audio velocity is returned scaled by the +schedule map's derivative d(sigma_a)/d(sigma_v). +""" + +import math + +import torch +import torch.nn as nn + +import comfy.ldm.common_dit +import comfy.model_management +import comfy.model_prefetch +import comfy.ops +import comfy.patcher_extension +import comfy.quant_ops +from comfy.ldm.modules.attention import optimized_attention + +FRAME_PER_TOKEN = (1, 4, 4, 4, 4) +FRAME_RESCALE = 5.0 / 3.0 +VISUAL_COND_TIMESTEP = 0.999 +AUDIO_COND_TIMESTEP = 1.0 + + +def time_shift_sigma(sigma, from_shift, to_shift): + # invert sigma = s*b/(1+(s-1)*b) to the base grid, re-apply the other shift + base = sigma / (from_shift + sigma * (1.0 - from_shift)) + return to_shift * base / (1.0 + (to_shift - 1.0) * base) + + +def time_shift_slope(sigma, from_shift, to_shift): + """d(sigma_to)/d(sigma_from) at the same base-grid point. + + Scaling a stream's returned velocity by this slope makes the flat ODE that + any sampler integrates on the from-schedule equal to that stream's true ODE + on its own schedule. + """ + base = sigma / (from_shift + sigma * (1.0 - from_shift)) + return (to_shift * (1.0 + (from_shift - 1.0) * base) ** 2) / (from_shift * (1.0 + (to_shift - 1.0) * base) ** 2) + + +def patchify_video(latent, patch_size=(1, 2, 2)): + # [B, C, T, H, W] -> [B*t*h*w, C*pt*ph*pw] + b, c, t_full, h_full, w_full = latent.shape + pt, ph, pw = patch_size + t, h, w = t_full // pt, h_full // ph, w_full // pw + x = latent.reshape(b, c, t, pt, h, ph, w, pw) + x = torch.einsum("nctrhpwq->nthwcrpq", x) + return x.reshape(b * t * h * w, c * pt * ph * pw) + + +def unpatchify_video(rows, t, h, w, c=24, patch_size=(1, 2, 2)): + pt, ph, pw = patch_size + x = rows.reshape(-1, t, h, w, c, pt, ph, pw) + x = torch.einsum("nthwcrpq->nctrhpwq", x) + return x.reshape(-1, c, t * pt, h * ph, w * pw) + + +def pack_audio(latent): + # [B, C=32, ch=2, T] -> [ch*T, 32] channel-major (ch0 t0..T-1, ch1 t0..T-1) + b, c, ch, t = latent.shape + return latent[0].permute(1, 2, 0).reshape(ch * t, c) + + +def unpack_audio(rows, ch=2): + t = rows.shape[0] // ch + return rows.reshape(ch, t, rows.shape[-1]).permute(2, 0, 1).unsqueeze(0) + + +def _axis_from_sqrt_area(dim, patch, sqrt_area): + # linspace((1 - ratio) / 2, (1 + ratio) / 2, dim // patch, endpoint=False) * 32 + ratio = dim / sqrt_area + n = dim // patch + return (torch.arange(n, dtype=torch.float64) * (ratio / n) + (1.0 - ratio) / 2.0) * 32.0 + + +def _frame_grid(h, w): + # area-normalized (h, w) coordinates of one latent frame's 2x2-patch rows + area = math.sqrt(h * w) + hh, ww = torch.meshgrid(_axis_from_sqrt_area(h, 2, area), _axis_from_sqrt_area(w, 2, area), indexing="ij") + return torch.stack([hh.reshape(-1), ww.reshape(-1)], dim=-1), _axis_from_sqrt_area(w, 2, area) + + +def _video_t_spans(n): + return [FRAME_RESCALE * FRAME_PER_TOKEN[k % 5] for k in range(n)] + + +def _video_t_grid(n, origin): + # origin + exclusive cumsum + spans = torch.tensor(_video_t_spans(n), dtype=torch.float64) + return float(origin) + torch.cat([torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(0)]) + + +def _audio_grid(cursor, t, w_low, w_high): + # channel-major stereo rows: t advances per latent frame, w pinned to the grid extremes per stereo channel, h stays 0 + g = torch.zeros(t * 2, 3, dtype=torch.float64) + g[:, 0] = (cursor + torch.arange(t, dtype=torch.float64)).repeat(2) + g[:t, 2] = w_low + g[t:, 2] = w_high + return g + + +def _video_grid(vt, frame, cursor): + g = torch.empty(vt, frame.shape[0], 3, dtype=torch.float64) + g[:, :, 0] = _video_t_grid(vt, cursor)[:, None] + g[:, :, 1:] = frame[None] + return g.reshape(-1, 3) + + +class TimeEmbedder(nn.Module): + def __init__(self, freq_dim, hidden, out, dtype=None, device=None, operations=None): + super().__init__() + self.freq_dim = freq_dim + self.proj_in = operations.Linear(freq_dim, hidden, bias=True, dtype=dtype, device=device) + self.proj_out = operations.Linear(hidden, out, bias=True, dtype=dtype, device=device) + + def forward(self, t): + # t: [M] in [0, 1]; fp32 throughout, cos before sin + half = self.freq_dim // 2 + freqs = torch.exp(-math.log(10000.0) * torch.arange(half, dtype=torch.float32, device=t.device) / half) + args = t.to(torch.float32)[:, None] * freqs[None] + emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + return self.proj_out(nn.functional.silu(self.proj_in(emb))) + + +def rope_rotation_table(angles, dtype): + """[S, rot_dim] pair angles -> [1, S, 1, rot_dim/2, 2, 2] rotation matrices.""" + half = angles.shape[-1] // 2 + ang = angles[:, :half] # duplicated halves: [:, :half] == [:, half:] + c, s = torch.cos(ang), torch.sin(ang) + table = torch.stack([c, -s, s, c], dim=-1).reshape(1, angles.shape[0], 1, half, 2, 2) + return table.to(dtype) + + +class Attention(nn.Module): + def __init__(self, hidden, heads, head_dim, eps, dtype=None, device=None, operations=None): + super().__init__() + self.heads = heads + self.head_dim = head_dim + inner = heads * head_dim + self.qkv_proj = operations.Linear(hidden, inner * 3, bias=False, dtype=dtype, device=device) + self.q_norm = operations.RMSNorm(head_dim, eps=eps, dtype=dtype, device=device) + self.k_norm = operations.RMSNorm(head_dim, eps=eps, dtype=dtype, device=device) + self.out_proj = operations.Linear(inner, hidden, bias=False, dtype=dtype, device=device) + + def forward(self, x, rope_freqs=None, transformer_options={}): + s = x.shape[0] + q, k, v = self.qkv_proj(x).split(self.heads * self.head_dim, dim=-1) + v = v.view(s, self.heads, self.head_dim) + if rope_freqs is not None: + # fused per-head RMSNorm + partial split-half rope, in place on the qkv buffer + q = q.view(1, s, self.heads, self.head_dim) + k = k.view(1, s, self.heads, self.head_dim) + qw = comfy.model_management.cast_to(self.q_norm.weight, device=x.device) + kw = comfy.model_management.cast_to(self.k_norm.weight, device=x.device) + rot = rope_freqs.shape[-3] * 2 + if comfy.model_management.in_training: + q, k = comfy.quant_ops.ck.rms_rope_split_half( + q, k, rope_freqs, qw, kw, epsilon=self.q_norm.eps, rot_dim=rot) + else: + comfy.quant_ops.ck.rms_rope_split_half_( + q, k, rope_freqs, qw, kw, epsilon=self.q_norm.eps, rot_dim=rot) + q = q[0] + k = k[0] + else: + q = self.q_norm(q.view(s, self.heads, self.head_dim)) + k = self.k_norm(k.view(s, self.heads, self.head_dim)) + q = q.transpose(0, 1).unsqueeze(0) + k = k.transpose(0, 1).unsqueeze(0) + v = v.transpose(0, 1).unsqueeze(0) + out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) + return self.out_proj(out.squeeze(0)) + + +class MLP(nn.Module): + def __init__(self, hidden, ffn, dtype=None, device=None, operations=None): + super().__init__() + self.fc1 = operations.Linear(hidden, ffn * 2, bias=False, dtype=dtype, device=device) + self.fc2 = operations.Linear(ffn, hidden, bias=False, dtype=dtype, device=device) + + def forward(self, x): + return comfy.ops.linear_input_act(self.fc2, self.fc1(x), "swiglu") + + +class AdalnProj(nn.Module): + def __init__(self, t_dim, hidden, expand, modalities, apply_silu=True, + dtype=None, device=None, operations=None): + super().__init__() + self.expand = expand + self.modalities = modalities + self.hidden = hidden + self.apply_silu = apply_silu + self.linear = operations.Linear(t_dim, expand * hidden * modalities, bias=True, dtype=dtype, device=device) + + def forward(self, t_emb): + # [M, t_dim] -> expand tensors of [M*modalities, hidden] + x = self.linear(nn.functional.silu(t_emb) if self.apply_silu else t_emb) + x = x.view(x.shape[0] * self.modalities, self.expand * self.hidden) + return x.chunk(self.expand, dim=-1) + + +def _mod_scale_shift(h, shift, scale, segments): + # segments: [(start, stop, mod_row)] covering h contiguously. + for a, b, row in segments: + h[a:b].mul_(1.0 + scale[row].to(h.dtype)).add_(shift[row].to(h.dtype)) + return h + + +def _mod_gate(x, gate, other, segments): + # other is the fresh attn/mlp output: accumulate the gated residual into the stream in place, one fused kernel per segment + for a, b, row in segments: + x[a:b].addcmul_(other[a:b], gate[row].to(x.dtype)) + return x + + +class RefinerBlock(nn.Module): + def __init__(self, hidden, heads, head_dim, ffn, eps, qk_eps, dtype=None, device=None, operations=None): + super().__init__() + self.norm1 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.norm2 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.attn = Attention(hidden, heads, head_dim, qk_eps, dtype=dtype, device=device, operations=operations) + self.mlp = MLP(hidden, ffn, dtype=dtype, device=device, operations=operations) + + def forward(self, x, transformer_options={}): + # attn/mlp outputs are fresh: accumulate residuals in place + x = self.attn(self.norm1(x), transformer_options=transformer_options).add_(x) + return self.mlp(self.norm2(x)).add_(x) + + +class TokenRefiner(nn.Module): + def __init__(self, num_layers, hidden, heads, head_dim, ffn, eps, qk_eps, final_eps, + dtype=None, device=None, operations=None): + super().__init__() + self.blocks = nn.ModuleList([ + RefinerBlock(hidden, heads, head_dim, ffn, eps, qk_eps, dtype=dtype, device=device, operations=operations) + for _ in range(num_layers)]) + self.final_norm = operations.RMSNorm(hidden, eps=final_eps, dtype=dtype, device=device) + + def forward(self, x, transformer_options={}): + for block in self.blocks: + x = block(x, transformer_options=transformer_options) + return self.final_norm(x) + + +class DiTBlock(nn.Module): + def __init__(self, hidden, heads, head_dim, ffn, t_dim, eps, qk_eps, + apply_silu=True, adaln_dtype=None, dtype=None, device=None, operations=None): + super().__init__() + self.norm1 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.norm2 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.attn = Attention(hidden, heads, head_dim, qk_eps, dtype=dtype, device=device, operations=operations) + self.mlp = MLP(hidden, ffn, dtype=dtype, device=device, operations=operations) + self.adaln_proj = AdalnProj(t_dim, hidden, 6, 3, apply_silu=apply_silu, + dtype=adaln_dtype if adaln_dtype is not None else dtype, + device=device, operations=operations) + + def forward(self, x, t_emb, mod_segments, rope_freqs, transformer_options={}): + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaln_proj(t_emb) + h = _mod_scale_shift(self.norm1(x), shift_msa, scale_msa, mod_segments) + x = _mod_gate(x, gate_msa, self.attn(h, rope_freqs=rope_freqs, transformer_options=transformer_options), mod_segments) + h = _mod_scale_shift(self.norm2(x), shift_mlp, scale_mlp, mod_segments) + return _mod_gate(x, gate_mlp, self.mlp(h), mod_segments) + + +class FinalLayer(nn.Module): + def __init__(self, hidden, t_dim, video_dim, audio_dim, eps, apply_silu=True, adaln_dtype=None, + dtype=None, device=None, operations=None): + super().__init__() + self.norm = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.adaln_proj = AdalnProj(t_dim, hidden, 2, 1, apply_silu=apply_silu, + dtype=adaln_dtype if adaln_dtype is not None else dtype, + device=device, operations=operations) + # output heads are the checkpoint's fp32 island; norm/adaln are stored at model dtype + self.video_out = operations.Linear(hidden, video_dim, bias=True, dtype=torch.float32, device=device) + self.audio_out = operations.Linear(hidden, audio_dim, bias=True, dtype=torch.float32, device=device) + + def forward(self, x, t_emb, video_seg, audio_seg): + # video_seg / audio_seg: (start, stop, timestep_row) of the target streams + shift, scale = self.adaln_proj(t_emb) + va, vb, vrow = video_seg + aa, ab, arow = audio_seg + hv = (self.norm(x[va:vb]) * (1.0 + scale[vrow]) + shift[vrow]).to(torch.float32) + ha = (self.norm(x[aa:ab]) * (1.0 + scale[arow]) + shift[arow]).to(torch.float32) + return self.video_out(hv), self.audio_out(ha) + + +class PackedLayout: + """Static packed-sequence structure for one shape/conditioning signature.""" + + def __init__(self, text_len, latent_t, latent_h, latent_w, audio_t, keyframes=None, refs=None, frame_count=None): + frame, w_grid = _frame_grid(latent_h, latent_w) + frame_rows = frame.shape[0] + + segments = [("text", text_len)] # (kind, n_rows) + g = torch.zeros(text_len, 3, dtype=torch.float64) + g[:, 0] = torch.arange(text_len, dtype=torch.float64) + pos = [g] # per segment: [n, 3] float64 (t, h, w) + + img_pos, img_update = [], [] + audio_pos, audio_update = [], [] + cursor = text_len + row = text_len + + if keyframes: + # fl2va: keyframe cond rows right after text, sharing the target spatial grid + for kf in keyframes: + pixel_index = kf["resolved_frame_index"] + if pixel_index == 0: + cond_t = float(text_len) + elif frame_count is not None and pixel_index == frame_count - 1: + cond_t = float(text_len) + sum(_video_t_spans(latent_t)) - FRAME_RESCALE + else: + raise ValueError("only first/last keyframe anchors are supported") + g = torch.empty(frame_rows, 3, dtype=torch.float64) + g[:, 0] = cond_t + g[:, 1:] = frame + segments.append(("cond", frame_rows)) + pos.append(g) + img_pos.append(torch.arange(row, row + frame_rows)) + img_update.append(torch.zeros(frame_rows, dtype=torch.bool)) + row += frame_rows + + target_audio_w = (float(w_grid[0]), float(w_grid[-1])) + if refs: + cursor = float(text_len) + for blk in refs: + kind = blk["kind"] + if kind == "image": + r_frame, _ = _frame_grid(blk["latent_h"], blk["latent_w"]) + n = r_frame.shape[0] + g = torch.empty(n, 3, dtype=torch.float64) + g[:, 0] = cursor + g[:, 1:] = r_frame + segments.append(("ref_img", n)) + pos.append(g) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + cursor += 1.0 + elif kind == "audio": + rt = blk["ref_audio_t"] + if rt > 0: + segments.append(("ref_audio", rt * 2)) + pos.append(_audio_grid(cursor, rt, *target_audio_w)) + audio_pos.append(torch.arange(row, row + rt * 2)) + audio_update.append(torch.zeros(rt * 2, dtype=torch.bool)) + row += rt * 2 + cursor += float(rt) + elif kind in ("video", "video_audio"): + # the block's audio rows pack immediately before its video + # rows, both sharing the cursor origin + rt = blk["ref_audio_t"] + vt = blk["latent_t"] + r_frame, r_w_grid = _frame_grid(blk["latent_h"], blk["latent_w"]) + if rt > 0: + segments.append(("ref_audio", rt * 2)) + pos.append(_audio_grid(cursor, rt, float(r_w_grid[0]), float(r_w_grid[-1]))) + audio_pos.append(torch.arange(row, row + rt * 2)) + audio_update.append(torch.zeros(rt * 2, dtype=torch.bool)) + row += rt * 2 + n = vt * r_frame.shape[0] + segments.append(("ref_img", n)) + pos.append(_video_grid(vt, r_frame, cursor)) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + cursor += max(float(rt), sum(_video_t_spans(vt))) + + # target audio then target video, always the last two segments + segments.append(("audio", audio_t * 2)) + pos.append(_audio_grid(cursor, audio_t, *target_audio_w)) + audio_pos.append(torch.arange(row, row + audio_t * 2)) + audio_update.append(torch.ones(audio_t * 2, dtype=torch.bool)) + row += audio_t * 2 + + n_video = latent_t * frame_rows + segments.append(("video", n_video)) + pos.append(_video_grid(latent_t, frame, cursor)) + img_pos.append(torch.arange(row, row + n_video)) + img_update.append(torch.ones(n_video, dtype=torch.bool)) + row += n_video + + self.seq_len = row + self.position_ids = torch.cat(pos) # [S, 3] float64 + self.img_pos = torch.cat(img_pos) + self.img_update = torch.cat(img_update) + self.audio_pos = torch.cat(audio_pos) + self.audio_update = torch.cat(audio_update) + self.signature = (text_len, latent_t, latent_h, latent_w, audio_t) + # contiguous segment table (start, stop, kind) + # kinds: text / cond / ref_img / ref_audio / audio / video + # the packed sequence is uniform per segment in (modality tag, timestep class), + # except the text span (tag runs resolved at forward time from the presentation tags) + seg_abs = [] + off = 0 + for kind, n in segments: + seg_abs.append((off, off + n, kind)) + off += n + self.segments = seg_abs + + +class MiniMaxH3Model(nn.Module): + def __init__(self, hidden_size=5376, num_layers=50, token_refiner_num_layers=2, + num_attention_heads=56, attention_head_dim=128, ffn_hidden_size=14336, + latents_dim=24, audio_latents_dim=32, patch_size=(1, 2, 2), text_dim=5120, + timestep_input_dim=256, time_embed_hidden_size=5376, time_embed_dim=2688, + rope_inv_freq_len=16, norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5, + sigma_shift_video=12.0, sigma_shift_audio=3.0, + adaln_curve_grid=None, + image_model=None, dtype=None, device=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.hidden_size = hidden_size + self.patch_size = tuple(patch_size) + self.latents_dim = latents_dim + self.audio_latents_dim = audio_latents_dim + self.sigma_shift_video = sigma_shift_video + self.sigma_shift_audio = sigma_shift_audio + self.use_adaln_curves = adaln_curve_grid is not None + # curve-form checkpoints replace the time embedder and full-width adaln weights with a small shared basis of the time-embedding curve + curve = {"apply_silu": not self.use_adaln_curves, + "adaln_dtype": torch.float32 if self.use_adaln_curves else dtype} + video_patch_dim = latents_dim * self.patch_size[0] * self.patch_size[1] * self.patch_size[2] + + self.video_patch_proj = operations.Linear(video_patch_dim, hidden_size, bias=True, dtype=torch.float32, device=device) + self.audio_patch_proj = operations.Linear(audio_latents_dim, hidden_size, bias=True, dtype=torch.float32, device=device) + self.condition_proj = operations.Linear(text_dim, hidden_size, bias=True, dtype=dtype, device=device) + if self.use_adaln_curves: + self.register_buffer("adaln_t_table", torch.empty(adaln_curve_grid, time_embed_dim, dtype=torch.float32)) + else: + self.time_embedder = TimeEmbedder(timestep_input_dim, time_embed_hidden_size, time_embed_dim, + dtype=torch.float32, device=device, operations=operations) + self.rope = nn.Module() + self.rope.register_buffer("inv_freq", torch.empty(rope_inv_freq_len, dtype=torch.float32)) + self.token_refiner = TokenRefiner(token_refiner_num_layers, hidden_size, num_attention_heads, + attention_head_dim, ffn_hidden_size, norm_eps, qk_norm_eps, + final_norm_eps, dtype=dtype, device=device, operations=operations) + self.blocks = nn.ModuleList([ + DiTBlock(hidden_size, num_attention_heads, attention_head_dim, ffn_hidden_size, + time_embed_dim, norm_eps, qk_norm_eps, **curve, dtype=dtype, device=device, operations=operations) + for _ in range(num_layers)]) + self.final_layer = FinalLayer(hidden_size, time_embed_dim, video_patch_dim, audio_latents_dim, + final_norm_eps, **curve, dtype=dtype, device=device, operations=operations) + + def preprocess_text_embeds(self, text_states): + """[B, L, text_dim] Qwen states -> [B, L, hidden] refined text embeds.""" + if text_states.shape[-1] == self.hidden_size: + return text_states + return self.token_refiner(self.condition_proj(text_states[0])).unsqueeze(0) + + def rope_freqs(self, position_ids, device): + # [S, 3] float64 -> [S, 96] fp32 + pos = position_ids.to(torch.float32).to(device) + inv = comfy.model_management.cast_to(self.rope.inv_freq, device=device) + per_axis = pos.unsqueeze(-1) * inv.view(1, 1, -1) # [S, 3, 16] + t_f, h_f, w_f = per_axis.unbind(dim=1) + half = torch.cat((t_f, h_f, w_f), dim=-1) # [S, 48] + return torch.cat((half, half), dim=-1) # [S, 96] + + def _cond_video_rows(self, payload, device): + """Concatenated visual condition rows (normalized latents -> patchified), with condition noise augmentation.""" + rows = [] + aug = payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP) + seed = int(payload.get("seed", 0)) + # every condition intentionally restarts the same RNG stream + for z in payload.get("cond_video_latents", []): + r = patchify_video(z.to(torch.float32), self.patch_size) + if aug < 1.0: + gen = torch.Generator("cpu").manual_seed(seed) + noise = torch.randn(r.shape, generator=gen, dtype=torch.float32) + r = aug * r + (1.0 - aug) * noise.to(r.device) + rows.append(r.to(device)) + return torch.cat(rows, dim=0) if rows else None + + def _cond_audio_rows(self, payload, device): + rows = [] + aug = payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP) + seed = int(payload.get("seed", 0)) + 1 + for z in payload.get("cond_audio_latents", []): + r = pack_audio(z.to(torch.float32)) + if aug < 1.0: + gen = torch.Generator("cpu").manual_seed(seed) + noise = torch.randn(r.shape, generator=gen, dtype=torch.float32) + r = aug * r + (1.0 - aug) * noise.to(r.device) + rows.append(r.to(device)) + return torch.cat(rows, dim=0) if rows else None + + def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload, **kwargs) + + def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs): + video_x, audio_x = x[0], x[1] + orig_t, orig_h, orig_w = video_x.shape[2], video_x.shape[3], video_x.shape[4] + video_x = comfy.ldm.common_dit.pad_to_patch_size(video_x, self.patch_size) + if video_x.shape[0] != 1: + raise ValueError("MiniMax H3 supports batch size 1") + payload = minimax_payload or {} + device = video_x.device + dtype = context.dtype # compute dtype + + latent_t, lat_h, lat_w = video_x.shape[2], video_x.shape[3], video_x.shape[4] + audio_t = audio_x.shape[-1] + text_len = context.shape[1] + # extra_conds prebuilds the layout once per sampling run + layout = payload.get("layout") + if layout is None or layout.signature != (text_len, latent_t, lat_h, lat_w, audio_t): + layout = PackedLayout(text_len, latent_t, lat_h, lat_w, audio_t, + keyframes=payload.get("keyframes"), + refs=payload.get("refs"), + frame_count=payload.get("frame_count")) + + # model_base passes model_sampling.timestep(sigma) = sigma * 1000 + shift_v = float(transformer_options.get("minimax_h3_sigma_shift_video", self.sigma_shift_video)) + shift_a = float(transformer_options.get("minimax_h3_sigma_shift_audio", self.sigma_shift_audio)) + sigma_v = (timestep.flatten()[0] / 1000.0).float().clamp(min=1e-6) + t_v = float(1.0 - sigma_v) + t_a = float(1.0 - time_shift_sigma(sigma_v, shift_v, shift_a)) + + # distinct timesteps are known analytically: text/pad follow video, cond rows pin near 1 + vis_aug = float(payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP)) + aud_aug = float(payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP)) + has_vis_cond = any(k in ("cond", "ref_img") for _, _, k in layout.segments) + has_aud_cond = any(k == "ref_audio" for _, _, k in layout.segments) + seg_t = {"text": t_v, "video": t_v, "audio": t_a, + "cond": max(t_v, vis_aug), "ref_img": max(t_v, vis_aug), + "ref_audio": max(t_a, aud_aug)} + unique_t = sorted({t_v, t_a} | ({seg_t["cond"]} if has_vis_cond else set()) + | ({seg_t["ref_audio"]} if has_aud_cond else set())) + t_row = {t: i for i, t in enumerate(unique_t)} + seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "ref_audio": 2} + + text_tags = payload.get("text_token_tags") + mod_segments = [] + for a, b, kind in layout.segments: + row_base = t_row[seg_t[kind]] * 3 + if kind == "text" and text_tags is not None: + # the presentation text span mixes tags (vision pads carry the video modality) split into tag runs + tags = text_tags.view(-1).tolist() + run_start = 0 + for i in range(1, b - a + 1): + if i == b - a or tags[i] != tags[run_start]: + mod_segments.append((a + run_start, a + i, row_base + int(tags[run_start]))) + run_start = i + else: + mod_segments.append((a, b, row_base + seg_tag[kind])) + + # embed + img_update = layout.img_update.to(device) + audio_update = layout.audio_update.to(device) + video_rows = patchify_video(video_x.to(torch.float32), self.patch_size) + audio_rows = pack_audio(audio_x.to(torch.float32)) + cond_video_rows = self._cond_video_rows(payload, device) + cond_audio_rows = self._cond_audio_rows(payload, device) + + all_video_rows = video_rows + if cond_video_rows is not None: + all_video_rows = torch.empty(img_update.shape[0], video_rows.shape[1], dtype=torch.float32, device=device) + all_video_rows[~img_update] = cond_video_rows + all_video_rows[img_update] = video_rows + all_audio_rows = audio_rows + if cond_audio_rows is not None: + all_audio_rows = torch.empty(audio_update.shape[0], audio_rows.shape[1], dtype=torch.float32, device=device) + all_audio_rows[~audio_update] = cond_audio_rows + all_audio_rows[audio_update] = audio_rows + + video_embed = self.video_patch_proj(all_video_rows).to(dtype) + audio_embed = self.audio_patch_proj(all_audio_rows).to(dtype) + text_states = context[0] + if text_states.shape[-1] != self.hidden_size: + text_states = self.token_refiner(self.condition_proj(text_states), + transformer_options=transformer_options) + + # segments are contiguous: assemble by slices, embed rows follow segment order + h = torch.empty(layout.seq_len, self.hidden_size, dtype=dtype, device=device) + voff = aoff = 0 + for a, b, kind in layout.segments: + n = b - a + if kind == "text": + h[a:b] = text_states + elif kind in ("cond", "ref_img", "video"): + h[a:b] = video_embed[voff:voff + n] + voff += n + else: # ref_audio / audio + h[a:b] = audio_embed[aoff:aoff + n] + aoff += n + + t_vals = torch.tensor(unique_t, dtype=torch.float32, device=device) + if self.use_adaln_curves: + # adaln projections consume interpolated coordinates of the time-embedding curve + table = comfy.model_management.cast_to(self.adaln_t_table, device=device) + pos = t_vals.clamp(0.0, 1.0) * (table.shape[0] - 1) # t in [0,1] -> fractional grid index, out-of-range t clamps to the curve ends + i0 = pos.floor().long().clamp(max=table.shape[0] - 2) # lower grid row, max-clamp keeps t=1.0 on the last interval instead of reading past the table + t_emb = torch.lerp(table[i0], table[i0 + 1], (pos - i0).unsqueeze(1)) # blend the two rows by the fractional part + else: + t_emb = self.time_embedder(t_vals).to(dtype) + + # rotation table computed once per forward, consumed by the kitchen split-half rope + rope_freqs = rope_rotation_table(self.rope_freqs(layout.position_ids, device), dtype) + + # blocks + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.blocks), device, transformer_options) + for i, block in enumerate(self.blocks): + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block) + if ("double_block", i) in blocks_replace: + def block_wrap(args): + return {"img": block(args["img"], args["t_emb"], args["mod_segments"], args["rope_freqs"], + transformer_options=args["transformer_options"])} + h = blocks_replace[("double_block", i)]( + {"img": h, "t_emb": t_emb, "mod_segments": mod_segments, "rope_freqs": rope_freqs, + "transformer_options": transformer_options}, + {"original_block": block_wrap})["img"] + else: + h = block(h, t_emb, mod_segments, rope_freqs, transformer_options=transformer_options) + if prefetch_queue is not None: + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, None) + + # target streams are single contiguous segments (audio then video, last two) + video_seg = next((a, b, t_row[seg_t["video"]]) for a, b, k in layout.segments if k == "video") + audio_seg = next((a, b, t_row[seg_t["audio"]]) for a, b, k in layout.segments if k == "audio") + v, a = self.final_layer(h, t_emb, video_seg, audio_seg) + + video_out = unpatchify_video(v, latent_t, lat_h // 2, lat_w // 2, self.latents_dim, self.patch_size) + video_out = video_out[:, :, :orig_t, :orig_h, :orig_w] + audio_out = unpack_audio(a) + + # The sampler integrates the flat ODE dX/dsigma_v = (X - denoised)/sigma_v. + # Scaling the audio velocity by d(sigma_a)/d(sigma_v) makes that ODE equal + # to the audio stream's true ODE on its own shifted schedule. + slope_a = time_shift_slope(sigma_v, shift_v, shift_a).to(audio_out.dtype) + return [-video_out.to(video_x.dtype), (-slope_a) * audio_out.to(audio_x.dtype)] diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py new file mode 100644 index 000000000..aeb3421a2 --- /dev/null +++ b/comfy/ldm/minimax/vae.py @@ -0,0 +1,694 @@ +# MiniMax H3 video VAE: 3D causal CNN encoder + ViT3D decoder. + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops +import comfy.quant_ops +import comfy.rmsnorm +from comfy.ldm.modules.attention import optimized_attention + +ops = comfy.ops.disable_weight_init + +IMAGENET_MEAN = (0.485, 0.456, 0.406) +IMAGENET_STD = (0.229, 0.224, 0.225) + +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.68317747116088865, 1.7549455165863037, + 1.5636216402053833, 2.194143533706665, 0.96531379222869875, 1.05698859691619875, + 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, + 0.7996809482574463, 0.44988900423049925, 0.7197399735450745, 0.69362932443618775, + 2.961095094680786, 2.7694199085235595, 3.0496184825897215, 2.1088054180145265, + 3.276226282119751, 3.1627357006073, 2.28168129920959475, 2.6127843856811525, +] + + +# 3D causal CNN encoder + +class CausalConv3d(ops.Conv3d): + # Reflect spatial padding, causal (zeros, front-only) temporal padding. + def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0): + super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride) + self.causal_padding = (padding,) * 3 if isinstance(padding, int) else tuple(padding) + + def forward(self, x): + if sum(self.causal_padding) == 0: + return super().forward(x) + + x = F.pad(x, (self.causal_padding[2], self.causal_padding[2], self.causal_padding[1], self.causal_padding[1], 0, 0), mode="reflect") + if x.shape[2] == 1: + # single frame: the causal front padding is all zeros truncate the temporal taps instead of convolving zero frames + return super().forward(x, autopad="causal_zero") + x = F.pad(x, (0, 0, 0, 0, self.causal_padding[0] * 2, 0), mode="constant") + return super().forward(x) + + +class TemporalIsolatedGroupNorm(ops.GroupNorm): + # GroupNorm with statistics computed per frame (time merged into batch). + def forward(self, x): + if x.dim() == 5: + b, c, t, h, w = x.shape + x = x.permute(0, 2, 1, 3, 4).contiguous().view(b * t, c, 1, h, w) + x = super().forward(x) + return x.view(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + return super().forward(x) + + +def group_norm_3d(num_channels): + return TemporalIsolatedGroupNorm(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True) + + +class Downsample3D(nn.Module): + def __init__(self, in_channels, out_channels, time_stride=1, space_stride=2): + super().__init__() + self.space_stride = space_stride + self.conv = CausalConv3d( + in_channels, + out_channels, + kernel_size=3, + padding=(1, 0, 0), + stride=(time_stride, space_stride, space_stride), + ) + + def forward(self, x): + if self.space_stride == 2: + x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect") + return self.conv(x) + + +class ResnetBlock3D(nn.Module): + def __init__(self, in_channels, out_channels=None): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + + self.norm1 = group_norm_3d(in_channels) + self.norm2 = group_norm_3d(out_channels) + self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, padding=1) + self.conv2 = CausalConv3d(out_channels, out_channels, kernel_size=3, padding=1) + if in_channels != out_channels: + self.nin_shortcut = CausalConv3d(in_channels, out_channels, kernel_size=1) + + def forward(self, x): + h = self.conv1(F.silu(self.norm1(x), inplace=True)) + h = self.conv2(F.silu(self.norm2(h), inplace=True)) + if self.in_channels != self.out_channels: + x = self.nin_shortcut(x) + return h.add_(x) + + +class EncoderFCN3D(nn.Module): + def __init__(self, ch, ch_mult, space_down, time_down, num_res_blocks, in_channels, z_channels, double_z=True): + super().__init__() + self.num_levels = len(ch_mult) + if isinstance(num_res_blocks, int): + num_res_blocks = [num_res_blocks] * self.num_levels + self.num_res_blocks = num_res_blocks + + block_mid = [ch * ch_mult[i] for i in range(self.num_levels)] + block_in = [block_mid[0]] + block_mid[:-1] + block_out = block_mid + + self.conv_in = CausalConv3d(in_channels, block_in[0], kernel_size=3, padding=1) + + self.down = nn.ModuleList() + for i_level in range(self.num_levels): + down = nn.Module() + down.block = nn.ModuleList() + for i in range(self.num_res_blocks[i_level]): + down.block.append( + ResnetBlock3D( + in_channels=block_in[i_level] if i == 0 else block_mid[i_level], + out_channels=block_mid[i_level], + ) + ) + if space_down[i_level] * time_down[i_level] > 1: + down.downsample = Downsample3D( + block_mid[i_level], + block_out[i_level], + time_stride=time_down[i_level], + space_stride=space_down[i_level], + ) + self.down.append(down) + + self.norm_out = group_norm_3d(block_out[-1]) + self.conv_out = CausalConv3d( + block_out[-1], + 2 * z_channels if double_z else z_channels, + kernel_size=3, + padding=1, + ) + + def forward(self, x): + h = self.conv_in(x) + for i_level in range(self.num_levels): + for i_block in range(self.num_res_blocks[i_level]): + h = self.down[i_level].block[i_block](h) + if hasattr(self.down[i_level], "downsample"): + h = self.down[i_level].downsample(h) + h = F.silu(self.norm_out(h)) + return self.conv_out(h) + + +# ViT3D decoder + +def create_token_ids(patch_dims, device, dtype): + coords_list = [] + for dim_size in patch_dims: + coords = torch.arange(0.5, dim_size, dtype=dtype, device=device) + coords = coords / dim_size + coords = 2.0 * coords - 1.0 + coords_list.append(coords) + coords = torch.stack(torch.meshgrid(*coords_list, indexing="ij"), dim=-1) + return coords.flatten(0, len(patch_dims) - 1).unsqueeze(0) + + +class RotaryEmbeddingND(nn.Module): + def __init__(self, dim, rotary_base=100.0, n_dim=3): + super().__init__() + self.n_dim = n_dim + self.angle_scale = 2.0 * math.pi + inv_freq = 1 / rotary_base ** torch.arange(0, 1, 2 * n_dim / dim, dtype=torch.float32) + self.register_buffer("inv_freq", inv_freq, persistent=False) + + def forward(self, img_ids): + # [B, S, n_dim] -> [B, S, 1, pairs, 2, 2] rotation table for the kitchen split-half rope + angles = ( + self.angle_scale + * img_ids[:, :, :, None].float() + * self.inv_freq.to(img_ids.device)[None, None, None, :] + ) + angles = angles.flatten(2, 3) + c, s = torch.cos(angles), torch.sin(angles) + table = torch.stack([c, -s, s, c], dim=-1).reshape(*angles.shape[:2], 1, angles.shape[-1], 2, 2) + return table.to(img_ids.dtype) + + +class FeedForward(nn.Module): + # Gated SiLU FFN. + def __init__(self, dim, mult=4, bias=True): + super().__init__() + inner_dim = dim * mult + self.w1 = ops.Linear(dim, inner_dim * 2, bias=bias) + self.w2 = ops.Linear(inner_dim, dim, bias=bias) + + def forward(self, x): + gate, x = self.w1(x).chunk(2, dim=-1) + return self.w2(F.silu(gate).mul_(x)) + + +class Attention(nn.Module): + def __init__(self, heads, dim_head, bias=True, eps=1e-5): + super().__init__() + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + self.norm_q = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) + self.norm_k = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) + self.to_qkv = ops.Linear(inner_dim, inner_dim * 3, bias=bias) + self.to_out = ops.Linear(inner_dim, inner_dim, bias=bias) + + def forward(self, x, rotary_pos_emb=None): + batch_size, seq_len, _ = x.shape + + qkv = self.to_qkv(x) + qkv = qkv.view(batch_size, seq_len, -1, 3 * self.dim_head) + query, key, value = torch.chunk(qkv, 3, dim=-1) + + query = comfy.rmsnorm.rms_norm(query, self.norm_q.weight, self.norm_q.eps) + key = comfy.rmsnorm.rms_norm(key, self.norm_k.weight, self.norm_k.eps) + + if rotary_pos_emb is not None: + rot = rotary_pos_emb.shape[-3] * 2 + query[..., :rot], key[..., :rot] = comfy.quant_ops.ck.apply_rope_split_half( + query[..., :rot], key[..., :rot], rotary_pos_emb) + + out = optimized_attention(query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), + self.heads, skip_reshape=True).nan_to_num_(0.0) + return self.to_out(out) + + +class TransformerBlock(nn.Module): + def __init__(self, heads, dim_head, bias=True, eps=1e-5): + super().__init__() + dim = heads * dim_head + self.norm1 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) + self.attn = Attention(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + self.scale1 = nn.Parameter(torch.empty(dim)) + self.norm2 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) + self.ff = FeedForward(dim=dim, bias=bias) + self.scale2 = nn.Parameter(torch.empty(dim)) + + def forward(self, x, rotary_pos_emb=None): + x = x.addcmul_(self.attn(comfy.rmsnorm.rms_norm(x, self.norm1.weight, self.norm1.eps), rotary_pos_emb), self.scale1) + return x.addcmul_(self.ff(comfy.rmsnorm.rms_norm(x, self.norm2.weight, self.norm2.eps)), self.scale2) + + +class ViT3DDecoder(nn.Module): + def __init__(self, patch_size=16, patch_size_t=4, in_channels=24, out_channels=3, num_layers=36, heads=32, dim_head=64, rope_theta=100.0, + rope_dim_ratio=0.75, bias=True, eps=1e-5, num_register_tokens=4): + super().__init__() + dim = heads * dim_head + self.patch_size = patch_size + self.patch_size_t = patch_size_t + self.out_channels = out_channels + self.num_register_tokens = num_register_tokens + + self.pos_embed = RotaryEmbeddingND(int(dim_head * rope_dim_ratio), rope_theta, n_dim=3) + self.x_embedder = ops.Linear(in_channels, dim) + self.register_tokens = nn.Parameter(torch.empty(1, num_register_tokens, dim)) + # unused at inference; kept so the checkpoint loads without leftover keys + self.register_buffer("mask_token", torch.empty(1, 1, dim)) + + self.transformer_blocks = nn.ModuleList( + [TransformerBlock(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + for _ in range(num_layers)] + ) + + self.norm_out = ops.LayerNorm(dim, elementwise_affine=True, eps=eps) + self.proj_out = ops.Linear(dim, out_channels * patch_size_t * patch_size * patch_size) + + def forward(self, x): + B, C, latent_T, latent_H, latent_W = x.shape + + h = self.x_embedder(x.flatten(2).transpose(1, 2)) # [B, T*H*W, C] + + num_patches = h.shape[1] + num_suffix = 1 + self.num_register_tokens + + h = torch.cat([h, self.register_tokens.expand(B, -1, -1), torch.zeros_like(h[:, 0:1, :])], dim=1) + + img_ids = create_token_ids((latent_T, latent_H, latent_W), x.device, x.dtype).expand(B, -1, -1) + suffix_ids = torch.zeros((B, num_suffix, 3), device=x.device, dtype=img_ids.dtype) + img_ids = torch.cat([img_ids, suffix_ids], dim=1) + + rotary_pos_emb = self.pos_embed(img_ids) + + for block in self.transformer_blocks: + h = block(h, rotary_pos_emb) + + output = self.proj_out(self.norm_out(h)) + + output = output[:, :num_patches, :] + + output = output.view( + B, latent_T, latent_H, latent_W, + self.out_channels, self.patch_size_t, self.patch_size, self.patch_size, + ) + output = output.permute(0, 4, 1, 5, 2, 6, 3, 7).contiguous() + output = output.reshape( + B, self.out_channels, + latent_T * self.patch_size_t, + latent_H * self.patch_size, + latent_W * self.patch_size, + ) + return output + + +# Full VAE + +class MiniMaxH3VideoVAE(nn.Module): + def __init__( + self, + in_channels=3, + out_ch=3, + ch=128, + embed_dim=24, + z_channels=24, + ch_mult=(1, 2, 2, 4, 4, 8), + num_res_blocks=2, + space_down=(2, 2, 2, 2, 1, 1), + time_down=(1, 2, 2, 1, 1, 1), + clip_length=17, + token_drop=3, + tile_size=256, + tile_overlap_min=64, + tiling=True, + ): + super().__init__() + self.vae_ratio = int(math.prod(space_down)) + self.vae_ratio_t = int(math.prod(time_down)) + + # temporal chunking parameters + self.clip_length = clip_length + self.token_drop = token_drop + self.frame_pre_padding = (-clip_length) % self.vae_ratio_t + self.tokens_chunk_size = math.ceil(clip_length / self.vae_ratio_t) + self.token_overlap = (-token_drop) % self.tokens_chunk_size + self.frame_overlap = max(self.token_overlap * self.vae_ratio_t - self.frame_pre_padding, 0) + + # spatial tiling parameters + self.tiling = tiling + self.tile_size = tile_size + self.tile_overlap_min = tile_overlap_min + + self.encoder = EncoderFCN3D( + ch=ch, + ch_mult=list(ch_mult), + space_down=list(space_down), + time_down=list(time_down), + num_res_blocks=num_res_blocks, + in_channels=in_channels, + z_channels=z_channels, + double_z=True, + ) + self.quant_conv = ops.Conv3d(z_channels * 2, 2 * embed_dim, 1) + self.post_quant_conv = ops.Conv3d(embed_dim, z_channels, 1) + self.decoder = ViT3DDecoder( + patch_size=self.vae_ratio, + patch_size_t=self.vae_ratio_t, + in_channels=z_channels, + out_channels=out_ch, + ) + + self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN)) + self.register_buffer("latents_std", torch.tensor(LATENTS_STD)) + self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1, 1), persistent=False) + self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False) + + # single-shot forward + + def _encode_moments(self, x): + return self.quant_conv(self.encoder(x)) + + def _decode_pixels(self, z): + return self.decoder(self.post_quant_conv(z)) + + def _adaptive_encode(self, x): + if self.tiling: + return self.tiled_encode(x) + return self._encode_moments(x) + + def _adaptive_decode(self, z): + if self.tiling: + return self.tiled_decode(z) + return self._decode_pixels(z) + + # spatial tiling + + def split_tiles(self, input_len): + tile_size = self.tile_size + if tile_size >= input_len: + return [0], [input_len], [] + + N = math.ceil(input_len / tile_size) + while True: + overlaps = [self.tile_overlap_min] * (N - 1) + remaining = tile_size * N - sum(overlaps) - input_len + if remaining < 0: + N += 1 + else: + break + + remaining_units = remaining // self.vae_ratio + for i in range(remaining_units): + overlaps[i % (N - 1)] += self.vae_ratio + + tile_start_idx = [0] + for i in range(N - 1): + tile_start_idx.append(tile_start_idx[-1] + tile_size - overlaps[i]) + + return tile_start_idx, [tile_size] * N, overlaps + + def blend(self, a, b, blend_extent, dim): + blend_extent = min(a.shape[dim], b.shape[dim], blend_extent) + + positions = torch.arange(blend_extent, device=b.device, dtype=b.dtype) + weight_a = 1 - positions / blend_extent + weight_b = positions / blend_extent + + shape = [1] * a.ndim + shape[dim] = blend_extent + weight_a = weight_a.view(shape) + weight_b = weight_b.view(shape) + + slice_a = [slice(None)] * a.ndim + slice_a[dim] = slice(-blend_extent, None) + slice_b = [slice(None)] * b.ndim + slice_b[dim] = slice(0, blend_extent) + + blended = a[tuple(slice_a)] * weight_a + b[tuple(slice_b)] * weight_b + + if blend_extent < b.shape[dim]: + slice_b_rest = [slice(None)] * b.ndim + slice_b_rest[dim] = slice(blend_extent, None) + return torch.cat([blended, b[tuple(slice_b_rest)]], dim=dim) + return blended + + def tiled_encode(self, x): + height, width = x.shape[-2], x.shape[-1] + y_idx, y_len, y_overlap = self.split_tiles(height) + x_idx, x_len, x_overlap = self.split_tiles(width) + + rows = [] + for i_pos, i_len in zip(y_idx, y_len): + row = [] + for j_pos, j_len in zip(x_idx, x_len): + tile = x[..., i_pos:i_pos + i_len, j_pos:j_pos + j_len] + row.append(self._encode_moments(tile)) + rows.append(row) + + latent_y_overlap = [o // self.vae_ratio for o in y_overlap] + latent_x_overlap = [o // self.vae_ratio for o in x_overlap] + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2) + if j > 0: + tile = self.blend(row[j - 1], tile, latent_x_overlap[j - 1], dim=-1) + if i < len(rows) - 1: + tile = tile[..., :-latent_y_overlap[i], :] + if j < len(row) - 1: + tile = tile[..., :, :-latent_x_overlap[j]] + result_row.append(tile) + result_rows.append(torch.cat(result_row, dim=-1)) + return torch.cat(result_rows, dim=-2) + + def tiled_decode(self, z): + height, width = z.shape[-2] * self.vae_ratio, z.shape[-1] * self.vae_ratio + y_idx, y_len, y_overlap = self.split_tiles(height) + x_idx, x_len, x_overlap = self.split_tiles(width) + + # Blended tiles are written straight into a pre-allocated canvas. + canvas = None + row_tails = [] + out_y = 0 + for i, (i_pos, i_len) in enumerate(zip(y_idx, y_len)): + zi, zl = i_pos // self.vae_ratio, i_len // self.vae_ratio + new_tails = [] + left_tail = None + out_x = 0 + for j, (j_pos, j_len) in enumerate(zip(x_idx, x_len)): + zj, zw = j_pos // self.vae_ratio, j_len // self.vae_ratio + tile = self._decode_pixels(z[..., zi:zi + zl, zj:zj + zw]) + if i < len(y_idx) - 1: + new_tails.append(tile[..., -y_overlap[i]:, :].clone()) + next_left_tail = tile[..., :, -x_overlap[j]:].clone() if j < len(x_idx) - 1 else None + if i > 0: + tile = self.blend(row_tails[j], tile, y_overlap[i - 1], dim=-2) + if j > 0: + tile = self.blend(left_tail, tile, x_overlap[j - 1], dim=-1) + left_tail = next_left_tail + if i < len(y_idx) - 1: + tile = tile[..., :-y_overlap[i], :] + if j < len(x_idx) - 1: + tile = tile[..., :, :-x_overlap[j]] + if canvas is None: + canvas = torch.empty(*tile.shape[:-2], height, width, dtype=tile.dtype, device=tile.device) + canvas[..., out_y:out_y + tile.shape[-2], out_x:out_x + tile.shape[-1]].copy_(tile) + out_x += tile.shape[-1] + row_tails = new_tails + out_y += tile.shape[-2] + return canvas + + # temporal chunking + + def encode_temporal(self, x): + if x.shape[2] % self.clip_length != 0: + pad_size = (-x.shape[2]) % self.clip_length + pad_frames = x[:, :, -1:].repeat(1, 1, pad_size, 1, 1) + x = torch.cat([x, pad_frames], dim=2) + + num_chunks = x.shape[2] // self.clip_length + + z_list = [] + for i in range(num_chunks): + clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :] + z_list.append(self._adaptive_encode(clip_x)) + + z = torch.cat(z_list, dim=2) + if self.token_drop > 0: + z = z[:, :, :-self.token_drop] + return z + + def _decode_temporal_pad_frames(self, z_len, pad_tokens): + if pad_tokens <= 0: + return 0 + intra_tail = self.clip_length % self.vae_ratio_t + if intra_tail == 0: + return pad_tokens * self.vae_ratio_t + + z_len_before_pad = z_len - pad_tokens + return sum( + (intra_tail if (z_len_before_pad + k) % self.tokens_chunk_size == 0 + else self.vae_ratio_t) + for k in range(pad_tokens) + ) + + def _decode_temporal_frame_plan(self, z_len, num_chunks, pad_tokens): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + total_frames = 0 + final_overlap_frames = 0 + + for i in range(num_chunks): + t_start_idx = i * self.tokens_chunk_size + t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap + clip_token_len = max(0, min(t_end_idx, z_len) - min(t_start_idx, z_len)) + clip_frame_len = clip_token_len * self.vae_ratio_t + + for j in range(split_count): + f_start_idx = j * chunk_dec + f_end_idx = min(f_start_idx + chunk_dec, clip_frame_len) + chunk_frames = max(0, f_end_idx - f_start_idx - self.frame_pre_padding) + if j == 0: + total_frames += chunk_frames + else: + final_overlap_frames = chunk_frames + + total_frames += final_overlap_frames + return total_frames - self._decode_temporal_pad_frames(z_len, pad_tokens) + + def decode_temporal(self, z): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + + pseudo_total_tokens = z.shape[2] + self.token_drop + + pad_tokens = 0 + remainder = pseudo_total_tokens % self.tokens_chunk_size + if remainder != 0: + pad_tokens = self.tokens_chunk_size - remainder + pseudo_total_tokens += pad_tokens + + num_chunks = pseudo_total_tokens // self.tokens_chunk_size - int(self.token_drop > 0) + if num_chunks < 1: + # too few tokens for one chunk (e.g. T_lat == 2): pad one extra chunk + pad_tokens += self.tokens_chunk_size + num_chunks += 1 + + if pad_tokens > 0: + pad_z = z[:, :, -1:, :, :].repeat(1, 1, pad_tokens, 1, 1) + z = torch.cat([z, pad_z], dim=2) + + output_frames = self._decode_temporal_frame_plan(z.shape[2], num_chunks, pad_tokens) + + dec = None + dec_overlap = None + write_pos = 0 + + def write_part(part): + nonlocal dec, write_pos + part_frames = part.shape[2] + if part_frames <= 0: + return + if dec is None: + out_shape = list(part.shape) + out_shape[2] = output_frames + dec = torch.empty(out_shape, dtype=part.dtype, device=part.device) + copy_frames = min(part_frames, max(0, dec.shape[2] - write_pos)) + if copy_frames > 0: + dec[:, :, write_pos:write_pos + copy_frames, :, :].copy_( + part[:, :, :copy_frames, :, :] + ) + write_pos += copy_frames + + for i in range(num_chunks): + t_start_idx = i * self.tokens_chunk_size + t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap + clip_z = z[:, :, t_start_idx:t_end_idx, :, :] + + clip_dec = self._adaptive_decode(clip_z) + + for j in range(split_count): + f_start_idx = j * chunk_dec + f_end_idx = min(f_start_idx + chunk_dec, clip_dec.shape[2]) + clip_dec_chunk = clip_dec[:, :, f_start_idx:f_end_idx, :, :] + clip_dec_chunk = clip_dec_chunk[:, :, self.frame_pre_padding:, :, :] + + if j == 0: + if dec_overlap is not None: + clip_dec_chunk = self.blend( + dec_overlap, clip_dec_chunk, self.frame_overlap, dim=-3 + ) + dec_overlap = None + write_part(clip_dec_chunk) + else: + dec_overlap = clip_dec_chunk.contiguous() + + if i == num_chunks - 1 and dec_overlap is not None: + write_part(dec_overlap) + dec_overlap = None + + del clip_dec, clip_z + + return dec + + + def encode(self, x): + # x: [B, 3, T, H, W] in [-1, 1] -> normalized latents [B, 24, T_lat, H/16, W/16] + if x.ndim == 4: + x = x.unsqueeze(2) + + x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + + if x.shape[2] == 1: + moments = self._adaptive_encode(x) + moments = moments[:, :, -1:, :, :] + else: + moments = self.encode_temporal(x) + + mean = torch.chunk(moments.float(), 2, dim=1)[0] + + latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(mean) + latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(mean) + return (mean - latents_mean) / latents_std + + def encode_tiled(self, x, **kwargs): + # tiling is always on internally with the reference's semantic tile sizes, ignore tiling fallbacks + return self.encode(x) + + def decode_tiled(self, z, **kwargs): + return self.decode(z) + + def decode(self, z): + # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> pixels [B, 3, T, H, W] in [-1, 1] + latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(z) + latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(z) + z = z * latents_std + latents_mean + + if z.shape[2] == 1: + dec = self._adaptive_decode(z) + dec = dec[:, :, -1:, :, :] + else: + dec = self.decode_temporal(z) + + dec = dec.float() + dec.mul_(self.pixel_std.to(dec)).add_(self.pixel_mean.to(dec)).clamp_(0.0, 1.0).mul_(2.0).sub_(1.0) + return dec diff --git a/comfy/model_base.py b/comfy/model_base.py index 7c68c6eba..6631c9eb0 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -21,6 +21,7 @@ import comfy.ldm.hunyuan3dv2_1.hunyuandit import torch import logging import comfy.ldm.lightricks.av_model +import comfy.ldm.minimax.model import comfy.ldm.lightricks.symmetric_patchifier import comfy.context_windows from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep @@ -2063,6 +2064,57 @@ class Hunyuan3Dv2_1(BaseModel): out['guidance'] = comfy.conds.CONDRegular(torch.FloatTensor([guidance])) return out +class MiniMaxH3(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.minimax.model.MiniMaxH3Model) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + # run condition_proj + token refiner once per sampling instead of per step + cross_attn = self.diffusion_model.preprocess_text_embeds( + cross_attn.to(device=kwargs["device"], dtype=self.get_dtype_inference())) + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + + latent_shapes = kwargs.get("latent_shapes", None) + if latent_shapes is not None: + out['latent_shapes'] = comfy.conds.CONDConstant(latent_shapes) + + # Everything H3-specific rides in one dict so _apply_model's dtype cast + # (which would flatten fp32 cond latents and long tags to bf16) skips it. + payload = {} + tags = kwargs.get("minimax_token_tags", None) + if tags is not None: + payload["text_token_tags"] = tags + keyframes = kwargs.get("minimax_keyframes", None) + if keyframes is not None: + payload["keyframes"] = keyframes + payload["frame_count"] = kwargs.get("minimax_frame_count", None) + payload["cond_video_latents"] = [kf["latent"] for kf in keyframes] + refs = kwargs.get("minimax_refs", None) + if refs is not None: + payload["refs"] = refs + payload["cond_video_latents"] = [r["latent"] for r in refs if "latent" in r] + payload["cond_audio_latents"] = [r["audio_latent"] for r in refs if r.get("audio_latent") is not None] + if kwargs.get("minimax_visual_cond_noise_aug", None) is not None: + payload["visual_cond_noise_aug"] = kwargs["minimax_visual_cond_noise_aug"] + if kwargs.get("minimax_audio_cond_noise_aug", None) is not None: + payload["audio_cond_noise_aug"] = kwargs["minimax_audio_cond_noise_aug"] + payload["seed"] = kwargs.get("seed", 0) + if cross_attn is not None and latent_shapes is not None and len(latent_shapes) > 1: + # packed layout built once per sampling run, h/w rounded up to the DiT's 2x2 patch + vs = latent_shapes[0] + payload["layout"] = comfy.ldm.minimax.model.PackedLayout( + cross_attn.shape[1], vs[2], (vs[3] + 1) // 2 * 2, (vs[4] + 1) // 2 * 2, + latent_shapes[1][-1], keyframes=payload.get("keyframes"), + refs=payload.get("refs"), frame_count=payload.get("frame_count")) + out['minimax_payload'] = comfy.conds.CONDConstant(payload) + return out + + def scale_latent_inpaint(self, sigma, noise, latent_image, **kwargs): + return latent_image + class TripoSplat(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.triposplat.model.LatentSeqMMFlowModel) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 39e973d36..103680fd1 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -359,6 +359,35 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): # PixArt diffusers return None + if '{}video_patch_proj.weight'.format(key_prefix) in state_dict_keys and '{}audio_patch_proj.weight'.format(key_prefix) in state_dict_keys: # MiniMax H3 + dit_config = {} + dit_config["image_model"] = "minimax_h3" + dit_config["num_layers"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') + dit_config["token_refiner_num_layers"] = count_blocks(state_dict_keys, '{}token_refiner.blocks.'.format(key_prefix) + '{}.') + dit_config["hidden_size"] = state_dict['{}video_patch_proj.weight'.format(key_prefix)].shape[0] + dit_config["latents_dim"] = state_dict['{}final_layer.video_out.weight'.format(key_prefix)].shape[0] // 4 # patch 1x2x2 + dit_config["audio_latents_dim"] = state_dict['{}final_layer.audio_out.weight'.format(key_prefix)].shape[0] + dit_config["attention_head_dim"] = state_dict['{}blocks.0.attn.q_norm.weight'.format(key_prefix)].shape[0] + qkv = state_dict['{}blocks.0.attn.qkv_proj.weight'.format(key_prefix)] + dit_config["num_attention_heads"] = qkv.shape[0] // (3 * dit_config["attention_head_dim"]) + dit_config["ffn_hidden_size"] = state_dict['{}blocks.0.mlp.fc1.weight'.format(key_prefix)].shape[0] // 2 + dit_config["text_dim"] = state_dict['{}condition_proj.weight'.format(key_prefix)].shape[1] + table_key = '{}adaln_t_table'.format(key_prefix) + if table_key in state_dict_keys: + # adaln shipped over a precomputed curve basis: the adaln linears span a small shared basis of the time-embedding curve (no time embedder) + table = state_dict[table_key].shape # [grid, k] + dit_config["adaln_curve_grid"] = table[0] + dit_config["time_embed_dim"] = table[1] + else: + te = state_dict['{}time_embedder.proj_in.weight'.format(key_prefix)] + dit_config["timestep_input_dim"] = te.shape[1] + dit_config["time_embed_hidden_size"] = te.shape[0] + dit_config["time_embed_dim"] = state_dict['{}time_embedder.proj_out.weight'.format(key_prefix)].shape[0] + dit_config["rope_inv_freq_len"] = state_dict['{}rope.inv_freq'.format(key_prefix)].shape[0] + if metadata is not None and "config" in metadata: + dit_config.update(json.loads(metadata["config"]).get("transformer", {})) + return dit_config + if '{}adaln_single.emb.timestep_embedder.linear_1.bias'.format(key_prefix) in state_dict_keys: #Lightricks ltxv dit_config = {} dit_config["image_model"] = "ltxav" if f'{key_prefix}audio_adaln_single.linear.weight' in state_dict_keys else "ltxv" diff --git a/comfy/ops.py b/comfy/ops.py index 6c3845eef..077a42351 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -943,13 +943,61 @@ if CUBLAS_IS_AVAILABLE: # ============================================================================== # Mixed Precision Operations # ============================================================================== +from . import quant_ops from .quant_ops import ( QuantizedTensor, QUANT_ALGOS, TensorCoreFP8Layout, + TensorWiseINT8Layout, get_layout_class, ) +def _swiglu_eager(x): + gate, up = x.chunk(2, dim=-1) + return torch.nn.functional.silu(gate).mul_(up) + + +INPUT_ACT_EAGER = { + "gelu_tanh": lambda x: torch.nn.functional.gelu(x, approximate="tanh"), + "swiglu": _swiglu_eager, +} + + +def linear_input_act(linear, x, input_act): + """``linear(act(x))``, with ``act`` folded into an INT8 activation quantizer. + + An INT8 linear quantizes its input anyway, so an elementwise activation can + ride along inside that kernel instead of writing a full-size intermediate to + HBM and reading it straight back. Worth it for an MLP's down-projection, + where the intermediate is several times the hidden size. + + """ + weight = linear.weight + if (comfy.model_management.in_training + or not isinstance(weight, QuantizedTensor) + or weight._layout_cls != "TensorWiseINT8Layout" + or getattr(weight._params, "transposed", False)): + return linear(INPUT_ACT_EAGER[input_act](x)) + + # want_requant keeps a vbar-streamed layer on the INT8 path when a LoRA is + # patched in on the fly; without it the cast hands back a dequantized weight. + weight, bias, offload_stream = cast_bias_weight( + linear, x, offloadable=True, compute_dtype=x.dtype, want_requant=True) + try: + if not isinstance(weight, QuantizedTensor): + # A LoRA weight_function, or activations whose dtype differs from the + # weight's, make the cast hand back a dequantized tensor. + return torch.nn.functional.linear(INPUT_ACT_EAGER[input_act](x), weight, bias) + qdata, scale = TensorWiseINT8Layout.get_plain_tensors(weight) + return quant_ops.ck.int8_linear( + x, qdata, scale, bias, x.dtype, + convrot=getattr(weight._params, "convrot", False), + convrot_groupsize=getattr(weight._params, "convrot_groupsize", 256), + input_act=input_act, + ) + finally: + uncast_bias_weight(linear, weight, bias, offload_stream) + class QuantLinearFunc(torch.autograd.Function): """Custom autograd function for quantized linear: quantized forward, optionally FP8 backward. @@ -1259,7 +1307,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def state_dict(self, *args, destination=None, prefix="", **kwargs): sd = destination if destination is not None else {} - return _quantized_weight_state_dict(self, sd, prefix, extra_quant_params=("input_scale",)) + return _quantized_weight_state_dict(self, sd, prefix, extra_quant_params=("input_scale", "pre_quant_scale")) def _forward(self, input, weight, bias): return torch.nn.functional.linear(input, weight, bias) @@ -1298,6 +1346,11 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def forward(self, input, *args, **kwargs): run_every_op() + # ModelOpt AWQ-style smoothing + pre_quant_scale = getattr(self, 'pre_quant_scale', None) + if pre_quant_scale is not None: + input = input * comfy.model_management.cast_to_device(pre_quant_scale, input.device, input.dtype) + input_shape = input.shape reshaped_nd = False #If cast needs to apply lora, it should be done in the compute dtype diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 15f9b1fdb..53586956a 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -240,7 +240,7 @@ QUANT_ALGOS = { }, "nvfp4": { "storage_t": torch.uint8, - "parameters": {"weight_scale", "weight_scale_2", "input_scale"}, + "parameters": {"weight_scale", "weight_scale_2", "input_scale", "pre_quant_scale"}, "comfy_tensor_layout": "TensorCoreNVFP4Layout", "group_size": 16, }, diff --git a/comfy/sd.py b/comfy/sd.py index caf78222d..8d670106e 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -72,6 +72,9 @@ import comfy.text_encoders.ace15 import comfy.text_encoders.longcat_image import comfy.text_encoders.qwen35 import comfy.text_encoders.qwen3vl +import comfy.text_encoders.minimax +import comfy.ldm.minimax.vae +import comfy.ldm.minimax.audio_vae import comfy.text_encoders.boogu import comfy.text_encoders.ernie import comfy.text_encoders.gemma4 @@ -936,6 +939,50 @@ class VAE: #Force cast it for --disable-dynamic-vram users until there is a true core fix. if not comfy.memory_management.aimdo_enabled: self.disable_offload = True + elif "decoder.transformer_blocks.0.scale1" in sd and "encoder.down.5.block.0.conv1.weight" in sd: # MiniMax H3 video VAE + self.first_stage_model = comfy.ldm.minimax.vae.MiniMaxH3VideoVAE() + self.latent_channels = 24 + self.latent_dim = 3 + # frames 17k+5 <-> latents 5k+2, 16x spatial + self.upscale_ratio = (lambda a: max(1, (a - 2) // 5 * 17 + 5), 16, 16) + self.upscale_index_formula = (4, 16, 16) + self.downscale_ratio = (lambda a: max(1, (a - 5) // 17 * 5 + 2) if a > 1 else 1, 16, 16) + self.downscale_index_formula = (4, 16, 16) + self.working_dtypes = [torch.float16, torch.float32] + # the model tiles internally (256px spatial, 17-frame temporal chunks) + self.handles_tiling = True + def estimate_encode_memory(frames, height, width, dtype): + fixed = 110_000_000 if frames == 1 else 1_300_000_000 + elements_per_pixel = 7 if frames == 1 else 9.5 + return (elements_per_pixel * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 + + def estimate_decode_memory(frames, height, width, dtype): + fixed = 110_000_000 if frames <= 22 else 270_000_000 + return (9.5 * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 + + self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], shape[3], shape[4], dtype) + self.memory_used_decode = lambda shape, dtype: estimate_decode_memory(self.upscale_ratio[0](shape[2]), shape[3] * self.upscale_ratio[1], shape[4] * self.upscale_ratio[2], dtype) + elif "pre_block.attn.zero_k_bias" in sd: # MiniMax H3 audio VAE (DAC encoder + BigVGAN decoder) + self.first_stage_model = comfy.ldm.minimax.audio_vae.MiniMaxH3AudioVAE() + self.latent_channels = 32 + self.output_channels = 2 + self.pad_channel_value = "replicate" + self.audio_sample_rate = 32000 + self.upscale_ratio = 800 + self.downscale_ratio = 800 + self.latent_dim = 2 # [B, 32, stereo 2, T] + self.process_output = lambda audio: audio + self.process_input = lambda audio: audio + self.working_dtypes = [torch.float32] + # encode gets the waveform shape [B, 2, samples], decode the latent shape [B, 32, 2, T] + def estimate_encode_memory(samples, dtype): + return (900 * samples + 105_000_000) * model_management.dtype_size(dtype) * 1.03 + + def estimate_decode_memory(samples, dtype): + return max(42_000_000, 220 * samples + 20_000_000) * model_management.dtype_size(dtype) * 1.03 + + self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], dtype) + self.memory_used_decode = lambda shape, dtype: estimate_decode_memory(shape[-1] * self.upscale_ratio, dtype) elif "gs.base_offset_scale" in sd and "octree.out_proj.weight" in sd: # TripoSplat octree gaussian decoder self.first_stage_model = comfy.ldm.triposplat.vae.OctreeGaussianDecoder() self.latent_channels = 16 @@ -1393,6 +1440,7 @@ class CLIPType(Enum): KREA2 = 32 JOYIMAGE = 33 MAGE = 34 + MINIMAX = 35 @@ -1449,6 +1497,7 @@ class TEModel(Enum): QWEN3VL_4B = 34 QWEN3VL_8B = 35 GEMMA_4_12B = 36 + QWEN3VL_32B = 37 def detect_te_model(sd): @@ -1515,6 +1564,9 @@ def detect_te_model(sd): return TEModel.QWEN35_2B if "model.visual.deepstack_merger_list.0.norm.weight" in sd: # DeepStack is unique to Qwen3-VL return TEModel.QWEN3VL_4B if sd["model.visual.merger.linear_fc2.weight"].shape[0] == 2560 else TEModel.QWEN3VL_8B + if "visual.deepstack_merger_list.0.norm.weight" in sd and "model.layers.49.self_attn.q_proj.weight" in sd: + # MiniMax H3 conditioning encoder: Qwen3-VL-32B, truncated to 50 layers + return TEModel.QWEN3VL_32B if "model.layers.0.post_attention_layernorm.weight" in sd: weight = sd['model.layers.0.post_attention_layernorm.weight'] if 'model.layers.0.self_attn.q_norm.weight' in sd: @@ -1744,6 +1796,9 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip qwen3vl_type = {TEModel.QWEN3VL_4B: "qwen3vl_4b", TEModel.QWEN3VL_8B: "qwen3vl_8b"}[te_model] clip_target.clip = comfy.text_encoders.qwen3vl.te(**llama_detect(clip_data), model_type=qwen3vl_type) clip_target.tokenizer = comfy.text_encoders.qwen3vl.tokenizer(model_type=qwen3vl_type) + elif te_model == TEModel.QWEN3VL_32B: + clip_target.clip = comfy.text_encoders.minimax.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.minimax.MiniMaxH3Tokenizer elif te_model == TEModel.QWEN3_06B: clip_target.clip = comfy.text_encoders.anima.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.anima.AnimaTokenizer diff --git a/comfy/supported_models.py b/comfy/supported_models.py index ca89850a5..51b58ed1e 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -15,6 +15,7 @@ import comfy.text_encoders.flux import comfy.text_encoders.genmo import comfy.text_encoders.lt import comfy.text_encoders.hunyuan_video +import comfy.text_encoders.minimax import comfy.text_encoders.cosmos import comfy.text_encoders.lumina2 import comfy.text_encoders.wan @@ -955,6 +956,33 @@ class LTXAV(LTXV): out = model_base.LTXAV(self, device=device) return out +class MiniMaxH3(supported_models_base.BASE): + unet_config = { + "image_model": "minimax_h3", + } + + sampling_settings = { + "shift": 12.0, + } + + unet_extra_config = {} + latent_format = latent_formats.MiniMaxH3AV + + memory_usage_factor = 0.114 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + return model_base.MiniMaxH3(self, device=device) + + def clip_target(self, state_dict={}, prefix=""): + pref = self.text_encoder_key_prefix[0] + detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_32b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.minimax.MiniMaxH3Tokenizer, comfy.text_encoders.minimax.te(**detect)) + class HunyuanVideo(supported_models_base.BASE): unet_config = { "image_model": "hunyuan_video", @@ -2407,6 +2435,7 @@ models = [ GenmoMochi, LTXV, LTXAV, + MiniMaxH3, HunyuanVideo15_SR_Distilled, HunyuanVideo15, HunyuanImage21Refiner, diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index 40d04007e..f5c5597ef 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -264,6 +264,17 @@ class Qwen3VL_4BConfig(Qwen3VL_8BConfig): intermediate_size: int = 9728 lm_head: bool = False # 4B ties word embeddings +@dataclass +class Qwen3VL_32BConfig(Qwen3VL_8BConfig): + # MiniMax H3 conditioning checkpoint: truncated to the first 50 of 64 layers, + # consumed as the unnormalized hidden state after layer 50 (no final norm, no lm_head) + hidden_size: int = 5120 + intermediate_size: int = 25600 + num_hidden_layers: int = 50 + num_attention_heads: int = 64 + lm_head: bool = False + final_norm: bool = False + @dataclass class Ovis25_2BConfig: vocab_size: int = 151936 diff --git a/comfy/text_encoders/minimax.py b/comfy/text_encoders/minimax.py new file mode 100644 index 000000000..c2dc47f7f --- /dev/null +++ b/comfy/text_encoders/minimax.py @@ -0,0 +1,201 @@ +"""MiniMax H3 text/vision conditioning: Qwen3-VL-32B (truncated to 50 layers). + +The H3 presentation is NOT chat-templated: token ids are raw prompt/label text +(no special tokens) with explicit vision blocks spliced in: + + t2va: + fl2va: ": " [": " ] + ref2va: per condition in request order (1-based ordinals per type): + image -> ": " + audio -> "