Merge remote-tracking branch 'upstream/master' into pixal3d

This commit is contained in:
kijai 2026-07-30 13:42:56 +03:00
commit 87e34c5c17
74 changed files with 3451 additions and 650 deletions

View File

@ -23,9 +23,9 @@ jobs:
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
# upstream changes; keep `workflows_ref` matching so prompts/scripts load
# from the same commit as the workflow definition.
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
with:
workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a
workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
diff_excludes: >-
:!**/.claude/**
:!**/dist/**

View File

@ -20,7 +20,7 @@ jobs:
git_tag: ${{ inputs.git_tag }}
cache_tag: "cu130"
python_minor: "13"
python_patch: "12"
python_patch: "14"
rel_name: "nvidia"
rel_extra_name: ""
test_release: true
@ -71,7 +71,7 @@ jobs:
git_tag: ${{ inputs.git_tag }}
cache_tag: "xpu"
python_minor: "13"
python_patch: "12"
python_patch: "14"
rel_name: "intel"
rel_extra_name: ""
test_release: true

View File

@ -162,8 +162,26 @@
adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`,
`comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and
`comfy-kitchen` helpers where they already solve the problem.
- Use optimized comfy-kitchen ops in places where they improve performance
without changing the expected dtype, device, memory, or interface behavior.
- Model implementations must use an existing optimized Comfy Kitchen or
ComfyUI operation whenever one supports the required math and tensor layout
without changing expected dtype, device, memory, or interface behavior. This
is the default implementation requirement, not an optional follow-up
optimization.
- Before implementing model math, inspect the operations already exposed by
Comfy Kitchen, `comfy.quant_ops`, and existing ComfyUI model helpers. Check
for optimized single, paired, fused, layout-specific, and quantized variants
before writing a local implementation or composing lower-level torch ops.
- Use the compatible optimized operation first and adapt the model's inputs to
its documented layout while preserving the model's exact math. If several
optimized variants apply, benchmark representative model shapes and select
the fastest valid path.
- Add or retain a local implementation only when no existing optimized
operation supports the required math, layout, dtype, device, autograd, or
patch contract. Keep differentiable or patch-compatible fallbacks when the
optimized inference operation does not provide those contracts.
- Use the existing ComfyUI cast, offload, and cleanup helpers for parameters
passed to optimized operations. Preserve model-specific epsilon, scaling,
layout, dtype, device, and output-shape behavior.
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
handwritten implementations of the same operation. Remove duplicate local
kernels and adapt inputs to the shared operation's documented layout while

View File

@ -1,5 +1,6 @@
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
/CODEOWNERS @comfyanonymous
/AGENTS.md @comfyanonymous
/.ci/ @comfyanonymous
/.github/ @comfyanonymous

View File

@ -315,15 +315,29 @@ async def download_asset_content(request: web.Request) -> web.Response:
404, "FILE_NOT_FOUND", "Underlying file not found on disk."
)
# User-controlled asset content must never render inline in the app origin
# User-controlled asset content must not render inline in the app origin
# (stored XSS via SVG/HTML/XML). Force dangerous types to download and
# override any requested inline disposition. Centralised through
# folder_paths.is_dangerous_content_type so this can't drift from /view and
# /userdata (the previous inline set here omitted image/svg+xml and missed
# the charset/casing/+xml-dialect bypasses).
# override any requested inline disposition; SVG loaded into an <img> is
# exempt, see renders_safely_as_image. Centralised through folder_paths so
# this can't drift from /view and /userdata (the previous inline set here
# omitted image/svg+xml and missed the charset/casing/+xml-dialect bypasses).
extra_headers = {}
sec_fetch_dest = request.headers.get("Sec-Fetch-Dest")
if folder_paths.is_dangerous_content_type(content_type):
content_type = "application/octet-stream"
disposition = "attachment"
# This response now depends on a request header, so it must not be
# reused across destinations by a browser or intermediary cache: an
# inline SVG primed by an <img> fetch and replayed to a document
# navigation of the same URL would re-enable the stored XSS.
extra_headers["Vary"] = "Sec-Fetch-Dest"
extra_headers["Cache-Control"] = "no-store"
if not folder_paths.renders_safely_as_image(content_type, sec_fetch_dest):
content_type = "application/octet-stream"
disposition = "attachment"
# mime_type is uploader-supplied and unvalidated, so it can carry
# parameters. aiohttp rejects a charset in the content_type argument with
# ValueError, which would turn a valid inline SVG into a 500.
content_type = content_type.split(";", 1)[0].strip() or "application/octet-stream"
safe_name = (filename or "").replace("\r", "").replace("\n", "")
encoded = urllib.parse.quote(safe_name)
@ -356,6 +370,7 @@ async def download_asset_content(request: web.Request) -> web.Response:
"Content-Disposition": cd,
"Content-Length": str(file_size),
"X-Content-Type-Options": "nosniff",
**extra_headers,
},
)

View File

@ -2,9 +2,12 @@ from collections import deque
from datetime import datetime
import io
import logging
import os
import sys
import threading
import comfy.logging
ANSI_NAMED_COLORS = {
'black': '\033[30m',
'red': '\033[31m',
@ -18,6 +21,7 @@ ANSI_NAMED_COLORS = {
ANSI_LEVEL_COLORS = {
'DEBUG': ANSI_NAMED_COLORS['cyan'],
'DETAIL': ANSI_NAMED_COLORS['blue'],
'INFO': ANSI_NAMED_COLORS['green'],
'WARNING': ANSI_NAMED_COLORS['yellow'],
'ERROR': ANSI_NAMED_COLORS['red'],
@ -85,7 +89,12 @@ def on_flush(callback):
if stderr_interceptor is not None:
stderr_interceptor.on_flush(callback)
def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool = False):
def get_log_level(level):
return comfy.logging.DETAIL if level == "DETAIL" else logging.getLevelName(level)
def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300, use_stdout: bool = False):
global logs
if logs:
return
@ -99,13 +108,18 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool
stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr)
# Setup default global logger
if file_outputs is None:
file_outputs = [('DETAIL', 'comfyui_detail.log')]
logger = logging.getLogger()
logger.setLevel(log_level)
console_level = get_log_level(log_level)
file_levels = [get_log_level(level) for level, _ in file_outputs]
logger.setLevel(min(console_level, *file_levels))
formatter = ColoredFormatter("%(message)s")
stream_handler = logging.StreamHandler()
stream_handler.setFormatter(formatter)
stream_handler.setLevel(console_level)
if use_stdout:
# Only errors and critical to stderr
@ -114,11 +128,24 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool
# Lesser to stdout
stdout_handler = logging.StreamHandler(sys.stdout)
stdout_handler.setFormatter(formatter)
stdout_handler.setLevel(console_level)
stdout_handler.addFilter(lambda record: record.levelno < logging.ERROR)
logger.addHandler(stdout_handler)
logger.addHandler(stream_handler)
for output_level, output_path in file_outputs:
output_path = os.path.abspath(output_path)
try:
output_handler = logging.FileHandler(output_path, encoding="utf-8")
except OSError as e:
logging.warning("Could not open %s log %s: %s", output_level, output_path, e)
continue
output_handler.setLevel(get_log_level(output_level))
output_handler.setFormatter(logging.Formatter("[%(asctime)s] [%(levelname)s] %(message)s"))
logger.addHandler(output_handler)
logging.info("%s log: %s", output_level.title(), output_path)
STARTUP_WARNINGS = []

View File

@ -343,13 +343,22 @@ class UserManager():
# XSS). Content-Disposition: attachment is the load-bearing guard;
# the content-type override and nosniff are defence in depth.
content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream'
if folder_paths.is_dangerous_content_type(content_type):
content_type = 'application/octet-stream'
user_root = self.get_request_user_filepath(request, None, create_dir=False)
is_user_css = path == os.path.abspath(os.path.join(user_root, "user.css"))
if is_user_css:
content_type = "text/css"
disposition = "inline"
else:
if folder_paths.is_dangerous_content_type(content_type):
content_type = 'application/octet-stream'
disposition = "attachment"
return web.FileResponse(path, headers={
"Content-Type": content_type,
"X-Content-Type-Options": "nosniff",
"Content-Disposition": "attachment",
"Content-Disposition": disposition,
})
@routes.post("/userdata/{file}")

View File

@ -33,6 +33,31 @@ class EnumAction(argparse.Action):
setattr(namespace, self.dest, value)
LOG_LEVELS = ('DEBUG', 'DETAIL', 'INFO', 'WARNING', 'ERROR', 'CRITICAL')
class VerboseAction(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None):
if len(values) == 0:
output = ('DEBUG', None)
elif len(values) == 1 and values[0] in LOG_LEVELS:
output = (values[0], None)
elif len(values) == 2 and values[0] in LOG_LEVELS:
output = tuple(values)
else:
parser.error(f"{option_string} expects no values, a console LEVEL, or LEVEL FILE")
setattr(namespace, self.dest, [*getattr(namespace, self.dest, []), output])
def get_console_log_level(outputs):
console_levels = [level for level, path in outputs if path is None]
return min(console_levels, key=LOG_LEVELS.index, default='INFO')
def get_file_log_outputs(outputs):
return [(level, path) for level, path in outputs if path is not None]
parser = argparse.ArgumentParser()
parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)")
@ -112,7 +137,7 @@ parser.add_argument("--preview-method", type=LatentPreviewMethod, default=Latent
parser.add_argument("--preview-size", type=int, default=512, help="Sets the maximum preview size for sampler nodes.")
cache_group = parser.add_mutually_exclusive_group()
cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 96GB).")
cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 128GB).")
cache_group.add_argument("--cache-classic", action="store_true", help="Use the old style (aggressive) caching.")
cache_group.add_argument("--cache-lru", type=int, default=0, help="Use LRU caching with a maximum of N node results cached. May use more RAM/VRAM.")
cache_group.add_argument("--cache-none", action="store_true", help="Reduced RAM/VRAM usage at the expense of executing every node for each run.")
@ -187,7 +212,7 @@ parser.add_argument("--disable-api-nodes", action="store_true", help="Disable lo
parser.add_argument("--multi-user", action="store_true", help="Enables per-user storage.")
parser.add_argument("--verbose", default='INFO', const='DEBUG', nargs="?", choices=['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'], help='Set the logging level')
parser.add_argument("--verbose", action=VerboseAction, nargs='*', default=[], metavar='LEVEL FILE', help='Set console logging with no values or LEVEL, or add a LEVEL FILE log output. May be repeated.')
parser.add_argument("--log-stdout", action="store_true", help="Send normal process output to stdout instead of stderr (default).")

View File

@ -5,6 +5,7 @@ import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention
import comfy.model_management
import comfy.ops
import comfy.quant_ops
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module):
query = q_flat.view(B, S, self.heads, self.head_dim)
key = k_flat.view(B, S, self.heads, self.head_dim)
query = self.norm_q(query)
key = self.norm_k(key)
if image_rotary_emb is not None:
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
if image_rotary_emb is not None and not comfy.model_management.in_training:
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True)
query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps)
comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if image_rotary_emb is not None:
query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
q_flat = query.reshape(B, S, -1)
k_flat = key.reshape(B, S, -1)

View File

@ -12,10 +12,13 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
import comfy.model_management
import comfy.ops
import comfy.patcher_extension
import comfy.quant_ops
from comfy.ldm.lumina.model import FeedForward
from comfy.ldm.modules.attention import optimized_attention_masked
from comfy.text_encoders.llama import apply_rope, precompute_freqs_cis
from comfy.text_encoders.llama import precompute_freqs_cis
# Per-token role indicators
SEQUENCE_PADDING_INDICATOR = -1
@ -25,6 +28,22 @@ LLM_TOKEN_INDICATOR = 3
IMAGE_POSITION_OFFSET = 65536
def _split_half_rope_matrix(freqs_cis):
cos, sin, neg_sin = freqs_cis
half_dim = sin.shape[-1]
matrix = torch.stack(
(cos[..., :half_dim], neg_sin, sin, cos[..., half_dim:]), dim=-1
)
return matrix.reshape(*matrix.shape[:-1], 2, 2).unsqueeze(2)
def _apply_rope_split_half1(x, freqs_cis):
x_dtype = x.dtype
x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(freqs_cis.dtype)
output = freqs_cis[..., 0] * x[..., 0] + freqs_cis[..., 1] * x[..., 1]
return output.movedim(-1, -2).reshape(*x.shape[:-3], -1).to(x_dtype)
class Ideogram4Attention(nn.Module):
def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None):
super().__init__()
@ -42,16 +61,23 @@ class Ideogram4Attention(nn.Module):
qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2)
q = self.norm_q(q)
k = self.norm_k(k)
if comfy.model_management.in_training:
q = _apply_rope_split_half1(self.norm_q(q), freqs_cis)
k = _apply_rope_split_half1(self.norm_k(k), freqs_cis)
else:
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, q, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, k, offloadable=True)
q, k = comfy.quant_ops.ck.rms_rope_split_half(
q, k, freqs_cis, q_scale, k_scale, self.norm_q.eps
)
comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
# (B, heads, L, head_dim)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
q, k = apply_rope(q, k, freqs_cis)
out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options)
return self.o(out)
@ -181,6 +207,7 @@ class Ideogram4Transformer(nn.Module):
self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta,
rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device,
)
freqs_cis = _split_half_rope_matrix(freqs_cis)
if attn_mask is not None and attn_mask.dtype == torch.bool:
attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max)

View File

@ -94,12 +94,21 @@ class JoyImageAttention(nn.Module):
txt_k = txt_k.unflatten(-1, (heads, -1))
txt_v = txt_v.unflatten(-1, (heads, -1))
img_q = self.img_attn_q_norm(img_q)
img_k = self.img_attn_k_norm(img_k)
txt_q = self.txt_attn_q_norm(txt_q)
txt_k = self.txt_attn_k_norm(txt_k)
img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb)
img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True)
img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True)
img_q, img_k = comfy_kitchen.rms_rope(
img_q,
img_k,
image_rotary_emb,
img_q_scale,
img_k_scale,
self.img_attn_q_norm.eps,
)
comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream)
comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream)
joint_q = torch.cat([img_q, txt_q], dim=1)
joint_k = torch.cat([img_k, txt_k], dim=1)

View File

@ -6,9 +6,8 @@ import torch
from comfy.ldm.lightricks.model import (
CrossAttention,
FeedForward,
freqs_cis_matrix,
generate_freq_grid_np,
interleaved_freqs_cis,
split_freqs_cis,
)
from torch import nn
@ -244,12 +243,15 @@ class Embeddings1DConnector(nn.Module):
expected_freqs = dim // 2
current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs
cos_freq, sin_freq = split_freqs_cis(
freqs, pad_size, self.num_attention_heads
)
else:
cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem)
return cos_freq.to(dtype=out_dtype), sin_freq.to(dtype=out_dtype), self.split_rope
pad_size = dim % n_elem
return freqs_cis_matrix(
freqs,
pad_size,
self.split_rope,
self.num_attention_heads,
out_dtype,
)
def forward(
self,

View File

@ -97,11 +97,11 @@ class SpatialRationalResampler(nn.Module):
For dims==3, work per-frame for spatial scaling (temporal axis untouched).
"""
def __init__(self, mid_channels: int, scale: float):
def __init__(self, mid_channels: int, scale: float, operations):
super().__init__()
self.scale = float(scale)
self.num, self.den = _rational_for_scale(self.scale)
self.conv = nn.Conv2d(
self.conv = operations.Conv2d(
mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1
)
self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num))
@ -119,18 +119,18 @@ class SpatialRationalResampler(nn.Module):
class ResBlock(nn.Module):
def __init__(
self, channels: int, mid_channels: Optional[int] = None, dims: int = 3
self, channels: int, operations, mid_channels: Optional[int] = None, dims: int = 3
):
super().__init__()
if mid_channels is None:
mid_channels = channels
Conv = nn.Conv2d if dims == 2 else nn.Conv3d
Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1)
self.norm1 = nn.GroupNorm(32, mid_channels)
self.norm1 = operations.GroupNorm(32, mid_channels)
self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1)
self.norm2 = nn.GroupNorm(32, channels)
self.norm2 = operations.GroupNorm(32, channels)
self.activation = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
@ -159,6 +159,7 @@ class LatentUpsampler(nn.Module):
def __init__(
self,
operations,
in_channels: int = 128,
mid_channels: int = 512,
num_blocks_per_stage: int = 4,
@ -179,34 +180,34 @@ class LatentUpsampler(nn.Module):
self.spatial_scale = float(spatial_scale)
self.rational_resampler = rational_resampler
Conv = nn.Conv2d if dims == 2 else nn.Conv3d
Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1)
self.initial_norm = nn.GroupNorm(32, mid_channels)
self.initial_norm = operations.GroupNorm(32, mid_channels)
self.initial_activation = nn.SiLU()
self.res_blocks = nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
[ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
)
if spatial_upsample and temporal_upsample:
self.upsampler = nn.Sequential(
nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
operations.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(3),
)
elif spatial_upsample:
if rational_resampler:
self.upsampler = SpatialRationalResampler(
mid_channels=mid_channels, scale=self.spatial_scale
mid_channels=mid_channels, scale=self.spatial_scale, operations=operations
)
else:
self.upsampler = nn.Sequential(
nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
operations.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(2),
)
elif temporal_upsample:
self.upsampler = nn.Sequential(
nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
operations.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(1),
)
else:
@ -215,11 +216,14 @@ class LatentUpsampler(nn.Module):
)
self.post_upsample_res_blocks = nn.ModuleList(
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
[ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
)
self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1)
def get_dtype(self):
return getattr(self.initial_conv, "weight_comfy_model_dtype", self.initial_conv.weight.dtype)
def forward(self, latent: torch.Tensor) -> torch.Tensor:
b, c, f, h, w = latent.shape
@ -266,7 +270,7 @@ class LatentUpsampler(nn.Module):
return x
@classmethod
def from_config(cls, config):
def from_config(cls, config, operations):
return cls(
in_channels=config.get("in_channels", 4),
mid_channels=config.get("mid_channels", 128),
@ -276,6 +280,7 @@ class LatentUpsampler(nn.Module):
temporal_upsample=config.get("temporal_upsample", False),
spatial_scale=config.get("spatial_scale", 2.0),
rational_resampler=config.get("rational_resampler", False),
operations=operations,
)
def config(self):

View File

@ -12,6 +12,8 @@ from torch import nn
import comfy.patcher_extension
import comfy.ldm.modules.attention
import comfy.ldm.common_dit
import comfy.model_management
import comfy.quant_ops
from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords
@ -322,40 +324,42 @@ class FeedForward(nn.Module):
return self.net(x)
def apply_rotary_emb(input_tensor, freqs_cis):
cos_freqs, sin_freqs = freqs_cis[0], freqs_cis[1]
split_pe = freqs_cis[2] if len(freqs_cis) > 2 else False
return (
apply_split_rotary_emb(input_tensor, cos_freqs, sin_freqs)
if split_pe else
apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs)
rotation_matrix, split_pe = freqs_cis
original_shape = input_tensor.shape
input_tensor = input_tensor.reshape(
input_tensor.shape[0], input_tensor.shape[1], rotation_matrix.shape[2], -1
)
def apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs): # TODO: remove duplicate funcs and pick the best/fastest one
t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2)
t1, t2 = t_dup.unbind(dim=-1)
t_dup = torch.stack((-t2, t1), dim=-1)
input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)")
if comfy.model_management.in_training:
if split_pe:
t = input_tensor.reshape(*input_tensor.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2)
else:
t = input_tensor.reshape(*input_tensor.shape[:-1], -1, 1, 2)
t = t.to(rotation_matrix.dtype)
output = rotation_matrix[..., 0] * t[..., 0] + rotation_matrix[..., 1] * t[..., 1]
if split_pe:
output = output.movedim(-1, -2)
output = output.reshape(input_tensor.shape).type_as(input_tensor)
elif split_pe:
output = comfy.quant_ops.ck.apply_rope_split_half1(input_tensor, rotation_matrix)
else:
output = comfy.quant_ops.ck.apply_rope1(input_tensor, rotation_matrix)
return output.reshape(original_shape)
out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs
def apply_rotary_emb_qk(q, k, freqs_cis):
if comfy.model_management.in_training:
return apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis)
return out
def apply_split_rotary_emb(input_tensor, cos, sin):
needs_reshape = False
if input_tensor.ndim != 4 and cos.ndim == 4:
B, H, T, _ = cos.shape
input_tensor = input_tensor.reshape(B, T, H, -1).swapaxes(1, 2)
needs_reshape = True
split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2)
first_half_input = split_input[..., :1, :]
second_half_input = split_input[..., 1:, :]
output = split_input * cos.unsqueeze(-2)
first_half_output = output[..., :1, :]
second_half_output = output[..., 1:, :]
first_half_output.addcmul_(-sin.unsqueeze(-2), second_half_input)
second_half_output.addcmul_(sin.unsqueeze(-2), first_half_input)
output = rearrange(output, "... d r -> ... (d r)")
return output.swapaxes(1, 2).reshape(B, T, -1) if needs_reshape else output
rotation_matrix, split_pe = freqs_cis
q_shape = q.shape
k_shape = k.shape
q = q.reshape(q.shape[0], q.shape[1], rotation_matrix.shape[2], -1)
k = k.reshape(k.shape[0], k.shape[1], rotation_matrix.shape[2], -1)
if split_pe:
q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rotation_matrix)
else:
q, k = comfy.quant_ops.ck.apply_rope(q, k, rotation_matrix)
return q.reshape(q_shape), k.reshape(k_shape)
class GuideAttentionMask:
@ -461,9 +465,13 @@ class CrossAttention(nn.Module):
q = self.q_norm(q)
k = self.k_norm(k)
# These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent.
if pe is not None:
q = apply_rotary_emb(q, pe)
k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
if k_pe is None and q.shape == k.shape:
q, k = apply_rotary_emb_qk(q, k, pe)
else:
q = apply_rotary_emb(q, pe)
k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
if mask is None:
out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options)
@ -653,36 +661,23 @@ def generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid):
)
return freqs
def interleaved_freqs_cis(freqs, pad_size):
cos_freq = freqs.cos().repeat_interleave(2, dim=-1)
sin_freq = freqs.sin().repeat_interleave(2, dim=-1)
if pad_size != 0:
cos_padding = torch.ones_like(cos_freq[:, :, : pad_size])
sin_padding = torch.zeros_like(cos_freq[:, :, : pad_size])
cos_freq = torch.cat([cos_padding, cos_freq], dim=-1)
sin_freq = torch.cat([sin_padding, sin_freq], dim=-1)
return cos_freq, sin_freq
def split_freqs_cis(freqs, pad_size, num_attention_heads):
cos_freq = freqs.cos()
sin_freq = freqs.sin()
if pad_size != 0:
cos_padding = torch.ones_like(cos_freq[:, :, :pad_size])
sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size])
cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1)
sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1)
# Reshape freqs to be compatible with multi-head attention
B , T, half_HD = cos_freq.shape
def freqs_cis_matrix(freqs, pad_size, split_mode, num_attention_heads, out_dtype):
cos_freq = freqs.cos().to(out_dtype)
sin_freq = freqs.sin().to(out_dtype)
if pad_size:
matrix_pad_size = pad_size if split_mode else pad_size // 2
cos_padding = torch.ones_like(cos_freq[:, :, :matrix_pad_size])
sin_padding = torch.zeros_like(sin_freq[:, :, :matrix_pad_size])
cos_freq = torch.cat((cos_padding, cos_freq), dim=-1)
sin_freq = torch.cat((sin_padding, sin_freq), dim=-1)
B, T, half_HD = cos_freq.shape
cos_freq = cos_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads)
sin_freq = sin_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads)
cos_freq = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2)
sin_freq = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2)
return cos_freq, sin_freq
rotation_matrix = torch.stack(
(cos_freq, -sin_freq, sin_freq, cos_freq), dim=-1
)
return rotation_matrix.reshape(*rotation_matrix.shape[:-1], 2, 2), split_mode
class LTXBaseModel(torch.nn.Module, ABC):
"""
@ -885,12 +880,17 @@ class LTXBaseModel(torch.nn.Module, ABC):
expected_freqs = dim // 2
current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs
cos_freq, sin_freq = split_freqs_cis(freqs, pad_size, num_attention_heads)
else:
# 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only
n_elem = 2 * indices_grid.shape[1]
cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem)
return cos_freq.to(out_dtype), sin_freq.to(out_dtype), split_mode
pad_size = dim % n_elem
return freqs_cis_matrix(
freqs,
pad_size,
split_mode,
num_attention_heads,
out_dtype,
)
def _prepare_positional_embeddings(self, pixel_coords, frame_rate, x_dtype):
"""Prepare positional embeddings."""

View File

@ -185,7 +185,7 @@ class AudioVAE(torch.nn.Module):
self.autoencoder.mel_bins,
)
def num_of_latents_from_frames(self, frames_number: int, frame_rate: int) -> int:
def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int:
return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second)
def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor:

View File

@ -49,6 +49,12 @@ class CausalConv3d(nn.Module):
)
self.temporal_cache_state={}
def _empty_output(self, x):
# empty (0 frame) outputs must still have the conv's output channels and spatial dims
h = (x.shape[3] + 2 * self.conv.padding[1] - self.conv.kernel_size[1]) // self.conv.stride[1] + 1
w = (x.shape[4] + 2 * self.conv.padding[2] - self.conv.kernel_size[2]) // self.conv.stride[2] + 1
return x.new_empty((x.shape[0], self.out_channels, 0, h, w))
def forward(self, x, causal: bool = True):
tid = threading.get_ident()
@ -58,7 +64,7 @@ class CausalConv3d(nn.Module):
if not causal:
padding_length = padding_length // 2
if x.shape[2] == 0:
return x
return self._empty_output(x)
cached = x[:, :, :1, :, :].repeat((1, 1, padding_length, 1, 1))
pieces = [ cached, x ]
if is_end and not causal:
@ -83,7 +89,7 @@ class CausalConv3d(nn.Module):
elif is_end:
self.temporal_cache_state[tid] = (None, True)
return self.conv(x) if x.shape[2] >= self.time_kernel_size else x[:, :, :0, :, :]
return self.conv(x) if x.shape[2] >= self.time_kernel_size else self._empty_output(x)
@property
def weight(self):

View File

@ -390,10 +390,10 @@ class Decoder(nn.Module):
# Compute output channel to be product of all channel-multiplier blocks
output_channel = base_channels
for block_name, block_params in list(reversed(blocks)):
for block_name, block_params in blocks:
block_params = block_params if isinstance(block_params, dict) else {}
if block_name == "res_x_y":
output_channel = output_channel * block_params.get("multiplier", 2)
output_channel = block_params.get("in_channels", output_channel * block_params.get("multiplier", 2))
if block_name == "compress_all":
output_channel = output_channel * block_params.get("multiplier", 1)
if block_name == "compress_space":
@ -432,7 +432,7 @@ class Decoder(nn.Module):
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "res_x_y":
output_channel = output_channel // block_params.get("multiplier", 2)
output_channel = block_params.get("out_channels", output_channel // block_params.get("multiplier", 2))
block = ResnetBlock3D(
dims=dims,
in_channels=input_channel,

View File

@ -6,6 +6,9 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
import comfy.ldm.common_dit
import comfy.model_management
import comfy.ops
import comfy.quant_ops
from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder
from comfy.ldm.modules.attention import optimized_attention_masked
@ -97,6 +100,7 @@ class JointAttention(nn.Module):
self.n_local_kv_heads = self.n_kv_heads
self.n_rep = self.n_local_heads // self.n_local_kv_heads
self.head_dim = dim // n_heads
self.qk_norm = qk_norm
self.qkv = operation_settings.get("operations").Linear(
dim,
@ -151,10 +155,21 @@ class JointAttention(nn.Module):
xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
xq = self.q_norm(xq)
xk = self.k_norm(xk)
xq, xk = apply_rope(xq, xk, freqs_cis)
if self.qk_norm and not comfy.model_management.in_training:
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True)
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True)
epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps
if self.n_local_heads == self.n_local_kv_heads:
xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon)
else:
xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon)
xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon)
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
else:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
xq, xk = apply_rope(xq, xk, freqs_cis)
n_rep = self.n_local_heads // self.n_local_kv_heads
if n_rep >= 1:

View File

@ -0,0 +1,186 @@
# Mage-Flow (https://github.com/microsoft/Mage) native-resolution MMDiT (MIT)
# Architecture is a 12-layer variant of the Qwen-Image double-stream block with
# patch_size=1 (no 2x2 packing), unrotated text tokens and a bf16-rounded
# timestep frequency table.
import math
import torch
import torch.nn as nn
from typing import Optional, Tuple
from comfy.ldm.lightricks.model import TimestepEmbedding
from comfy.ldm.flux.layers import EmbedND
from comfy.ldm.qwen_image.model import QwenImageTransformerBlock, LastLayer
import comfy.patcher_extension
class MageTimestepProjEmbeddings(nn.Module):
def __init__(self, embedding_dim, dtype=None, device=None, operations=None):
super().__init__()
self.timestep_embedder = TimestepEmbedding(
in_channels=256, time_embed_dim=embedding_dim,
dtype=dtype, device=device, operations=operations
)
def forward(self, timestep, hidden_states):
half_dim = 128
exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim
emb = torch.exp(exponent).to(timestep.dtype)
emb = timestep[:, None].float() * emb[None, :]
emb = 1000.0 * emb
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) # flip_sin_to_cos
return self.timestep_embedder(emb.to(dtype=hidden_states.dtype))
class MageFlowTransformer2DModel(nn.Module):
def __init__(
self,
in_channels: int = 128,
out_channels: Optional[int] = 128,
num_layers: int = 12,
attention_head_dim: int = 128,
num_attention_heads: int = 24,
joint_attention_dim: int = 2560,
axes_dims_rope: Tuple[int, int, int] = (16, 56, 56),
image_model=None,
dtype=None,
device=None,
operations=None,
):
super().__init__()
self.dtype = dtype
self.patch_size = 1
self.in_channels = in_channels
self.out_channels = out_channels or in_channels
self.inner_dim = num_attention_heads * attention_head_dim
self.pe_embedder = EmbedND(dim=attention_head_dim, theta=10000, axes_dim=list(axes_dims_rope))
self.time_text_embed = MageTimestepProjEmbeddings(embedding_dim=self.inner_dim, dtype=dtype, device=device, operations=operations)
self.txt_norm = operations.RMSNorm(joint_attention_dim, eps=1e-6, dtype=dtype, device=device)
self.img_in = operations.Linear(in_channels, self.inner_dim, dtype=dtype, device=device)
self.txt_in = operations.Linear(joint_attention_dim, self.inner_dim, dtype=dtype, device=device)
self.transformer_blocks = nn.ModuleList([
QwenImageTransformerBlock(
dim=self.inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
dtype=dtype,
device=device,
operations=operations
)
for _ in range(num_layers)
])
self.norm_out = LastLayer(self.inner_dim, self.inner_dim, dtype=dtype, device=device, operations=operations)
self.proj_out = operations.Linear(self.inner_dim, self.out_channels, bias=True, dtype=dtype, device=device)
def process_img(self, x, index=0):
# patch_size=1: tokens are raw latent pixels, no 2x2 packing.
bs, c, h, w = x.shape
hidden_states = x.movedim(1, -1).reshape(bs, h * w, c)
img_ids = torch.zeros((h, w, 3), device=x.device)
# Frame axis: positive image index (0 = target, 1..N = reference images).
img_ids[:, :, 0] = index
# Mage scale_rope centering: positions [-ceil(n/2), floor(n/2)), i.e.
# offset by (n - n//2). Differs from Qwen-Image's -(n//2) for odd sizes.
img_ids[:, :, 1] = img_ids[:, :, 1] + torch.arange(h, device=x.device)[:, None] - (h - h // 2)
img_ids[:, :, 2] = img_ids[:, :, 2] + torch.arange(w, device=x.device)[None, :] - (w - w // 2)
return hidden_states, img_ids.reshape(h * w, 3).unsqueeze(0).expand(bs, -1, -1), (h, w)
def forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
self._forward,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
).execute(x, timestep, context, attention_mask, ref_latents, transformer_options, **kwargs)
def _forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, control=None, **kwargs):
if attention_mask is not None and not torch.is_floating_point(attention_mask):
attention_mask = (attention_mask - 1).to(x.dtype) * torch.finfo(x.dtype).max
hidden_states, img_ids, orig_shape = self.process_img(x)
num_embeds = hidden_states.shape[1]
if ref_latents is not None:
ref_num_tokens = []
index = 0
for ref in ref_latents:
index += 1
kontext, kontext_ids, _ = self.process_img(ref, index=index)
hidden_states = torch.cat([hidden_states, kontext], dim=1)
img_ids = torch.cat([img_ids, kontext_ids], dim=1)
ref_num_tokens.append(kontext.shape[1])
transformer_options = transformer_options.copy()
transformer_options["reference_image_num_tokens"] = ref_num_tokens
# Text tokens are not rotated in Mage-Flow: RoPE at position 0 is the
# identity rotation.
txt_ids = torch.zeros((x.shape[0], context.shape[1], 3), device=x.device)
hidden_states = self.img_in(hidden_states)
context = self.txt_norm(context)
context = self.txt_in(context)
temb = self.time_text_embed(timestep, hidden_states)
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
if "post_input" in patches:
for p in patches["post_input"]:
out = p({"img": hidden_states, "txt": context, "img_ids": img_ids, "txt_ids": txt_ids, "transformer_options": transformer_options})
hidden_states = out["img"]
context = out["txt"]
img_ids = out["img_ids"]
txt_ids = out["txt_ids"]
ids = torch.cat((txt_ids, img_ids), dim=1)
image_rotary_emb = self.pe_embedder(ids).contiguous()
del ids, txt_ids, img_ids
transformer_options["total_blocks"] = len(self.transformer_blocks)
transformer_options["block_type"] = "double"
for i, block in enumerate(self.transformer_blocks):
transformer_options["block_index"] = i
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["txt"], out["img"] = block(hidden_states=args["img"], encoder_hidden_states=args["txt"], encoder_hidden_states_mask=attention_mask, temb=args["vec"], image_rotary_emb=args["pe"], transformer_options=args["transformer_options"])
return out
out = blocks_replace[("double_block", i)]({"img": hidden_states, "txt": context, "vec": temb, "pe": image_rotary_emb, "transformer_options": transformer_options}, {"original_block": block_wrap})
hidden_states = out["img"]
context = out["txt"]
else:
context, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=context,
encoder_hidden_states_mask=attention_mask,
temb=temb,
image_rotary_emb=image_rotary_emb,
transformer_options=transformer_options,
)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": hidden_states, "txt": context, "x": x, "block_index": i, "transformer_options": transformer_options})
hidden_states = out["img"]
context = out["txt"]
if control is not None: # Controlnet
control_i = control.get("input")
if i < len(control_i):
add = control_i[i]
if add is not None:
hidden_states[:, :add.shape[1]] += add
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states[:, :num_embeds]
h, w = orig_shape
return hidden_states.reshape(x.shape[0], h, w, self.out_channels).movedim(-1, 1)

477
comfy/ldm/mage_flow/vae.py Normal file
View File

@ -0,0 +1,477 @@
# Mage-VAE (https://github.com/microsoft/Mage) (MIT)
# Symmetric one-step diffusion codec: DConvEncoder (image -> 128ch latent) and
# DConvDenoiser + CoD Decoder (latent -> image). 16x downsample, latents in the
# Flux.2-VAE-anchored space (no patch packing, no BN normalization).
# Both encode and decode are single forward passes at t=0.
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import comfy.ops
from comfy.ldm.modules.diffusionmodules.model import vae_attention
ops = comfy.ops.disable_weight_init
def nonlinearity(x):
return torch.nn.functional.silu(x)
def Normalize(in_channels):
return ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
def modulate(x, shift, scale):
if x.dim() == 4:
b, c = x.shape[:2]
return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1)
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
class LayerNorm2d(ops.LayerNorm):
def __init__(self, num_channels, eps=1e-6, affine=True):
super().__init__(num_channels, eps=eps, elementwise_affine=affine)
def forward(self, x):
x = x.permute(0, 2, 3, 1).contiguous()
x = super().forward(x)
return x.permute(0, 3, 1, 2).contiguous()
class TimestepEmbedder(nn.Module):
"""DConv-style timestep MLP (max_period=10000, freq_size=256)."""
def __init__(self, hidden_size, frequency_embedding_size=256):
super().__init__()
self.mlp = nn.Sequential(
ops.Linear(frequency_embedding_size, hidden_size, bias=True),
nn.SiLU(),
ops.Linear(hidden_size, hidden_size, bias=True),
)
self.frequency_embedding_size = frequency_embedding_size
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half
).to(t.device)
args = t[:, None].float() * freqs[None]
emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
return emb
def forward(self, t, dtype):
emb = self.timestep_embedding(t, self.frequency_embedding_size)
return self.mlp(emb.to(dtype))
class BottleneckPatchEmbed(nn.Module):
"""Image patch embed concatenated with a per-patch conditioning vector."""
def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True):
super().__init__()
self.proj1 = ops.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
self.proj2 = ops.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias)
def forward(self, x, cond):
return self.proj2(torch.cat([self.proj1(x), cond], dim=1))
class DiCoBlock(nn.Module):
"""DConv block with adaLN modulation."""
def __init__(self, hidden_size, mlp_ratio=4.0):
super().__init__()
self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.ca = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
nn.Sigmoid(),
)
ffn = int(mlp_ratio * hidden_size)
self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
self.norm1 = LayerNorm2d(hidden_size, affine=False)
self.norm2 = LayerNorm2d(hidden_size, affine=False)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
ops.Linear(hidden_size, 6 * hidden_size, bias=True),
)
def forward(self, inp, c):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
x = modulate(self.norm1(inp), shift_msa, scale_msa)
x = F.gelu(self.conv2(self.conv1(x)))
x = x * self.ca(x)
x = self.conv3(x)
x = inp + gate_msa[..., None, None] * x
x = x + gate_mlp[..., None, None] * self.conv5(
F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp)))
)
return x
class EncoderDiCoBlock(nn.Module):
"""DiCoBlock without adaLN, for the encoder head pathway."""
def __init__(self, hidden_size, mlp_ratio=4.0):
super().__init__()
self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
self.ca = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
nn.Sigmoid(),
)
ffn = int(mlp_ratio * hidden_size)
self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
self.norm1 = LayerNorm2d(hidden_size)
self.norm2 = LayerNorm2d(hidden_size)
def forward(self, inp):
x = self.norm1(inp)
x = F.gelu(self.conv2(self.conv1(x)))
x = x * self.ca(x)
x = self.conv3(x)
x = inp + x
return x + self.conv5(F.gelu(self.conv4(self.norm2(x))))
class NerfEmbedder(nn.Module):
"""Patch-position embedder used by the DConv decoder x-pathway."""
def __init__(self, in_channels, hidden_size_input, max_freqs=8):
super().__init__()
self.max_freqs = max_freqs
self.embedder = nn.Sequential(
ops.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True),
)
def fetch_pos(self, patch_size, device, dtype):
pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype)
pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij")
pos_x = pos_x.reshape(-1, 1, 1)
pos_y = pos_y.reshape(-1, 1, 1)
freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device)
fx = freqs[None, :, None]
fy = freqs[None, None, :]
coeffs = (1 + fx * fy) ** -1
dct_x = torch.cos(pos_x * fx * torch.pi)
dct_y = torch.cos(pos_y * fy * torch.pi)
return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2)
def forward(self, x):
B, P2, _ = x.shape
ps = int(P2 ** 0.5)
dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1)
return self.embedder(torch.cat([x, dct], dim=-1))
class NerfFinalLayer(nn.Module):
def __init__(self, hidden_size, out_channels):
super().__init__()
self.norm = ops.RMSNorm(hidden_size, eps=1e-6)
self.linear = ops.Linear(hidden_size, out_channels, bias=True)
def forward(self, x):
return self.linear(self.norm(x))
class MLPResBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.in_ln = ops.LayerNorm(channels, eps=1e-6)
self.mlp = nn.Sequential(
ops.Linear(channels, channels, bias=True),
nn.SiLU(),
ops.Linear(channels, channels, bias=True),
)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
ops.Linear(channels, 3 * channels, bias=True),
)
def forward(self, x, y):
shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1)
h = self.in_ln(x) * (1 + scale) + shift
return x + gate * self.mlp(h)
class SimpleMLPAdaLN(nn.Module):
"""Final small MLP that maps NerfEmbedder features to per-patch RGB."""
def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size):
super().__init__()
self.in_channels = in_channels
self.model_channels = model_channels
self.out_channels = out_channels
self.num_res_blocks = num_res_blocks
self.patch_size = patch_size
self.cond_embed = ops.Linear(z_channels, patch_size ** 2 * model_channels)
self.input_proj = ops.Linear(in_channels, model_channels)
self.res_blocks = nn.ModuleList(MLPResBlock(model_channels) for _ in range(num_res_blocks))
def forward(self, x, c):
x = self.input_proj(x)
c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1)
for block in self.res_blocks:
x = block(x, c)
return x
class ResnetBlock(nn.Module):
"""GroupNorm + Conv ResBlock used by the CoD Decoder."""
def __init__(self, *, in_channels, out_channels=None):
super().__init__()
out_channels = out_channels or in_channels
self.in_channels = in_channels
self.out_channels = out_channels
self.norm1 = Normalize(in_channels)
self.conv1 = ops.Conv2d(in_channels, out_channels, 3, padding=1)
self.norm2 = Normalize(out_channels)
self.conv2 = ops.Conv2d(out_channels, out_channels, 3, padding=1)
if in_channels != out_channels:
self.nin_shortcut = ops.Conv2d(in_channels, out_channels, 1)
def forward(self, x):
h = self.conv1(nonlinearity(self.norm1(x)))
h = self.conv2(nonlinearity(self.norm2(h)))
if self.in_channels != self.out_channels:
x = self.nin_shortcut(x)
return x + h
class AttnBlock(nn.Module):
"""Patched (windowed) self-attention used by the CoD Decoder."""
def __init__(self, in_channels, patch_size=32):
super().__init__()
self.in_channels = in_channels
self.patch_size = patch_size
self.norm = Normalize(in_channels)
self.q = ops.Conv2d(in_channels, in_channels, 1)
self.k = ops.Conv2d(in_channels, in_channels, 1)
self.v = ops.Conv2d(in_channels, in_channels, 1)
self.proj_out = ops.Conv2d(in_channels, in_channels, 1)
# VAE attention selection: full-precision backends only (no sage/quantized attention)
self.optimized_attention = vae_attention()
def forward(self, x):
h_ = self.norm(x)
Q = self.q(h_)
K = self.k(h_)
V = self.v(h_)
d = self.patch_size
b, c, H, W = Q.shape
pad_h = (d - H % d) % d
pad_w = (d - W % d) % d
if pad_h or pad_w:
Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate")
K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate")
V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate")
_, _, H_pad, W_pad = Q.shape
nph, npw = H_pad // d, W_pad // d
np_ = nph * npw
def to_patches(t):
return (t.reshape(b, c, nph, d, npw, d)
.permute(0, 2, 4, 1, 3, 5)
.reshape(b * np_, c, d * d))
# [b*np, c, d*d]: attention over the d*d spatial positions of each window
Q = to_patches(Q)
K = to_patches(K)
V = to_patches(V)
h_ = self.optimized_attention(Q, K, V)
h_ = h_.reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad)
if pad_h or pad_w:
h_ = h_[:, :, :H, :W]
return x + self.proj_out(h_)
class CoDDecoder(nn.Module):
"""CoD Decoder: latent -> conditioning features for the denoiser (ds=16, light)."""
def __init__(self, out_ch=384, z_ch=128):
super().__init__()
self.conv_in = ops.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1)
self.block = nn.Sequential(
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
AttnBlock(out_ch, patch_size=32),
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
AttnBlock(out_ch, patch_size=32),
ResnetBlock(in_channels=out_ch, out_channels=out_ch),
)
self.norm_out = Normalize(out_ch)
self.conv_out = ops.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1)
self.ada = nn.Identity()
def forward(self, z):
h = self.block(self.conv_in(z))
h = self.conv_out(nonlinearity(self.norm_out(h)))
return self.ada(h)
class DConvEncoder(nn.Module):
"""DConvEncoder: image -> packed (mean, logvar) latent."""
def __init__(
self,
z_ch=128,
hidden_size=384,
num_blocks=21,
patch_size=16,
mlp_ratio=4.0,
head_size=768,
num_head_blocks=2,
out_ch_mult=2,
):
super().__init__()
self.z_ch = z_ch
self.patch_size = patch_size
self.patch_cond_embed = ops.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True)
self.head_blocks = nn.ModuleList([
EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks)
])
self.proj_down = ops.Conv2d(head_size, hidden_size, kernel_size=1, bias=True)
self.z_proj = ops.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True)
self.fuse_proj = ops.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True)
self.t_embedder = TimestepEmbedder(hidden_size)
self.blocks = nn.ModuleList([
DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks)
])
self.norm_out = LayerNorm2d(hidden_size)
self.proj_out = ops.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True)
def forward_pred(self, z_t, t, y):
cond = self.patch_cond_embed(y)
for block in self.head_blocks:
cond = block(cond)
cond = self.proj_down(cond)
s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1))
c = self.t_embedder(t.view(-1), y.dtype)
for block in self.blocks:
s = block(s, c)
return self.proj_out(self.norm_out(s))
class YEmbedder(nn.Module):
"""Holds only the CoD decoder (the original Flux2-VAE encoder side is dropped at load)."""
def __init__(self, ch=384, z_ch=128):
super().__init__()
self.decoder = CoDDecoder(out_ch=ch, z_ch=z_ch)
class DConvDenoiser(nn.Module):
"""One-step DConv denoiser: latent (via cond) + zero noise -> reconstructed image."""
def __init__(
self,
patch_size=16,
in_channels=3,
hidden_size=384,
hidden_size_x=32,
mlp_ratio=4.0,
num_blocks=24,
num_cond_blocks=21,
bottleneck_dim=128,
):
super().__init__()
self.in_channels = in_channels
self.patch_size = patch_size
self.hidden_size = hidden_size
self.num_cond_blocks = num_cond_blocks
self.t_embedder = TimestepEmbedder(hidden_size)
self.y_embedder_x = ops.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0)
self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8)
self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
self.blocks = nn.ModuleList([
DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks)
])
self.dec_net = SimpleMLPAdaLN(
in_channels=hidden_size_x,
model_channels=hidden_size_x,
out_channels=in_channels,
z_channels=hidden_size,
num_res_blocks=num_blocks - num_cond_blocks,
patch_size=patch_size,
)
self.final_layer = NerfFinalLayer(hidden_size_x, in_channels)
self.y_embedder = YEmbedder(ch=hidden_size, z_ch=bottleneck_dim)
def forward(self, x, t, cond):
b, _, h, w = x.shape
c = self.t_embedder(t.view(-1), x.dtype)
s = self.s_embedder(x, cond)
for block in self.blocks:
s = block(s, c)
length = s.shape[-2] * s.shape[-1]
s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size)
x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size)
x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1)
x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1)
x = self.x_embedder(x)
x = self.dec_net(x, s)
x = self.final_layer(x)
x = x.transpose(1, 2).reshape(b, length, -1)
return torch.nn.functional.fold(
x.transpose(1, 2).contiguous(), (h, w),
kernel_size=self.patch_size, stride=self.patch_size,
)
class MageVAE(nn.Module):
"""
Encode: DConvEncoder (one-step at t=0) -> posterior mean [B, 128, H/16, W/16]
Decode: DConvDenoiser + CoD Decoder -> image [B, 3, H, W] in [-1, 1]
"""
latent_channels = 128
downsample_factor = 16
def __init__(self):
super().__init__()
self.dconv_encoder = DConvEncoder()
self.decoder_model = DConvDenoiser()
def encode(self, x):
B, _, H, W = x.shape
ps = self.dconv_encoder.patch_size
z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype)
t = torch.zeros(B, device=x.device, dtype=x.dtype)
out = self.dconv_encoder.forward_pred(z_t, t, x)
return out[:, : self.latent_channels] # posterior mean (sample_posterior=False)
def decode(self, z):
cond = self.decoder_model.y_embedder.decoder(z)
B = z.shape[0]
H = z.shape[2] * self.downsample_factor
W = z.shape[3] * self.downsample_factor
noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype)
t = torch.zeros(B, device=z.device, dtype=z.dtype)
return self.decoder_model.forward(noise, t, cond)

View File

@ -552,6 +552,7 @@ class WanModel(torch.nn.Module):
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@ -564,11 +565,13 @@ class WanModel(torch.nn.Module):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# In-context reference (Bernini)
context_latents = kwargs.get("context_latents", None)
@ -589,6 +592,7 @@ class WanModel(torch.nn.Module):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -604,6 +608,11 @@ class WanModel(torch.nn.Module):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@ -777,6 +786,7 @@ class VaceWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@ -807,6 +817,7 @@ class VaceWanModel(WanModel):
x_orig = x
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -822,6 +833,11 @@ class VaceWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
ii = self.vace_layers_mapping.get(i, None)
if ii is not None:
for iii in range(len(c)):
@ -887,6 +903,7 @@ class CameraWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if self.control_adapter is not None and camera_conditions is not None:
x = x + self.control_adapter(camera_conditions).to(x.dtype)
@ -909,6 +926,7 @@ class CameraWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -924,6 +942,11 @@ class CameraWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel):
# embeddings
bs, _, time, height, width = x.shape
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if control_video is not None:
x = x + self.cond_encoder(control_video)
@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel):
context = self.text_embedding(context)
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None:
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
# head
@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel):
bs, _, time, height, width = x.shape
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel):
audio = None
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)
@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
x_input = x
img_offset = 0
if reference_latent is not None:
x = torch.cat((reference_latent, x), dim=2)
img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
(reference_latent.shape[3] // self.patch_size[1]) * \
(reference_latent.shape[4] // self.patch_size[2])
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
# head
x = self.head(x, e)

View File

@ -493,6 +493,7 @@ class AnimateWanModel(WanModel):
**kwargs,
):
# embeddings
x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
grid_sizes = x.shape[2:]
@ -505,11 +506,13 @@ class AnimateWanModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@ -522,6 +525,7 @@ class AnimateWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -537,6 +541,11 @@ class AnimateWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if i % 5 == 0 and motion_vec is not None:
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)

View File

@ -111,6 +111,7 @@ class WanDancerModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
# embeddings
x_input = x
if int(fps + 0.5) != 30:
x = self.patch_embedding_global(x.float()).to(x.dtype)
else:
@ -128,11 +129,13 @@ class WanDancerModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
img_offset = 0
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@ -163,6 +166,7 @@ class WanDancerModel(WanModel):
context_img_len += clip_fea_ref.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@ -177,6 +181,12 @@ class WanDancerModel(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
if "double_block" in patches:
for p in patches["double_block"]:
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
x = out["img"]
if audio_emb is not None:
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)

149
comfy/ldm/wan/uni3c.py Normal file
View File

@ -0,0 +1,149 @@
# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
# Converted from the original diffusers based implementation.
import torch
import torch.nn as nn
from comfy.ldm.flux.layers import EmbedND
from .model import WanSelfAttention
class Uni3CLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim,
embedding_dim,
eps=1e-5,
device=None, dtype=None, operations=None
):
super().__init__()
self.silu = nn.SiLU()
self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
def forward(self, x, temb):
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
return x, gate[:, None, :]
class Uni3CAttentionBlock(nn.Module):
def __init__(
self,
dim,
ffn_dim,
num_heads,
time_embed_dim=5120,
eps=1e-6,
device=None, dtype=None, operations=None
):
super().__init__()
operation_settings = {"operations": operations, "device": device, "dtype": dtype}
self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
self.ffn = nn.Sequential(
operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
def forward(self, x, temb, freqs):
norm_x, gate_msa = self.norm1(x, temb)
x = x + gate_msa * self.self_attn(norm_x, freqs)
norm_x, gate_ff = self.norm2(x, temb)
x = x + gate_ff * self.ffn(norm_x)
return x
class MaskCamEmbed(nn.Module):
def __init__(
self,
add_channels=7,
mid_channels=256,
conv_out_dim=5120,
device=None, dtype=None, operations=None
):
super().__init__()
self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
self.mask_proj = nn.Sequential(
operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
nn.SiLU())
self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
def forward(self, add_inputs):
add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
add_embeds = self.mask_proj(add_padded)
add_embeds = self.mask_zero_proj(add_embeds)
add_embeds = add_embeds.flatten(2).transpose(1, 2)
return add_embeds
class WanUni3CControlnet(nn.Module):
def __init__(
self,
in_channels=36,
conv_out_dim=5120,
dim=1024,
ffn_dim=8192,
num_heads=16,
num_layers=20,
time_embed_dim=5120,
out_proj_dim=5120,
add_channels=7,
mid_channels=256,
device=None, dtype=None, operations=None
):
super().__init__()
patch_size = (1, 2, 2)
self.num_layers = num_layers
self.controlnet_patch_embedding = operations.Conv3d(
in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
if conv_out_dim != dim:
self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
else:
self.proj_in = nn.Identity()
self.controlnet_blocks = nn.ModuleList([
Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
for _ in range(num_layers)])
self.proj_out = nn.ModuleList([
operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
for _ in range(num_layers)])
head_dim = dim // num_heads
self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids).movedim(1, 2)
return freqs
def process_input(self, control_input, render_mask=None, camera_embedding=None):
# render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
t_len, h_len, w_len = hidden.shape[2:]
freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
hidden = hidden.flatten(2).transpose(1, 2)
add_inputs = None
if camera_embedding is not None and render_mask is not None:
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
elif render_mask is not None:
add_inputs = render_mask
if add_inputs is not None:
hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
hidden = self.proj_in(hidden)
return hidden, freqs
def forward_block(self, block_index, hidden, temb, freqs):
hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
residual = self.proj_out[block_index](hidden)
return hidden, residual

10
comfy/logging.py Normal file
View File

@ -0,0 +1,10 @@
import logging
DETAIL = 15
logging.addLevelName(DETAIL, "DETAIL")
def detail(message, *args, **kwargs):
kwargs.setdefault("stacklevel", 2)
logging.log(DETAIL, message, *args, **kwargs)

View File

@ -58,6 +58,7 @@ import comfy.ldm.omnigen.omnigen2
import comfy.ldm.seedvr.model
import comfy.ldm.boogu.model
import comfy.ldm.qwen_image.model
import comfy.ldm.mage_flow.model
import comfy.ldm.joyimage.model
import comfy.ldm.ideogram4.model
import comfy.ldm.krea2.model
@ -2046,11 +2047,11 @@ class WAN22_WanDancer(WAN21):
fps = kwargs.get("fps", None)
if fps is not None:
out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps]))
out['fps'] = comfy.conds.CONDConstant(fps)
audio_inject_scale = kwargs.get("audio_inject_scale", None)
if audio_inject_scale is not None:
out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale]))
out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale)
return out
class Hunyuan3Dv2(BaseModel):
@ -2265,8 +2266,8 @@ class Boogu(Omnigen2):
self.memory_usage_factor_conds = ("ref_latents",)
class QwenImage(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel)
def __init__(self, model_config, model_type=ModelType.FLUX, device=None, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel):
super().__init__(model_config, model_type, device=device, unet_model=unet_model)
self.memory_usage_factor_conds = ("ref_latents",)
def extra_conds(self, **kwargs):
@ -2296,6 +2297,21 @@ class QwenImage(BaseModel):
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
return out
class MageFlow(QwenImage):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel)
def process_timestep(self, timestep, **kwargs):
# Mage runs in bf16 and rounds its timestep frequency table to the timestep dtype, keep that on fp32 devices.
return timestep.to(torch.bfloat16)
def extra_conds_shapes(self, **kwargs):
out = {}
ref_latents = kwargs.get("reference_latents", None)
if ref_latents is not None:
out['ref_latents'] = list([1, 128, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 128])
return out
class JoyImage(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel)

View File

@ -903,6 +903,13 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
"selected_layer_index": selected_layer_index,
}
if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys and '{}proj_out.weight'.format(key_prefix) in state_dict_keys and state_dict['{}txt_norm.weight'.format(key_prefix)].shape[0] == 2560 and state_dict['{}proj_out.weight'.format(key_prefix)].shape[0] == 128: # Mage-Flow (Qwen Image txt_norm/proj_out are 3584/64)
dit_config = {}
dit_config["image_model"] = "mage_flow"
dit_config["in_channels"] = 128
dit_config["num_layers"] = count_blocks(state_dict_keys, '{}transformer_blocks.'.format(key_prefix) + '{}.')
return dit_config
if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys: # Qwen Image
dit_config = {}
dit_config["image_model"] = "qwen_image"

View File

@ -34,6 +34,7 @@ import comfy.utils
import comfy.quant_ops
import comfy_aimdo.host_buffer
import comfy_aimdo.vram_buffer
from comfy.logging import detail
from typing import TYPE_CHECKING
if TYPE_CHECKING:
@ -473,7 +474,7 @@ except:
SUPPORT_FP8_OPS = args.supports_fp8_compute
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN'
try:
@ -632,18 +633,50 @@ def mark_mmap_dirty(storage):
if mmap_refs is not None:
DIRTY_MMAPS.add(mmap_refs[0])
def free_pins(size, evict_active=False):
PIN_SUBSETS = [ "weights", "patches" ]
LOADED_PIN_SUBSETS = [ "weights-loaded", "patches-loaded" ]
def models_for_pin_eviction(active, current_prompt=None):
for loaded_model in current_loaded_models:
model = loaded_model.model
if model is None or not model.is_dynamic():
continue
pin_state = model.model.dynamic_pins[model.load_device]
if ((active is None or pin_state["active"] == active) and
(current_prompt is None or pin_state["current_prompt"] == current_prompt)):
yield model
def free_model_pins(size, subsets, current_prompt, active, registrations=False):
freed_total = 0
for loaded_model in reversed(current_loaded_models):
for model in models_for_pin_eviction(active, current_prompt=current_prompt):
if size <= 0:
return freed_total
model = loaded_model.model
if model is not None and model.is_dynamic() and (evict_active or not model.model.dynamic_pins[model.load_device]["active"]):
freed = model.partially_unload_ram(size)
freed_total += freed
size -= freed
if registrations:
freed = model.unregister_inactive_pins(size, subsets=subsets)
else:
freed = model.partially_unload_ram(size, subsets=subsets)
freed_total += freed
size -= freed
return freed_total
def pin_eviction_tiers(loaded, evict_active):
tiers = [
(PIN_SUBSETS, False, None),
(LOADED_PIN_SUBSETS, False, None),
(LOADED_PIN_SUBSETS, True, None),
]
if not loaded:
tiers.append((PIN_SUBSETS, True, False))
if evict_active:
tiers.append((PIN_SUBSETS, True, True))
return tiers
def free_pins(size, evict_active=False, loaded=False):
freed = 0
for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active):
freed += free_model_pins(size - freed, subsets, current_prompt, active)
return freed
def should_free_pins_for_ram_pressure(shortfall):
if shortfall <= 0:
return False
@ -653,7 +686,7 @@ def should_free_pins_for_ram_pressure(shortfall):
return True
return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT
def ensure_pin_budget(size, evict_active=False):
def ensure_pin_budget(size, evict_active=False, loaded=False):
if args.high_ram:
return True
if args.fast_disk:
@ -664,32 +697,21 @@ def ensure_pin_budget(size, evict_active=False):
return True
to_free = shortfall + PIN_PRESSURE_HYSTERESIS
return free_pins(to_free, evict_active=evict_active) >= shortfall
return free_pins(to_free, evict_active=evict_active, loaded=loaded) >= shortfall
def free_registrations(shortfall, evict_active=True):
def free_registrations(shortfall, evict_active=True, loaded=False):
if MAX_PINNED_MEMORY <= 0:
return False
if shortfall <= 0:
return True
shortfall += REGISTERABLE_PIN_HYSTERESIS
for loaded_model in reversed(current_loaded_models):
model = loaded_model.model
if model is not None and model.is_dynamic() and not model.model.dynamic_pins[model.load_device]["active"]:
shortfall -= model.unregister_inactive_pins(shortfall)
if shortfall <= 0:
return True
if evict_active:
for loaded_model in current_loaded_models:
model = loaded_model.model
if model is not None and model.is_dynamic() and model.model.dynamic_pins[model.load_device]["active"]:
shortfall -= model.unregister_inactive_pins(shortfall)
if shortfall <= 0:
return True
for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active):
shortfall -= free_model_pins(shortfall, subsets, current_prompt, active, registrations=True)
return shortfall <= REGISTERABLE_PIN_HYSTERESIS
def ensure_pin_registerable(size, evict_active=True):
return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active)
def ensure_pin_registerable(size, evict_active=True, loaded=False):
return free_registrations(TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY, evict_active=evict_active, loaded=loaded)
class LoadedModel:
def __init__(self, model: ModelPatcher):
@ -815,6 +837,8 @@ def minimum_inference_memory():
def free_memory(memory_required, device, keep_loaded=[], for_dynamic=False, pins_required=0, ram_required=0):
cleanup_models_gc()
if not for_dynamic:
detail("Non dynamic memory free called! memory_required=%s pins_required=%s ram_required=%s", memory_required, pins_required, ram_required)
unloaded_model = []
can_unload = []
unloaded_models = []
@ -953,6 +977,9 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu
lowvram_model_memory = 0.1
loaded_model.model_load(lowvram_model_memory, force_patch_weights=force_patch_weights)
vram_used = 0 if is_device_cpu(torch_dev) else loaded_model.model_loaded_memory()
ram_used = model.loaded_ram_size() if model.is_dynamic() else loaded_model.model_memory() - vram_used
detail("Model loaded: patcher=%s model=%s ram_mb=%.1f vram_mb=%.1f", model.__class__.__name__, model.model.__class__.__name__, ram_used / (1024 ** 2), vram_used / (1024 ** 2))
current_loaded_models.insert(0, loaded_model)
return
@ -1379,15 +1406,17 @@ def reset_cast_buffers():
pin_state = model.model.dynamic_pins[model.load_device]
if pin_state["active"]:
*_, buckets = pin_state["weights"]
for size, bucket in list(buckets.items()):
bucket[:] = [ entry for entry in bucket if entry[-1] is not None ]
if not bucket:
del buckets[size]
for subset in ("weights", "weights-loaded"):
*_, buckets = pin_state[subset]
for size, bucket in list(buckets.items()):
bucket[:] = [ entry for entry in bucket if entry[-1] is not None ]
if not bucket:
del buckets[size]
pin_state["active"] = False
model.partially_unload_ram(1e30, subsets=[ "patches" ])
model.model.dynamic_pins[model.load_device]["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {})
model.partially_unload_ram(1e30, subsets=[ "patches", "patches-loaded" ])
for subset in ("patches", "patches-loaded"):
pin_state[subset] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {})
STREAM_CAST_BUFFERS.clear()
STREAM_AIMDO_CAST_BUFFERS.clear()

View File

@ -22,6 +22,7 @@ import collections
import inspect
import logging
import math
import time
import uuid
from typing import Callable, Optional
@ -37,11 +38,58 @@ import comfy.patcher_extension
import comfy.utils
import comfy_aimdo.host_buffer
from comfy.comfy_types import UnetWrapperFunction
from comfy.logging import detail
from comfy.quant_ops import QuantizedTensor
from comfy.patcher_extension import CallbacksMP, PatcherInjection, WrappersMP
import comfy_aimdo.model_vbar
def is_model_patcher_output(output):
return isinstance(output, ModelPatcher) or isinstance(getattr(output, "patcher", None), ModelPatcher)
class PromptModelTracker:
def __init__(self):
self.models = {}
def start(self):
self.end()
def add(self, outputs):
if isinstance(outputs, collections.abc.Mapping):
outputs = outputs.values()
elif not isinstance(outputs, (list, tuple)):
outputs = (outputs,)
for output in outputs:
if isinstance(output, (collections.abc.Mapping, list, tuple)):
self.add(output)
continue
models = []
if isinstance(output, ModelPatcher):
models.append(output)
models.extend(output.model_patches_models())
models.extend(output.get_nested_additional_models())
else:
patcher = getattr(output, "patcher", None)
if isinstance(patcher, ModelPatcher):
models.append(patcher)
get_models = getattr(output, "get_models", None)
if callable(get_models):
models.extend(get_models())
for model in models:
if not isinstance(model, ModelPatcher) or not model.is_dynamic():
continue
key = (id(model.model), model.load_device)
self.models[key] = model
model.set_in_use_by_current_prompt(True)
def end(self):
for model in self.models.values():
model.set_in_use_by_current_prompt(False)
self.models.clear()
def set_model_options_patch_replace(model_options, patch, name, block_name, number, transformer_index=None):
to = model_options["transformer_options"].copy()
@ -1724,14 +1772,20 @@ class ModelPatcherDynamic(ModelPatcher):
self.model.dynamic_pins[device] = {
"weights": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}),
"patches": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}),
"weights-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}),
"patches-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}),
"hostbufs_initialized": False,
"failed": False,
"active": False,
"current_prompt": False,
}
def is_dynamic(self):
return True
def set_in_use_by_current_prompt(self, in_use):
self.model.dynamic_pins[self.load_device]["current_prompt"] = in_use
def _vbar_get(self, create=False):
if self.load_device == torch.device("cpu"):
return None
@ -1802,6 +1856,8 @@ class ModelPatcherDynamic(ModelPatcher):
hostbuf_size = comfy.model_management.pinned_hostbuf_size(self.model_size())
pin_state["weights"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {})
pin_state["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {})
pin_state["weights-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {})
pin_state["patches-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {})
pin_state["hostbufs_initialized"] = True
pin_state["failed"] = False
pin_state["active"] = True
@ -1935,20 +1991,37 @@ class ModelPatcherDynamic(ModelPatcher):
assert self.load_device != torch.device("cpu")
vbar = self._vbar_get()
freed = 0 if vbar is None else vbar.free_memory(memory_to_free)
vbar_freed = 0 if vbar is None else vbar.free_memory(memory_to_free)
freed = vbar_freed
backup_freed = 0
if freed < memory_to_free:
freed += self.restore_loaded_backups()
backup_freed = self.restore_loaded_backups()
freed += backup_freed
method = "vbar+backups" if vbar_freed and backup_freed else "vbar" if vbar_freed else "backups" if backup_freed else "none"
free_methods = getattr(self, "_free_methods", {})
free_methods[method] = free_methods.get(method, 0) + 1
self._free_methods = free_methods
now = time.monotonic()
if now - getattr(self, "_last_free_log_time", 0) >= 5:
requested = "all" if memory_to_free >= 1e30 else f"{memory_to_free / (1024 ** 2):.1f}MB"
prevailing_method = max(free_methods, key=free_methods.get)
detail("AIMDO free: model=%s device=%s prevailing_method=%s methods=%s requested=%s vbar_mb=%.1f backups_mb=%.1f", self.model.__class__.__name__, self.load_device, prevailing_method, free_methods, requested, vbar_freed / (1024 ** 2), backup_freed / (1024 ** 2))
self._free_methods = {}
self._last_free_log_time = now
return freed
def loaded_ram_size(self):
return (self.model.dynamic_pins[self.load_device]["weights"][0].size)
pin_state = self.model.dynamic_pins[self.load_device]
return pin_state["weights"][0].size + pin_state["weights-loaded"][0].size
def pinned_memory_size(self):
return (self.model.dynamic_pins[self.load_device]["weights"][3][0])
pin_state = self.model.dynamic_pins[self.load_device]
return pin_state["weights"][3][0] + pin_state["weights-loaded"][3][0]
def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights", "patches" ]):
def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]):
freed = 0
pin_state = self.model.dynamic_pins[self.load_device]
for subset in subsets:
@ -1956,15 +2029,17 @@ class ModelPatcherDynamic(ModelPatcher):
split = stack_split[0]
while split >= 0:
module, offset = stack[split]
module_pin = module._pins[subset]
split -= 1
stack_split[0] = split
if not module._pin_registered:
if not module_pin["registered"]:
continue
size = module._pin.numel() * module._pin.element_size()
if torch.cuda.cudart().cudaHostUnregister(module._pin.data_ptr()) != 0:
pin = module_pin["pin"]
size = pin.numel() * pin.element_size()
if torch.cuda.cudart().cudaHostUnregister(pin.data_ptr()) != 0:
comfy.model_management.discard_cuda_async_error()
continue
module._pin_registered = False
module_pin["registered"] = False
comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size)
pinned_size[0] = max(0, pinned_size[0] - size)
freed += size
@ -1973,20 +2048,23 @@ class ModelPatcherDynamic(ModelPatcher):
return freed
return freed
def partially_unload_ram(self, ram_to_unload, subsets=[ "weights", "patches" ]):
def partially_unload_ram(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]):
freed = 0
pin_state = self.model.dynamic_pins[self.load_device]
for subset in subsets:
hostbuf, stack, stack_split, pinned_size, *_ = pin_state[subset]
while len(stack) > 0:
module, offset = stack.pop()
size = module._pin.numel() * module._pin.element_size()
module._pin_balancer_entry[-1] = None
del module._pin_balancer_entry
del module._pin
hostbuf.truncate(offset, do_unregister=module._pin_registered)
module_pin = module._pins[subset]
pin = module_pin["pin"]
size = pin.numel() * pin.element_size()
module_pin["balancer_entry"][-1] = None
del module_pin["balancer_entry"]
del module_pin["pin"]
registered = module_pin["registered"]
hostbuf.truncate(offset, do_unregister=registered)
stack_split[0] = min(stack_split[0], len(stack) - 1)
if module._pin_registered:
if registered:
comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size)
pinned_size[0] = max(0, pinned_size[0] - size)
freed += size

View File

@ -41,7 +41,7 @@ def scaled_dot_product_attention(q, k, v, *args, **kwargs):
try:
if torch.cuda.is_available() and comfy.model_management.WINDOWS:
if torch.cuda.is_available():
from torch.nn.attention import SDPBackend, sdpa_kernel
import inspect
if "set_priority" in inspect.signature(sdpa_kernel).parameters:
@ -51,7 +51,10 @@ try:
SDPBackend.MATH,
]
SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
if comfy.model_management.WINDOWS:
SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
else:
SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION)
def scaled_dot_product_attention(q, k, v, *args, **kwargs):
if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower
@ -144,8 +147,13 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin
needs_cast = False
xfer_source = [ s.weight, s.bias ]
pin = comfy.pinned_memory.get_pin(s)
subset = "weights"
pin = comfy.pinned_memory.get_pin(s, subset=subset)
if pin is None and not args.fast_disk:
loaded_pin = comfy.pinned_memory.get_pin(s, subset="weights-loaded")
if loaded_pin is not None or signature is not None:
subset = "weights-loaded"
pin = loaded_pin
if pin is not None:
xfer_source = [ pin ]
@ -182,12 +190,12 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin
if pin is not None:
cast_maybe_lowvram_patch([pin], dest, offload_stream)
return
if signature is None or args.high_ram:
if signature is None or not args.fast_disk or args.high_ram:
comfy.pinned_memory.pin_memory(m, subset=subset, size=size)
pin = comfy.pinned_memory.get_pin(m, subset=subset)
cast_maybe_lowvram_patch(source, pin, offload_stream, xfer_dest2=dest)
handle_pin(s, pin, xfer_source, xfer_dest, size=dest_size)
handle_pin(s, pin, xfer_source, xfer_dest, subset=subset, size=dest_size)
for param_key in ("weight", "bias"):
lowvram_source = getattr(s, param_key + "_lowvram_function", None)
@ -197,8 +205,16 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin
lowvram_dest = get_cast_buffer(lowvram_size)
lowvram_source.prepare(lowvram_dest, None, copy=False, commit=True)
pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches")
handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset="patches", size=lowvram_size)
subset = "patches"
pin = comfy.pinned_memory.get_pin(lowvram_source, subset=subset)
if pin is None:
loaded_pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches-loaded")
if loaded_pin is not None:
subset = "patches-loaded"
pin = loaded_pin
elif signature is not None and not args.fast_disk:
subset = "patches-loaded"
handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset=subset, size=lowvram_size)
prefetch["xfer_dest"] = xfer_dest
@ -1469,12 +1485,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
if layer_conf is not None:
layer_conf = json.loads(layer_conf.numpy().tobytes())
# Only fp8 makes sense for embeddings (per-row dequant via index select).
# Only fp8 and int8_tensorwise support per-row dequant via index select.
# Block-scaled formats (NVFP4, MXFP8) can't do per-row lookup efficiently.
quant_format = layer_conf.get("format") if layer_conf is not None else None
manually_loaded_keys = []
if quant_format in ("float8_e4m3fn", "float8_e5m2") and weight_key in state_dict:
if quant_format in ("float8_e4m3fn", "float8_e5m2", "int8_tensorwise") and weight_key in state_dict:
self.quant_format = quant_format
qconfig = QUANT_ALGOS[quant_format]
self.layout_type = qconfig["comfy_tensor_layout"]
@ -1488,10 +1504,16 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
scale = scale.float()
manually_loaded_keys.append(scale_key)
extra = {}
if quant_format == "int8_tensorwise" and layer_conf.get("convrot", False):
# rotated embedding table: record it so the forward un-rotates after lookup
extra["convrot"] = True
extra["convrot_groupsize"] = int(layer_conf.get("convrot_groupsize", 256))
params = layout_cls.Params(
scale=scale if scale is not None else torch.ones((), dtype=torch.float32),
orig_dtype=MixedPrecisionOps._compute_dtype,
orig_shape=(self.num_embeddings, self.embedding_dim),
**extra,
)
self.weight = torch.nn.Parameter(
QuantizedTensor(weight.to(dtype=qconfig["storage_t"]), qconfig["comfy_tensor_layout"], params),
@ -1513,15 +1535,23 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
def forward_comfy_cast_weights(self, input, out_dtype=None):
weight = self.weight
# Optimized path: lookup in fp8, dequantize only the selected rows.
# Optimized path: lookup in fp8/int8, dequantize only the selected rows.
if isinstance(weight, QuantizedTensor) and len(self.weight_function) == 0:
qdata, _, offload_stream = cast_bias_weight(self, device=input.device, dtype=weight.dtype, offloadable=True)
if isinstance(qdata, QuantizedTensor):
scale = qdata._params.scale
params = qdata._params
scale = params.scale
qdata = qdata._qdata
else:
params = weight._params
scale = None
# int8: per-row scale possible ConvRot, so let the layout do the gather
if self.quant_format == "int8_tensorwise":
x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input)
uncast_bias_weight(self, qdata, None, offload_stream)
return x if out_dtype is None else x.to(dtype=out_dtype)
x = torch.nn.functional.embedding(
input, qdata, self.padding_idx, self.max_norm,
self.norm_type, self.scale_grad_by_freq, self.sparse)

View File

@ -9,14 +9,14 @@ import torch
from comfy.cli_args import args
def _add_to_bucket(module, buckets, size, priority):
def _add_to_bucket(module, module_pin, buckets, size, priority):
bucket = buckets.setdefault(size, [])
entry = [-priority, 0, module]
entry[1] = id(entry)
bisect.insort(bucket, entry)
module._pin_balancer_entry = entry
module_pin["balancer_entry"] = entry
def _steal_pin(module, stack, buckets, size, priority):
def _steal_pin(module, stack, buckets, size, priority, subset):
bucket = buckets.get(size)
if bucket is None:
return False
@ -31,34 +31,39 @@ def _steal_pin(module, stack, buckets, size, priority):
return False
*_, victim = bucket.pop()
module._pin = victim._pin
module._pin_registered = victim._pin_registered
module._pin_stack_index = victim._pin_stack_index
stack[module._pin_stack_index] = (module, stack[module._pin_stack_index][1])
module_pin = module._pins[subset]
victim_pin = victim._pins[subset]
module_pin["pin"] = victim_pin["pin"]
module_pin["registered"] = victim_pin["registered"]
module_pin["stack_index"] = victim_pin["stack_index"]
stack_index = module_pin["stack_index"]
stack[stack_index] = (module, stack[stack_index][1])
victim._pin_registered = False
del victim._pin
del victim._pin_stack_index
del victim._pin_balancer_entry
victim_pin["registered"] = False
del victim_pin["pin"]
del victim_pin["stack_index"]
del victim_pin["balancer_entry"]
_add_to_bucket(module, buckets, size, priority)
_add_to_bucket(module, module_pin, buckets, size, priority)
return True
def get_pin(module, subset="weights"):
pin = getattr(module, "_pin", None)
if pin is None or module._pin_registered or args.disable_pinned_memory:
pins = module.__dict__.get("_pins")
module_pin = None if pins is None else pins.get(subset)
pin = None if module_pin is None else module_pin.get("pin")
if pin is None or module_pin["registered"] or args.disable_pinned_memory:
return pin
_, _, stack_split, pinned_size, *_ = module._pin_state[subset]
size = pin.nbytes
comfy.model_management.ensure_pin_registerable(size)
comfy.model_management.ensure_pin_registerable(size, loaded=subset.endswith("-loaded"))
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
return pin
module._pin_registered = True
stack_split[0] = max(stack_split[0], module._pin_stack_index)
module_pin["registered"] = True
stack_split[0] = max(stack_split[0], module_pin["stack_index"])
comfy.model_management.TOTAL_PINNED_MEMORY += size
pinned_size[0] += size
return pin
@ -72,23 +77,26 @@ def pin_memory(module, subset="weights", size=None):
if pin is not None:
return
pins = module.__dict__.setdefault("_pins", {})
module_pin = pins.setdefault(subset, {})
hostbuf, stack, stack_split, pinned_size, counter, buckets = pin_state[subset]
if size is None:
size = comfy.memory_management.vram_aligned_size([ module.weight, module.bias ])
offset = hostbuf.size
registerable_size = size
priority = getattr(module, "_pin_balancer_priority", None)
loaded = subset.endswith("-loaded")
priority = module_pin.get("balancer_priority")
if priority is None:
priority = comfy.utils.bit_reverse_range(counter[0], 16)
counter[0] += 1
module._pin_balancer_priority = priority
module_pin["balancer_priority"] = priority
comfy.memory_management.extra_ram_release(comfy.memory_management.RAM_CACHE_HEADROOM)
if (not comfy.model_management.ensure_pin_budget(size) or
not comfy.model_management.ensure_pin_registerable(registerable_size)):
return _steal_pin(module, stack, buckets, size, priority)
if (not comfy.model_management.ensure_pin_budget(size, loaded=loaded) or
not comfy.model_management.ensure_pin_registerable(registerable_size, loaded=loaded)):
return _steal_pin(module, stack, buckets, size, priority, subset)
offset = hostbuf.size
extended = False
try:
hostbuf.extend(size=size, register=False)
@ -97,23 +105,23 @@ def pin_memory(module, subset="weights", size=None):
pin.untyped_storage()._comfy_hostbuf = hostbuf
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
comfy.model_management.free_registrations(size)
comfy.model_management.free_registrations(size, loaded=loaded)
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
del pin
hostbuf.truncate(offset, do_unregister=False)
return _steal_pin(module, stack, buckets, size, priority)
return _steal_pin(module, stack, buckets, size, priority, subset)
except RuntimeError:
if extended:
hostbuf.truncate(offset, do_unregister=False)
return _steal_pin(module, stack, buckets, size, priority)
return _steal_pin(module, stack, buckets, size, priority, subset)
module._pin = pin
module_pin["pin"] = pin
stack.append((module, offset))
module._pin_registered = True
module._pin_stack_index = len(stack) - 1
stack_split[0] = max(stack_split[0], module._pin_stack_index)
module_pin["registered"] = True
module_pin["stack_index"] = len(stack) - 1
stack_split[0] = max(stack_split[0], module_pin["stack_index"])
comfy.model_management.TOTAL_PINNED_MEMORY += size
pinned_size[0] += size
_add_to_bucket(module, buckets, size, priority)
_add_to_bucket(module, module_pin, buckets, size, priority)
return True

View File

@ -20,6 +20,7 @@ import comfy.hooks
import comfy.context_windows
import comfy.multigpu
import comfy.utils
from comfy.logging import detail
import scipy.stats
import numpy
@ -991,10 +992,15 @@ class KSAMPLER(Sampler):
noise = model_wrap.inner_model.model_sampling.noise_scaling(sigmas[0], noise, latent_image, self.max_denoise(model_wrap, sigmas))
k_callback = None
total_steps = len(sigmas) - 1
if callback is not None:
k_callback = lambda x: callback(x["i"], x["denoised"], x["x"], total_steps)
first_step = True
def k_callback(x):
nonlocal first_step
if first_step:
detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, self.sampler_function.__name__, x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape))
first_step = False
if callback is not None:
callback(x["i"], x["denoised"], x["x"], total_steps)
samples = self.sampler_function(model_k, noise, sigmas, extra_args=extra_args, callback=k_callback, disable=disable_pbar, **self.extra_options)
samples = model_wrap.inner_model.model_sampling.inverse_noise_scaling(sigmas[-1], samples)
@ -1270,10 +1276,13 @@ class CFGGuider:
return latent_image
if latent_image.is_nested:
sampler_shapes = [tuple(x.shape) for x in latent_image.unbind()]
latent_image, latent_shapes = comfy.utils.pack_latents(latent_image.unbind())
noise, _ = comfy.utils.pack_latents(noise.unbind())
else:
latent_shapes = [latent_image.shape]
sampler_shapes = [tuple(latent_image.shape)]
detail("Sampler: model=%s latent_shapes=%s", self.model_patcher.model.__class__.__name__, sampler_shapes)
if denoise_mask is not None:
if denoise_mask.is_nested:

View File

@ -18,6 +18,7 @@ import comfy.ldm.trellis2.vae
import comfy.ldm.wan.vae2_2
import comfy.ldm.hunyuan3d.vae
import comfy.ldm.seedvr.vae
import comfy.ldm.mage_flow.vae
import comfy.ldm.triposplat.vae
import comfy.ldm.ace.vae.music_dcae_pipeline
import comfy.ldm.cogvideo.vae
@ -61,6 +62,7 @@ import comfy.text_encoders.qwen_image
import comfy.text_encoders.hunyuan_image
import comfy.text_encoders.z_image
import comfy.text_encoders.krea2
import comfy.text_encoders.mage_flow
import comfy.text_encoders.ideogram4
import comfy.text_encoders.ovis
import comfy.text_encoders.kandinsky5
@ -578,6 +580,17 @@ class VAE:
self.upscale_index_formula = (4, 8, 8)
self.process_input = lambda image: image * 2.0 - 1.0
self.crop_input = False
elif "student.dconv_encoder.proj_out.weight" in sd: # Mage-VAE (one-step diffusion codec, Flux2-anchored 128ch/16x latents)
sd = comfy.utils.state_dict_prefix_replace(sd, {"student.dconv_encoder.": "dconv_encoder.", "pipeline.": "decoder_model."})
# Drop the unused Flux2-VAE anchor encoder carried in the checkpoint.
sd = {k: v for k, v in sd.items() if not k.startswith("decoder_model.y_embedder.encoder.") and not k.startswith("decoder_model.y_embedder.bottleneck.")}
self.first_stage_model = comfy.ldm.mage_flow.vae.MageVAE()
self.latent_channels = 128
self.downscale_ratio = 16
self.upscale_ratio = 16
self.working_dtypes = [torch.bfloat16, torch.float32]
self.memory_used_encode = lambda shape, dtype: (400 * shape[2] * shape[3]) * model_management.dtype_size(dtype)
self.memory_used_decode = lambda shape, dtype: (1000 * shape[2] * shape[3] * 16 * 16) * model_management.dtype_size(dtype)
elif "decoder.conv_in.weight" in sd:
if sd['decoder.conv_in.weight'].shape[1] == 64:
ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True}
@ -1399,6 +1412,7 @@ class CLIPType(Enum):
BOOGU = 31
KREA2 = 32
JOYIMAGE = 33
MAGE = 34
@ -1454,6 +1468,7 @@ class TEModel(Enum):
GPT_OSS_20B = 33
QWEN3VL_4B = 34
QWEN3VL_8B = 35
GEMMA_4_12B = 36
def detect_te_model(sd):
@ -1483,6 +1498,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:
@ -1638,10 +1656,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel
clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer
tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None)
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B):
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B):
variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B,
TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B,
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model]
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B,
TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model]
clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant)
clip_target.tokenizer = variant.tokenizer
tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None)
@ -1728,6 +1747,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer
elif clip_type == CLIPType.MAGE and te_model == TEModel.QWEN3VL_4B: # Mage-Flow: full Qwen3-VL-4B, last hidden state, Qwen-Image-style templates.
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
clip_target.clip = comfy.text_encoders.mage_flow.te(**llama_detect(clip_data))
clip_target.tokenizer = comfy.text_encoders.mage_flow.MageFlowTokenizer
elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx.
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data))

View File

@ -27,6 +27,7 @@ import comfy.text_encoders.z_image
import comfy.text_encoders.ideogram4
import comfy.text_encoders.boogu
import comfy.text_encoders.krea2
import comfy.text_encoders.mage_flow
import comfy.text_encoders.joyimage
import comfy.text_encoders.anima
import comfy.text_encoders.ace15
@ -1908,6 +1909,35 @@ class Krea2(supported_models_base.BASE):
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect))
class MageFlow(supported_models_base.BASE):
unet_config = {
"image_model": "mage_flow",
}
sampling_settings = {
"multiplier": 1.0,
"shift": 6.0,
}
memory_usage_factor = 6.5
unet_extra_config = {}
latent_format = latent_formats.Flux2
supported_inference_dtypes = [torch.bfloat16, torch.float32]
vae_key_prefix = ["vae."]
text_encoder_key_prefix = ["text_encoders."]
def get_model(self, state_dict, prefix="", device=None):
out = model_base.MageFlow(self, device=device)
return out
def clip_target(self, state_dict={}):
pref = self.text_encoder_key_prefix[0]
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref))
return supported_models_base.ClipTarget(comfy.text_encoders.mage_flow.MageFlowTokenizer, comfy.text_encoders.mage_flow.te(**hunyuan_detect))
class QwenImage(supported_models_base.BASE):
unet_config = {
"image_model": "qwen_image",
@ -2446,6 +2476,7 @@ models = [
ACEStep15,
Omnigen2,
Boogu,
MageFlow,
QwenImage,
JoyImage,
Ideogram4,

View File

@ -1,11 +1,15 @@
import torch
import torch.nn as nn
import torchaudio.functional as AF
import torchvision.transforms.functional as TVF
import numpy as np
from tokenizers import Tokenizer
from dataclasses import dataclass
import math
from comfy import sd1_clip
import comfy.model_management
import comfy.ops
from comfy.ldm.modules.attention import optimized_attention_for_device
from comfy.rmsnorm import rms_norm
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
@dataclass
class Gemma4Config:
vocab_size: int = 262144
@ -35,6 +43,9 @@ class Gemma4Config:
transformer_type: str = "gemma4"
head_dim = 256
global_head_dim = 512
num_global_key_value_heads = None
attention_k_eq_v = False
vision_bidirectional = False
rms_norm_add = False
mlp_activation = "gelu_pytorch_tanh"
qkv_bias = False
@ -51,6 +62,7 @@ class Gemma4Config:
num_kv_shared_layers: int = 18
use_double_wide_mlp: bool = False
stop_tokens = [1, 50, 106]
suppress_tokens = []
vision_config = GEMMA4_VISION_CONFIG
audio_config = GEMMA4_AUDIO_CONFIG
mm_tokens_per_image = 280
@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config):
num_hidden_layers: int = 60
num_attention_heads: int = 32
num_key_value_heads: int = 16
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = None
vision_config = GEMMA4_VISION_31B_CONFIG
@dataclass
class Gemma4_12B_Config(Gemma4Config):
hidden_size: int = 3840
intermediate_size: int = 15360
num_hidden_layers: int = 48
num_attention_heads: int = 16
num_key_value_heads: int = 8
num_global_key_value_heads = 1
attention_k_eq_v = True
vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
vision_config = GEMMA4_UNIFIED_VISION_CONFIG
suppress_tokens = [258883, 258882]
# unfused RoPE as addcmul_ RoPE diverges from reference code
def _apply_rotary_pos_emb(x, freqs_cis):
@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis):
return out
class Gemma4Attention(nn.Module):
def __init__(self, config, head_dim, device=None, dtype=None, ops=None):
def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
super().__init__()
self.num_heads = config.num_attention_heads
self.num_kv_heads = config.num_key_value_heads
self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
self.hidden_size = config.hidden_size
self.head_dim = head_dim
self.inner_size = self.num_heads * head_dim
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
# k_eq_v: V reuses the K projection (no separate v_proj weight)
self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
self.q_norm = None
@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module):
shareable_kv = None
else:
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
if self.v_proj is not None:
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
else:
xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
if self.k_norm is not None:
xk = self.k_norm(xk)
xv = rms_norm(xv)
@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module):
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops)
# k_eq_v only on global layers, which then use num_global_key_value_heads
k_eq_v = config.attention_k_eq_v and not self.sliding_attention
num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
num_kv_shared = config.num_kv_shared_layers
first_kv_shared = config.num_hidden_layers - num_kv_shared
@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module):
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
else:
self.layer_scalar = None
# layer_scalar exists on every gemma4 variant, independent of per-layer input
self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
sliding_window = None
@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module):
x = self.post_per_layer_input_norm(x)
x = residual + x
if self.layer_scalar is not None:
x = x * self.layer_scalar
x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
return x, present_key_value, shareable_kv
@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module):
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
mask = mask + causal_mask if mask is not None else causal_mask
# Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
if self.config.vision_bidirectional and past_len == 0 and embeds_info:
block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
group = 0
for info in embeds_info:
if info.get("type") == "image":
start = info["index"]
block_ids[start:start + info["size"]] = group
group += 1
if group > 0:
same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
mask = mask.masked_fill(same_block, 0.0)
# Per-layer inputs
per_layer_inputs = None
if self.hidden_size_per_layer_input:
@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module):
shared_global_kv = None # KV from last non-shared global layer
intermediate = None
all_intermediate = None
only_layers = None
if intermediate_output is not None:
if isinstance(intermediate_output, list):
all_intermediate = []
only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
elif intermediate_output == "all":
all_intermediate = []
intermediate_output = None
elif intermediate_output < 0:
intermediate_output = len(self.layers) + intermediate_output
next_key_values = []
for i, layer in enumerate(self.layers):
if all_intermediate is not None:
if only_layers is None or (i in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
layer_kwargs = {}
@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module):
if self.norm is not None:
x = self.norm(x)
if len(next_key_values) > 0:
if all_intermediate is not None:
if only_layers is None or (len(self.layers) in only_layers):
all_intermediate.append(x.unsqueeze(1).clone())
if len(all_intermediate) > 0:
intermediate = torch.cat(all_intermediate, dim=1)
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
intermediate = self.norm(intermediate)
# Only hand back the KV cache when caching was actually requested; SDClipModel reads
# outputs[2] as the pooled output.
if past_key_values is not None and len(next_key_values) > 0:
return x, intermediate, next_key_values
return x, intermediate
@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
cap = self.model.config.final_logit_softcapping
if cap:
logits = cap * torch.tanh(logits / cap)
if self.model.config.suppress_tokens:
logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
return logits
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
@ -441,6 +519,28 @@ class Gemma4AudioMixin:
return None, None
class Gemma4UnifiedBase(Gemma4Base):
"""Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
def _init_model(self, config, dtype, device, operations):
self.num_layers = config.num_hidden_layers
self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
self.dtype = dtype
self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
def preprocess_embed(self, embed, device):
if embed["type"] == "image":
pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
patches, positions = self.vision_model.patchify(pixels)
vision_out = self.vision_model(patches, positions)
return self.multi_modal_projector(vision_out), None
if embed["type"] == "audio":
audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
return self.audio_projector(audio), None
return None, None
# Vision Encoder
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
def _patches_merge(patches, positions_xy, length):
patch_size = math.isqrt(patches.shape[-1] // 3)
k = math.isqrt(patches.shape[-2] // length)
batch = patches.shape[:-2]
max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
kidx = torch.div(positions_xy, k, rounding_mode="floor")
rem = torch.remainder(positions_xy, k)
order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
perm = order.long().argsort(dim=-1)
merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
pad = (positions_xy == -1).all(dim=-1, keepdim=True)
pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
return merged, pos
class Gemma4UnifiedVisionEmbedder(nn.Module):
"""Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
def __init__(self, config, device=None, dtype=None, ops=None):
super().__init__()
self.patch_size = config["patch_size"]
self.pooling_kernel_size = config["pooling_kernel_size"]
patch_dim = config["model_patch_size"] ** 2 * 3
mm_embed_dim = config["mm_embed_dim"]
self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
def patchify(self, pixels):
"""pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
ps, k = self.patch_size, self.pooling_kernel_size
out_patches, out_positions = [], []
for img in pixels:
ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
n_model = teacher.shape[0] // (k * k)
mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
out_patches.append(mp.squeeze(0))
out_positions.append(mpos.squeeze(0))
return torch.stack(out_patches), torch.stack(out_positions)
def forward(self, pixel_values, image_position_ids):
x = self.patch_ln1(pixel_values)
x = self.patch_dense(x)
x = self.patch_ln2(x)
clamped = image_position_ids.clamp(min=0).long()
valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
axes = torch.arange(2, device=image_position_ids.device)
pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
pos_embs = (pos[clamped, axes] * valid).sum(-2)
x = x + pos_embs
return self.pos_norm(x)
# Audio Encoder
class Gemma4AudioConvSubsampler(nn.Module):
@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
# Tokenizer and Wrappers
def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
target_px = max_patches * patch_size ** 2
factor = math.sqrt(target_px / (height * width))
side_mult = pooling_kernel_size * patch_size
target_height = math.floor(factor * height / side_mult) * side_mult
target_width = math.floor(factor * width / side_mult) * side_mult
if target_height == 0 and target_width == 0:
raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
if target_height == 0:
target_height = side_mult
target_width = min(math.floor(width / height) * side_mult, max_side_length)
elif target_width == 0:
target_width = side_mult
target_height = min(math.floor(height / width) * side_mult, max_side_length)
if target_height * target_width > target_px:
raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
return target_height, target_width
class Gemma4_Tokenizer():
tokenizer_json_data = None
@ -998,25 +1189,35 @@ class Gemma4_Tokenizer():
return {"tokenizer_json": self.tokenizer_json_data}
return {}
def _extract_mel_spectrogram(self, waveform, sample_rate):
"""Extract 128-bin log mel spectrogram.
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
"""
# Mix to mono first, then resample to 16kHz
def _audio_token_count(self, num_samples):
# Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
return min(_t, 750)
@staticmethod
def _resample_16k(waveform, sample_rate):
"""Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
if waveform.dim() > 1 and waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
audio = waveform.squeeze(0).float().numpy()
audio = waveform.float()
if sample_rate != 16000:
# Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match)
from scipy.signal import resample_poly, firwin
from math import gcd
g = gcd(sample_rate, 16000)
up, down = 16000 // g, sample_rate // g
L = max(up, down)
h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5))
audio = resample_poly(audio, up, down, window=h).astype(np.float32)
audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
return audio.squeeze(0).contiguous()
def _extract_audio_features(self, waveform, sample_rate):
"""Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
"""
audio = self._resample_16k(waveform, sample_rate).numpy()
n = len(audio)
# Pad to multiple of 128, build sample-level mask
@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer():
if audio is not None:
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate)
audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T])
feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
# Process image/video frames
is_video = video is not None
@ -1090,13 +1291,8 @@ class Gemma4_Tokenizer():
pooling_k = 3
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
max_patches = max_soft_tokens * pooling_k * pooling_k
target_px = max_patches * patch_size * patch_size
factor = (target_px / (h * w)) ** 0.5
side_mult = pooling_k * patch_size
target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
import torchvision.transforms.functional as TVF
for i in range(num_frames):
# rescaling to match reference code
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer():
llama_text = llama_template.format(text)
else:
# Build template from modalities present
system = "<|turn>system\n<|think|><turn|>\n" if thinking else ""
system = "<|turn>system\n<|think|>\n<turn|>\n" if thinking else ""
media = ""
if len(images) > 0:
if is_video:
@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer():
if len(audio_features) > 0:
# Compute audio token count (always at 16kHz)
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
_t = _nmel
for _ in range(2):
_t = (_t + 2 - 3) // 2 + 1
n_audio_tokens = min(_t, 750)
n_audio_tokens = self._audio_token_count(num_samples)
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>"
llama_text = f"{system}<|turn>user\n{media}{text}<turn|>\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<channel|>"
llama_text = f"{system}<|turn>user\n{text}{media}<turn|>\n<|turn>model\n{model_open}"
text_tokens = super().tokenize_with_weights(llama_text, return_word_ids)
@ -1178,7 +1370,6 @@ class Gemma4_Tokenizer():
class _Gemma4Tokenizer:
"""Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)"""
def __init__(self, tokenizer_json_bytes=None, **kwargs):
from tokenizers import Tokenizer
if isinstance(tokenizer_json_bytes, torch.Tensor):
tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist())
self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8"))
@ -1224,6 +1415,30 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class)
class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer):
"""Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram."""
embedding_size = 3840
def _extract_audio_features(self, waveform, sample_rate):
audio = self._resample_16k(waveform, sample_rate)
spt = 640 # audio_samples_per_token (40ms at 16kHz)
pad = (-audio.shape[0]) % spt
if pad:
audio = torch.nn.functional.pad(audio, (0, pad))
num_tokens = audio.shape[0] // spt
feats = audio[:num_tokens * spt].reshape(num_tokens, spt)
feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s)
mask = torch.ones(feats.shape[0], dtype=torch.bool)
return feats, mask
def _audio_token_count(self, num_samples):
return min((num_samples + 639) // 640, 750)
class Gemma4UnifiedTokenizer(Gemma4Tokenizer):
tokenizer_class = Gemma4UnifiedSDTokenizer
# Model wrappers
class Gemma4Model(sd1_clip.SDClipModel):
model_class = None
@ -1256,7 +1471,7 @@ class Gemma4Model(sd1_clip.SDClipModel):
expanded_idx += 1
initial_token_ids = [ids]
input_ids = torch.tensor(initial_token_ids, device=self.execution_device)
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids)
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info)
def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None):
@ -1296,3 +1511,11 @@ def _make_variant(config_cls):
Gemma4_E4B = _make_variant(Gemma4Config)
Gemma4_E2B = _make_variant(Gemma4_E2B_Config)
Gemma4_31B = _make_variant(Gemma4_31B_Config)
# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant).
class Gemma4_12B(Gemma4UnifiedBase):
def __init__(self, config_dict, dtype, device, operations):
super().__init__()
self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations)
Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer

View File

@ -876,7 +876,7 @@ class BaseGenerate:
torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0))
return past_key_values
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None):
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None):
device = embeds.device
if stop_tokens is None:
@ -911,7 +911,7 @@ class BaseGenerate:
if step == 0 and deepstack_embeds is not None:
extra["deepstack_embeds"] = deepstack_embeds
extra["visual_pos_masks"] = visual_pos_masks
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra)
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None))
logits = self.logits(x)[:, -1]
next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty)
token_id = next_token[0].item()

View File

@ -0,0 +1,94 @@
"""Mage-Flow text encoder: Qwen3-VL-4B, last hidden state (2560-dim).
Mage-Flow conditions on the final hidden state of Qwen3-VL-4B with the leading
system + user-opening template tokens stripped (reference start_idx 34 for t2i,
64 for edit). The t2i template is identical to Qwen-Image's; the edit template
uses the same system prompt as Qwen-Image-Edit with "Image N: " reference
prefixes and no <think> block.
"""
import numbers
import torch
import comfy.text_encoders.qwen3vl
from comfy import sd1_clip
MAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>"
MAGE_T2I_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
MAGE_EDIT_TEMPLATE = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
class MageFlowTokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer):
def __init__(self, embedding_directory=None, tokenizer_data={}):
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b")
self.llama_template = MAGE_T2I_TEMPLATE
self.llama_template_images = MAGE_EDIT_TEMPLATE
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs):
image = kwargs.get("image", None)
if image is not None and len(images) == 0:
images = [image[i:i + 1] for i in range(image.shape[0])]
if llama_template is None:
if len(images) > 0:
# Training-time multi-reference body: "Image 1: <ph>Image 2: <ph>...{instruction}"
prefix = "".join("Image {}: {}".format(j + 1, MAGE_VISION_BLOCK) for j in range(len(images)))
llama_template = self.llama_template_images.replace("{}", prefix + "{}", 1)
else:
llama_template = self.llama_template
# thinking=True: Mage templates end at "<|im_start|>assistant\n" with no <think> block.
return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs)
class MageFlowQwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel):
def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}, model_type="qwen3vl_4b"):
super().__init__(device=device, dtype=dtype, attention_mask=attention_mask, model_options=model_options, model_type=model_type)
# apply the final RMSNorm to the tapped last layer (HF last_hidden_state)
self.layer_norm_hidden_state = True
class MageFlowTEModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
clip_model = lambda **kw: MageFlowQwen3VLClipModel(**kw, model_type="qwen3vl_4b") # noqa: E731
super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=clip_model, model_options=model_options)
def encode_token_weights(self, token_weight_pairs, template_end=-1):
# Strip the system + user-opening prefix (reference drop_idx: 34 t2i / 64 edit).
out, pooled, extra = super().encode_token_weights(token_weight_pairs)
tok_pairs = token_weight_pairs["qwen3vl_4b"][0]
count_im_start = 0
if template_end == -1:
for i, v in enumerate(tok_pairs):
elem = v[0]
if not torch.is_tensor(elem):
if isinstance(elem, numbers.Integral):
if elem == 151644 and count_im_start < 2: # <|im_start|>
template_end = i
count_im_start += 1
if out.shape[1] > (template_end + 3):
if tok_pairs[template_end + 1][0] == 872: # "user"
if tok_pairs[template_end + 2][0] == 198: # "\n"
template_end += 3
out = out[:, template_end:]
if "attention_mask" in extra:
extra["attention_mask"] = extra["attention_mask"][:, template_end:]
if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]):
extra.pop("attention_mask") # attention mask is useless if no masked elements
return out, pooled, extra
def te(dtype_llama=None, llama_quantization_metadata=None):
class MageFlowTEModel_(MageFlowTEModel):
def __init__(self, device="cpu", dtype=None, model_options={}):
if dtype_llama is not None:
dtype = dtype_llama
if llama_quantization_metadata is not None:
model_options = model_options.copy()
model_options["quantization_metadata"] = llama_quantization_metadata
super().__init__(device=device, dtype=dtype, model_options=model_options)
return MageFlowTEModel_

View File

@ -158,12 +158,12 @@ class Qwen3VLTokenizer(sd1_clip.SD1Tokenizer):
self.llama_template = "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
self.llama_template_images = "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, **kwargs):
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=False, skip_template=False, **kwargs):
image = kwargs.get("image", None)
if image is not None and len(images) == 0:
images = [image[i:i + 1] for i in range(image.shape[0])]
skip_template = text.startswith('<|im_start|>')
skip_template = skip_template or text.startswith('<|im_start|>')
if prevent_empty_text and text == '':
text = ' '

View File

@ -244,10 +244,10 @@ RECRAFT_V4_PRO_SIZES = [
"2304x1792",
"1792x2304",
"1664x2688",
"1434x1024",
"1024x1434",
"2560x1792",
"1792x2560",
"2688x1536",
"1536x2688",
]

View File

@ -28,6 +28,10 @@ ANTHROPIC_IMAGE_MAX_PIXELS = 1568 * 1568
CLAUDE_MAX_IMAGES = 20
CLAUDE_MODELS: dict[str, str] = {
"Opus 5": "claude-opus-5",
"Opus 4.8": "claude-opus-4-8",
"Fable 5": "claude-fable-5",
"Sonnet 5": "claude-sonnet-5",
"Opus 4.7": "claude-opus-4-7",
"Opus 4.6": "claude-opus-4-6",
"Sonnet 4.6": "claude-sonnet-4-6",
@ -36,9 +40,12 @@ CLAUDE_MODELS: dict[str, str] = {
}
_THINKING_UNSUPPORTED = {"Haiku 4.5"}
# Models that use the newer "adaptive" thinking mode (Opus 4.7 requires it; older models keep the explicit budget API).
# Models that use the newer "adaptive" thinking mode (Opus 4.7+ require it; older models keep the explicit budget API).
# Anthropic decides the actual budget when adaptive is used, based on the `output_config.effort` hint.
_ADAPTIVE_THINKING_MODELS = {"Opus 4.7", "Opus 4.6", "Sonnet 4.6"}
_ADAPTIVE_THINKING_MODELS = {"Opus 4.8", "Sonnet 5", "Opus 4.7", "Opus 4.6", "Sonnet 4.6"}
_ALWAYS_THINKING_MODELS = {"Opus 5", "Fable 5"}
_EXPLICIT_THINKING_OFF_MODELS = {"Sonnet 5"}
_NO_TEMPERATURE_MODELS = {"Opus 5", "Opus 4.8", "Fable 5", "Sonnet 5"}
# Budget mode (Sonnet 4.5): effort -> reasoning budget in tokens. Must be < max_tokens.
# Sized so even the "high" budget fits comfortably under the default max_tokens=32768.
@ -60,20 +67,33 @@ def _claude_model_inputs(model_label: str):
tooltip="Maximum number of tokens to generate (includes reasoning tokens when enabled).",
advanced=True,
),
IO.Float.Input(
"temperature",
default=1.0,
min=0.0,
max=1.0,
step=0.01,
tooltip=(
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
),
advanced=True,
),
]
if model_label not in _THINKING_UNSUPPORTED:
if model_label not in _NO_TEMPERATURE_MODELS:
inputs.append(
IO.Float.Input(
"temperature",
default=1.0,
min=0.0,
max=1.0,
step=0.01,
tooltip=(
"Controls randomness. 0.0 is deterministic, 1.0 is most random. "
"Ignored for Opus 4.7 and any model when reasoning_effort is set."
),
advanced=True,
)
)
if model_label in _ALWAYS_THINKING_MODELS:
inputs.append(
IO.Combo.Input(
"reasoning_effort",
options=[e for e in _REASONING_EFFORTS if e != "off"],
default="high",
tooltip="Extended thinking effort. Reasoning is always enabled for this model.",
advanced=True,
)
)
elif model_label not in _THINKING_UNSUPPORTED:
inputs.append(
IO.Combo.Input(
"reasoning_effort",
@ -86,43 +106,6 @@ def _claude_model_inputs(model_label: str):
return inputs
def _model_price_per_million(model: str) -> tuple[float, float] | None:
"""Return (input_per_1M, output_per_1M) USD for a Claude model, or None if unknown."""
if "opus-4-7" in model or "opus-4-6" in model or "opus-4-5" in model:
return 5.0, 25.0
if "sonnet-4" in model:
return 3.0, 15.0
if "haiku-4-5" in model:
return 1.0, 5.0
return None
def calculate_tokens_price(response: AnthropicMessagesResponse) -> float | None:
"""Compute approximate USD price from response usage. Server-side billing is authoritative."""
if not response.usage or not response.model:
return None
rates = _model_price_per_million(response.model)
if rates is None:
return None
input_rate, output_rate = rates
input_tokens = response.usage.input_tokens or 0
output_tokens = response.usage.output_tokens or 0
cache_read = response.usage.cache_read_input_tokens or 0
cache_5m = 0
cache_1h = 0
if response.usage.cache_creation:
cache_5m = response.usage.cache_creation.ephemeral_5m_input_tokens or 0
cache_1h = response.usage.cache_creation.ephemeral_1h_input_tokens or 0
total = (
input_tokens * input_rate
+ output_tokens * output_rate
+ cache_read * input_rate * 0.1
+ cache_5m * input_rate * 1.25
+ cache_1h * input_rate * 2.0
)
return total / 1_000_000.0
def _get_text_from_response(response: AnthropicMessagesResponse) -> str:
if not response.content:
return ""
@ -213,7 +196,27 @@ class ClaudeNode(IO.ComfyNode):
expr="""
(
$m := widgets.model;
$contains($m, "opus") ? {
$contains($m, "fable") ? {
"type": "list_usd",
"usd": [0.0143, 0.0715],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "opus 4.8") ? {
"type": "list_usd",
"usd": [0.00715, 0.03575],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "sonnet 5") ? {
"type": "list_usd",
"usd": [0.00286, 0.0143],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "opus 5") ? {
"type": "list_usd",
"usd": [0.00715, 0.03575],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
}
: $contains($m, "opus") ? {
"type": "list_usd",
"usd": [0.005, 0.025],
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
@ -247,18 +250,23 @@ class ClaudeNode(IO.ComfyNode):
model_label = model["model"]
max_tokens = model.get("max_tokens", 32768)
reasoning_effort = model.get("reasoning_effort", "off")
thinking_enabled = reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED
always_thinking = model_label in _ALWAYS_THINKING_MODELS
thinking_enabled = always_thinking or (
reasoning_effort not in ("off", None) and model_label not in _THINKING_UNSUPPORTED
)
# Anthropic requires temperature to be unset (defaults to 1.0) when thinking is enabled.
# Opus 4.7 also rejects user-supplied temperature.
if thinking_enabled or model_label == "Opus 4.7":
if model_label in _NO_TEMPERATURE_MODELS or thinking_enabled or model_label == "Opus 4.7":
temperature = None
else:
temperature = model.get("temperature", 1.0)
thinking_cfg: AnthropicThinkingConfig | None = None
output_cfg: AnthropicOutputConfig | None = None
if thinking_enabled:
if always_thinking:
output_cfg = AnthropicOutputConfig(effort=reasoning_effort)
elif thinking_enabled:
if model_label in _ADAPTIVE_THINKING_MODELS:
# Adaptive mode - Anthropic chooses the budget based on effort hint
thinking_cfg = AnthropicThinkingConfig(type="adaptive")
@ -268,6 +276,8 @@ class ClaudeNode(IO.ComfyNode):
budget = _REASONING_BUDGET[reasoning_effort]
budget = min(budget, max(1024, max_tokens - 1024))
thinking_cfg = AnthropicThinkingConfig(type="enabled", budget_tokens=budget)
elif model_label in _EXPLICIT_THINKING_OFF_MODELS:
thinking_cfg = AnthropicThinkingConfig(type="disabled")
image_tensors: list[Input.Image] = [t for t in (images or {}).values() if t is not None]
if sum(get_number_of_images(t) for t in image_tensors) > CLAUDE_MAX_IMAGES:
@ -291,8 +301,12 @@ class ClaudeNode(IO.ComfyNode):
thinking=thinking_cfg,
output_config=output_cfg,
),
price_extractor=calculate_tokens_price,
)
if response.stop_reason == "refusal":
raise ValueError(
"Claude declined to answer this request for safety reasons. "
"Rephrase the prompt or try a different model."
)
return IO.NodeOutput(_get_text_from_response(response) or "Empty response from Claude model.")

View File

@ -2690,7 +2690,8 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
"with ByteDance Seed Audio 1.0. Describe the voice(s), emotion, ambience, background music "
"and sound effects in the prompt, and include the lines to speak. Optionally pick a built-in "
"preset voice, clone voices from up to 3 reference clips (tagged @Audio1-3 in the prompt), "
"or derive a voice from a character image. Up to 2 minutes of audio per run."
"or derive a voice from a character image. Up to 2 minutes of audio per run. "
"The multilingual model supports 20 languages and timestamp-based timing control."
),
inputs=[
IO.String.Input(
@ -2701,7 +2702,9 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
"Describe the voice(s), emotion, pacing, ambience, background music and sound "
"effects, and include the lines to speak (name characters inline for dialogue). "
"In 'audio reference' mode, refer to connected clips by order as @Audio1, @Audio2, "
"@Audio3. Maximum 3000 characters."
"@Audio3. With the multilingual model, a quoted line can start with a timestamp "
'range that controls when and how long it is spoken, e.g. "[5.5s:8.0s] Wait for me!". '
"Write the prompt in the same language as the lines to speak. Maximum 3000 characters."
),
),
IO.DynamicCombo.Input(
@ -2796,6 +2799,19 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
tooltip="Seed controls whether the node should re-run; "
"results are non-deterministic regardless of seed.",
),
IO.Combo.Input(
"model",
options=["seed-audio-1.0-multilingual", "seed-audio-1.0"],
default="seed-audio-1.0-multilingual",
optional=True,
tooltip=(
"seed-audio-1.0-multilingual: 20 languages (English, Chinese, Japanese, Korean, "
"Mexican & Castilian Spanish, Indonesian, German, Brazilian Portuguese, French, "
"Thai, Vietnamese, Malay, Filipino, Italian, Russian, Dutch, Polish, Turkish, "
'Swedish) plus per-sentence timing control via "[5.5s:8.0s] ..." timestamps. '
"seed-audio-1.0: English and Chinese only, no timing control."
),
),
],
outputs=[IO.Audio.Output()],
hidden=[
@ -2819,6 +2835,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
loudness_rate: int,
pitch_rate: int,
seed: int,
model: str = "seed-audio-1.0-multilingual",
) -> IO.NodeOutput:
mode = reference_mode["reference_mode"]
audio_indices = connected_audio_indices(reference_mode)
@ -2845,6 +2862,7 @@ class ByteDanceSeedAudioNode(IO.ComfyNode):
ApiEndpoint(path="/proxy/byteplus/api/v3/tts/create", method="POST"),
response_model=SeedAudioResponse,
data=SeedAudioRequest(
model=model,
text_prompt=text_prompt,
references=references,
audio_config=SeedAudioConfig(

View File

@ -34,13 +34,6 @@ SEED_MODELS: dict[str, str] = {
"Seed 2.0 Mini": "seed-2-0-mini-260215",
}
# USD per 1M tokens: (input, cache_hit_input, output)
_SEED_PRICES_PER_MILLION: dict[str, tuple[float, float, float]] = {
"seed-2-0-pro-260328": (0.50, 0.10, 3.00),
"seed-2-0-lite-260228": (0.25, 0.05, 2.00),
"seed-2-0-mini-260215": (0.10, 0.02, 0.40),
}
def _seed_model_inputs(max_images: int = SEED_MAX_IMAGES, max_videos: int = SEED_MAX_VIDEOS):
return [
@ -74,24 +67,6 @@ def _seed_model_inputs(max_images: int = SEED_MAX_IMAGES, max_videos: int = SEED
]
def _calculate_price(model_id: str, response: BytePlusResponseObject) -> float | None:
"""Compute approximate USD price from response usage."""
if not response.usage:
return None
rates = _SEED_PRICES_PER_MILLION.get(model_id)
if rates is None:
return None
input_rate, cache_hit_rate, output_rate = rates
input_tokens = response.usage.input_tokens or 0
output_tokens = response.usage.output_tokens or 0
cached = 0
if response.usage.input_tokens_details:
cached = response.usage.input_tokens_details.cached_tokens or 0
fresh_input = max(0, input_tokens - cached)
total = fresh_input * input_rate + cached * cache_hit_rate + output_tokens * output_rate
return total / 1_000_000.0
def _get_text_from_response(response: BytePlusResponseObject) -> str:
"""Extract concatenated text from all assistant message output_text blocks."""
if not response.output:
@ -251,7 +226,6 @@ class ByteDanceSeedNode(IO.ComfyNode):
store=False,
stream=False,
),
price_extractor=lambda r: _calculate_price(model_id, r),
)
if response.error:
raise ValueError(f"Seed API error ({response.error.code}): {response.error.message}")

View File

@ -35,7 +35,6 @@ from comfy_api_nodes.apis.gemini import (
GeminiSystemInstructionContent,
GeminiTextPart,
GeminiThinkingConfig,
Modality,
)
from comfy_api_nodes.util import (
ApiEndpoint,
@ -60,6 +59,7 @@ GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions"
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
GEMINI_URL_INPUT_BUDGET = 10
GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024
GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB
GEMINI_IMAGE_SYS_PROMPT = (
"You are an expert image-generation engine. You must ALWAYS produce an image.\n"
"Interpret all user input—regardless of "
@ -237,60 +237,6 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug
return torch.cat(image_tensors, dim=0)
def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None:
if not response.modelVersion:
return None
# Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing
if response.modelVersion == "gemini-2.5-pro":
input_tokens_price = 1.25
output_text_tokens_price = 10.0
output_image_tokens_price = 0.0
elif response.modelVersion == "gemini-2.5-flash":
input_tokens_price = 0.30
output_text_tokens_price = 2.50
output_image_tokens_price = 0.0
elif response.modelVersion == "gemini-2.5-flash-image":
input_tokens_price = 0.30
output_text_tokens_price = 2.50
output_image_tokens_price = 30.0
elif response.modelVersion in ("gemini-3-pro-preview", "gemini-3.1-pro-preview"):
input_tokens_price = 2
output_text_tokens_price = 12.0
output_image_tokens_price = 0.0
elif response.modelVersion in ("gemini-3.1-flash-lite-preview", "gemini-3.1-flash-lite"):
input_tokens_price = 0.25
output_text_tokens_price = 1.50
output_image_tokens_price = 0.0
elif response.modelVersion == "gemini-3.5-flash":
input_tokens_price = 1.50
output_text_tokens_price = 9.0
output_image_tokens_price = 0.0
elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"):
input_tokens_price = 2
output_text_tokens_price = 12.0
output_image_tokens_price = 120.0
elif response.modelVersion in ("gemini-3.1-flash-image-preview", "gemini-3.1-flash-image"):
input_tokens_price = 0.5
output_text_tokens_price = 3.0
output_image_tokens_price = 60.0
elif response.modelVersion == "gemini-3.1-flash-lite-image":
input_tokens_price = 0.25
output_text_tokens_price = 1.50
output_image_tokens_price = 30.0
else:
return None
final_price = response.usageMetadata.promptTokenCount * input_tokens_price
if response.usageMetadata.candidatesTokensDetails:
for i in response.usageMetadata.candidatesTokensDetails:
if i.modality == Modality.IMAGE:
final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models
else:
final_price += output_text_tokens_price * i.tokenCount
if response.usageMetadata.thoughtsTokenCount:
final_price += output_text_tokens_price * response.usageMetadata.thoughtsTokenCount
return final_price / 1_000_000.0
def get_text_from_interaction(interaction: GeminiInteraction) -> str:
"""Extract and concatenate all model output text from an Interactions API response."""
texts = []
@ -325,24 +271,6 @@ async def get_video_from_interaction(
)
def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None:
if interaction.usage is None:
return None
input_tokens_price = 1.5
output_tokens_prices = {"text": 9.0, "video": 17.5}
thoughts_tokens_price = 9.0
final_price = 0.0
for i in interaction.usage.input_tokens_by_modality or []:
if i.tokens:
final_price += input_tokens_price * i.tokens
for i in interaction.usage.output_tokens_by_modality or []:
if i.tokens and i.modality in output_tokens_prices:
final_price += output_tokens_prices[i.modality] * i.tokens
if interaction.usage.total_thought_tokens:
final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens
return final_price / 1_000_000.0
def create_video_parts(video_input: Input.Video) -> list[GeminiPart]:
"""Convert a single video input to Gemini API compatible parts (inline MP4/H.264)."""
base_64_string = video_to_base64_string(
@ -469,9 +397,10 @@ async def build_gemini_media_parts(
part, nbytes = _media_inline_part(kind, payload)
inline_bytes += nbytes
if inline_bytes > max_inline_bytes:
detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else ""
raise ValueError(
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first "
f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media."
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). "
"Reduce the number or size of attached media."
)
parts.append(part)
return parts
@ -655,7 +584,6 @@ class GeminiNode(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
output_text = get_text_from_response(response)
@ -870,7 +798,6 @@ class GeminiNodeV2(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
output_text = get_text_from_response(response)
@ -1083,7 +1010,6 @@ class GeminiImage(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
return IO.NodeOutput(await get_image_from_response(response), get_text_from_response(response))
@ -1223,7 +1149,6 @@ class GeminiImage2(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
return IO.NodeOutput(await get_image_from_response(response), get_text_from_response(response))
@ -1383,7 +1308,6 @@ class GeminiNanoBanana2(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
return IO.NodeOutput(
await get_image_from_response(response),
@ -1608,7 +1532,6 @@ class GeminiNanoBanana2V2(IO.ComfyNode):
systemInstruction=gemini_system_prompt,
),
response_model=GeminiGenerateContentResponse,
price_extractor=calculate_tokens_price,
)
return IO.NodeOutput(
await get_image_from_response(response),
@ -1738,7 +1661,14 @@ class GeminiVideoOmni(IO.ComfyNode):
parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = []
if images or videos:
media_parts = await build_gemini_media_parts(cls, images, [], videos)
# The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL.
media_parts = await build_gemini_media_parts(
cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES
)
video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts)
media_parts += await build_gemini_media_parts(
cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes
)
parts.extend(to_interaction_media_part(p) for p in media_parts)
parts.append(GeminiInteractionTextPart(text=prompt))
interaction = await sync_op(
@ -1753,7 +1683,6 @@ class GeminiVideoOmni(IO.ComfyNode):
),
),
response_model=GeminiInteraction,
price_extractor=calculate_interaction_tokens_price,
)
if interaction.status != "completed":
model_message = get_text_from_interaction(interaction).strip()

View File

@ -155,7 +155,6 @@ class GrokImageNode(IO.ComfyNode):
resolution=resolution.lower(),
),
response_model=ImageGenerationResponse,
price_extractor=_extract_grok_price,
)
if len(response.data) == 1:
return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url))
@ -351,7 +350,6 @@ class GrokImageEditNode(IO.ComfyNode):
aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio,
),
response_model=ImageGenerationResponse,
price_extractor=_extract_grok_price,
)
if len(response.data) == 1:
return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url))
@ -488,7 +486,6 @@ class GrokImageEditNodeV2(IO.ComfyNode):
aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio,
),
response_model=ImageGenerationResponse,
price_extractor=_extract_grok_price,
)
if len(response.data) == 1:
return IO.NodeOutput(await download_url_to_image_tensor(response.data[0].url))

View File

@ -364,19 +364,6 @@ class OpenAIDalle3(IO.ComfyNode):
return IO.NodeOutput(await validate_and_cast_response(response))
def calculate_tokens_price_image_1(response: OpenAIImageGenerationResponse) -> float | None:
# https://platform.openai.com/docs/pricing
return ((response.usage.input_tokens * 10.0) + (response.usage.output_tokens * 40.0)) / 1_000_000.0
def calculate_tokens_price_image_1_5(response: OpenAIImageGenerationResponse) -> float | None:
return ((response.usage.input_tokens * 8.0) + (response.usage.output_tokens * 32.0)) / 1_000_000.0
def calculate_tokens_price_image_2_0(response: OpenAIImageGenerationResponse) -> float | None:
return ((response.usage.input_tokens * 8.0) + (response.usage.output_tokens * 30.0)) / 1_000_000.0
class OpenAIGPTImage1(IO.ComfyNode):
@classmethod
@ -570,15 +557,10 @@ class OpenAIGPTImage1(IO.ComfyNode):
if size not in ("auto", "1024x1024", "1024x1536", "1536x1024"):
raise ValueError(f"Resolution {size} is only supported by GPT Image 2 model")
if model == "gpt-image-1":
price_extractor = calculate_tokens_price_image_1
elif model == "gpt-image-1.5":
price_extractor = calculate_tokens_price_image_1_5
elif model == "gpt-image-2":
price_extractor = calculate_tokens_price_image_2_0
if model == "gpt-image-2":
if background == "transparent":
raise ValueError("Transparent background is not supported for GPT Image 2 model")
else:
elif model not in ("gpt-image-1", "gpt-image-1.5"):
raise ValueError(f"Unknown model: {model}")
if image is not None:
@ -633,7 +615,6 @@ class OpenAIGPTImage1(IO.ComfyNode):
),
content_type="multipart/form-data",
files=files,
price_extractor=price_extractor,
)
else:
response = await sync_op(
@ -650,7 +631,6 @@ class OpenAIGPTImage1(IO.ComfyNode):
size=size,
moderation="low",
),
price_extractor=price_extractor,
)
return IO.NodeOutput(await validate_and_cast_response(response))
@ -879,13 +859,7 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode):
)
size = f"{custom_width}x{custom_height}"
if model_id == "gpt-image-1":
price_extractor = calculate_tokens_price_image_1
elif model_id == "gpt-image-1.5":
price_extractor = calculate_tokens_price_image_1_5
elif model_id == "gpt-image-2":
price_extractor = calculate_tokens_price_image_2_0
else:
if model_id not in ("gpt-image-1", "gpt-image-1.5", "gpt-image-2"):
raise ValueError(f"Unknown model: {model_id}")
if image_tensors:
@ -944,7 +918,6 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode):
),
content_type="multipart/form-data",
files=files,
price_extractor=price_extractor,
)
else:
response = await sync_op(
@ -960,7 +933,6 @@ class OpenAIGPTImageNodeV2(IO.ComfyNode):
size=size,
moderation="low",
),
price_extractor=price_extractor,
)
return IO.NodeOutput(await validate_and_cast_response(response))

View File

@ -45,27 +45,40 @@ class _ModelSpec:
MODELS: list[_ModelSpec] = [
_ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.000005, 0.000025, max_images=20),
_ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.00003, 0.00018, max_images=20),
_ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.000005, 0.00003, max_images=20),
_ModelSpec("google/gemini-3.5-flash", "reasoning", 0.0000015, 0.000009, max_images=20, max_videos=4),
_ModelSpec("x-ai/grok-4.20", "reasoning", 0.00000125, 0.0000025, max_images=20),
_ModelSpec("x-ai/grok-4.3", "reasoning", 0.00000125, 0.0000025, max_images=20),
_ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.000000435, 0.00000087),
_ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.000000112, 0.000000224),
_ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.000000252, 0.000000378),
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.00000104, 0.00000624),
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.000000325, 0.00000195, max_images=10, max_videos=4),
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.0000001875, 0.000001125, max_images=10, max_videos=4),
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.0000005, 0.0000015, max_images=8),
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.0000015, 0.0000075, max_images=8),
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.00000043, 0.00000174),
_ModelSpec("z-ai/glm-5", "reasoning", 0.0000006, 0.00000192),
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.00000073, 0.00000349, max_images=10),
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.0000006, 0.0000025),
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.000003, 0.000015),
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.000002, 0.000008),
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.000002, 0.000008),
_ModelSpec("anthropic/claude-opus-5", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("anthropic/claude-opus-4.8", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("anthropic/claude-opus-4.7", "frontier_reasoning", 0.00000715, 0.00003575, max_images=20),
_ModelSpec("anthropic/claude-fable-5", "frontier_reasoning", 0.0000143, 0.0000715, max_images=20),
_ModelSpec("anthropic/claude-sonnet-5", "frontier_reasoning", 0.00000286, 0.0000143, max_images=20),
_ModelSpec("anthropic/claude-haiku-4.5", "frontier_reasoning", 0.00000143, 0.00000715, max_images=20),
_ModelSpec("openai/gpt-5.6-sol-pro", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("openai/gpt-5.6-sol", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("openai/gpt-5.6-terra-pro", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
_ModelSpec("openai/gpt-5.6-terra", "frontier_reasoning", 0.000003575, 0.00002145, max_images=20),
_ModelSpec("openai/gpt-5.6-luna-pro", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
_ModelSpec("openai/gpt-5.6-luna", "frontier_reasoning", 0.00000143, 0.00000858, max_images=20),
_ModelSpec("openai/gpt-5.5-pro", "frontier_reasoning", 0.0000429, 0.0002574, max_images=20),
_ModelSpec("openai/gpt-5.5", "frontier_reasoning", 0.00000715, 0.0000429, max_images=20),
_ModelSpec("google/gemini-3.5-flash", "reasoning", 0.000002145, 0.00001287, max_images=20, max_videos=4),
_ModelSpec("x-ai/grok-4.5", "reasoning", 0.00000286, 0.00000858, max_images=20),
_ModelSpec("x-ai/grok-4.20", "reasoning", 0.0000017875, 0.000003575, max_images=20),
_ModelSpec("x-ai/grok-4.3", "reasoning", 0.0000017875, 0.000003575, max_images=20),
_ModelSpec("deepseek/deepseek-v4-pro", "reasoning", 0.00000062205, 0.0000012441),
_ModelSpec("deepseek/deepseek-v4-flash", "reasoning", 0.00000016016, 0.00000032032),
_ModelSpec("deepseek/deepseek-v3.2", "reasoning", 0.00000036036, 0.00000054054),
_ModelSpec("qwen/qwen3.6-max-preview", "reasoning", 0.0000014872, 0.0000089232),
_ModelSpec("qwen/qwen3.6-plus", "reasoning", 0.00000046475, 0.0000027885, max_images=10, max_videos=4),
_ModelSpec("qwen/qwen3.6-flash", "reasoning", 0.000000268125, 0.00000160875, max_images=10, max_videos=4),
_ModelSpec("mistralai/mistral-large-2512", "standard", 0.000000715, 0.000002145, max_images=8),
_ModelSpec("mistralai/mistral-medium-3-5", "reasoning", 0.000002145, 0.000010725, max_images=8),
_ModelSpec("z-ai/glm-4.6", "reasoning", 0.0000006149, 0.0000024882),
_ModelSpec("z-ai/glm-5", "reasoning", 0.000000858, 0.0000027456),
_ModelSpec("moonshotai/kimi-k3", "reasoning", 0.00000429, 0.00002145, max_images=10),
_ModelSpec("moonshotai/kimi-k2.6", "reasoning", 0.0000010439, 0.0000049907, max_images=10),
_ModelSpec("moonshotai/kimi-k2-thinking", "reasoning", 0.000000858, 0.000003575),
_ModelSpec("perplexity/sonar-pro", "perplexity", 0.00000429, 0.00002145),
_ModelSpec("perplexity/sonar-reasoning-pro", "perplexity_reasoning", 0.00000286, 0.00001144),
_ModelSpec("perplexity/sonar-deep-research", "perplexity_reasoning", 0.00000286, 0.00001144),
]
_MODELS_BY_SLUG: dict[str, _ModelSpec] = {m.slug: m for m in MODELS}
@ -146,12 +159,6 @@ def _build_model_options() -> list[IO.DynamicCombo.Option]:
return [IO.DynamicCombo.Option(spec.slug, _inputs_for_model(spec)) for spec in MODELS]
def _calculate_price(response: OpenRouterChatResponse) -> float | None:
if response.usage and response.usage.cost is not None:
return float(response.usage.cost)
return None
def _price_badge_jsonata() -> str:
rates_pairs = []
for spec in MODELS:
@ -269,8 +276,8 @@ class OpenRouterLLMNode(IO.ComfyNode):
essentials_category="Text Generation",
description=(
"Generate text responses through OpenRouter. Routes to a curated set of popular "
"models from xAI, DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and "
"Perplexity Sonar."
"models from Anthropic (Claude), OpenAI (GPT), Google (Gemini), xAI (Grok), "
"DeepSeek, Qwen, Mistral, Z.AI (GLM), Moonshot (Kimi), and Perplexity Sonar."
),
inputs=[
IO.String.Input(
@ -359,7 +366,6 @@ class OpenRouterLLMNode(IO.ComfyNode):
ApiEndpoint(path=OPENROUTER_CHAT_ENDPOINT, method="POST"),
response_model=OpenRouterChatResponse,
data=request,
price_extractor=_calculate_price,
)
return IO.NodeOutput(_extract_text(response))

View File

@ -399,7 +399,7 @@ class RecraftTextToImageNode(IO.ComfyNode):
def define_schema(cls):
return IO.Schema(
node_id="RecraftTextToImageNode",
display_name="Recraft Text to Image",
display_name="Recraft V3 Text to Image",
category="partner/image/Recraft",
description="Generates images synchronously based on prompt and resolution.",
inputs=[
@ -511,7 +511,7 @@ class RecraftImageToImageNode(IO.ComfyNode):
def define_schema(cls):
return IO.Schema(
node_id="RecraftImageToImageNode",
display_name="Recraft Image to Image",
display_name="Recraft V3 Image to Image",
category="partner/image/Recraft",
description="Modify image based on prompt and strength.",
inputs=[
@ -731,7 +731,7 @@ class RecraftTextToVectorNode(IO.ComfyNode):
def define_schema(cls):
return IO.Schema(
node_id="RecraftTextToVectorNode",
display_name="Recraft Text to Vector",
display_name="Recraft V3 Text to Vector",
category="partner/image/Recraft",
description="Generates SVG synchronously based on prompt and resolution.",
inputs=[
@ -1087,7 +1087,7 @@ class RecraftV4TextToImageNode(IO.ComfyNode):
node_id="RecraftV4TextToImageNode",
display_name="Recraft V4 Text to Image",
category="partner/image/Recraft",
description="Generates images using Recraft V4 or V4 Pro models.",
description="Generates images using Recraft V4 and V4.1 models.",
inputs=[
IO.String.Input(
"prompt",
@ -1097,11 +1097,56 @@ class RecraftV4TextToImageNode(IO.ComfyNode):
IO.String.Input(
"negative_prompt",
multiline=True,
tooltip="An optional text description of undesired elements on an image.",
tooltip="This input is ignored: negative prompt is not supported by "
"Recraft V4 and V4.1 models.",
),
IO.DynamicCombo.Input(
"model",
options=[
IO.DynamicCombo.Option(
"recraftv4_1",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_SIZES,
default="1024x1024",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_utility",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_SIZES,
default="1024x1024",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_pro",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_PRO_SIZES,
default="2048x2048",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_utility_pro",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_PRO_SIZES,
default="2048x2048",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4",
[
@ -1162,7 +1207,14 @@ class RecraftV4TextToImageNode(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]),
expr="""
(
$prices := {"recraftv4": 0.04, "recraftv4_pro": 0.25};
$prices := {
"recraftv4_1": 0.035,
"recraftv4_1_utility": 0.035,
"recraftv4_1_pro": 0.21,
"recraftv4_1_utility_pro": 0.21,
"recraftv4": 0.04,
"recraftv4_pro": 0.25
};
{"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n}
)
""",
@ -1179,14 +1231,13 @@ class RecraftV4TextToImageNode(IO.ComfyNode):
seed: int,
recraft_controls: RecraftControls | None = None,
) -> IO.NodeOutput:
validate_string(prompt, strip_whitespace=False, min_length=1, max_length=10000)
validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000)
response = await sync_op(
cls,
ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"),
response_model=RecraftImageGenerationResponse,
data=RecraftImageGenerationRequest(
prompt=prompt,
negative_prompt=negative_prompt if negative_prompt else None,
model=model["model"],
size=model["size"],
n=n,
@ -1211,7 +1262,7 @@ class RecraftV4TextToVectorNode(IO.ComfyNode):
node_id="RecraftV4TextToVectorNode",
display_name="Recraft V4 Text to Vector",
category="partner/image/Recraft",
description="Generates SVG using Recraft V4 or V4 Pro models.",
description="Generates SVG using Recraft V4 and V4.1 models.",
inputs=[
IO.String.Input(
"prompt",
@ -1221,11 +1272,56 @@ class RecraftV4TextToVectorNode(IO.ComfyNode):
IO.String.Input(
"negative_prompt",
multiline=True,
tooltip="An optional text description of undesired elements on an image.",
tooltip="This input is ignored: negative prompt is not supported by "
"Recraft V4 and V4.1 models.",
),
IO.DynamicCombo.Input(
"model",
options=[
IO.DynamicCombo.Option(
"recraftv4_1_vector",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_SIZES,
default="1024x1024",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_utility_vector",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_SIZES,
default="1024x1024",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_pro_vector",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_PRO_SIZES,
default="2048x2048",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4_1_utility_pro_vector",
[
IO.Combo.Input(
"size",
options=RECRAFT_V4_PRO_SIZES,
default="2048x2048",
tooltip="The size of the generated image.",
),
],
),
IO.DynamicCombo.Option(
"recraftv4",
[
@ -1286,7 +1382,14 @@ class RecraftV4TextToVectorNode(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]),
expr="""
(
$prices := {"recraftv4": 0.08, "recraftv4_pro": 0.30};
$prices := {
"recraftv4_1_vector": 0.08,
"recraftv4_1_utility_vector": 0.08,
"recraftv4_1_pro_vector": 0.30,
"recraftv4_1_utility_pro_vector": 0.30,
"recraftv4": 0.08,
"recraftv4_pro": 0.30
};
{"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n}
)
""",
@ -1303,18 +1406,17 @@ class RecraftV4TextToVectorNode(IO.ComfyNode):
seed: int,
recraft_controls: RecraftControls | None = None,
) -> IO.NodeOutput:
validate_string(prompt, strip_whitespace=False, min_length=1, max_length=10000)
validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000)
response = await sync_op(
cls,
ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"),
response_model=RecraftImageGenerationResponse,
data=RecraftImageGenerationRequest(
prompt=prompt,
negative_prompt=negative_prompt if negative_prompt else None,
model=model["model"],
size=model["size"],
n=n,
style="vector_illustration",
style=None if model["model"].endswith("_vector") else "vector_illustration",
substyle=None,
controls=recraft_controls.create_api_model() if recraft_controls else None,
),

View File

@ -62,13 +62,6 @@ def _postprocessing_inputs():
]
def _reve_price_extractor(headers: dict) -> float | None:
credits_used = headers.get("x-reve-credits-used")
if credits_used is not None:
return float(credits_used) / 524.48
return None
def _reve_response_header_validator(headers: dict) -> None:
error_code = headers.get("x-reve-error-code")
if error_code:
@ -180,7 +173,6 @@ class ReveImageCreateNode(IO.ComfyNode):
headers={"Accept": "image/webp"},
),
as_binary=True,
price_extractor=_reve_price_extractor,
response_header_validator=_reve_response_header_validator,
data=ReveImageCreateRequest(
prompt=prompt,
@ -279,7 +271,6 @@ class ReveImageEditNode(IO.ComfyNode):
headers={"Accept": "image/webp"},
),
as_binary=True,
price_extractor=_reve_price_extractor,
response_header_validator=_reve_response_header_validator,
data=ReveImageEditRequest(
edit_instruction=edit_instruction,
@ -396,7 +387,6 @@ class ReveImageRemixNode(IO.ComfyNode):
headers={"Accept": "image/webp"},
),
as_binary=True,
price_extractor=_reve_price_extractor,
response_header_validator=_reve_response_header_validator,
data=ReveImageRemixRequest(
prompt=prompt,

View File

@ -194,6 +194,7 @@ class RunwayImageToVideoNodeGen3a(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
),
is_deprecated=True,
)
@classmethod
@ -390,6 +391,7 @@ class RunwayFirstLastFrameNode(IO.ComfyNode):
depends_on=IO.PriceBadgeDepends(widgets=["duration"]),
expr="""{"type":"usd","usd": 0.0715 * widgets.duration}""",
),
is_deprecated=True,
)
@classmethod

View File

@ -2,8 +2,10 @@ import asyncio
import contextlib
import json
import logging
import math
import time
import uuid
import weakref
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from enum import Enum
@ -84,11 +86,37 @@ class _PollUIState:
_RETRY_STATUS = {408, 500, 502, 503, 504} # status 429 is handled separately
_MAX_RETRY_AFTER_WAIT = 150.0 # Cap a server Retry-After at this many seconds so a large hint can't block execution
PRICE_CREDITS_HEADER = "X-Comfy-Credits-Used"
"""Proxy response header with the actual cost in Comfy credits. When present on any successful proxied response,
it takes precedence over ``price_extractor``."""
_credits_used_by_execution: "weakref.WeakKeyDictionary[type, float]" = weakref.WeakKeyDictionary()
"""Last PRICE_CREDITS_HEADER value per node execution, keyed by the node's per-execution class clone."""
COMPLETED_STATUSES = ["succeeded", "succeed", "success", "completed", "finished", "done", "complete"]
FAILED_STATUSES = ["cancelled", "canceled", "canceling", "fail", "failed", "error"]
QUEUED_STATUSES = ["created", "queued", "queueing", "submitted", "initializing", "wait", "in_queue"]
def _maybe_remember_credits_used(node_cls: type[IO.ComfyNode], header_value: str | None) -> None:
"""Remember a PRICE_CREDITS_HEADER value from a successful proxied response."""
if not header_value:
return
try:
credits_used = float(header_value)
except (TypeError, ValueError):
logging.debug("Ignoring malformed %s header: %r", PRICE_CREDITS_HEADER, header_value)
return
if not math.isfinite(credits_used) or credits_used < 0:
logging.debug("Ignoring out-of-range %s header: %r", PRICE_CREDITS_HEADER, header_value)
return
_credits_used_by_execution[node_cls] = credits_used + 0.0 # normalize -0.0
def _get_remembered_credits_used(node_cls: type[IO.ComfyNode]) -> float | None:
return _credits_used_by_execution.get(node_cls)
async def sync_op(
cls: type[IO.ComfyNode],
endpoint: ApiEndpoint,
@ -450,10 +478,15 @@ def _display_text(
display_lines: list[str] = []
if status:
display_lines.append(f"Status: {status.capitalize() if isinstance(status, str) else status}")
if price is not None:
server_credits = _get_remembered_credits_used(node_cls)
if server_credits is not None:
p = f"{server_credits:,.2f}".rstrip("0").rstrip(".")
elif price is not None:
p = f"{float(price) * 211:,.1f}".rstrip("0").rstrip(".")
if p != "0":
display_lines.append(f"Price: {p} credits")
else:
p = None
if p is not None and p != "0":
display_lines.append(f"Price: {p} credits")
if text is not None:
display_lines.append(text)
if display_lines:
@ -606,7 +639,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool):
"""Core request with retries, per-second interruption monitoring, true cancellation, and friendly errors."""
url = cfg.endpoint.path
parsed_url = urlparse(url)
if not parsed_url.scheme and not parsed_url.netloc: # is URL relative?
is_comfy_api_request = not parsed_url.scheme and not parsed_url.netloc # is URL relative?
if is_comfy_api_request:
url = urljoin(default_base_url().rstrip("/") + "/", url.lstrip("/"))
method = cfg.endpoint.method
@ -644,7 +678,7 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool):
logging.debug("[DEBUG] HTTP %s %s (attempt %d)", method, url, attempt)
payload_headers = {"Accept": "*/*"} if expect_binary else {"Accept": "application/json"}
if not parsed_url.scheme and not parsed_url.netloc: # is URL relative?
if is_comfy_api_request:
payload_headers.update(get_comfy_api_headers(cfg.node_cls))
if cfg.endpoint.headers:
payload_headers.update(cfg.endpoint.headers)
@ -804,6 +838,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool):
)
bytes_payload = bytes(buff)
resp_headers = {k.lower(): v for k, v in resp.headers.items()}
if is_comfy_api_request:
_maybe_remember_credits_used(cfg.node_cls, resp.headers.get(PRICE_CREDITS_HEADER))
if cfg.price_extractor:
with contextlib.suppress(Exception):
extracted_price = cfg.price_extractor(resp_headers)
@ -831,6 +867,8 @@ async def _request_base(cfg: _RequestConfig, expect_binary: bool):
except json.JSONDecodeError:
payload = {"_raw": text}
response_content_to_log = payload if isinstance(payload, dict) else text
if is_comfy_api_request:
_maybe_remember_credits_used(cfg.node_cls, resp.headers.get(PRICE_CREDITS_HEADER))
with contextlib.suppress(Exception):
extracted_price = cfg.price_extractor(payload) if cfg.price_extractor else None
operation_succeeded = True

