Merge remote-tracking branch 'upstream/master' into pixal3d
This commit is contained in:
commit
87e34c5c17
|
|
@ -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/**
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
22
AGENTS.md
22
AGENTS.md
|
|
@ -162,8 +162,26 @@
|
|||
adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`,
|
||||
`comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and
|
||||
`comfy-kitchen` helpers where they already solve the problem.
|
||||
- Use optimized comfy-kitchen ops in places where they improve performance
|
||||
without changing the expected dtype, device, memory, or interface behavior.
|
||||
- Model implementations must use an existing optimized Comfy Kitchen or
|
||||
ComfyUI operation whenever one supports the required math and tensor layout
|
||||
without changing expected dtype, device, memory, or interface behavior. This
|
||||
is the default implementation requirement, not an optional follow-up
|
||||
optimization.
|
||||
- Before implementing model math, inspect the operations already exposed by
|
||||
Comfy Kitchen, `comfy.quant_ops`, and existing ComfyUI model helpers. Check
|
||||
for optimized single, paired, fused, layout-specific, and quantized variants
|
||||
before writing a local implementation or composing lower-level torch ops.
|
||||
- Use the compatible optimized operation first and adapt the model's inputs to
|
||||
its documented layout while preserving the model's exact math. If several
|
||||
optimized variants apply, benchmark representative model shapes and select
|
||||
the fastest valid path.
|
||||
- Add or retain a local implementation only when no existing optimized
|
||||
operation supports the required math, layout, dtype, device, autograd, or
|
||||
patch contract. Keep differentiable or patch-compatible fallbacks when the
|
||||
optimized inference operation does not provide those contracts.
|
||||
- Use the existing ComfyUI cast, offload, and cleanup helpers for parameters
|
||||
passed to optimized operations. Preserve model-specific epsilon, scaling,
|
||||
layout, dtype, device, and output-shape behavior.
|
||||
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
|
||||
handwritten implementations of the same operation. Remove duplicate local
|
||||
kernels and adapt inputs to the shared operation's documented layout while
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
|
||||
|
||||
/CODEOWNERS @comfyanonymous
|
||||
/AGENTS.md @comfyanonymous
|
||||
/.ci/ @comfyanonymous
|
||||
/.github/ @comfyanonymous
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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).")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
54
comfy/ops.py
54
comfy/ops.py
|
|
@ -41,7 +41,7 @@ def scaled_dot_product_attention(q, k, v, *args, **kwargs):
|
|||
|
||||
|
||||
try:
|
||||
if torch.cuda.is_available() and comfy.model_management.WINDOWS:
|
||||
if torch.cuda.is_available():
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
import inspect
|
||||
if "set_priority" in inspect.signature(sdpa_kernel).parameters:
|
||||
|
|
@ -51,7 +51,10 @@ try:
|
|||
SDPBackend.MATH,
|
||||
]
|
||||
|
||||
SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
|
||||
if comfy.model_management.WINDOWS:
|
||||
SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
|
||||
else:
|
||||
SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION)
|
||||
|
||||
def scaled_dot_product_attention(q, k, v, *args, **kwargs):
|
||||
if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
27
comfy/sd.py
27
comfy/sd.py
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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_
|
||||
|
|
@ -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 = ' '
|
||||
|
||||
|
|
|
|||
|
|
@ -244,10 +244,10 @@ RECRAFT_V4_PRO_SIZES = [
|
|||
"2304x1792",
|
||||
"1792x2304",
|
||||
"1664x2688",
|
||||
"1434x1024",
|
||||
"1024x1434",
|
||||
"2560x1792",
|
||||
"1792x2560",
|
||||
"2688x1536",
|
||||
"1536x2688",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
13
execution.py
13
execution.py
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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/
|
||||
|
|
|
|||
|
|
@ -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
20
main.py
|
|
@ -2,6 +2,7 @@ import comfy.options
|
|||
comfy.options.enable_args_parsing()
|
||||
|
||||
from comfy.cli_args import args
|
||||
from comfy.cli_args import get_console_log_level, get_file_log_outputs
|
||||
|
||||
if args.list_feature_flags:
|
||||
import json
|
||||
|
|
@ -17,7 +18,9 @@ import folder_paths
|
|||
import time
|
||||
from comfy.cli_args import enables_dynamic_vram
|
||||
from app.logger import setup_logger
|
||||
setup_logger(log_level=args.verbose, use_stdout=args.log_stdout)
|
||||
console_log_level = get_console_log_level(args.verbose)
|
||||
file_log_outputs = [('DETAIL', 'comfyui_detail.log'), *get_file_log_outputs(args.verbose)]
|
||||
setup_logger(log_level=console_log_level, file_outputs=file_log_outputs, use_stdout=args.log_stdout)
|
||||
|
||||
from app.assets.seeder import asset_seeder
|
||||
from app.assets.services import register_output_files
|
||||
|
|
@ -251,13 +254,18 @@ if args.enable_dynamic_vram or (enables_dynamic_vram() and comfy.model_managemen
|
|||
aimdo_initialized = comfy_aimdo.control.init_devices(d.index for d in comfy.model_management.get_all_torch_devices())
|
||||
|
||||
if aimdo_initialized:
|
||||
if args.verbose == 'DEBUG':
|
||||
if console_log_level == 'DEBUG':
|
||||
comfy_aimdo.control.set_log_debug()
|
||||
elif args.verbose == 'CRITICAL':
|
||||
elif console_log_level == 'DETAIL':
|
||||
try:
|
||||
comfy_aimdo.control.set_log_detail()
|
||||
except AttributeError:
|
||||
comfy_aimdo.control.set_log_info()
|
||||
elif console_log_level == 'CRITICAL':
|
||||
comfy_aimdo.control.set_log_critical()
|
||||
elif args.verbose == 'ERROR':
|
||||
elif console_log_level == 'ERROR':
|
||||
comfy_aimdo.control.set_log_error()
|
||||
elif args.verbose == 'WARNING':
|
||||
elif console_log_level == 'WARNING':
|
||||
comfy_aimdo.control.set_log_warning()
|
||||
else: #INFO
|
||||
comfy_aimdo.control.set_log_info()
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
3
nodes.py
3
nodes.py
|
|
@ -992,7 +992,7 @@ class CLIPLoader:
|
|||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage", "mage"], ),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["default", "cpu"], {"advanced": True}),
|
||||
|
|
@ -2462,6 +2462,7 @@ async def init_builtin_extra_nodes():
|
|||
"nodes_seedvr.py",
|
||||
"nodes_context_windows.py",
|
||||
"nodes_qwen.py",
|
||||
"nodes_mage.py",
|
||||
"nodes_joyimage.py",
|
||||
"nodes_boogu.py",
|
||||
"nodes_chroma_radiance.py",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
34
server.py
34
server.py
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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()"""
|
||||
|
|
|
|||
Loading…
Reference in New Issue