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 1/4] 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 9d7fa731f84..e15e0a9fd2a 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 @@ class EmptyClass: 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 0bba8341b7f..5163c1676de 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_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 @@ def forward( 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 @@ def __init__(self, config, index, device=None, dtype=None, ops=None): 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 @@ def __init__(self, config, index, device=None, dtype=None, ops=None): 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 @@ def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, p 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 @@ def forward(self, x, attention_mask=None, embeds=None, num_tokens=None, intermed 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 @@ def forward(self, x, attention_mask=None, embeds=None, num_tokens=None, intermed 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 @@ def forward(self, x, attention_mask=None, embeds=None, num_tokens=None, intermed 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 @@ def logits(self, x): 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 @@ def preprocess_embed(self, embed, device): 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 @@ def __init__(self, config, dtype=None, device=None, ops=None): 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 @@ def __init__(self, config, dtype=None, device=None, ops=None): # 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 @@ def state_dict(self): 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 @@ def tokenize_with_weights(self, text, return_word_ids=False, image=None, audio=N 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 @@ def tokenize_with_weights(self, text, return_word_ids=False, image=None, audio=N 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 @@ def tokenize_with_weights(self, text, return_word_ids=False, image=None, audio=N 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 @@ def tokenize_with_weights(self, text, return_word_ids=False, image=None, audio=N 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 @@ def _replace_placeholders(token_list, token_id, embeds): 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 @@ def __init__(self, embedding_directory=None, tokenizer_data={}): 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 @@ def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min 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 @@ class Tokenizer(Gemma4Tokenizer): 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 3f98fb0a594..40d04007e2e 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -876,7 +876,7 @@ def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): 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 @@ def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, 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 2/4] 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 222005b6f4d..766e9ea89cb 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -473,7 +473,7 @@ def amd_min_version(device=None, min_rdna_version=0): 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 3/4] 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 043c0ec75f9..634927dd646 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 4/4] [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 73fe75b7fca..d7e4652cf8f 100644 --- a/comfy_extras/nodes_dataset.py +++ b/comfy_extras/nodes_dataset.py @@ -2,6 +2,7 @@ import os import json +import av import numpy as np import torch from PIL import Image @@ -9,7 +10,7 @@ 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 @@ def execute(cls, folder): 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 @@ def define_schema(cls): @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 @@ def execute(cls, images, **kwargs): 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 @@ def execute(cls, images, texts, seed): 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 @@ async def get_node_list(self) -> list[type[io.ComfyNode]]: 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 @@ async def get_node_list(self) -> list[type[io.ComfyNode]]: 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 a27217b804b..0dde97fc9a5 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,