View File

@ -5,7 +5,7 @@ import psutil
import time
import torch
from typing import Sequence, Mapping, Dict
from comfy.model_patcher import ModelPatcher
from comfy.model_patcher import is_model_patcher_output
from comfy_execution.graph import DynamicPrompt
from abc import ABC, abstractmethod
@ -524,6 +524,13 @@ class RAMPressureCache(LRUCache):
def __init__(self, key_class, enable_providers=False):
super().__init__(key_class, 0, enable_providers=enable_providers)
self.timestamps = {}
self.active_evictions = False
self.full_evictions = False
async def set_prompt(self, dynprompt, node_ids, is_changed_cache):
self.active_evictions = False
self.full_evictions = False
await super().set_prompt(dynprompt, node_ids, is_changed_cache)
def clean_unused(self):
self._clean_subcaches()
@ -567,7 +574,7 @@ class RAMPressureCache(LRUCache):
elif isinstance(output, torch.Tensor) and output.device.type == 'cpu':
ram_usage += output.numel() * output.element_size()
oom_ram_usage += output.numel() * output.element_size()
elif isinstance(output, ModelPatcher) and self.used_generation[key] != self.generation:
elif is_model_patcher_output(output) and self.used_generation[key] != self.generation:
#old ModelPatchers are the first to go
oom_ram_usage = 1e30
scan_list_for_ram_usage(cache_entry.outputs)
@ -588,4 +595,8 @@ class RAMPressureCache(LRUCache):
self.timestamps.pop(key, None)
self.children.pop(key, None)
freed += ram_usage
if freed and free_active:
self.active_evictions = True
if min_entry_size == 0:
self.full_evictions = True
return freed

