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 01/49] 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 02/49] 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 03/49] 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 04/49] 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 05/49] [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 06/49] 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 07/49] 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 08/49] [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 09/49] 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 10/49] 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 11/49] [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 12/49] [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 13/49] 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 14/49] 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 15/49] 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 16/49] [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 17/49] [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 18/49] 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 19/49] 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 20/49] 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 21/49] 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 22/49] 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 23/49] 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 24/49] 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 25/49] [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 26/49] 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 27/49] 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 28/49] 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 29/49] 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 30/49] [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 31/49] 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 32/49] 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 33/49] [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 34/49] [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 35/49] 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 36/49] 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 37/49] 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 38/49] 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 39/49] 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 40/49] 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 41/49] 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 42/49] 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 43/49] [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 44/49] 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 45/49] 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 46/49] 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 47/49] 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 48/49] 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 49/49] 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()