From a779de4d89fd21d90c300f6a40ca3ed2dc04e2ee Mon Sep 17 00:00:00 2001 From: Comfy Org PR Bot Date: Fri, 14 Aug 2026 08:55:28 +0900 Subject: [PATCH 1/7] Bump comfyui-frontend-package to 1.49.6 (#15526) Co-authored-by: Alexis Rolland --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index e180e7884..18ce16a56 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,4 @@ -comfyui-frontend-package==1.48.7 +comfyui-frontend-package==1.49.6 comfyui-workflow-templates==0.11.41 comfyui-embedded-docs==0.5.9 torch From 7fe8a6138504f90ff7be82f3babf416da32876b1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Fri, 14 Aug 2026 04:12:55 +0300 Subject: [PATCH 2/7] Speedup Gemma4 text generation (CORE-371) (#15054) --- comfy/model_prefetch.py | 16 ++- comfy/ops.py | 2 +- comfy/text_encoders/gemma4.py | 263 +++++++++++++++++++++++++++++++--- 3 files changed, 260 insertions(+), 21 deletions(-) diff --git a/comfy/model_prefetch.py b/comfy/model_prefetch.py index 7aedab530..bdde5137a 100644 --- a/comfy/model_prefetch.py +++ b/comfy/model_prefetch.py @@ -1,4 +1,5 @@ import torch +import warnings import weakref import comfy_aimdo.model_vbar @@ -28,6 +29,18 @@ def cleanup_prefetched_modules(module, comfy_modules): comfy_aimdo.model_vbar.vbar_unpin(module._v_block) del module._v_block_faulted +def _drop_graph(module): + graph = getattr(module, "_comfy_graph", None) + if graph is None: + return + # reset() through the bound method surfaces the allocator's benign + # "uncaptured free of a captured allocation" as catchable Python warnings; + # a plain del frees from the C++ dealloc path and spams stderr instead + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + graph["graph"].reset() + del module._comfy_graph + def cleanup_prefetch_queues(): global PREFETCH_QUEUES, GRAPH_CAPTURE_STREAMS @@ -41,7 +54,7 @@ def cleanup_prefetch_queues(): cleanup_prefetched_modules(prefetched_module, comfy_modules) PREFETCH_QUEUES = [] for module in GRAPH_MODULES: - del module._comfy_graph + _drop_graph(module) GRAPH_MODULES.clear() GRAPH_WARMED_MODULES.clear() GRAPH_CAPTURE_STREAMS = {} @@ -117,6 +130,7 @@ def prefetch_queue_pop(queue, device, module, dtype=None, core=None, enable_grap if signature is not None: module._v_block_faulted = True if signature is not None: + _drop_graph(module) graph = torch.cuda.CUDAGraph() if generator is not None: graph.register_generator_state(generator) diff --git a/comfy/ops.py b/comfy/ops.py index 73ae46674..ff64aad59 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -1639,7 +1639,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec self.norm_type, self.scale_grad_by_freq, self.sparse) target_dtype = out_dtype if out_dtype is not None else weight._params.orig_dtype x = x.to(dtype=target_dtype) - if scale is not None and scale != 1.0: + if scale is not None: x = x * scale.to(dtype=target_dtype) return x diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 606f8993e..61bc3a3f0 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -9,10 +9,12 @@ import math from comfy import sd1_clip import comfy.model_management +import comfy.model_prefetch import comfy.ops +import comfy.quant_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 +from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, FixedKV, _make_scaled_embedding # Intentional minor divergences from transformers -reference implementation: @@ -109,7 +111,28 @@ class Gemma4_12B_Config(Gemma4Config): suppress_tokens = [258883, 258882] -# unfused RoPE as addcmul_ RoPE diverges from reference code +class RingKV(FixedKV): + # sliding-window ring: writes wrap at capacity, validity saturates + def prepare(self, num_tokens): + capacity = self.key.shape[2] + self.position.fill_(self.index % capacity) + self.seqlen.fill_(min(self.index + num_tokens, capacity)) + + +def _fixed_kv_decode_mask(mask, cache, min_val): + capacity = cache.key.shape[2] + valid = min(cache.index + 1, capacity) + output = mask.new_full((*mask.shape[:-1], capacity), min_val) + if isinstance(cache, RingKV): + positions = torch.arange(cache.index + 1 - valid, cache.index + 1, device=mask.device) % capacity + output.index_copy_(-1, positions, mask[..., -valid:]) + else: + output[..., :valid] = mask[..., :valid] + return output + + +# unfused RoPE as addcmul_ RoPE diverges from reference code (vision only; text +# layers use the kitchen split-half kernel, bitwise-equal to this with bf16 freqs) def _apply_rotary_pos_emb(x, freqs_cis): cos, sin = freqs_cis[0], freqs_cis[1] half = x.shape[-1] // 2 @@ -140,6 +163,23 @@ class Gemma4Attention(nn.Module): if config.k_norm == "gemma3": self.k_norm = RMSNorm(head_dim, eps=config.rms_norm_eps, device=device, dtype=dtype) + def _decode_attention(self, xq, cache, bias): + if bias is None: + # eager decode: slice the cache to the valid length (python-side index, + # no mask needed; a full ring is order-invariant under softmax) + n = min(cache.index + 1, cache.key.shape[2]) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + attention = optimized_attention_for_device(xq.device, mask=False, small_input=True) + return attention(xq, cache.key[:, :, :n], cache.value[:, :, :n], self.num_heads, skip_reshape=True, scale=1.0, **gqa_kwargs) + # graph capture: fixed-length masked attention over the full capacity, explicit + # math (SDPA leaves its fast path on broadcast-bias + GQA and costs ~0.5ms/layer) + batch_size = xq.shape[0] + groups = self.num_heads // self.num_kv_heads + q = xq.reshape(batch_size, self.num_kv_heads, groups, self.head_dim) + scores = q @ cache.key.transpose(-1, -2) + bias + probs = torch.softmax(scores.float(), dim=-1).to(xq.dtype) + return (probs @ cache.value).reshape(batch_size, 1, self.inner_size) + def forward( self, hidden_states: torch.Tensor, @@ -156,10 +196,16 @@ class Gemma4Attention(nn.Module): if self.q_norm is not None: xq = self.q_norm(xq) + if isinstance(shared_kv, FixedKV): + # decode on a KV-shared layer: attend the source layer's fixed cache + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) + output = self._decode_attention(xq, shared_kv, attention_mask) + return self.o_proj(output), None, None + if shared_kv is not None: xk, xv = shared_kv # Apply RoPE to Q only (K already has RoPE from source layer) - xq = _apply_rotary_pos_emb(xq, freqs_cis) + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) present_key_value = None shareable_kv = None else: @@ -173,11 +219,39 @@ class Gemma4Attention(nn.Module): xv = rms_norm(xv) xk = xk.transpose(1, 2) xv = xv.transpose(1, 2) - xq = _apply_rotary_pos_emb(xq, freqs_cis) - xk = _apply_rotary_pos_emb(xk, freqs_cis) + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) + xk = comfy.quant_ops.ck.apply_rope_split_half1(xk, freqs_cis) present_key_value = None - if past_key_value is not None: + fixed_cache = past_key_value if isinstance(past_key_value, FixedKV) else None + if fixed_cache is not None: + if seq_length == 1: + # CUDA-graphable decode: write at the device-side ring/linear position + fixed_cache.key.index_copy_(2, fixed_cache.position, xk) + fixed_cache.value.index_copy_(2, fixed_cache.position, xv) + output = self._decode_attention(xq, fixed_cache, attention_mask) + return self.o_proj(output), fixed_cache, None + + # prefill: attend the local sequence, persist the tail into the cache + capacity = fixed_cache.key.shape[2] + index = fixed_cache.index + if index + seq_length <= capacity: + fixed_cache.key[:, :, index:index + seq_length] = xk + fixed_cache.value[:, :, index:index + seq_length] = xv + if index > 0: + xk = fixed_cache.key[:, :, :index + seq_length] + xv = fixed_cache.value[:, :, :index + seq_length] + elif index == 0: + # prefill longer than the sliding ring: attend the full local K/V + # (per-query windows come from the prefill sliding mask), cache only + # the last `capacity` keys at their wrapped slots (position % capacity) + slots = torch.arange(seq_length - capacity, seq_length, device=xk.device) % capacity + fixed_cache.key.index_copy_(2, slots, xk[:, :, -capacity:]) + fixed_cache.value.index_copy_(2, slots, xv[:, :, -capacity:]) + else: + raise RuntimeError("gemma4: chunked prefill past the sliding window is not supported") + present_key_value = fixed_cache + elif past_key_value is not None: cumulative_len = 0 if len(past_key_value) > 0: past_key, past_value, cumulative_len = past_key_value @@ -245,6 +319,7 @@ class TransformerBlockGemma4(nn.Module): 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): + output = x sliding_window = None if self.sliding_attention: sliding_window = self.sliding_attention @@ -281,7 +356,8 @@ class TransformerBlockGemma4(nn.Module): x = self.post_per_layer_input_norm(x) x = residual + x - x = x * comfy.ops.cast_to_input(self.layer_scalar, x) + # in-place into the input buffer so CUDA-graph replays land in the static x + x = torch.mul(x, comfy.ops.cast_to_input(self.layer_scalar, x), out=output) return x, present_key_value, shareable_kv @@ -290,6 +366,9 @@ class Gemma4Transformer(nn.Module): def __init__(self, config, device=None, dtype=None, ops=None): super().__init__() self.config = config + self.fixed_kv = True + self.prefetch_dynamic_vbars = True + self.graph_dynamic_vbar_blocks = True self.embed_tokens = _make_scaled_embedding(ops, config.vocab_size, config.hidden_size, config.hidden_size ** 0.5, device, dtype) @@ -298,6 +377,19 @@ class Gemma4Transformer(nn.Module): for i in range(config.num_hidden_layers) ]) + # KV-shared layers never run k_proj/v_proj/k_norm: their never-resolved vbar + # signatures would block layer graph capture, so prefetch only what executes + first_kv_shared = config.num_hidden_layers - config.num_kv_shared_layers if config.num_kv_shared_layers > 0 else config.num_hidden_layers + self._prefetch_units = [] + for i, layer in enumerate(self.layers): + if i >= first_kv_shared: + dead = {layer.self_attn.k_proj, layer.self_attn.v_proj, layer.self_attn.k_norm} + self._prefetch_units.append([ + m for m in layer.modules() if next(m.children(), None) is None and m not in dead + ]) + else: + self._prefetch_units.append(layer) + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype) if config.final_norm else None # Precompute RoPE inv_freq on CPU to match reference code's exact value @@ -311,6 +403,9 @@ class Gemma4Transformer(nn.Module): sliding_inv = 1.0 / (config.rope_theta[1] ** (torch.arange(0, config.head_dim, 2).float() / config.head_dim)) self.register_buffer("_sliding_inv_freq", sliding_inv, persistent=False) + if config.suppress_tokens: + self.register_buffer("_suppress_tokens", torch.tensor(config.suppress_tokens, dtype=torch.long), persistent=False) + # Per-layer input mechanism self.hidden_size_per_layer_input = config.hidden_size_per_layer_input if self.hidden_size_per_layer_input: @@ -322,19 +417,26 @@ class Gemma4Transformer(nn.Module): self.hidden_size_per_layer_input, eps=config.rms_norm_eps, device=device, dtype=dtype) + def get_dynamic_vram__units(self): + return (list(self.layers), []) if self.graph_dynamic_vbar_blocks else ([], []) + def get_past_len(self, past_key_values): for kv in past_key_values: + if isinstance(kv, FixedKV): + return kv.index if len(kv) >= 3: return kv[2] return 0 def _freqs_from_inv(self, inv_freq, position_ids, device, dtype): - """Compute cos/sin from stored inv_freq""" + """Compute per-pair 2x2 rotation matrices [B, 1, S, d/2, 2, 2] from stored inv_freq""" inv_exp = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(device) pos_exp = position_ids[:, None, :].float() freqs = (inv_exp @ pos_exp).transpose(1, 2) - emb = torch.cat((freqs, freqs), dim=-1) - return emb.cos().unsqueeze(1).to(dtype), emb.sin().unsqueeze(1).to(dtype) + cos, sin = freqs.cos(), freqs.sin() + mat = torch.stack((torch.stack((cos, -sin), dim=-1), + torch.stack((sin, cos), dim=-1)), dim=-2) + return mat.unsqueeze(1).to(dtype) def compute_freqs_cis(self, position_ids, device, dtype=None): global_freqs = self._freqs_from_inv(self._global_inv_freq, position_ids, device, dtype) @@ -401,6 +503,71 @@ class Gemma4Transformer(nn.Module): first_kv_shared = self.config.num_hidden_layers - num_kv_shared if num_kv_shared > 0 else self.config.num_hidden_layers shared_sliding_kv = None # KV from last non-shared sliding layer shared_global_kv = None # KV from last non-shared global layer + share_source = {} + if num_kv_shared > 0: + for i in range(first_kv_shared): + share_source[bool(self.layers[i].sliding_attention)] = i + + prefetch_queue = comfy.model_prefetch.make_prefetch_queue( + list(self._prefetch_units), x.device, + {"prefetch_dynamic_vbars": self.prefetch_dynamic_vbars and past_key_values is not None}) + + fixed_kv = (past_key_values is not None and len(past_key_values) > 0 + and isinstance(past_key_values[0], FixedKV)) + decode = fixed_kv and seq_len == 1 + # mirror the conditions under which prefetch_queue_pop can actually capture, so + # eager fallbacks keep the sliced decode path instead of the full-capacity one + enable_graph = (decode and mask is None and self.graph_dynamic_vbar_blocks + and prefetch_queue is not None + and hasattr(self.layers[0], "_v_block") + and not comfy.model_management.args.disable_cuda_graphs + and comfy.model_management.is_device_cuda(x.device)) + decode_bias = None + decode_masks = None + if decode: + prepared = set() + for kv in past_key_values: + if isinstance(kv, FixedKV) and id(kv.position) not in prepared: + kv.prepare(seq_len) + prepared.add(id(kv.position)) + if mask is not None: + decode_masks = {} + for kv in past_key_values: + if isinstance(kv, FixedKV) and id(kv.position) not in decode_masks: + decode_masks[id(kv.position)] = _fixed_kv_decode_mask(mask, kv, min_val) + if enable_graph: + # static buffers + per-capacity attention biases: layer graphs replay against + # stable storage, refreshed eagerly each step + capacities = tuple(sorted({kv.key.shape[2] for kv in past_key_values if isinstance(kv, FixedKV)})) + state_key = (x.shape, x.dtype, x.device, tuple(t.shape for t in freqs_cis), capacities, + None if per_layer_inputs is None else per_layer_inputs.shape) + state = getattr(self, "_comfy_cross_step_state", None) + if state is None or state["key"] != state_key: + state = {"key": state_key, + "x": torch.empty_like(x), + "freqs_cis": [torch.empty_like(t) for t in freqs_cis], + "bias": {c: torch.empty((1, 1, 1, c), dtype=x.dtype, device=x.device) for c in capacities}, + "per_layer": None if per_layer_inputs is None else torch.empty_like(per_layer_inputs), + "bias_valid": -1} + self._comfy_cross_step_state = state + comfy.model_management._register_cross_step(self) + state["x"].copy_(x) + for source, target in zip(freqs_cis, state["freqs_cis"]): + target.copy_(source) + x = state["x"] + freqs_cis = state["freqs_cis"] + if per_layer_inputs is not None: + state["per_layer"].copy_(per_layer_inputs) + per_layer_inputs = state["per_layer"] + valid = past_len + 1 + for capacity, bias in state["bias"].items(): + if state["bias_valid"] != past_len: + bias.fill_(min_val) + bias[..., :min(valid, capacity)] = 0 + elif past_len < capacity: + bias[..., past_len:valid] = 0 + state["bias_valid"] = valid + decode_bias = state["bias"] intermediate = None all_intermediate = None @@ -429,12 +596,36 @@ class Gemma4Transformer(nn.Module): is_sliding = hasattr(layer, 'sliding_attention') and layer.sliding_attention if i >= first_kv_shared and num_kv_shared > 0: - shared = shared_sliding_kv if is_sliding else shared_global_kv - if shared is not None: - layer_kwargs['shared_kv'] = shared + if decode: + layer_kwargs['shared_kv'] = past_key_values[share_source[bool(is_sliding)]] + else: + shared = shared_sliding_kv if is_sliding else shared_global_kv + if shared is not None: + layer_kwargs['shared_kv'] = shared - x, current_kv, shareable_kv = layer(x=x, attention_mask=mask, freqs_cis=freqs_cis, past_key_value=past_kv, **layer_kwargs) + if enable_graph: + bias_cache = layer_kwargs.get('shared_kv', past_kv) + layer_mask = decode_bias[bias_cache.key.shape[2]] + elif decode: + bias_cache = layer_kwargs.get('shared_kv', past_kv) + layer_mask = None if decode_masks is None else decode_masks[id(bias_cache.position)] + else: + layer_mask = mask + result = [] + + def core(): + nonlocal x + x, current_kv, shareable_kv = layer(x=x, attention_mask=layer_mask, freqs_cis=freqs_cis, past_key_value=past_kv, **layer_kwargs) + result.append((current_kv, shareable_kv)) + + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, layer, x.dtype, core=core, enable_graph=enable_graph) + + if result: + current_kv, shareable_kv = result[0] + else: + # graph replay: the cache already holds this step's write + current_kv, shareable_kv = past_kv, None next_key_values.append(current_kv if current_kv is not None else ()) # Only track the last sliding/global before the sharing boundary @@ -447,6 +638,14 @@ class Gemma4Transformer(nn.Module): if i == intermediate_output: intermediate = x.clone() + if prefetch_queue is not None: + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, None) + + if fixed_kv: + for kv in past_key_values: + if isinstance(kv, FixedKV): + kv.advance(seq_len) + if self.norm is not None: x = self.norm(x) @@ -481,14 +680,37 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module): 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 + logits.index_fill_(-1, self.model._suppress_tokens, torch.finfo(logits.dtype).min) return logits def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): - past_key_values = [] - for _ in range(self.model.config.num_hidden_layers): - past_key_values.append(()) - return past_key_values + cfg = self.model.config + num_layers = cfg.num_hidden_layers + if not self.model.fixed_kv: + return [() for _ in range(num_layers)] + first_shared = num_layers - cfg.num_kv_shared_layers if cfg.num_kv_shared_layers > 0 else num_layers + # position/seqlen device tensors are shared per cache geometry and filled once per step + trackers = {} + caches = [] + for i in range(num_layers): + if i >= first_shared: + caches.append(()) + continue + sliding = cfg.sliding_attention[i % len(cfg.sliding_attention)] if cfg.sliding_attention else False + head_dim = cfg.head_dim if sliding else cfg.global_head_dim + k_eq_v = cfg.attention_k_eq_v and not sliding + kv_heads = cfg.num_global_key_value_heads if k_eq_v else cfg.num_key_value_heads + length = min(sliding, max_cache_len) if sliding else max_cache_len + cache_cls = RingKV if sliding else FixedKV + tracker = trackers.get((cache_cls, length)) + if tracker is None: + tracker = (torch.empty((1,), device=device, dtype=torch.int64), + torch.empty((batch,), device=device, dtype=torch.int32)) + trackers[(cache_cls, length)] = tracker + # zero-init: decode attends full capacity with masked tails, 0*0 stays finite + key = torch.zeros((batch, kv_heads, length, head_dim), device=device, dtype=execution_dtype) + caches.append(cache_cls(key, torch.zeros_like(key), 0, tracker[0], tracker[1])) + return caches def preprocess_embed(self, embed, device): if embed["type"] == "image": @@ -1497,6 +1719,9 @@ def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=No if dtype_llama is not None: dtype = dtype_llama super().__init__(device=device, dtype=dtype, name="gemma4", clip_model=clip_model, model_options=model_options) + + def get_dynamic_vram__units(self): + return getattr(self, self.clip).transformer.model.get_dynamic_vram__units() return Gemma4TEModel_ From 1c6d8d45b3693bfbb32385b410d813a7fd6be216 Mon Sep 17 00:00:00 2001 From: "Daxiong (Lin)" Date: Sat, 15 Aug 2026 01:31:38 +0800 Subject: [PATCH 3/7] chore: update embedded docs to v0.5.10 (#15613) --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 18ce16a56..61577a6fd 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ comfyui-frontend-package==1.49.6 comfyui-workflow-templates==0.11.41 -comfyui-embedded-docs==0.5.9 +comfyui-embedded-docs==0.5.10 torch torchsde torchvision From 55b6a9b11dffecdd65a3ccd5eb6a1b3a178c96dc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Fri, 14 Aug 2026 22:17:18 +0300 Subject: [PATCH 4/7] Fix thinking handling (#15611) --- comfy/text_encoders/gemma4.py | 9 ++++++--- comfy_extras/nodes_textgen.py | 10 ++++------ 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 61bc3a3f0..da116ab27 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -6,6 +6,7 @@ import numpy as np from tokenizers import Tokenizer from dataclasses import dataclass import math +import re from comfy import sd1_clip import comfy.model_management @@ -1624,11 +1625,13 @@ class Gemma4SDTokenizer(Gemma4_Tokenizer, sd1_clip.SDTokenizer): def decode(self, token_ids, **kwargs): text = super().decode(token_ids, skip_special_tokens=False) - # Translate thinking channel markers to standard / tags + # Only a close that ends a thought channel becomes : generation primed with + # another channel leaves its opener in the prompt, so its close is not reasoning. + text = re.sub(r"<\|channel>thought\n(.*?)", r"\n\1", text, flags=re.DOTALL) text = text.replace("<|channel>thought\n", "\n") - text = text.replace("", "") # Strip remaining special tokens - text = text.replace("", "").replace("", "").strip() + text = re.sub(r"<\|channel>\w*\n?||<\|turn>\w*\n?|", "", text) + text = text.replace("", "").strip() return text diff --git a/comfy_extras/nodes_textgen.py b/comfy_extras/nodes_textgen.py index 40004652c..ef5858a7d 100644 --- a/comfy_extras/nodes_textgen.py +++ b/comfy_extras/nodes_textgen.py @@ -256,12 +256,10 @@ class TextGenerateLTX2Prompt(TextGenerate): out = super().execute(clip, formatted_prompt, max_length, sampling_mode, image=image, thinking=thinking, use_default_template=use_default_template, video=video, audio=audio) - text = out.args[0] - text = re.sub(r".*?", "", text, flags=re.DOTALL) - if "" in text: # unclosed/truncated reasoning: keep what follows the last close - text = text.rsplit("", 1)[-1] - text = re.sub(r"|<\|channel>\w*\n?||<\|turn>\w*\n?", "", text).strip() - return io.NodeOutput(text) + # Drop reasoning, including a block left unclosed by max_length. Both system prompts ask + # for the original prompt back when there is nothing to give; empty conditions on nothing. + text = re.sub(r".*?(?:|$)", "", out.args[0], flags=re.DOTALL).strip() + return io.NodeOutput(text or prompt) class TextgenExtension(ComfyExtension): From a7365071e47175fb06572d0a56d1bf4116c2f581 Mon Sep 17 00:00:00 2001 From: comfyanonymous <121283862+comfyanonymous@users.noreply.github.com> Date: Fri, 14 Aug 2026 17:34:40 -0700 Subject: [PATCH 5/7] Enable dynamic vram by default on ROCm 7.14 and higher. (#15633) --- main.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index eb64726ea..b6f7d6cef 100644 --- a/main.py +++ b/main.py @@ -248,7 +248,17 @@ import hook_breaker_ac10a0 import comfy.memory_management import comfy.model_patcher -if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_management.is_nvidia()): + +def dynamic_vram_supported(): + if comfy.model_management.is_nvidia(): + return True + if comfy.model_management.is_amd(): + if comfy.model_management.rocm_version >= (7, 14): + return True + return False + + +if args.enable_dynamic_vram or (enables_dynamic_vram() and dynamic_vram_supported()): if (not args.enable_dynamic_vram) and (comfy.model_management.torch_version_numeric < (2, 8)): logging.warning("Unsupported Pytorch detected. DynamicVRAM support requires Pytorch version 2.8 or later. Falling back to legacy ModelPatcher. VRAM estimates may be unreliable especially on Windows") else: From 0f1fa67ad8a68b62c65ebc97a7bf485df2459c3a Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Sat, 15 Aug 2026 08:37:17 +0400 Subject: [PATCH 6/7] [Partner Nodes] feat(ByteDance): add Seedance 2.5 task_type for video extension (#15579) Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/bytedance.py | 1 + comfy_api_nodes/nodes_bytedance.py | 171 ++++++++++++++++++++++++++--- 2 files changed, 154 insertions(+), 18 deletions(-) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 7ee83e5f3..64f8cde37 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -116,6 +116,7 @@ class Seedance2TaskCreationRequest(BaseModel): seed: int | None = Field(None, ge=0, le=2147483647) watermark: bool | None = Field(None) output_format: str | None = Field(None) + omni_reference_task_type: str | None = Field(None, description="One of: auto, reference, edit, extend.") class TaskCreationResponse(BaseModel): diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 09fc445d8..265f94d13 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -2069,7 +2069,7 @@ def _seedance2_text_inputs(resolutions: list[str], default_ratio: str = "16:9"): ] -def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = False): +def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = False, with_task_type: bool = False): return [ IO.String.Input( "prompt", @@ -2124,6 +2124,29 @@ def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = if with_video_editing else [] ), + *( + [ + IO.Combo.Input( + "task_type", + options=["auto", "reference", "edit", "extend"], + default="auto", + tooltip="What to do with the reference media. Every value except auto is " + "validated when the task is submitted, so mismatched settings fail before " + "generation starts. auto: the model infers the task from the prompt and " + "inputs, and settings that conflict with its reading fail only after " + "generation has started. reference: generate a new video guided by the " + "reference images, videos, and audio. edit: change a connected reference " + "video (add, remove, replace); the output keeps the source clip's own length " + "and aspect ratio, and the duration and ratio widgets are ignored. extend: " + "continue a connected reference video forward or backward; the prompt should " + "say 'extend forward', 'extend backward', or 'continue', the aspect ratio " + "follows the source clip, and the output contains only the newly generated " + "segment of the duration you set, not the source clip.", + ) + ] + if with_task_type + else [] + ), IO.Combo.Input( "output_format", options=["mp4"], @@ -2133,9 +2156,9 @@ def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = ] -def _seedance25_reference_inputs(): +def _seedance25_reference_inputs(with_video_editing: bool = False, with_task_type: bool = False): return [ - *_seedance25_text_inputs(with_video_editing=True), + *_seedance25_text_inputs(with_video_editing=with_video_editing, with_task_type=with_task_type), IO.Autogrow.Input( "reference_images", template=IO.Autogrow.TemplateNames( @@ -2196,17 +2219,23 @@ def _seedance2_build_request( watermark: bool, ratio: str, ) -> Seedance2TaskCreationRequest: - video_editing = bool(model.get("video_editing")) + task_type = model.get("task_type", "auto") + duration = model["duration"] + if model.get("video_editing") or task_type == "edit": + ratio, duration = "adaptive", -1 + elif task_type == "extend": + ratio = "adaptive" return Seedance2TaskCreationRequest( model=model_id, content=content, generate_audio=model["generate_audio"], resolution=model["resolution"], - ratio="adaptive" if video_editing else ratio, - duration=-1 if video_editing else model["duration"], + ratio=ratio, + duration=duration, seed=seed, watermark=watermark, output_format=model.get("output_format"), + omni_reference_task_type=None if task_type == "auto" else task_type, ) @@ -2216,7 +2245,7 @@ _SEEDANCE2_PRICE_EXPR_TEMPLATE = """ $res := $lookup(widgets, "model.resolution"); $ratio := $lookup(widgets, "model.ratio"); $dur := $lookup(widgets, "model.duration"); - $auto := $lookup(widgets, "model.video_editing") = true; + $auto := __IS_EDIT__; $hasVideo := __HAS_VIDEO__; $ready := $type($m) = "string" and $type($res) = "string" and ($auto or $type($dur) = "number"); $ready ? ( @@ -2261,6 +2290,7 @@ _SEEDANCE2_PRICE_EXPR_TEMPLATE = """ _SEEDANCE_AUDIO_POLICY_CODE = "OutputAudioSensitiveContentDetected.PolicyViolation" _SEEDANCE_TASK_TYPE_CONSTRAINT_CODE = "InvalidParameter.TaskTypeConstraint" +_SEEDANCE_TASK_TYPE_MISMATCH_CODE = "InvalidParameter.TaskTypeMismatch" async def _seedance2_poll_video_task( @@ -2269,6 +2299,7 @@ async def _seedance2_poll_video_task( model_id: str, resolution: str, has_video_input: bool, + task_type: str | None = None, ) -> TaskStatusResponse: try: return await poll_op( @@ -2289,19 +2320,48 @@ async def _seedance2_poll_video_task( "to get a silent video, or adjust the prompt and try again." ) from exc if _SEEDANCE_TASK_TYPE_CONSTRAINT_CODE in str(exc): + if task_type is None: + raise ValueError( + "Seedance read this prompt as editing the reference video, and an edit always " + "takes its duration and aspect ratio from that video. Enable video_editing on " + "this node and run again, or reword the prompt so it describes a new video " + "rather than a change to the reference one." + ) from exc + if task_type == "edit": + raise ValueError( + "The request does not satisfy the 'edit' constraints: the clip being edited " + "must be 4 to 30 seconds long." + ) from exc + if task_type == "extend": + raise ValueError( + "The request does not satisfy the 'extend' constraints: the clip being " + "extended must be 1.9 to 30 seconds long." + ) from exc raise ValueError( - "Seedance read this prompt as editing the reference video, and an edit always " - "takes its duration and aspect ratio from that video. Enable video_editing on " - "this node and run again, or reword the prompt so it describes a new video " - "rather than a change to the reference one." + "Seedance decided from the prompt that this task's duration or aspect ratio " + "must come from the reference video, and the current settings conflict with " + "that. Set task_type to the task you mean ('edit' or 'extend') and run again, " + "or reword the prompt so it describes a new video rather than a change to the " + "reference one." + ) from exc + if _SEEDANCE_TASK_TYPE_MISMATCH_CODE in str(exc): + raise ValueError( + f"Seedance read this prompt as a different task than the selected task_type " + f"'{task_type}'. Reword the prompt so it matches: an extend prompt should say " + "'extend forward', 'extend backward', or 'continue'; an edit prompt should use " + "words like add, remove, replace, or change. Or set task_type to auto." ) from exc raise -def _seedance2_price_badge(with_reference_videos: bool) -> IO.PriceBadge: +def _seedance2_price_badge(with_reference_videos: bool, legacy_video_editing: bool = False) -> IO.PriceBadge: widgets = ["model", "model.resolution", "model.ratio", "model.duration"] + if legacy_video_editing: + is_edit = '$lookup(widgets, "model.video_editing") = true' + else: + is_edit = '$lookup(widgets, "model.task_type") = "edit"' if with_reference_videos: - widgets.append("model.video_editing") + widgets.append("model.video_editing" if legacy_video_editing else "model.task_type") has_video = ( '$exists(inputGroups) and $lookup(inputGroups, "model.reference_videos") > 0' if with_reference_videos @@ -2312,7 +2372,7 @@ def _seedance2_price_badge(with_reference_videos: bool) -> IO.PriceBadge: widgets=widgets, input_groups=["model.reference_videos"] if with_reference_videos else [], ), - expr=_SEEDANCE2_PRICE_EXPR_TEMPLATE.replace("__HAS_VIDEO__", has_video), + expr=_SEEDANCE2_PRICE_EXPR_TEMPLATE.replace("__HAS_VIDEO__", has_video).replace("__IS_EDIT__", is_edit), ) @@ -2662,12 +2722,12 @@ def _seedance2_reference_inputs(resolutions: list[str], default_ratio: str = "16 ] -class ByteDance2ReferenceNode(IO.ComfyNode): +class ByteDance2ReferenceNodeV2(IO.ComfyNode): @classmethod def define_schema(cls): return IO.Schema( - node_id="ByteDance2ReferenceNode", + node_id="ByteDance2ReferenceNodeV2", display_name="ByteDance Seedance 2.5 Reference to Video", category="partner/video/ByteDance", description="Generate, edit, or extend video using Seedance 2.5 or 2.0 with reference " @@ -2676,7 +2736,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs()), + IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs(with_task_type=True)), IO.DynamicCombo.Option( "Seedance 2.0", _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), @@ -2761,6 +2821,13 @@ class ByteDance2ReferenceNode(IO.ComfyNode): f"(videos={len(reference_videos)}, video assets={len(reference_video_assets)}). " f"Maximum is {limits['max_videos']}." ) + task_type = model.get("task_type") + if task_type in ("edit", "extend") and total_videos == 0: + raise ValueError( + f"A '{task_type}' task needs at least one reference video. Connect the video " + f"you want to {'change' if task_type == 'edit' else 'continue'}, or set " + "task_type to 'reference' to generate a new video from the references you have." + ) total_audios = len(reference_audios) + len(reference_audio_assets) if total_audios > limits["max_audios"]: raise ValueError( @@ -2893,11 +2960,78 @@ class ByteDance2ReferenceNode(IO.ComfyNode): response_model=TaskCreationResponse, ) response = await _seedance2_poll_video_task( - cls, initial_response.id, model_id, model["resolution"], has_video_input=has_video_input + cls, + initial_response.id, + model_id, + model["resolution"], + has_video_input=has_video_input, + task_type=task_type, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) +class ByteDance2ReferenceNode(ByteDance2ReferenceNodeV2): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="ByteDance2ReferenceNode", + display_name="ByteDance Seedance 2.5 Reference to Video (Legacy)", + category="partner/video/ByteDance", + description="Generate, edit, or extend video using Seedance 2.5 or 2.0 with reference " + "images, videos, and audio. Supports multimodal reference, video editing, and video extension.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs(with_video_editing=True)), + IO.DynamicCombo.Option( + "Seedance 2.0", + _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), + ), + IO.DynamicCombo.Option( + "Seedance 2.0 Fast", + _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), + ), + IO.DynamicCombo.Option( + "Seedance 2.0 Mini", + _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), + ), + ], + tooltip=SEEDANCE_MODEL_TOOLTIP, + ), + IO.Int.Input( + "seed", + default=0, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add a watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + is_deprecated=True, + price_badge=_seedance2_price_badge(with_reference_videos=True, legacy_video_editing=True), + ) + + async def process_video_task( cls: type[IO.ComfyNode], payload: Text2VideoTaskCreationRequest | Image2VideoTaskCreationRequest, @@ -3413,6 +3547,7 @@ class ByteDanceExtension(ComfyExtension): ByteDance2TextToVideoNode, ByteDance2FirstLastFrameNode, ByteDance2ReferenceNode, + ByteDance2ReferenceNodeV2, ByteDanceCreateImageAsset, ByteDanceCreateVideoAsset, ByteDanceSeedAudioNode, From a9ab2b62dac17743c62fc1cab8468340496d302e Mon Sep 17 00:00:00 2001 From: Christian Byrne Date: Sat, 15 Aug 2026 10:27:24 -0700 Subject: [PATCH 7/7] [Partner Nodes] Stop adding an opaque alpha channel to API node images (#15369) * Stop adding an opaque alpha channel to API node images bytesio_to_image_tensor converted every downloaded image to RGBA, so nodes whose API returns no transparency still emitted a 4 channel IMAGE. Keep the alpha when the decoded image has one, stay RGB when it does not. --------- Signed-off-by: bigcat88 Co-authored-by: bigcat88 --- comfy_api_nodes/nodes_bytedance.py | 6 +- comfy_api_nodes/nodes_gemini.py | 5 +- comfy_api_nodes/nodes_recraft.py | 11 +-- comfy_api_nodes/util/__init__.py | 2 + comfy_api_nodes/util/conversions.py | 19 ++++- .../gemini_response_images_test.py | 57 +++++++++++++ .../image_conversions_test.py | 80 +++++++++++++++++++ 7 files changed, 170 insertions(+), 10 deletions(-) create mode 100644 tests-unit/comfy_api_nodes_test/gemini_response_images_test.py create mode 100644 tests-unit/comfy_api_nodes_test/image_conversions_test.py diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 265f94d13..6040075f8 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -56,6 +56,8 @@ from comfy_api_nodes.util import ( ApiEndpoint, audio_bytes_to_audio_input, audio_input_to_mp3, + bytesio_to_image_tensor, + download_url_as_bytesio, download_url_to_image_tensor, download_url_to_video_output, downscale_image_tensor_by_max_side, @@ -1315,7 +1317,9 @@ class ByteDanceSeedreamLayerSeparationNode(IO.ComfyNode): left, top, rect_w, rect_h = spec["left"], spec["top"], spec["rect_w"], spec["rect_h"] async with semaphore: try: - rgba = (await download_url_to_image_tensor(str(item["url"])))[0] + # the layer math below needs the alpha channel, and ByteDance encodes + # alpha-less images as plain RGB (the base plate is one), so force RGBA + rgba = bytesio_to_image_tensor(await download_url_as_bytesio(str(item["url"])), mode="RGBA")[0] except ProcessingInterrupted: raise except Exception as exc: diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 131590751..8343633e0 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -43,6 +43,7 @@ from comfy_api_nodes.util import ( download_url_to_image_tensor, download_url_to_video_output, get_number_of_images, + pad_images_to_common_channels, sync_op, tensor_to_base64_string, upload_audio_to_comfyapi, @@ -233,8 +234,8 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug "Try rephrasing your prompt or changing the response modality to 'IMAGE+TEXT' " "to see the model's reasoning." ) - return torch.zeros((1, 1024, 1024, 4)) - return torch.cat(image_tensors, dim=0) + return torch.zeros((1, 1024, 1024, 3)) + return torch.cat(pad_images_to_common_channels(image_tensors), dim=0) def get_text_from_interaction(interaction: GeminiInteraction) -> str: diff --git a/comfy_api_nodes/nodes_recraft.py b/comfy_api_nodes/nodes_recraft.py index 2605b9021..9f1823426 100644 --- a/comfy_api_nodes/nodes_recraft.py +++ b/comfy_api_nodes/nodes_recraft.py @@ -27,6 +27,7 @@ from comfy_api_nodes.util import ( ApiEndpoint, bytesio_to_image_tensor, download_url_as_bytesio, + pad_images_to_common_channels, resize_mask_to_image, sync_op, tensor_to_bytesio, @@ -621,7 +622,7 @@ class RecraftImageToImageNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftImageInpaintingNode(IO.ComfyNode): @@ -723,7 +724,7 @@ class RecraftImageInpaintingNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftTextToVectorNode(IO.ComfyNode): @@ -954,7 +955,7 @@ class RecraftReplaceBackgroundNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftRemoveBackgroundNode(IO.ComfyNode): @@ -995,7 +996,7 @@ class RecraftRemoveBackgroundNode(IO.ComfyNode): image=image[i], path="/proxy/recraft/images/removeBackground", ) - images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) + images.append(torch.cat([bytesio_to_image_tensor(x, mode="RGBA") for x in sub_bytes], dim=0)) pbar.update(1) images_tensor = torch.cat(images, dim=0) @@ -1047,7 +1048,7 @@ class RecraftCrispUpscaleNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftCreativeUpscaleNode(RecraftCrispUpscaleNode): diff --git a/comfy_api_nodes/util/__init__.py b/comfy_api_nodes/util/__init__.py index 1fb6b96cf..2bb4a1b04 100644 --- a/comfy_api_nodes/util/__init__.py +++ b/comfy_api_nodes/util/__init__.py @@ -18,6 +18,7 @@ from .conversions import ( downscale_image_tensor_by_max_side, downscale_video_to_max_pixels, image_tensor_pair_to_batch, + pad_images_to_common_channels, pil_to_bytesio, resize_mask_to_image, tensor_to_base64_string, @@ -92,6 +93,7 @@ __all__ = [ "downscale_image_tensor_by_max_side", "downscale_video_to_max_pixels", "image_tensor_pair_to_batch", + "pad_images_to_common_channels", "pil_to_bytesio", "resize_mask_to_image", "tensor_to_base64_string", diff --git a/comfy_api_nodes/util/conversions.py b/comfy_api_nodes/util/conversions.py index f46cac3f8..eb81447a0 100644 --- a/comfy_api_nodes/util/conversions.py +++ b/comfy_api_nodes/util/conversions.py @@ -16,12 +16,14 @@ from comfy_api.latest import Input, InputImpl, Types from ._helpers import mimetype_to_extension -def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor: +def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str | None = None) -> torch.Tensor: """Converts image data from BytesIO to a torch.Tensor. Args: image_bytesio: BytesIO object containing the image data. - mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA"). + mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA"). Defaults + to RGBA when the decoded image carries transparency and RGB when it + does not, so an API that returns no alpha does not get an opaque one. Returns: A torch.Tensor representing the image (1, H, W, C). @@ -31,6 +33,8 @@ def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch ValueError: If the specified mode is invalid. """ image = Image.open(image_bytesio) + if mode is None: + mode = "RGBA" if "A" in image.getbands() or "transparency" in image.info else "RGB" image = image.convert(mode) image_array = np.array(image).astype(np.float32) / 255.0 return torch.from_numpy(image_array).unsqueeze(0) @@ -53,6 +57,17 @@ def image_tensor_pair_to_batch(image1: torch.Tensor, image2: torch.Tensor) -> to return torch.cat((image1, image2), dim=0) +def pad_images_to_common_channels(images: list[torch.Tensor]) -> list[torch.Tensor]: + """Pads [B, H, W, C] image tensors with opaque alpha so they all share the largest channel count.""" + channels = max(image.shape[-1] for image in images) + return [ + torch.nn.functional.pad(image, (0, channels - image.shape[-1]), value=1.0) + if image.shape[-1] < channels + else image + for image in images + ] + + def tensor_to_bytesio( image: torch.Tensor, *, diff --git a/tests-unit/comfy_api_nodes_test/gemini_response_images_test.py b/tests-unit/comfy_api_nodes_test/gemini_response_images_test.py new file mode 100644 index 000000000..7b4708b43 --- /dev/null +++ b/tests-unit/comfy_api_nodes_test/gemini_response_images_test.py @@ -0,0 +1,57 @@ +import asyncio +import base64 +from io import BytesIO + +import torch +from PIL import Image + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +from comfy_api_nodes.apis.gemini import ( # noqa: E402 + GeminiCandidate, + GeminiContent, + GeminiGenerateContentResponse, + GeminiInlineData, + GeminiPart, +) +from comfy_api_nodes.nodes_gemini import get_image_from_response # noqa: E402 + + +def image_part(mode, color): + buffer = BytesIO() + Image.new(mode, (4, 4), color).save(buffer, format="PNG") + return GeminiPart( + inlineData=GeminiInlineData( + data=base64.b64encode(buffer.getvalue()).decode(), + mimeType="image/png", + ) + ) + + +def response(*parts): + return GeminiGenerateContentResponse( + candidates=[GeminiCandidate(content=GeminiContent(parts=list(parts), role="model"))] + ) + + +def test_rgb_only_response_stays_three_channels(): + out = asyncio.run(get_image_from_response(response(image_part("RGB", (10, 20, 30))))) + assert out.shape == (1, 4, 4, 3) + + +def test_mixed_rgb_and_rgba_parts_are_padded_to_the_same_width(): + out = asyncio.run( + get_image_from_response( + response( + image_part("RGB", (10, 20, 30)), + image_part("RGBA", (10, 20, 30, 0)), + ) + ) + ) + assert out.shape == (2, 4, 4, 4) + # the part that had no alpha is padded opaque, the transparent one is preserved + assert out[0, ..., 3].min() == 1.0 + assert out[1, ..., 3].max() == 0.0 diff --git a/tests-unit/comfy_api_nodes_test/image_conversions_test.py b/tests-unit/comfy_api_nodes_test/image_conversions_test.py new file mode 100644 index 000000000..3caff2818 --- /dev/null +++ b/tests-unit/comfy_api_nodes_test/image_conversions_test.py @@ -0,0 +1,80 @@ +from io import BytesIO + +import pytest +import torch +from PIL import Image + +from comfy.cli_args import args + +if not torch.cuda.is_available(): + args.cpu = True + +from comfy_api_nodes.util.conversions import bytesio_to_image_tensor, pad_images_to_common_channels # noqa: E402 + + +def encode(image: Image.Image, image_format: str = "PNG") -> BytesIO: + buffer = BytesIO() + image.save(buffer, format=image_format) + buffer.seek(0) + return buffer + + +def test_rgb_png_stays_three_channels(): + tensor = bytesio_to_image_tensor(encode(Image.new("RGB", (4, 4), (10, 20, 30)))) + assert tensor.shape == (1, 4, 4, 3) + + +def test_jpeg_stays_three_channels(): + tensor = bytesio_to_image_tensor(encode(Image.new("RGB", (4, 4), (10, 20, 30)), "JPEG")) + assert tensor.shape == (1, 4, 4, 3) + + +def test_grayscale_is_expanded_to_rgb(): + tensor = bytesio_to_image_tensor(encode(Image.new("L", (4, 4), 128))) + assert tensor.shape == (1, 4, 4, 3) + + +def test_rgba_png_keeps_its_alpha(): + tensor = bytesio_to_image_tensor(encode(Image.new("RGBA", (4, 4), (10, 20, 30, 0)))) + assert tensor.shape == (1, 4, 4, 4) + assert tensor[..., 3].max() == 0.0 + + +def test_palette_png_with_transparency_keeps_its_alpha(): + image = Image.new("P", (4, 4), 1) + image.putpalette([0, 0, 0, 255, 255, 255]) + image.info["transparency"] = 0 + image.putpixel((0, 0), 0) + tensor = bytesio_to_image_tensor(encode(image)) + assert tensor.shape == (1, 4, 4, 4) + assert tensor[0, 0, 0, 3] == 0.0 + assert tensor[0, 1, 1, 3] == 1.0 + + +@pytest.mark.parametrize("mode,channels", [("RGB", 3), ("RGBA", 4)]) +def test_explicit_mode_is_respected(mode, channels): + tensor = bytesio_to_image_tensor(encode(Image.new("RGBA", (4, 4), (10, 20, 30, 128))), mode=mode) + assert tensor.shape == (1, 4, 4, channels) + + +def test_pad_mixed_channels_concatenates(): + rgb = torch.rand(1, 4, 4, 3) + rgba = torch.rand(2, 4, 4, 4) + padded = pad_images_to_common_channels([rgb, rgba]) + result = torch.cat(padded, dim=0) + assert result.shape == (3, 4, 4, 4) + + +def test_pad_adds_opaque_alpha_and_keeps_rgb_values(): + rgb = torch.rand(1, 4, 4, 3) + rgba = torch.rand(1, 4, 4, 4) + padded_rgb, padded_rgba = pad_images_to_common_channels([rgb, rgba]) + assert torch.equal(padded_rgb[..., :3], rgb) + assert padded_rgb[..., 3].min() == 1.0 + assert padded_rgba is rgba + + +def test_pad_leaves_homogeneous_channels_unchanged(): + images = [torch.rand(1, 4, 4, 3), torch.rand(2, 4, 4, 3)] + padded = pad_images_to_common_channels(images) + assert all(p is i for p, i in zip(padded, images))