View File

@ -195,9 +195,10 @@ class ExecutionList(TopologicalSort):
ExecutionList implements a topological dissolve of the graph. After a node is staged for execution,
it can still be returned to the graph after having further dependencies added.
"""
def __init__(self, dynprompt, output_cache):
def __init__(self, dynprompt, output_cache, output_link_callback=None):
super().__init__(dynprompt)
self.output_cache = output_cache
self.output_link_callback = output_link_callback
self.staged_node_id = None
self.execution_cache = {}
self.execution_cache_listeners = {}
@ -205,13 +206,16 @@ class ExecutionList(TopologicalSort):
def is_cached(self, node_id):
return self.output_cache.get_local(node_id) is not None
def cache_link(self, from_node_id, to_node_id):
def cache_link(self, from_node_id, to_node_id, from_socket=None):
if to_node_id not in self.execution_cache:
self.execution_cache[to_node_id] = {}
self.execution_cache[to_node_id][from_node_id] = self.output_cache.get_local(from_node_id)
value = self.output_cache.get_local(from_node_id)
self.execution_cache[to_node_id][from_node_id] = value
if from_node_id not in self.execution_cache_listeners:
self.execution_cache_listeners[from_node_id] = set()
self.execution_cache_listeners[from_node_id].add(to_node_id)
self.execution_cache_listeners[from_node_id].add((to_node_id, from_socket))
if value is not None and from_socket is not None and self.output_link_callback is not None:
self.output_link_callback(value.outputs[from_socket])
def get_cache(self, from_node_id, to_node_id):
if to_node_id not in self.execution_cache:
@ -225,13 +229,15 @@ class ExecutionList(TopologicalSort):
def cache_update(self, node_id, value):
if node_id in self.execution_cache_listeners:
for to_node_id in self.execution_cache_listeners[node_id]:
for to_node_id, from_socket in self.execution_cache_listeners[node_id]:
if to_node_id in self.execution_cache:
self.execution_cache[to_node_id][node_id] = value
if from_socket is not None and self.output_link_callback is not None:
self.output_link_callback(value.outputs[from_socket])
def add_strong_link(self, from_node_id, from_socket, to_node_id):
super().add_strong_link(from_node_id, from_socket, to_node_id)
self.cache_link(from_node_id, to_node_id)
self.cache_link(from_node_id, to_node_id, from_socket)
async def stage_node_execution(self):
assert self.staged_node_id is None

View File

@ -170,6 +170,19 @@ def is_previewable(media_type: str, item: dict) -> bool:
return False
def is_text_preview(media_type: str, item: dict) -> bool:
"""
Check if a previewable output item is textual rather than visual media.
Saved text files (SaveText's .txt/.md/.json) are real outputs but must not
outrank visual media when picking the job preview.
"""
if media_type == 'text':
return True
filename = item.get('filename', '').lower()
return any(filename.endswith(ext) for ext in TEXT_EXTENSIONS)
def normalize_queue_item(item: tuple, status: str) -> dict:
"""Convert queue item tuple to unified job dict.
@ -259,8 +272,13 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
Returns (outputs_count, preview_output).
Preview priority (matching frontend):
1. type="output" with previewable media
2. Any previewable media
1. type="output" visual media (saved images/video/audio/3d)
2. any other previewable visual media (e.g. temp/preview images)
3. saved text file (e.g. SaveText's .txt/.md/.json)
4. raw text (only when the job produced nothing else previewable)
Text is kept in its own slots so node/execution order can't let a text
output mask a visual one (e.g. a text node that runs before an image).
Text content entries (strings under 'text') are preview-only metadata,
matching the frontend's METADATA_KEYS: they can serve as the fallback
@ -269,6 +287,8 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
count = 0
preview_output = None
fallback_preview = None
text_file_fallback = None
text_fallback = None
for node_id, node_outputs in outputs.items():
if not isinstance(node_outputs, dict):
@ -296,8 +316,8 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
'nodeId': node_id,
'mediaType': media_type
}
if fallback_preview is None:
fallback_preview = enriched
if text_fallback is None:
text_fallback = enriched
continue
# normalize_output_item returned a dict (e.g. 3D file)
item = normalized
@ -314,12 +334,15 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
}
if 'mediaType' not in item:
enriched['mediaType'] = media_type
if item.get('type') == 'output':
if is_text_preview(media_type, item):
if text_file_fallback is None:
text_file_fallback = enriched
elif item.get('type') == 'output':
preview_output = enriched
elif fallback_preview is None:
fallback_preview = enriched
return count, preview_output or fallback_preview
return count, preview_output or fallback_preview or text_file_fallback or text_fallback
def apply_sorting(jobs: list[dict], sort_by: str, sort_order: str) -> list[dict]:

View File

@ -2,6 +2,7 @@ import logging
import os
import json
import av
import numpy as np
import torch
from PIL import Image
@ -9,7 +10,7 @@ from typing_extensions import override
import folder_paths
import node_helpers
from comfy_api.latest import ComfyExtension, io
from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types
def load_and_process_images(image_files, input_dir):
@ -42,6 +43,130 @@ def load_and_process_images(image_files, input_dir):
return output_images
def secure_subfolder_path(base_dir, folder_name):
"""Resolve folder_name inside base_dir, rejecting anything that escapes it.
Blocks '..', absolute paths, drive letters and symlink escapes using the
same realpath containment check as the core file endpoints.
"""
target = os.path.abspath(os.path.join(base_dir, folder_name))
if not folder_paths.is_within_directory(base_dir, target):
raise ValueError(f"Invalid folder name {folder_name!r}: resolves outside of {base_dir}")
return target
def list_dataset_folders():
"""Relative paths of dataset folders found under all dataset roots.
Any subfolder containing a metadata.json or *.safetensors shard counts as
a dataset; the walk doesn't descend into matched folders.
Symlinked directories are followed, but symlink loops are avoided.
"""
found = set()
for root in folder_paths.get_folder_paths("datasets"):
if not os.path.isdir(root):
continue
root = os.path.abspath(root)
seen_dirs = set()
for dirpath, subdirs, filenames in os.walk(root, followlinks=True):
try:
st = os.stat(dirpath) # follows symlinks
except OSError:
subdirs[:] = []
continue
dir_key = (st.st_dev, st.st_ino)
if dir_key in seen_dirs:
subdirs[:] = []
continue
seen_dirs.add(dir_key)
if dirpath != root and (
"metadata.json" in filenames
or any(f.endswith(".safetensors") for f in filenames)
):
found.add(os.path.relpath(dirpath, root).replace(os.sep, "/"))
subdirs[:] = []
continue
kept_subdirs = []
for name in subdirs:
child = os.path.join(dirpath, name)
try:
child_st = os.stat(child) # follows symlinks
except OSError:
continue
child_key = (child_st.st_dev, child_st.st_ino)
if child_key not in seen_dirs:
kept_subdirs.append(name)
subdirs[:] = kept_subdirs
return sorted(found)
def get_dataset_save_dir(folder_name):
"""Resolve the folder to save a new dataset into, inside the default root.
The folder is not created here; callers makedirs after validation.
"""
root = folder_paths.get_folder_paths("datasets")[0]
target = secure_subfolder_path(root, folder_name)
if os.path.realpath(target) == os.path.realpath(root):
raise ValueError("folder_name must name a subfolder of the datasets directory, e.g. 'my_dataset'.")
return target
def get_dataset_dir(folder_name):
"""Find an existing dataset folder by relative name across all dataset roots."""
roots = folder_paths.get_folder_paths("datasets")
for root in roots:
target = secure_subfolder_path(root, folder_name)
if os.path.realpath(target) == os.path.realpath(root):
raise ValueError("folder_name must name a subfolder of the datasets directory, e.g. 'my_dataset'.")
if os.path.isdir(target):
return target
raise ValueError(f"Dataset folder {folder_name!r} not found in: {', '.join(roots)}")
VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"]
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 +282,116 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode):
return io.NodeOutput(output_tensor, captions)
class LoadVideoDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video (from Folder)",
category="video",
description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = sorted([
f for f in os.listdir(sub_input_dir)
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
])
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files]
logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}")
return io.NodeOutput(videos)
class LoadVideoTextDataSetFromFolderNode(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="LoadVideoTextDataSetFromFolder",
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
display_name="Load Video-Text (from Folder)",
category="video",
description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
is_experimental=True,
inputs=[
io.Combo.Input(
"folder",
options=folder_paths.get_input_subfolders(),
tooltip="The folder containing video files and .txt captions.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Lazy video references; frames are decoded only when needed downstream.",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="List of text captions.",
),
],
)
@classmethod
def execute(cls, folder):
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
video_files = []
for item in sorted(os.listdir(sub_input_dir)):
path = os.path.join(sub_input_dir, item)
if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS):
video_files.append(path)
elif os.path.isdir(path):
# Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/
repeat = 1
if item.split("_")[0].isdigit():
repeat = int(item.split("_")[0])
video_files.extend([
os.path.join(path, f)
for f in sorted(os.listdir(path))
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
] * repeat)
if not video_files:
raise ValueError(f"No video files found in {sub_input_dir}")
captions = []
for vf in video_files:
caption_path = os.path.splitext(vf)[0] + ".txt"
if os.path.exists(caption_path):
with open(caption_path, "r", encoding="utf-8") as f:
captions.append(f.read().strip())
else:
captions.append("")
videos = [InputImpl.VideoFromFile(vf) for vf in video_files]
logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}")
return io.NodeOutput(videos, captions)
def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True):
"""Utility function to save a list of image tensors to disk.
@ -252,7 +487,7 @@ class SaveImageDataSetToFolderNode(io.ComfyNode):
filename_prefix = filename_prefix[0]
mode = mode[0]
output_dir = os.path.join(folder_paths.get_output_directory(), folder_name)
output_dir = secure_subfolder_path(folder_paths.get_output_directory(), folder_name)
saved_files = save_images_to_folder(images, output_dir, filename_prefix, mode=='overwrite')
logging.info(f"Saved {len(saved_files)} images to {output_dir}.")
@ -306,7 +541,7 @@ class SaveImageTextDataSetToFolderNode(io.ComfyNode):
filename_prefix = filename_prefix[0]
mode = mode[0]
output_dir = os.path.join(folder_paths.get_output_directory(), folder_name)
output_dir = secure_subfolder_path(folder_paths.get_output_directory(), folder_name)
saved_files = save_images_to_folder(images, output_dir, filename_prefix, mode=='overwrite')
# Save captions
@ -470,7 +705,15 @@ class ImageProcessingNode(io.ComfyNode):
@classmethod
def execute(cls, images, **kwargs):
"""Execute the node. Routes to _process or _group_process based on mode."""
"""Execute the node. Routes to _process or _group_process based on mode.
For individual processing (_process), automatically handles multi-frame
inputs (video tensors [T, H, W, C]) by applying _process per-frame and
concatenating the results. This allows all spatial transform nodes to
work with video without modification. Nodes that natively handle batched
tensors (e.g. pure tensor math) can set per_frame_process = False to
skip the per-frame loop.
"""
is_group = cls._detect_processing_mode()
if is_group:
@ -489,7 +732,16 @@ class ImageProcessingNode(io.ComfyNode):
result = cls._group_process(images, **params)
else:
# Individual processing: images is single item, call _process
result = cls._process(images, **params)
# Auto-loop over frames for multi-frame inputs (video [T, H, W, C])
# so that PIL-based spatial transforms work per-frame automatically.
if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True):
results = []
for i in range(images.shape[0]):
frame_result = cls._process(images[i:i + 1], **params)
results.append(frame_result)
result = torch.cat(results, dim=0)
else:
result = cls._process(images, **params)
return io.NodeOutput(result)
@ -803,6 +1055,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 +1086,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 +1108,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 +1190,261 @@ class ShuffleImageTextDatasetNode(io.ComfyNode):
return io.NodeOutput(shuffled_images, shuffled_texts)
# ========== Video Processing Nodes ==========
class VideoFrameSampleNode(io.ComfyNode):
"""Sample a fixed number of frames from a video using various strategies.
For contiguous strategies ("head"/"tail") the result is a fully lazy
VideoInput (no frames decoded). For non-contiguous strategies
("uniform"/"random") only the selected indices are decoded.
"""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoFrameSample",
search_aliases=["sample frames", "extract frames"],
display_name="Sample Video Frame",
category="video",
description="Sample a fixed number of frames from a video using various strategies.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"num_frames",
default=16,
min=1,
max=9999,
tooltip="Number of frames to sample.",
),
io.Combo.Input(
"strategy",
options=["uniform", "head", "tail", "random"],
default="uniform",
tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed (only used with 'random' strategy).",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Sampled video."),
],
)
@classmethod
def execute(cls, video, num_frames, strategy, seed):
total_frames = video.get_frame_count()
num_frames = min(num_frames, total_frames)
fps = float(video.get_frame_rate())
if strategy == "head":
return io.NodeOutput(
video.as_trimmed(0.0, num_frames / fps, strict_duration=False)
)
if strategy == "tail":
start_t = (total_frames - num_frames) / fps
return io.NodeOutput(
video.as_trimmed(start_t, num_frames / fps, strict_duration=False)
)
if strategy == "uniform":
if num_frames == 1:
indices = [total_frames // 2]
else:
indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)]
elif strategy == "random":
rng = np.random.RandomState(seed % (2**32 - 1))
indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist())
else:
raise ValueError(f"Unknown strategy: {strategy}")
return io.NodeOutput(_decode_selected_frames(video, indices))
class VideoTemporalCropNode(io.ComfyNode):
"""Crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video"],
display_name="Crop Video (Temporal)",
category="video/transform",
description="Crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"start_frame",
default=0,
min=0,
max=99999,
tooltip="Starting frame index.",
),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, start_frame, length):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
start_frame = min(start_frame, max(total_frames - 1, 0))
length = min(length, total_frames - start_frame)
return io.NodeOutput(
video.as_trimmed(start_frame / fps, length / fps, strict_duration=False)
)
class VideoRandomTemporalCropNode(io.ComfyNode):
"""Randomly crop a continuous range of frames from a video (fully lazy)."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="VideoRandomTemporalCrop",
search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"],
display_name="Crop Video (Temporal Random)",
category="video/transform",
description="Randomly crop a continuous range of frames from a video.",
is_experimental=True,
inputs=[
io.Video.Input("video", tooltip="Input video."),
io.Int.Input(
"length",
default=16,
min=1,
max=99999,
tooltip="Number of frames to keep.",
),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
],
)
@classmethod
def execute(cls, video, length, seed):
total_frames = video.get_frame_count()
fps = float(video.get_frame_rate())
length = min(length, total_frames)
max_start = total_frames - length
rng = np.random.RandomState(seed % (2**32 - 1))
start = rng.randint(0, max_start + 1) if max_start > 0 else 0
return io.NodeOutput(
video.as_trimmed(start / fps, length / fps, strict_duration=False)
)
class ShuffleVideoDatasetNode(io.ComfyNode):
"""Randomly shuffle the order of videos in the dataset."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Videos List",
category="video/batch",
description="Randomly shuffle the order of videos in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.Int.Input(
"seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed."
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
],
)
@classmethod
def execute(cls, videos, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput([videos[i] for i in indices])
class ShuffleVideoTextDatasetNode(io.ComfyNode):
"""Shuffle videos and their captions together, preserving pairs."""
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ShuffleVideoTextDataset",
search_aliases=["shuffle", "randomize", "mix"],
display_name="Shuffle Pairs of Video-Text",
category="dataset/video",
description="Randomly shuffle the order of pairs of video-text in a list.",
is_experimental=True,
is_input_list=True,
inputs=[
io.Video.Input("videos", tooltip="List of videos to shuffle."),
io.String.Input("texts", tooltip="List of texts to shuffle."),
io.Int.Input(
"seed",
default=0,
min=0,
max=0xFFFFFFFFFFFFFFFF,
tooltip="Random seed.",
),
],
outputs=[
io.Video.Output(
display_name="videos",
is_output_list=True,
tooltip="Shuffled videos",
),
io.String.Output(
display_name="texts",
is_output_list=True,
tooltip="Shuffled texts",
),
],
)
@classmethod
def execute(cls, videos, texts, seed):
seed = seed[0] if isinstance(seed, list) else seed
np.random.seed(seed % (2**32 - 1))
indices = np.random.permutation(len(videos))
return io.NodeOutput(
[videos[i] for i in indices],
[texts[i] for i in indices],
)
# ========== Text Transform Nodes ==========
@ -1443,7 +1953,7 @@ class SaveTrainingDataset(io.ComfyNode):
io.String.Input(
"folder_name",
default="training_dataset",
tooltip="Name of folder to save dataset (inside output directory).",
tooltip="Name of folder to save the dataset into, inside the datasets directory. Subfolders like 'project/run1' are allowed.",
),
io.Int.Input(
"shard_size",
@ -1473,8 +1983,8 @@ class SaveTrainingDataset(io.ComfyNode):
f"Something went wrong in dataset preparation."
)
# Create output directory
output_dir = os.path.join(folder_paths.get_output_directory(), folder_name)
# Create output directory (inside the datasets root, traversal-safe)
output_dir = get_dataset_save_dir(folder_name)
os.makedirs(output_dir, exist_ok=True)
# Prepare data pairs
@ -1533,10 +2043,10 @@ class LoadTrainingDataset(io.ComfyNode):
description="Load encoded training dataset (latents + conditioning) from disk for use in training.",
is_experimental=True,
inputs=[
io.String.Input(
io.Combo.Input(
"folder_name",
default="training_dataset",
tooltip="Name of folder containing the saved dataset (inside output directory).",
options=list_dataset_folders(),
tooltip="Saved dataset to load, from the datasets directory.",
),
],
outputs=[
@ -1555,11 +2065,8 @@ class LoadTrainingDataset(io.ComfyNode):
@classmethod
def execute(cls, folder_name):
# Get dataset directory
dataset_dir = os.path.join(folder_paths.get_output_directory(), folder_name)
if not os.path.exists(dataset_dir):
raise ValueError(f"Dataset directory not found: {dataset_dir}")
# Get dataset directory (searched across all dataset roots, traversal-safe)
dataset_dir = get_dataset_dir(folder_name)
# Find all shard files
shard_files = sorted(
@ -1608,7 +2115,10 @@ class DatasetExtension(ComfyExtension):
LoadImageTextDataSetFromFolderNode,
SaveImageDataSetToFolderNode,
SaveImageTextDataSetToFolderNode,
# Image transform nodes
# Video data loading nodes
LoadVideoDataSetFromFolderNode,
LoadVideoTextDataSetFromFolderNode,
# Image transform nodes (auto-handle video via per-frame processing)
ResizeImagesByShorterEdgeNode,
ResizeImagesByLongerEdgeNode,
CenterCropImagesNode,
@ -1618,6 +2128,12 @@ class DatasetExtension(ComfyExtension):
AdjustContrastNode,
ShuffleDatasetNode,
ShuffleImageTextDatasetNode,
# Video processing nodes (lazy VideoInput in/out)
VideoFrameSampleNode,
VideoTemporalCropNode,
VideoRandomTemporalCropNode,
ShuffleVideoDatasetNode,
ShuffleVideoTextDatasetNode,
# Text transform nodes
TextToLowercaseNode,
TextToUppercaseNode,

View File

@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters:
x: Input tensor of shape (B, C, H, W)
x: Input tensor of shape (..., H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
# Initialize mask with high-frequency scaling factor
mask = torch.ones(x_freq.shape, device=device) * scale_high
m = mask
for d in range(len(x_freq.shape) - 2):
dim = d + 2
for d in range(2):
dim = len(x_freq.shape) - 2 + d
cc = x_freq.shape[dim] // 2
f_c = min(freq_cutoff, cc)
m = m.narrow(dim, cc - f_c, f_c * 2)

View File

@ -2,6 +2,8 @@ import nodes
import node_helpers
import torch
import comfy.model_management
import comfy.model_patcher
import comfy.ops
from typing_extensions import override
from comfy_api.latest import ComfyExtension, io
from comfy.ldm.hunyuan_video.upsampler import HunyuanVideo15SRModel
@ -217,8 +219,11 @@ class LatentUpscaleModelLoader(io.ComfyNode):
model.load_sd(sd)
elif "post_upsample_res_blocks.0.conv2.bias" in sd:
config = json.loads(metadata["config"])
model = LatentUpsampler.from_config(config).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32]))
model.load_state_dict(sd)
model = LatentUpsampler.from_config(config, operations=comfy.ops.disable_weight_init).to(dtype=comfy.model_management.vae_dtype(allowed_dtypes=[torch.bfloat16, torch.float32]))
comfy.model_management.archive_model_dtypes(model)
model_patcher = comfy.model_patcher.CoreModelPatcher(model, load_device=comfy.model_management.get_torch_device(), offload_device=comfy.model_management.unet_offload_device())
model.load_state_dict(sd, assign=model_patcher.is_dynamic())
model = model_patcher
return io.NodeOutput(model)

View File

@ -50,8 +50,8 @@ class GetICLoRAParameters(io.ComfyNode):
factor = 1
if metadata:
try:
factor = max(1, round(float(metadata.get("reference_downscale_factor", 1))))
except (TypeError, ValueError):
factor = max(1, round(float(next(v for k, v in metadata.items() if k.endswith("reference_downscale_factor")))))
except (StopIteration, TypeError, ValueError):
factor = 1
parameters = {"reference_downscale_factor": factor}
return io.NodeOutput(parameters)

View File

@ -107,14 +107,17 @@ class LTXVEmptyLatentAudio(io.ComfyNode):
display_mode=io.NumberDisplay.number,
tooltip="Number of frames.",
),
io.Int.Input(
"frame_rate",
default=25,
min=1,
max=1000,
step=1,
display_mode=io.NumberDisplay.number,
tooltip="Number of frames per second.",
io.MultiType.Input(
io.Float.Input(
"frame_rate",
default=25.0,
min=1.0,
max=1000.0,
step=0.01,
display_mode=io.NumberDisplay.number,
tooltip="Number of frames per second.",
),
[io.Int],
),
io.Int.Input(
"batch_size",
@ -137,7 +140,7 @@ class LTXVEmptyLatentAudio(io.ComfyNode):
def execute(
cls,
frames_number: int,
frame_rate: int,
frame_rate: float,
batch_size: int,
audio_vae,
) -> io.NodeOutput:

View File

@ -38,26 +38,20 @@ class LTXVLatentUpsampler(IO.ComfyNode):
Returns:
tuple: Tuple containing the upsampled latent
"""
device = model_management.get_torch_device()
memory_required = model_management.module_size(upscale_model)
model_dtype = next(upscale_model.parameters()).dtype
device = upscale_model.load_device
model = upscale_model.model
model_dtype = upscale_model.model_dtype()
latents = samples["samples"]
input_dtype = latents.dtype
memory_required += math.prod(latents.shape) * 3000.0 # TODO: more accurate
model_management.free_memory(memory_required, device)
memory_required = math.prod(latents.shape) * 3000.0 # TODO: more accurate
model_management.load_models_gpu([upscale_model], memory_required=memory_required)
try:
upscale_model.to(device) # TODO: use the comfy model management system.
latents = latents.to(dtype=model_dtype, device=device)
latents = latents.to(dtype=model_dtype, device=device)
"""Upsample latents without tiling."""
latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents)
upsampled_latents = upscale_model(latents)
finally:
upscale_model.cpu()
"""Upsample latents without tiling."""
latents = vae.first_stage_model.per_channel_statistics.un_normalize(latents)
upsampled_latents = model(latents)
upsampled_latents = vae.first_stage_model.per_channel_statistics.normalize(
upsampled_latents

103
comfy_extras/nodes_mage.py Normal file
View File

@ -0,0 +1,103 @@
from typing_extensions import override
import comfy.utils
import node_helpers
import torch
import comfy.model_management
from comfy_api.latest import ComfyExtension, io
class TextEncodeMageFlowEdit(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="TextEncodeMageFlowEdit",
category="model/conditioning/mage",
description="Encode an edit instruction with one or more reference images for Mage-Flow-Edit. Reference latents are resized to the output resolution (width/height, or the first image's size when 0). Use the latent output for sampling so the sizes always match.",
inputs=[
io.Clip.Input("clip"),
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
io.String.Input("negative_prompt", multiline=True, dynamic_prompts=True, advanced=True),
io.Vae.Input("vae", optional=True),
io.Autogrow.Input(
"images",
template=io.Autogrow.TemplateNames(
io.Image.Input("image"),
names=[f"image_{i}" for i in range(1, 17)],
min=0,
),
tooltip="Reference image(s) to edit. All references are resized to the output resolution before encoding.",
),
io.Int.Input("width", default=0, min=0, max=8192, step=16, tooltip="Output width. 0 = use the first reference image's size."),
io.Int.Input("height", default=0, min=0, max=8192, step=16, tooltip="Output height. 0 = use the first reference image's size."),
io.Int.Input("batch_size", default=1, min=1, max=4096),
],
outputs=[
io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative"),
io.Latent.Output(display_name="latent"),
],
)
@classmethod
def execute(cls, clip, prompt, negative_prompt="", vae=None, images: io.Autogrow.Type = None, width=0, height=0, batch_size=1) -> io.NodeOutput:
ref_latents = []
images = images or {}
images = [images[name] for name in sorted(images, key=lambda n: int(n.rsplit("_", 1)[-1])) if images[name] is not None]
images_vl = []
# Output resolution: explicit width/height, else the primary reference's own size, floored to /16.
# Each dimension falls back independently so a 0 on one axis keeps an explicit value on the other.
if width == 0 or height == 0:
if len(images) > 0:
ref_h, ref_w = images[0].shape[1], images[0].shape[2]
else:
ref_h, ref_w = 1024, 1024
height = height or ref_h
width = width or ref_w
width = max(16, (width // 16) * 16)
height = max(16, (height // 16) * 16)
for image in images:
samples = image.movedim(-1, 1)
# VL conditioning copy: cap the long edge at 384 (training preprocessing).
long_edge = max(samples.shape[3], samples.shape[2])
if long_edge > 384:
scale_by = 384 / long_edge
s = comfy.utils.common_upscale(samples, max(1, round(samples.shape[3] * scale_by)), max(1, round(samples.shape[2] * scale_by)), "bicubic", "disabled")
images_vl.append(s.movedim(1, -1))
else:
images_vl.append(image)
if vae is not None:
# All references are resized to the output resolution before encoding, because Mage's RoPE aligns reference and target content by position
if samples.shape[3] != width or samples.shape[2] != height:
s = comfy.utils.common_upscale(samples, width, height, "bicubic", "disabled")
else:
s = samples
ref_latents.append(vae.encode(s.movedim(1, -1)[:, :, :, :3]))
# Negative branch keeps the same reference images (VL tokens + ref latents), only the instruction differs.
positive = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=images_vl))
negative = clip.encode_from_tokens_scheduled(clip.tokenize(negative_prompt if negative_prompt else " ", images=images_vl))
if len(ref_latents) > 0:
positive = node_helpers.conditioning_set_values(positive, {"reference_latents": ref_latents}, append=True)
negative = node_helpers.conditioning_set_values(negative, {"reference_latents": ref_latents}, append=True)
latent = torch.zeros([batch_size, 128, height // 16, width // 16], device=comfy.model_management.intermediate_device())
return io.NodeOutput(positive, negative, {"samples": latent})
class MageExtension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
TextEncodeMageFlowEdit,
]
async def comfy_entrypoint() -> MageExtension:
return MageExtension()

View File

@ -9,6 +9,7 @@ import comfy.latent_formats
import comfy.ldm.lumina.controlnet
import comfy.ldm.supir.supir_modules
import comfy.ldm.anima.lllite
import comfy.ldm.wan.uni3c
from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel
from comfy_api.latest import io
from comfy.ldm.supir.supir_patch import SUPIRPatch
@ -264,6 +265,37 @@ class ModelPatchLoader:
if torch.count_nonzero(ref_weight) == 0:
config['broken'] = True
model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config)
elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan
attn_key_replace = {".self_attn.to_q.": ".self_attn.q.",
".self_attn.to_k.": ".self_attn.k.",
".self_attn.to_v.": ".self_attn.v.",
".self_attn.to_out.0.": ".self_attn.o."}
converted_sd = {}
for k, w in sd.items():
for r, rr in attn_key_replace.items():
k = k.replace(r, rr)
converted_sd[k] = w
sd = converted_sd
num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight"))
conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0]
if "proj_in.weight" in sd:
dim = sd["proj_in.weight"].shape[0]
else:
dim = conv_out_dim
model = comfy.ldm.wan.uni3c.WanUni3CControlnet(
in_channels=sd["controlnet_patch_embedding.weight"].shape[1],
conv_out_dim=conv_out_dim,
dim=dim,
ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0],
num_layers=num_layers,
time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1],
out_proj_dim=sd["proj_out.0.weight"].shape[0],
add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1],
mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0],
device=comfy.model_management.unet_offload_device(),
dtype=dtype,
operations=comfy.ops.manual_cast)
elif "audio_proj.proj1.weight" in sd:
model = MultiTalkModelPatch(
audio_window=5, context_tokens=32, vae_scale=4,
@ -561,6 +593,150 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet):
CATEGORY = "model/patch/z-image"
class WanUni3CCnetPatch:
def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end):
self.model_patch = model_patch
self.render_video = render_video
self.vae = vae
self.latent_format = latent_format
self.strength = strength
self.sigma_start = sigma_start
self.sigma_end = sigma_end
self.prepared_render = None
self.temp_data = None
def encode_render_video(self, target_latent_shape):
t_len, h_len, w_len = target_latent_shape
temporal_compression = self.vae.temporal_compression_decode() or 1
spatial_compression = self.vae.spacial_compression_encode()
target_frames = (t_len - 1) * temporal_compression + 1
target_height = h_len * spatial_compression
target_width = w_len * spatial_compression
frames = self.render_video
if frames.shape[0] > target_frames:
frames = frames[:target_frames]
elif frames.shape[0] < target_frames:
last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1)
frames = torch.cat([frames, last_frame], dim=0)
if frames.shape[1] != target_height or frames.shape[2] != target_width:
frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1)
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
render_latent = self.vae.encode(frames)
comfy.model_management.load_models_gpu(loaded_models)
return self.latent_format.process_in(render_latent)
def build_controlnet_input(self, x, dtype, samples_per_cond):
# first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V)
hidden = x[:samples_per_cond, :20].to(dtype)
if hidden.shape[1] < 20:
pad_shape = list(hidden.shape)
pad_shape[1] = 20 - hidden.shape[1]
hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1)
render = self.prepared_render
if render is None or render.shape[2:] != hidden.shape[2:]:
render = self.encode_render_video(hidden.shape[2:])
render = render.to(device=hidden.device, dtype=dtype)
self.prepared_render = render
if render.shape[0] != hidden.shape[0]:
render = render.expand(hidden.shape[0], -1, -1, -1, -1)
return torch.cat([hidden, render], dim=1)
def __call__(self, kwargs):
img = kwargs.get("img")
block_index = kwargs.get("block_index")
transformer_options = kwargs.get("transformer_options", {})
if block_index == 0:
self.temp_data = None
active = True
sigmas = transformer_options.get("sigmas", None)
if sigmas is not None:
sigma = sigmas[0].item()
if sigma > self.sigma_start or sigma < self.sigma_end:
active = False
if active:
x = kwargs.get("x")
# cond and uncond chunks share latents, so we can reuse residuals
num_conds = len(transformer_options.get("cond_or_uncond", [0]))
samples_per_cond = x.shape[0]
if num_conds > 0 and x.shape[0] % num_conds == 0:
samples_per_cond = x.shape[0] // num_conds
temb = kwargs.get("vec")[:samples_per_cond]
if temb.ndim == 3:
temb = temb[:, 0]
model = self.model_patch.model
controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond)
hidden, freqs = model.process_input(controlnet_input)
self.temp_data = (hidden, temb.to(img.dtype), freqs)
num_layers = self.model_patch.model.num_layers
if self.temp_data is not None and block_index < num_layers:
hidden, temb, freqs = self.temp_data
hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs)
residual = residual.to(img.dtype) * self.strength
if residual.shape[0] != img.shape[0]:
residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1)
img_offset = kwargs.get("img_offset", 0)
img[:, img_offset:img_offset + residual.shape[1]] += residual
if block_index >= num_layers - 1:
self.temp_data = None
else:
self.temp_data = (hidden, temb, freqs)
return kwargs
def to(self, device_or_dtype):
if isinstance(device_or_dtype, torch.device):
if self.prepared_render is not None:
self.prepared_render = self.prepared_render.to(device_or_dtype)
self.temp_data = None
return self
def models(self):
return [self.model_patch]
class WanUni3CControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model": ("MODEL",),
"model_patch": ("MODEL_PATCH",),
"vae": ("VAE",),
"render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply_patch"
EXPERIMENTAL = True
CATEGORY = "model/patch/wan"
def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent):
if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet):
raise ValueError("The connected model patch is not a Uni3C ControlNet.")
cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features
model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None)
if model_dim is None:
raise ValueError("The Uni3C ControlNet only works with Wan models.")
if model_dim != cnet_dim:
raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim))
model_patched = model.clone()
model_sampling = model.get_model_object("model_sampling")
sigma_start = model_sampling.percent_to_sigma(start_percent)
sigma_end = model_sampling.percent_to_sigma(end_percent)
latent_format = model.get_model_object("latent_format")
patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end)
model_patched.set_model_double_block_patch(patch)
return (model_patched,)
class UsoStyleProjectorPatch:
def __init__(self, model_patch, encoded_image):
self.model_patch = model_patch
@ -719,6 +895,7 @@ NODE_CLASS_MAPPINGS = {
"ModelPatchLoader": ModelPatchLoader,
"QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet,
"ZImageFunControlnet": ZImageFunControlnet,
"WanUni3CControlnetApply": WanUni3CControlnetApply,
"USOStyleReference": USOStyleReference,
"SUPIRApply": SUPIRApply,
"AnimaLLLiteApply": AnimaLLLiteApply,
@ -728,6 +905,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ModelPatchLoader": "Load Model Patch",
"QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet",
"ZImageFunControlnet": "Apply Z-Image Fun ControlNet",
"WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet",
"USOStyleReference": "Apply USO Style Reference",
"SUPIRApply": "Apply SUPIR Patch",
"AnimaLLLiteApply": "Apply Anima LLLite",

View File

@ -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,

View File

@ -7,6 +7,7 @@ import folder_paths
from typing_extensions import override
from comfy_api.latest import ComfyExtension, io
import comfy.model_management
import comfy.model_patcher
try:
from spandrel_extra_arches import EXTRA_REGISTRY
@ -42,6 +43,7 @@ class UpscaleModelLoader(io.ComfyNode):
if not isinstance(out, ImageModelDescriptor):
raise Exception("Upscale model must be a single-image model.")
out.patcher = comfy.model_patcher.CoreModelPatcher(out.model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device())
return io.NodeOutput(out)
load_model = execute # TODO: remove
@ -66,14 +68,12 @@ class ImageUpscaleWithModel(io.ComfyNode):
@classmethod
def execute(cls, upscale_model, image) -> io.NodeOutput:
device = model_management.get_torch_device()
device = upscale_model.patcher.load_device
memory_required = model_management.module_size(upscale_model.model)
memory_required += (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate
memory_required = (512 * 512 * 3) * image.element_size() * max(upscale_model.scale, 1.0) * 384.0 #The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate
memory_required += image.nelement() * image.element_size()
model_management.free_memory(memory_required, device)
model_management.load_models_gpu([upscale_model.patcher], memory_required=memory_required)
upscale_model.to(device)
in_img = image.movedim(-1,-3).to(device)
tile = 512
@ -82,20 +82,17 @@ class ImageUpscaleWithModel(io.ComfyNode):
output_device = comfy.model_management.intermediate_device()
oom = True
try:
while oom:
try:
steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap)
pbar = comfy.utils.ProgressBar(steps)
s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device)
oom = False
except Exception as e:
model_management.raise_non_oom(e)
tile //= 2
if tile < 128:
raise e
finally:
upscale_model.to("cpu")
while oom:
try:
steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap)
pbar = comfy.utils.ProgressBar(steps)
s = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a.float()), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar, output_device=output_device)
oom = False
except Exception as e:
model_management.raise_non_oom(e)
tile //= 2
if tile < 128:
raise e
s = torch.clamp(s.movedim(-3,-1), min=0, max=1.0).to(comfy.model_management.intermediate_dtype())
return io.NodeOutput(s)

View File

@ -1,3 +1,3 @@
# This file is automatically generated by the build process when version is
# updated in pyproject.toml.
__version__ = "0.28.0"
__version__ = "0.29.0"

View File

@ -13,11 +13,13 @@ import asyncio
import torch
from comfy.cli_args import args
from comfy.cli_args import args, get_console_log_level
import comfy.memory_management
import comfy.model_management
import comfy.model_patcher
import comfy.model_prefetch
import comfy_aimdo.model_vbar
from comfy.logging import detail
from latent_preview import set_preview_method
import nodes
@ -543,7 +545,7 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
output_data, output_ui, has_subgraph, has_pending_tasks = await get_output_data(prompt_id, unique_id, obj, input_data_all, execution_block_cb=execution_block_cb, pre_execute_cb=pre_execute_cb, v3_data=v3_data)
finally:
if comfy.memory_management.aimdo_enabled:
if args.verbose == "DEBUG":
if get_console_log_level(args.verbose) == "DEBUG":
comfy_aimdo.control.analyze()
comfy.model_management.reset_cast_buffers()
comfy.model_prefetch.cleanup_prefetch_queues()
@ -664,6 +666,7 @@ class PromptExecutor:
self.cache_args = cache_args
self.cache_type = cache_type
self.server = server
self.prompt_model_tracker = comfy.model_patcher.PromptModelTracker()
self.reset()
def reset(self):
@ -728,6 +731,7 @@ class PromptExecutor:
set_preview_method(extra_data.get("preview_method"))
nodes.interrupt_processing(False)
self.prompt_model_tracker.start()
if "client_id" in extra_data:
self.server.client_id = extra_data["client_id"]
@ -770,7 +774,7 @@ class PromptExecutor:
pending_async_nodes = {} # TODO - Unify this with pending_subgraph_results
ui_node_outputs = {}
executed = set()
execution_list = ExecutionList(dynamic_prompt, self.caches.outputs)
execution_list = ExecutionList(dynamic_prompt, self.caches.outputs, self.prompt_model_tracker.add)
current_outputs = self.caches.outputs.all_node_ids()
for node_id in list(execute_outputs):
execution_list.add_node(node_id)
@ -832,7 +836,10 @@ class PromptExecutor:
if comfy.model_management.DISABLE_SMART_MEMORY:
comfy.model_management.unload_all_models()
finally:
if self.cache_type == CacheType.RAM_PRESSURE:
detail("RAM cache evictions: prompt=%s active=%s full=%s", prompt_id, self.caches.outputs.active_evictions, self.caches.outputs.full_evictions)
comfy.memory_management.set_ram_cache_release_state(None, 0)
self.prompt_model_tracker.end()
self._notify_prompt_lifecycle("end", prompt_id)

View File

@ -29,6 +29,7 @@
# upscale_models: models/upscale_models/
# latent_upscale_models: models/latent_upscale_models/
# custom_nodes: custom_nodes/
# datasets: datasets/
# hypernetworks: models/hypernetworks/
# photomaker: models/photomaker/
# classifiers: models/classifiers/

View File

@ -44,6 +44,8 @@ folder_names_and_paths["latent_upscale_models"] = ([os.path.join(models_dir, "la
folder_names_and_paths["custom_nodes"] = ([os.path.join(base_path, "custom_nodes")], set())
folder_names_and_paths["datasets"] = ([os.path.join(base_path, "datasets")], set())
folder_names_and_paths["hypernetworks"] = ([os.path.join(models_dir, "hypernetworks")], supported_pt_extensions)
folder_names_and_paths["photomaker"] = ([os.path.join(models_dir, "photomaker")], supported_pt_extensions)
@ -304,6 +306,23 @@ def is_dangerous_content_type(content_type: str | None) -> bool:
return normalized.endswith('+xml') or normalized.endswith('/xml')
def renders_safely_as_image(content_type: str | None, sec_fetch_dest: str | None) -> bool:
"""Return True if a dangerous `content_type` is safe to serve inline anyway.
An SVG referenced by an ``<img>`` is loaded in secure static mode: scripts
and external references are disabled, so the stored XSS that
``is_dangerous_content_type`` guards against cannot fire. The attack needs
the SVG to become a document, which is a separate ``Sec-Fetch-Dest``.
Browsers set that header themselves and script cannot override it (the
``Sec-`` prefix makes it a forbidden header name), so it is trustworthy for
this decision. Anything else, including a missing header from a non-browser
client or a proxy that strips it, fails closed.
"""
if sec_fetch_dest != 'image':
return False
return (content_type or '').split(';', 1)[0].strip().lower() == 'image/svg+xml'
def is_within_directory(directory: str, target: str) -> bool:
"""Return True if `target` resolves to a path inside `directory`.

20
main.py
View File

@ -2,6 +2,7 @@ import comfy.options
comfy.options.enable_args_parsing()
from comfy.cli_args import args
from comfy.cli_args import get_console_log_level, get_file_log_outputs
if args.list_feature_flags:
import json
@ -17,7 +18,9 @@ import folder_paths
import time
from comfy.cli_args import enables_dynamic_vram
from app.logger import setup_logger
setup_logger(log_level=args.verbose, use_stdout=args.log_stdout)
console_log_level = get_console_log_level(args.verbose)
file_log_outputs = [('DETAIL', 'comfyui_detail.log'), *get_file_log_outputs(args.verbose)]
setup_logger(log_level=console_log_level, file_outputs=file_log_outputs, use_stdout=args.log_stdout)
from app.assets.seeder import asset_seeder
from app.assets.services import register_output_files
@ -251,13 +254,18 @@ if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_managemen
aimdo_initialized = comfy_aimdo.control.init_devices(d.index for d in comfy.model_management.get_all_torch_devices())
if aimdo_initialized:
if args.verbose == 'DEBUG':
if console_log_level == 'DEBUG':
comfy_aimdo.control.set_log_debug()
elif args.verbose == 'CRITICAL':
elif console_log_level == 'DETAIL':
try:
comfy_aimdo.control.set_log_detail()
except AttributeError:
comfy_aimdo.control.set_log_info()
elif console_log_level == 'CRITICAL':
comfy_aimdo.control.set_log_critical()
elif args.verbose == 'ERROR':
elif console_log_level == 'ERROR':
comfy_aimdo.control.set_log_error()
elif args.verbose == 'WARNING':
elif console_log_level == 'WARNING':
comfy_aimdo.control.set_log_warning()
else: #INFO
comfy_aimdo.control.set_log_info()
@ -319,7 +327,7 @@ def prompt_worker(q, server_instance):
cache_ram_inactive = 0
if not args.cache_classic and not args.cache_none and args.cache_lru <= 0:
cache_ram = min(10.0, max(2.0, comfy.model_management.total_ram * 0.10 / 1024.0))
cache_ram_inactive = min(96.0, comfy.model_management.total_ram / 1024.0)
cache_ram_inactive = min(128.0, comfy.model_management.total_ram / 1024.0)
if len(args.cache_ram) > 0:
cache_ram = args.cache_ram[0]
if len(args.cache_ram) > 1:

View File

@ -992,7 +992,7 @@ class CLIPLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ),
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ),
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage", "mage"], ),
},
"optional": {
"device": (["default", "cpu"], {"advanced": True}),
@ -2462,6 +2462,7 @@ async def init_builtin_extra_nodes():
"nodes_seedvr.py",
"nodes_context_windows.py",
"nodes_qwen.py",
"nodes_mage.py",
"nodes_joyimage.py",
"nodes_boogu.py",
"nodes_chroma_radiance.py",

View File

@ -1,6 +1,6 @@
[project]
name = "ComfyUI"
version = "0.28.0"
version = "0.29.0"
readme = "README.md"
license = { file = "LICENSE" }
requires-python = ">=3.10"

View File

@ -1,6 +1,6 @@
comfyui-frontend-package==1.45.21
comfyui-workflow-templates==0.11.12
comfyui-embedded-docs==0.5.8
comfyui-frontend-package==1.47.10
comfyui-workflow-templates==0.11.19
comfyui-embedded-docs==0.5.9
torch
torchsde
torchvision
@ -22,7 +22,7 @@ alembic
SQLAlchemy>=2.0.0
filelock
av>=16.0.0
comfy-kitchen==0.2.22
comfy-kitchen==0.2.24
comfy-aimdo==0.4.10
requests
simpleeval>=1.0.0

View File

@ -624,8 +624,9 @@ class PromptServer():
# For security, force renderable/active types (HTML, JS,
# CSS, SVG, XML — anything that can carry inline <script>
# and execute in the page origin) to download instead of
# displaying inline, preventing stored XSS. The
# attachment disposition is the load-bearing guard: a
# displaying inline, preventing stored XSS. SVG loaded
# into an <img> is exempt, see renders_safely_as_image.
# The attachment disposition is the load-bearing guard: a
# bare filename= hint does not force a download per
# RFC 6266, so we only attach it on the dangerous branch
# to avoid breaking inline display of legitimate images.
@ -635,18 +636,27 @@ class PromptServer():
# header's quoted-string and malform the disposition.
safe_filename = filename.replace("\\", "\\\\").replace('"', '\\"')
disposition = f"filename=\"{safe_filename}\""
headers = {"X-Content-Type-Options": "nosniff"}
sec_fetch_dest = request.headers.get('Sec-Fetch-Dest')
if folder_paths.is_dangerous_content_type(content_type):
content_type = 'application/octet-stream'
disposition = f"attachment; filename=\"{safe_filename}\""
# This response now depends on a request header, so
# it must not be reused across destinations.
# FileResponse emits Last-Modified/ETag and nothing
# sets Cache-Control on /view, which makes it
# heuristically cacheable: without these headers a
# cache could replay the inline SVG served to an
# <img> to a later document navigation of the same
# URL and re-enable the stored XSS, or replay the
# attachment to an <img> and re-break the preview.
headers["Vary"] = "Sec-Fetch-Dest"
headers["Cache-Control"] = "no-store"
if not folder_paths.renders_safely_as_image(content_type, sec_fetch_dest):
content_type = 'application/octet-stream'
disposition = f"attachment; filename=\"{safe_filename}\""
return web.FileResponse(
file,
headers={
"Content-Disposition": disposition,
"Content-Type": content_type,
"X-Content-Type-Options": "nosniff"
}
)
headers["Content-Disposition"] = disposition
headers["Content-Type"] = content_type
return web.FileResponse(file, headers=headers)
return web.Response(status=404)

View File

@ -0,0 +1,191 @@
"""CI unit guard for the Sec-Fetch-Dest exemption to FIX #5 of GHSA-779p-m5rp-r4h4.
FIX #5 forced every SVG served by /view and the assets download route to
application/octet-stream + Content-Disposition: attachment. That blocked the
stored XSS, but it also stopped SVG node outputs and Media Assets thumbnails
from rendering, because the frontend requests them with a plain <img>.
An SVG referenced by an <img> is loaded in secure static mode: scripting and
external references are disabled, so the payload from vuln #5 cannot execute.
The attack needs the SVG to load as a document, which arrives with a different
Sec-Fetch-Dest. Browsers set that header themselves and page script cannot
override it, so renders_safely_as_image() uses it to re-allow only the <img>
case. Everything else, including a missing header, keeps the forced download.
Because the decision reads a request header, the response must also carry
Vary: Sec-Fetch-Dest and Cache-Control: no-store. Without them a cache keyed on
the URL alone can replay the inline SVG served to an <img> to a later document
navigation of the same URL, which re-enables the very XSS the exemption is
built around. The route tests at the bottom pin those headers.
server.py cannot be imported in a unit test (importing it spins up the full
PromptServer/aiohttp app and its global side effects), so the /view side is
covered by pinning the helper its closure calls.
"""
import pytest
from aiohttp import web
import folder_paths
from app.assets.api import routes as asset_routes
from app.assets.services.schemas import DownloadResolutionResult
# Every Sec-Fetch-Dest that must keep the forced download. 'document' is the
# vuln #5 attack itself; the rest either cannot execute an SVG or have no
# reason to receive one inline. None covers curl and proxies that strip the
# header.
UNSAFE_DESTS = [
'document',
'iframe',
'object',
'embed',
'frame',
'script',
'style',
'empty',
'',
None,
]
def test_svg_renders_inline_for_image_dest():
assert folder_paths.renders_safely_as_image('image/svg+xml', 'image')
def test_svg_forced_to_download_for_every_other_dest():
for dest in UNSAFE_DESTS:
assert not folder_paths.renders_safely_as_image('image/svg+xml', dest), (
f"SVG must not be served inline for Sec-Fetch-Dest={dest!r}. Only "
"an <img> load is safe, everything else can reach document context"
)
def test_missing_header_fails_closed():
assert not folder_paths.renders_safely_as_image('image/svg+xml', None)
def test_exemption_normalises_parameters_and_casing():
for content_type in ('IMAGE/SVG+XML', 'image/svg+xml; charset=utf-8', ' image/svg+xml '):
assert folder_paths.renders_safely_as_image(content_type, 'image'), (
f"{content_type!r} is an SVG and must not be excluded from the "
"exemption by casing or a charset parameter"
)
def test_other_dangerous_types_are_never_exempt():
# Only SVG is safe as an <img>. Nothing else in the blocklist may ride the
# image dest back into an inline response.
for content_type in ('text/html', 'text/javascript', 'text/css',
'application/xml', 'text/xml', 'application/xhtml+xml',
'image/svg+xml.html', 'application/rss+xml'):
assert not folder_paths.renders_safely_as_image(content_type, 'image'), (
f"{content_type!r} must stay forced-download regardless of Sec-Fetch-Dest"
)
def test_exemption_does_not_widen_the_blocklist():
# The exemption is a call-site gate, not a change to what counts as
# dangerous. is_dangerous_content_type must still flag SVG on its own.
assert folder_paths.is_dangerous_content_type('image/svg+xml')
# --- Response-level guards on the assets content route -----------------------
#
# The helper above decides correctly, but the exemption is only safe if the
# response cannot be reused across Sec-Fetch-Dest values. These mount the real
# route and assert on the headers that actually ship.
SVG_PAYLOAD = b'<svg xmlns="http://www.w3.org/2000/svg"><script>alert(1)</script></svg>'
ASSET_ID = "00000000-0000-4000-8000-000000000001"
CONTENT_URL = f"/api/assets/{ASSET_ID}/content"
class _StubUserManager:
def get_request_user_id(self, request):
return "test-user"
@pytest.fixture
def asset_app(monkeypatch, tmp_path):
"""Mount the real /api/assets/{id}/content route over a stored SVG."""
def _factory(stored_mime_type):
svg = tmp_path / "thumb.svg"
svg.write_bytes(SVG_PAYLOAD)
monkeypatch.setattr(asset_routes, "_ASSETS_ENABLED", True)
monkeypatch.setattr(asset_routes, "USER_MANAGER", _StubUserManager())
monkeypatch.setattr(
asset_routes,
"resolve_asset_for_download",
lambda reference_id, owner_id: DownloadResolutionResult(
abs_path=str(svg),
content_type=stored_mime_type,
download_name="thumb.svg",
),
)
app = web.Application()
app.add_routes(asset_routes.ROUTES)
return app
return _factory
@pytest.mark.asyncio
async def test_inline_svg_response_is_not_cacheable_across_destinations(
aiohttp_client, asset_app
):
client = await aiohttp_client(asset_app("image/svg+xml"))
resp = await client.get(
CONTENT_URL, params={"disposition": "inline"}, headers={"Sec-Fetch-Dest": "image"}
)
assert resp.status == 200
assert "image/svg+xml" in resp.headers.get("Content-Type", "").lower()
# The load-bearing assertion: a cache must not be able to hand this inline
# SVG to a later document navigation of the same URL.
assert "sec-fetch-dest" in resp.headers.get("Vary", "").lower(), (
"The response varies on Sec-Fetch-Dest but does not say so, so a cache "
"keyed on the URL alone can replay the inline SVG into document context."
)
assert "no-store" in resp.headers.get("Cache-Control", "").lower()
@pytest.mark.asyncio
async def test_forced_download_response_also_declares_the_variance(
aiohttp_client, asset_app
):
# The attachment branch needs the same headers, in both directions: a
# cached octet-stream replayed to an <img> re-breaks the preview this fix
# exists to restore.
client = await aiohttp_client(asset_app("image/svg+xml"))
resp = await client.get(
CONTENT_URL,
params={"disposition": "inline"},
headers={"Sec-Fetch-Dest": "document"},
)
assert resp.status == 200
assert "application/octet-stream" in resp.headers.get("Content-Type", "").lower()
assert "attachment" in resp.headers.get("Content-Disposition", "").lower()
assert "sec-fetch-dest" in resp.headers.get("Vary", "").lower()
assert "no-store" in resp.headers.get("Cache-Control", "").lower()
@pytest.mark.asyncio
async def test_parameterised_svg_mime_type_does_not_500(aiohttp_client, asset_app):
# mime_type is uploader-supplied and unvalidated. aiohttp rejects a charset
# in the content_type argument with ValueError, so the exempt branch must
# strip parameters before building the response.
client = await aiohttp_client(asset_app("image/svg+xml; charset=utf-8"))
resp = await client.get(
CONTENT_URL, params={"disposition": "inline"}, headers={"Sec-Fetch-Dest": "image"}
)
assert resp.status == 200, (
"A charset parameter on the stored mime type must not turn a valid "
"inline SVG request into a 500."
)
assert "image/svg+xml" in resp.headers.get("Content-Type", "").lower()

View File

@ -280,6 +280,86 @@ class TestGetOutputsSummary:
assert preview['filename'] == 'model.glb'
assert preview['mediaType'] == '3d'
def test_media_preview_preferred_over_text(self):
"""A visual output wins the preview even when a text node is iterated
first (regression: text could mask a later temp/preview image)."""
outputs = {
'text_node': {'text': ['a caption']},
'image_node': {'images': [{'filename': 'preview.png', 'type': 'temp'}]},
}
count, preview = get_outputs_summary(outputs)
# Text is preview-only metadata and not counted; only the image counts.
assert count == 1
assert preview['filename'] == 'preview.png'
assert preview['mediaType'] == 'images'
def test_text_used_as_preview_when_no_media(self):
"""Text is the preview only when the job produced no media output."""
outputs = {
'text_node': {'text': ['hello world']},
}
count, preview = get_outputs_summary(outputs)
assert count == 0 # text entries are not counted as outputs
assert preview['mediaType'] == 'text'
assert preview['content'] == 'hello world'
def test_media_preview_preferred_over_saved_text_file(self):
"""A visual output wins the preview over a saved text file (SaveText),
even a temp/preview image iterated after the text node."""
outputs = {
'save_text': {
'text': ['the text'],
'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}],
},
'preview_image': {'images': [{'filename': 'preview.png', 'type': 'temp'}]},
}
count, preview = get_outputs_summary(outputs)
assert count == 2 # the .txt file and the image; raw text is metadata
assert preview['filename'] == 'preview.png'
assert preview['mediaType'] == 'images'
def test_saved_media_preferred_over_saved_text_file(self):
outputs = {
'save_text': {
'text': ['the text'],
'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}],
},
'save_image': {'images': [{'filename': 'result.png', 'type': 'output'}]},
}
count, preview = get_outputs_summary(outputs)
assert count == 2
assert preview['filename'] == 'result.png'
def test_mime_format_file_preferred_over_saved_text_file(self):
"""Custom-node outputs previewable via MIME format (e.g. VHS videos
under arbitrary keys) rank as visual media, above saved text files."""
outputs = {
'save_text': {
'files': [{'filename': 'notes.md', 'subfolder': '', 'type': 'output'}],
},
'video_node': {
'files': [{'filename': 'clip.webm', 'format': 'video/webm', 'type': 'output'}],
},
}
count, preview = get_outputs_summary(outputs)
assert count == 2
assert preview['filename'] == 'clip.webm'
def test_saved_text_file_preferred_over_raw_text(self):
"""With no media in the job, the saved text file (a real, counted
output) is the preview rather than the raw text metadata."""
outputs = {
'save_text': {
'text': ['the text'],
'files': [{'filename': 'ComfyUI_00001.txt', 'subfolder': '', 'type': 'output'}],
},
}
count, preview = get_outputs_summary(outputs)
assert count == 1
assert preview['filename'] == 'ComfyUI_00001.txt'
assert preview['mediaType'] == 'files'
class TestHas3DExtension:
"""Unit tests for has_3d_extension()"""