diff --git a/.github/workflows/ci-cursor-review.yml b/.github/workflows/ci-cursor-review.yml
index 2312c0ccd..a7a0692c9 100644
--- a/.github/workflows/ci-cursor-review.yml
+++ b/.github/workflows/ci-cursor-review.yml
@@ -23,9 +23,9 @@ jobs:
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
# upstream changes; keep `workflows_ref` matching so prompts/scripts load
# from the same commit as the workflow definition.
- uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9
+ uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
with:
- workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a
+ workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
diff_excludes: >-
:!**/.claude/**
:!**/dist/**
diff --git a/.github/workflows/release-stable-all.yml b/.github/workflows/release-stable-all.yml
index d7cf69fe2..10f1ccf96 100644
--- a/.github/workflows/release-stable-all.yml
+++ b/.github/workflows/release-stable-all.yml
@@ -20,7 +20,7 @@ jobs:
git_tag: ${{ inputs.git_tag }}
cache_tag: "cu130"
python_minor: "13"
- python_patch: "12"
+ python_patch: "14"
rel_name: "nvidia"
rel_extra_name: ""
test_release: true
@@ -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
diff --git a/AGENTS.md b/AGENTS.md
index 20014ce7e..bfe0976fd 100644
--- a/AGENTS.md
+++ b/AGENTS.md
@@ -162,8 +162,26 @@
adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`,
`comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and
`comfy-kitchen` helpers where they already solve the problem.
-- Use optimized comfy-kitchen ops in places where they improve performance
- without changing the expected dtype, device, memory, or interface behavior.
+- Model implementations must use an existing optimized Comfy Kitchen or
+ ComfyUI operation whenever one supports the required math and tensor layout
+ without changing expected dtype, device, memory, or interface behavior. This
+ is the default implementation requirement, not an optional follow-up
+ optimization.
+- Before implementing model math, inspect the operations already exposed by
+ Comfy Kitchen, `comfy.quant_ops`, and existing ComfyUI model helpers. Check
+ for optimized single, paired, fused, layout-specific, and quantized variants
+ before writing a local implementation or composing lower-level torch ops.
+- Use the compatible optimized operation first and adapt the model's inputs to
+ its documented layout while preserving the model's exact math. If several
+ optimized variants apply, benchmark representative model shapes and select
+ the fastest valid path.
+- Add or retain a local implementation only when no existing optimized
+ operation supports the required math, layout, dtype, device, autograd, or
+ patch contract. Keep differentiable or patch-compatible fallbacks when the
+ optimized inference operation does not provide those contracts.
+- Use the existing ComfyUI cast, offload, and cleanup helpers for parameters
+ passed to optimized operations. Preserve model-specific epsilon, scaling,
+ layout, dtype, device, and output-shape behavior.
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
handwritten implementations of the same operation. Remove duplicate local
kernels and adapt inputs to the shared operation's documented layout while
diff --git a/CODEOWNERS b/CODEOWNERS
index 043c0ec75..634927dd6 100644
--- a/CODEOWNERS
+++ b/CODEOWNERS
@@ -1,5 +1,6 @@
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
/CODEOWNERS @comfyanonymous
+/AGENTS.md @comfyanonymous
/.ci/ @comfyanonymous
/.github/ @comfyanonymous
diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py
index 43e60094c..e25b8a57f 100644
--- a/app/assets/api/routes.py
+++ b/app/assets/api/routes.py
@@ -315,15 +315,29 @@ async def download_asset_content(request: web.Request) -> web.Response:
404, "FILE_NOT_FOUND", "Underlying file not found on disk."
)
- # User-controlled asset content must never render inline in the app origin
+ # User-controlled asset content must not render inline in the app origin
# (stored XSS via SVG/HTML/XML). Force dangerous types to download and
- # override any requested inline disposition. Centralised through
- # folder_paths.is_dangerous_content_type so this can't drift from /view and
- # /userdata (the previous inline set here omitted image/svg+xml and missed
- # the charset/casing/+xml-dialect bypasses).
+ # override any requested inline disposition; SVG loaded into an
is
+ # exempt, see renders_safely_as_image. Centralised through folder_paths so
+ # this can't drift from /view and /userdata (the previous inline set here
+ # omitted image/svg+xml and missed the charset/casing/+xml-dialect bypasses).
+ extra_headers = {}
+ sec_fetch_dest = request.headers.get("Sec-Fetch-Dest")
if folder_paths.is_dangerous_content_type(content_type):
- content_type = "application/octet-stream"
- disposition = "attachment"
+ # This response now depends on a request header, so it must not be
+ # reused across destinations by a browser or intermediary cache: an
+ # inline SVG primed by an
fetch and replayed to a document
+ # navigation of the same URL would re-enable the stored XSS.
+ extra_headers["Vary"] = "Sec-Fetch-Dest"
+ extra_headers["Cache-Control"] = "no-store"
+ if not folder_paths.renders_safely_as_image(content_type, sec_fetch_dest):
+ content_type = "application/octet-stream"
+ disposition = "attachment"
+
+ # mime_type is uploader-supplied and unvalidated, so it can carry
+ # parameters. aiohttp rejects a charset in the content_type argument with
+ # ValueError, which would turn a valid inline SVG into a 500.
+ content_type = content_type.split(";", 1)[0].strip() or "application/octet-stream"
safe_name = (filename or "").replace("\r", "").replace("\n", "")
encoded = urllib.parse.quote(safe_name)
@@ -356,6 +370,7 @@ async def download_asset_content(request: web.Request) -> web.Response:
"Content-Disposition": cd,
"Content-Length": str(file_size),
"X-Content-Type-Options": "nosniff",
+ **extra_headers,
},
)
diff --git a/app/logger.py b/app/logger.py
index bde815822..1aed54e37 100644
--- a/app/logger.py
+++ b/app/logger.py
@@ -2,9 +2,12 @@ from collections import deque
from datetime import datetime
import io
import logging
+import os
import sys
import threading
+import comfy.logging
+
ANSI_NAMED_COLORS = {
'black': '\033[30m',
'red': '\033[31m',
@@ -18,6 +21,7 @@ ANSI_NAMED_COLORS = {
ANSI_LEVEL_COLORS = {
'DEBUG': ANSI_NAMED_COLORS['cyan'],
+ 'DETAIL': ANSI_NAMED_COLORS['blue'],
'INFO': ANSI_NAMED_COLORS['green'],
'WARNING': ANSI_NAMED_COLORS['yellow'],
'ERROR': ANSI_NAMED_COLORS['red'],
@@ -85,7 +89,12 @@ def on_flush(callback):
if stderr_interceptor is not None:
stderr_interceptor.on_flush(callback)
-def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool = False):
+
+def get_log_level(level):
+ return comfy.logging.DETAIL if level == "DETAIL" else logging.getLevelName(level)
+
+
+def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300, use_stdout: bool = False):
global logs
if logs:
return
@@ -99,13 +108,18 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool
stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr)
# Setup default global logger
+ if file_outputs is None:
+ file_outputs = [('DETAIL', 'comfyui_detail.log')]
logger = logging.getLogger()
- logger.setLevel(log_level)
+ console_level = get_log_level(log_level)
+ file_levels = [get_log_level(level) for level, _ in file_outputs]
+ logger.setLevel(min(console_level, *file_levels))
formatter = ColoredFormatter("%(message)s")
stream_handler = logging.StreamHandler()
stream_handler.setFormatter(formatter)
+ stream_handler.setLevel(console_level)
if use_stdout:
# Only errors and critical to stderr
@@ -114,11 +128,24 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool
# Lesser to stdout
stdout_handler = logging.StreamHandler(sys.stdout)
stdout_handler.setFormatter(formatter)
+ stdout_handler.setLevel(console_level)
stdout_handler.addFilter(lambda record: record.levelno < logging.ERROR)
logger.addHandler(stdout_handler)
logger.addHandler(stream_handler)
+ for output_level, output_path in file_outputs:
+ output_path = os.path.abspath(output_path)
+ try:
+ output_handler = logging.FileHandler(output_path, encoding="utf-8")
+ except OSError as e:
+ logging.warning("Could not open %s log %s: %s", output_level, output_path, e)
+ continue
+ output_handler.setLevel(get_log_level(output_level))
+ output_handler.setFormatter(logging.Formatter("[%(asctime)s] [%(levelname)s] %(message)s"))
+ logger.addHandler(output_handler)
+ logging.info("%s log: %s", output_level.title(), output_path)
+
STARTUP_WARNINGS = []
diff --git a/app/user_manager.py b/app/user_manager.py
index de261ad39..55e7e81e3 100644
--- a/app/user_manager.py
+++ b/app/user_manager.py
@@ -343,13 +343,22 @@ class UserManager():
# XSS). Content-Disposition: attachment is the load-bearing guard;
# the content-type override and nosniff are defence in depth.
content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream'
- if folder_paths.is_dangerous_content_type(content_type):
- content_type = 'application/octet-stream'
+
+ user_root = self.get_request_user_filepath(request, None, create_dir=False)
+ is_user_css = path == os.path.abspath(os.path.join(user_root, "user.css"))
+
+ if is_user_css:
+ content_type = "text/css"
+ disposition = "inline"
+ else:
+ if folder_paths.is_dangerous_content_type(content_type):
+ content_type = 'application/octet-stream'
+ disposition = "attachment"
return web.FileResponse(path, headers={
"Content-Type": content_type,
"X-Content-Type-Options": "nosniff",
- "Content-Disposition": "attachment",
+ "Content-Disposition": disposition,
})
@routes.post("/userdata/{file}")
diff --git a/comfy/cli_args.py b/comfy/cli_args.py
index e2e0d97ec..792148f0a 100644
--- a/comfy/cli_args.py
+++ b/comfy/cli_args.py
@@ -33,6 +33,31 @@ class EnumAction(argparse.Action):
setattr(namespace, self.dest, value)
+LOG_LEVELS = ('DEBUG', 'DETAIL', 'INFO', 'WARNING', 'ERROR', 'CRITICAL')
+
+
+class VerboseAction(argparse.Action):
+ def __call__(self, parser, namespace, values, option_string=None):
+ if len(values) == 0:
+ output = ('DEBUG', None)
+ elif len(values) == 1 and values[0] in LOG_LEVELS:
+ output = (values[0], None)
+ elif len(values) == 2 and values[0] in LOG_LEVELS:
+ output = tuple(values)
+ else:
+ parser.error(f"{option_string} expects no values, a console LEVEL, or LEVEL FILE")
+ setattr(namespace, self.dest, [*getattr(namespace, self.dest, []), output])
+
+
+def get_console_log_level(outputs):
+ console_levels = [level for level, path in outputs if path is None]
+ return min(console_levels, key=LOG_LEVELS.index, default='INFO')
+
+
+def get_file_log_outputs(outputs):
+ return [(level, path) for level, path in outputs if path is not None]
+
+
parser = argparse.ArgumentParser()
parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)")
@@ -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).")
diff --git a/comfy/ldm/ernie/model.py b/comfy/ldm/ernie/model.py
index f158ca1d2..88a3775d0 100644
--- a/comfy/ldm/ernie/model.py
+++ b/comfy/ldm/ernie/model.py
@@ -5,6 +5,7 @@ import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention
import comfy.model_management
+import comfy.ops
import comfy.quant_ops
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
@@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module):
query = q_flat.view(B, S, self.heads, self.head_dim)
key = k_flat.view(B, S, self.heads, self.head_dim)
- query = self.norm_q(query)
- key = self.norm_k(key)
-
- if image_rotary_emb is not None:
- query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
+ if image_rotary_emb is not None and not comfy.model_management.in_training:
+ q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True)
+ k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True)
+ query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps)
+ comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
+ comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
+ else:
+ query = self.norm_q(query)
+ key = self.norm_k(key)
+ if image_rotary_emb is not None:
+ query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb)
q_flat = query.reshape(B, S, -1)
k_flat = key.reshape(B, S, -1)
diff --git a/comfy/ldm/ideogram4/model.py b/comfy/ldm/ideogram4/model.py
index 4ea5b8aaf..12e1a14fb 100644
--- a/comfy/ldm/ideogram4/model.py
+++ b/comfy/ldm/ideogram4/model.py
@@ -12,10 +12,13 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
+import comfy.model_management
+import comfy.ops
import comfy.patcher_extension
+import comfy.quant_ops
from comfy.ldm.lumina.model import FeedForward
from comfy.ldm.modules.attention import optimized_attention_masked
-from comfy.text_encoders.llama import apply_rope, precompute_freqs_cis
+from comfy.text_encoders.llama import precompute_freqs_cis
# Per-token role indicators
SEQUENCE_PADDING_INDICATOR = -1
@@ -25,6 +28,22 @@ LLM_TOKEN_INDICATOR = 3
IMAGE_POSITION_OFFSET = 65536
+def _split_half_rope_matrix(freqs_cis):
+ cos, sin, neg_sin = freqs_cis
+ half_dim = sin.shape[-1]
+ matrix = torch.stack(
+ (cos[..., :half_dim], neg_sin, sin, cos[..., half_dim:]), dim=-1
+ )
+ return matrix.reshape(*matrix.shape[:-1], 2, 2).unsqueeze(2)
+
+
+def _apply_rope_split_half1(x, freqs_cis):
+ x_dtype = x.dtype
+ x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(freqs_cis.dtype)
+ output = freqs_cis[..., 0] * x[..., 0] + freqs_cis[..., 1] * x[..., 1]
+ return output.movedim(-1, -2).reshape(*x.shape[:-3], -1).to(x_dtype)
+
+
class Ideogram4Attention(nn.Module):
def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None):
super().__init__()
@@ -42,16 +61,23 @@ class Ideogram4Attention(nn.Module):
qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2)
- q = self.norm_q(q)
- k = self.norm_k(k)
+ if comfy.model_management.in_training:
+ q = _apply_rope_split_half1(self.norm_q(q), freqs_cis)
+ k = _apply_rope_split_half1(self.norm_k(k), freqs_cis)
+ else:
+ q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, q, offloadable=True)
+ k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, k, offloadable=True)
+ q, k = comfy.quant_ops.ck.rms_rope_split_half(
+ q, k, freqs_cis, q_scale, k_scale, self.norm_q.eps
+ )
+ comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream)
+ comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream)
# (B, heads, L, head_dim)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
- q, k = apply_rope(q, k, freqs_cis)
-
out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options)
return self.o(out)
@@ -181,6 +207,7 @@ class Ideogram4Transformer(nn.Module):
self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta,
rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device,
)
+ freqs_cis = _split_half_rope_matrix(freqs_cis)
if attn_mask is not None and attn_mask.dtype == torch.bool:
attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max)
diff --git a/comfy/ldm/joyimage/model.py b/comfy/ldm/joyimage/model.py
index bca12c391..9d6951e54 100644
--- a/comfy/ldm/joyimage/model.py
+++ b/comfy/ldm/joyimage/model.py
@@ -94,12 +94,21 @@ class JoyImageAttention(nn.Module):
txt_k = txt_k.unflatten(-1, (heads, -1))
txt_v = txt_v.unflatten(-1, (heads, -1))
- img_q = self.img_attn_q_norm(img_q)
- img_k = self.img_attn_k_norm(img_k)
txt_q = self.txt_attn_q_norm(txt_q)
txt_k = self.txt_attn_k_norm(txt_k)
- img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb)
+ img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True)
+ img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True)
+ img_q, img_k = comfy_kitchen.rms_rope(
+ img_q,
+ img_k,
+ image_rotary_emb,
+ img_q_scale,
+ img_k_scale,
+ self.img_attn_q_norm.eps,
+ )
+ comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream)
+ comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream)
joint_q = torch.cat([img_q, txt_q], dim=1)
joint_k = torch.cat([img_k, txt_k], dim=1)
diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py
index 2811080be..1a6ddcc8d 100644
--- a/comfy/ldm/lightricks/embeddings_connector.py
+++ b/comfy/ldm/lightricks/embeddings_connector.py
@@ -6,9 +6,8 @@ import torch
from comfy.ldm.lightricks.model import (
CrossAttention,
FeedForward,
+ freqs_cis_matrix,
generate_freq_grid_np,
- interleaved_freqs_cis,
- split_freqs_cis,
)
from torch import nn
@@ -244,12 +243,15 @@ class Embeddings1DConnector(nn.Module):
expected_freqs = dim // 2
current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs
- cos_freq, sin_freq = split_freqs_cis(
- freqs, pad_size, self.num_attention_heads
- )
else:
- cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem)
- return cos_freq.to(dtype=out_dtype), sin_freq.to(dtype=out_dtype), self.split_rope
+ pad_size = dim % n_elem
+ return freqs_cis_matrix(
+ freqs,
+ pad_size,
+ self.split_rope,
+ self.num_attention_heads,
+ out_dtype,
+ )
def forward(
self,
diff --git a/comfy/ldm/lightricks/latent_upsampler.py b/comfy/ldm/lightricks/latent_upsampler.py
index 78ed7653f..6a4beb1bf 100644
--- a/comfy/ldm/lightricks/latent_upsampler.py
+++ b/comfy/ldm/lightricks/latent_upsampler.py
@@ -97,11 +97,11 @@ class SpatialRationalResampler(nn.Module):
For dims==3, work per-frame for spatial scaling (temporal axis untouched).
"""
- def __init__(self, mid_channels: int, scale: float):
+ def __init__(self, mid_channels: int, scale: float, operations):
super().__init__()
self.scale = float(scale)
self.num, self.den = _rational_for_scale(self.scale)
- self.conv = nn.Conv2d(
+ self.conv = operations.Conv2d(
mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1
)
self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num))
@@ -119,18 +119,18 @@ class SpatialRationalResampler(nn.Module):
class ResBlock(nn.Module):
def __init__(
- self, channels: int, mid_channels: Optional[int] = None, dims: int = 3
+ self, channels: int, operations, mid_channels: Optional[int] = None, dims: int = 3
):
super().__init__()
if mid_channels is None:
mid_channels = channels
- Conv = nn.Conv2d if dims == 2 else nn.Conv3d
+ Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1)
- self.norm1 = nn.GroupNorm(32, mid_channels)
+ self.norm1 = operations.GroupNorm(32, mid_channels)
self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1)
- self.norm2 = nn.GroupNorm(32, channels)
+ self.norm2 = operations.GroupNorm(32, channels)
self.activation = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
@@ -159,6 +159,7 @@ class LatentUpsampler(nn.Module):
def __init__(
self,
+ operations,
in_channels: int = 128,
mid_channels: int = 512,
num_blocks_per_stage: int = 4,
@@ -179,34 +180,34 @@ class LatentUpsampler(nn.Module):
self.spatial_scale = float(spatial_scale)
self.rational_resampler = rational_resampler
- Conv = nn.Conv2d if dims == 2 else nn.Conv3d
+ Conv = operations.Conv2d if dims == 2 else operations.Conv3d
self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1)
- self.initial_norm = nn.GroupNorm(32, mid_channels)
+ self.initial_norm = operations.GroupNorm(32, mid_channels)
self.initial_activation = nn.SiLU()
self.res_blocks = nn.ModuleList(
- [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
+ [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
)
if spatial_upsample and temporal_upsample:
self.upsampler = nn.Sequential(
- nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
+ operations.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(3),
)
elif spatial_upsample:
if rational_resampler:
self.upsampler = SpatialRationalResampler(
- mid_channels=mid_channels, scale=self.spatial_scale
+ mid_channels=mid_channels, scale=self.spatial_scale, operations=operations
)
else:
self.upsampler = nn.Sequential(
- nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
+ operations.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(2),
)
elif temporal_upsample:
self.upsampler = nn.Sequential(
- nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
+ operations.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
PixelShuffleND(1),
)
else:
@@ -215,11 +216,14 @@ class LatentUpsampler(nn.Module):
)
self.post_upsample_res_blocks = nn.ModuleList(
- [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
+ [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)]
)
self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1)
+ def get_dtype(self):
+ return getattr(self.initial_conv, "weight_comfy_model_dtype", self.initial_conv.weight.dtype)
+
def forward(self, latent: torch.Tensor) -> torch.Tensor:
b, c, f, h, w = latent.shape
@@ -266,7 +270,7 @@ class LatentUpsampler(nn.Module):
return x
@classmethod
- def from_config(cls, config):
+ def from_config(cls, config, operations):
return cls(
in_channels=config.get("in_channels", 4),
mid_channels=config.get("mid_channels", 128),
@@ -276,6 +280,7 @@ class LatentUpsampler(nn.Module):
temporal_upsample=config.get("temporal_upsample", False),
spatial_scale=config.get("spatial_scale", 2.0),
rational_resampler=config.get("rational_resampler", False),
+ operations=operations,
)
def config(self):
diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py
index 9953b6679..f9de3a38e 100644
--- a/comfy/ldm/lightricks/model.py
+++ b/comfy/ldm/lightricks/model.py
@@ -12,6 +12,8 @@ from torch import nn
import comfy.patcher_extension
import comfy.ldm.modules.attention
import comfy.ldm.common_dit
+import comfy.model_management
+import comfy.quant_ops
from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords
@@ -322,40 +324,42 @@ class FeedForward(nn.Module):
return self.net(x)
def apply_rotary_emb(input_tensor, freqs_cis):
- cos_freqs, sin_freqs = freqs_cis[0], freqs_cis[1]
- split_pe = freqs_cis[2] if len(freqs_cis) > 2 else False
- return (
- apply_split_rotary_emb(input_tensor, cos_freqs, sin_freqs)
- if split_pe else
- apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs)
+ rotation_matrix, split_pe = freqs_cis
+ original_shape = input_tensor.shape
+ input_tensor = input_tensor.reshape(
+ input_tensor.shape[0], input_tensor.shape[1], rotation_matrix.shape[2], -1
)
-def apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs): # TODO: remove duplicate funcs and pick the best/fastest one
- t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2)
- t1, t2 = t_dup.unbind(dim=-1)
- t_dup = torch.stack((-t2, t1), dim=-1)
- input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)")
+ if comfy.model_management.in_training:
+ if split_pe:
+ t = input_tensor.reshape(*input_tensor.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2)
+ else:
+ t = input_tensor.reshape(*input_tensor.shape[:-1], -1, 1, 2)
+ t = t.to(rotation_matrix.dtype)
+ output = rotation_matrix[..., 0] * t[..., 0] + rotation_matrix[..., 1] * t[..., 1]
+ if split_pe:
+ output = output.movedim(-1, -2)
+ output = output.reshape(input_tensor.shape).type_as(input_tensor)
+ elif split_pe:
+ output = comfy.quant_ops.ck.apply_rope_split_half1(input_tensor, rotation_matrix)
+ else:
+ output = comfy.quant_ops.ck.apply_rope1(input_tensor, rotation_matrix)
+ return output.reshape(original_shape)
- out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs
+def apply_rotary_emb_qk(q, k, freqs_cis):
+ if comfy.model_management.in_training:
+ return apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis)
- return out
-
-def apply_split_rotary_emb(input_tensor, cos, sin):
- needs_reshape = False
- if input_tensor.ndim != 4 and cos.ndim == 4:
- B, H, T, _ = cos.shape
- input_tensor = input_tensor.reshape(B, T, H, -1).swapaxes(1, 2)
- needs_reshape = True
- split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2)
- first_half_input = split_input[..., :1, :]
- second_half_input = split_input[..., 1:, :]
- output = split_input * cos.unsqueeze(-2)
- first_half_output = output[..., :1, :]
- second_half_output = output[..., 1:, :]
- first_half_output.addcmul_(-sin.unsqueeze(-2), second_half_input)
- second_half_output.addcmul_(sin.unsqueeze(-2), first_half_input)
- output = rearrange(output, "... d r -> ... (d r)")
- return output.swapaxes(1, 2).reshape(B, T, -1) if needs_reshape else output
+ rotation_matrix, split_pe = freqs_cis
+ q_shape = q.shape
+ k_shape = k.shape
+ q = q.reshape(q.shape[0], q.shape[1], rotation_matrix.shape[2], -1)
+ k = k.reshape(k.shape[0], k.shape[1], rotation_matrix.shape[2], -1)
+ if split_pe:
+ q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rotation_matrix)
+ else:
+ q, k = comfy.quant_ops.ck.apply_rope(q, k, rotation_matrix)
+ return q.reshape(q_shape), k.reshape(k_shape)
class GuideAttentionMask:
@@ -461,9 +465,13 @@ class CrossAttention(nn.Module):
q = self.q_norm(q)
k = self.k_norm(k)
+ # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent.
if pe is not None:
- q = apply_rotary_emb(q, pe)
- k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
+ if k_pe is None and q.shape == k.shape:
+ q, k = apply_rotary_emb_qk(q, k, pe)
+ else:
+ q = apply_rotary_emb(q, pe)
+ k = apply_rotary_emb(k, pe if k_pe is None else k_pe)
if mask is None:
out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options)
@@ -653,36 +661,23 @@ def generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid):
)
return freqs
-def interleaved_freqs_cis(freqs, pad_size):
- cos_freq = freqs.cos().repeat_interleave(2, dim=-1)
- sin_freq = freqs.sin().repeat_interleave(2, dim=-1)
- if pad_size != 0:
- cos_padding = torch.ones_like(cos_freq[:, :, : pad_size])
- sin_padding = torch.zeros_like(cos_freq[:, :, : pad_size])
- cos_freq = torch.cat([cos_padding, cos_freq], dim=-1)
- sin_freq = torch.cat([sin_padding, sin_freq], dim=-1)
- return cos_freq, sin_freq
-
-def 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."""
diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py
index dd5320c8f..b4a8c7524 100644
--- a/comfy/ldm/lightricks/vae/audio_vae.py
+++ b/comfy/ldm/lightricks/vae/audio_vae.py
@@ -185,7 +185,7 @@ class AudioVAE(torch.nn.Module):
self.autoencoder.mel_bins,
)
- def num_of_latents_from_frames(self, frames_number: int, frame_rate: int) -> int:
+ def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int:
return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second)
def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor:
diff --git a/comfy/ldm/lightricks/vae/causal_conv3d.py b/comfy/ldm/lightricks/vae/causal_conv3d.py
index 7515f0d4e..bb1803f12 100644
--- a/comfy/ldm/lightricks/vae/causal_conv3d.py
+++ b/comfy/ldm/lightricks/vae/causal_conv3d.py
@@ -49,6 +49,12 @@ class CausalConv3d(nn.Module):
)
self.temporal_cache_state={}
+ def _empty_output(self, x):
+ # empty (0 frame) outputs must still have the conv's output channels and spatial dims
+ h = (x.shape[3] + 2 * self.conv.padding[1] - self.conv.kernel_size[1]) // self.conv.stride[1] + 1
+ w = (x.shape[4] + 2 * self.conv.padding[2] - self.conv.kernel_size[2]) // self.conv.stride[2] + 1
+ return x.new_empty((x.shape[0], self.out_channels, 0, h, w))
+
def forward(self, x, causal: bool = True):
tid = threading.get_ident()
@@ -58,7 +64,7 @@ class CausalConv3d(nn.Module):
if not causal:
padding_length = padding_length // 2
if x.shape[2] == 0:
- return x
+ return self._empty_output(x)
cached = x[:, :, :1, :, :].repeat((1, 1, padding_length, 1, 1))
pieces = [ cached, x ]
if is_end and not causal:
@@ -83,7 +89,7 @@ class CausalConv3d(nn.Module):
elif is_end:
self.temporal_cache_state[tid] = (None, True)
- return self.conv(x) if x.shape[2] >= self.time_kernel_size else x[:, :, :0, :, :]
+ return self.conv(x) if x.shape[2] >= self.time_kernel_size else self._empty_output(x)
@property
def weight(self):
diff --git a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py
index 5975015e2..5d0eec5b8 100644
--- a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py
+++ b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py
@@ -390,10 +390,10 @@ class Decoder(nn.Module):
# Compute output channel to be product of all channel-multiplier blocks
output_channel = base_channels
- for block_name, block_params in list(reversed(blocks)):
+ for block_name, block_params in blocks:
block_params = block_params if isinstance(block_params, dict) else {}
if block_name == "res_x_y":
- output_channel = output_channel * block_params.get("multiplier", 2)
+ output_channel = block_params.get("in_channels", output_channel * block_params.get("multiplier", 2))
if block_name == "compress_all":
output_channel = output_channel * block_params.get("multiplier", 1)
if block_name == "compress_space":
@@ -432,7 +432,7 @@ class Decoder(nn.Module):
spatial_padding_mode=spatial_padding_mode,
)
elif block_name == "res_x_y":
- output_channel = output_channel // block_params.get("multiplier", 2)
+ output_channel = block_params.get("out_channels", output_channel // block_params.get("multiplier", 2))
block = ResnetBlock3D(
dims=dims,
in_channels=input_channel,
diff --git a/comfy/ldm/lumina/model.py b/comfy/ldm/lumina/model.py
index d0ee97d33..cdf03b2b5 100644
--- a/comfy/ldm/lumina/model.py
+++ b/comfy/ldm/lumina/model.py
@@ -6,6 +6,9 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
import comfy.ldm.common_dit
+import comfy.model_management
+import comfy.ops
+import comfy.quant_ops
from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder
from comfy.ldm.modules.attention import optimized_attention_masked
@@ -97,6 +100,7 @@ class JointAttention(nn.Module):
self.n_local_kv_heads = self.n_kv_heads
self.n_rep = self.n_local_heads // self.n_local_kv_heads
self.head_dim = dim // n_heads
+ self.qk_norm = qk_norm
self.qkv = operation_settings.get("operations").Linear(
dim,
@@ -151,10 +155,21 @@ class JointAttention(nn.Module):
xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim)
- xq = self.q_norm(xq)
- xk = self.k_norm(xk)
-
- xq, xk = apply_rope(xq, xk, freqs_cis)
+ if self.qk_norm and not comfy.model_management.in_training:
+ q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True)
+ k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True)
+ epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps
+ if self.n_local_heads == self.n_local_kv_heads:
+ xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon)
+ else:
+ xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon)
+ xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon)
+ comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
+ comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
+ else:
+ xq = self.q_norm(xq)
+ xk = self.k_norm(xk)
+ xq, xk = apply_rope(xq, xk, freqs_cis)
n_rep = self.n_local_heads // self.n_local_kv_heads
if n_rep >= 1:
diff --git a/comfy/ldm/mage_flow/model.py b/comfy/ldm/mage_flow/model.py
new file mode 100644
index 000000000..ac29bb610
--- /dev/null
+++ b/comfy/ldm/mage_flow/model.py
@@ -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)
diff --git a/comfy/ldm/mage_flow/vae.py b/comfy/ldm/mage_flow/vae.py
new file mode 100644
index 000000000..e6e21b99f
--- /dev/null
+++ b/comfy/ldm/mage_flow/vae.py
@@ -0,0 +1,477 @@
+# Mage-VAE (https://github.com/microsoft/Mage) (MIT)
+# Symmetric one-step diffusion codec: DConvEncoder (image -> 128ch latent) and
+# DConvDenoiser + CoD Decoder (latent -> image). 16x downsample, latents in the
+# Flux.2-VAE-anchored space (no patch packing, no BN normalization).
+# Both encode and decode are single forward passes at t=0.
+import math
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+import comfy.ops
+from comfy.ldm.modules.diffusionmodules.model import vae_attention
+
+ops = comfy.ops.disable_weight_init
+
+
+def nonlinearity(x):
+ return torch.nn.functional.silu(x)
+
+
+def Normalize(in_channels):
+ return ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
+
+
+def modulate(x, shift, scale):
+ if x.dim() == 4:
+ b, c = x.shape[:2]
+ return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1)
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
+
+
+class LayerNorm2d(ops.LayerNorm):
+ def __init__(self, num_channels, eps=1e-6, affine=True):
+ super().__init__(num_channels, eps=eps, elementwise_affine=affine)
+
+ def forward(self, x):
+ x = x.permute(0, 2, 3, 1).contiguous()
+ x = super().forward(x)
+ return x.permute(0, 3, 1, 2).contiguous()
+
+
+class TimestepEmbedder(nn.Module):
+ """DConv-style timestep MLP (max_period=10000, freq_size=256)."""
+
+ def __init__(self, hidden_size, frequency_embedding_size=256):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ ops.Linear(frequency_embedding_size, hidden_size, bias=True),
+ nn.SiLU(),
+ ops.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.frequency_embedding_size = frequency_embedding_size
+
+ @staticmethod
+ def timestep_embedding(t, dim, max_period=10000):
+ half = dim // 2
+ freqs = torch.exp(
+ -math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half
+ ).to(t.device)
+ args = t[:, None].float() * freqs[None]
+ emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
+ return emb
+
+ def forward(self, t, dtype):
+ emb = self.timestep_embedding(t, self.frequency_embedding_size)
+ return self.mlp(emb.to(dtype))
+
+
+class BottleneckPatchEmbed(nn.Module):
+ """Image patch embed concatenated with a per-patch conditioning vector."""
+
+ def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True):
+ super().__init__()
+ self.proj1 = ops.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False)
+ self.proj2 = ops.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias)
+
+ def forward(self, x, cond):
+ return self.proj2(torch.cat([self.proj1(x), cond], dim=1))
+
+
+class DiCoBlock(nn.Module):
+ """DConv block with adaLN modulation."""
+
+ def __init__(self, hidden_size, mlp_ratio=4.0):
+ super().__init__()
+ self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
+ self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
+ self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
+
+ self.ca = nn.Sequential(
+ nn.AdaptiveAvgPool2d(1),
+ ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
+ nn.Sigmoid(),
+ )
+
+ ffn = int(mlp_ratio * hidden_size)
+ self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
+ self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
+
+ self.norm1 = LayerNorm2d(hidden_size, affine=False)
+ self.norm2 = LayerNorm2d(hidden_size, affine=False)
+
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ ops.Linear(hidden_size, 6 * hidden_size, bias=True),
+ )
+
+ def forward(self, inp, c):
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
+ x = modulate(self.norm1(inp), shift_msa, scale_msa)
+ x = F.gelu(self.conv2(self.conv1(x)))
+ x = x * self.ca(x)
+ x = self.conv3(x)
+ x = inp + gate_msa[..., None, None] * x
+ x = x + gate_mlp[..., None, None] * self.conv5(
+ F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp)))
+ )
+ return x
+
+
+class EncoderDiCoBlock(nn.Module):
+ """DiCoBlock without adaLN, for the encoder head pathway."""
+
+ def __init__(self, hidden_size, mlp_ratio=4.0):
+ super().__init__()
+ self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
+ self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True)
+ self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True)
+ self.ca = nn.Sequential(
+ nn.AdaptiveAvgPool2d(1),
+ ops.Conv2d(hidden_size, hidden_size, 1, bias=True),
+ nn.Sigmoid(),
+ )
+ ffn = int(mlp_ratio * hidden_size)
+ self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True)
+ self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True)
+ self.norm1 = LayerNorm2d(hidden_size)
+ self.norm2 = LayerNorm2d(hidden_size)
+
+ def forward(self, inp):
+ x = self.norm1(inp)
+ x = F.gelu(self.conv2(self.conv1(x)))
+ x = x * self.ca(x)
+ x = self.conv3(x)
+ x = inp + x
+ return x + self.conv5(F.gelu(self.conv4(self.norm2(x))))
+
+
+class NerfEmbedder(nn.Module):
+ """Patch-position embedder used by the DConv decoder x-pathway."""
+
+ def __init__(self, in_channels, hidden_size_input, max_freqs=8):
+ super().__init__()
+ self.max_freqs = max_freqs
+ self.embedder = nn.Sequential(
+ ops.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True),
+ )
+
+ def fetch_pos(self, patch_size, device, dtype):
+ pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype)
+ pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij")
+ pos_x = pos_x.reshape(-1, 1, 1)
+ pos_y = pos_y.reshape(-1, 1, 1)
+ freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device)
+ fx = freqs[None, :, None]
+ fy = freqs[None, None, :]
+ coeffs = (1 + fx * fy) ** -1
+ dct_x = torch.cos(pos_x * fx * torch.pi)
+ dct_y = torch.cos(pos_y * fy * torch.pi)
+ return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2)
+
+ def forward(self, x):
+ B, P2, _ = x.shape
+ ps = int(P2 ** 0.5)
+ dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1)
+ return self.embedder(torch.cat([x, dct], dim=-1))
+
+
+class NerfFinalLayer(nn.Module):
+ def __init__(self, hidden_size, out_channels):
+ super().__init__()
+ self.norm = ops.RMSNorm(hidden_size, eps=1e-6)
+ self.linear = ops.Linear(hidden_size, out_channels, bias=True)
+
+ def forward(self, x):
+ return self.linear(self.norm(x))
+
+
+class MLPResBlock(nn.Module):
+ def __init__(self, channels):
+ super().__init__()
+ self.in_ln = ops.LayerNorm(channels, eps=1e-6)
+ self.mlp = nn.Sequential(
+ ops.Linear(channels, channels, bias=True),
+ nn.SiLU(),
+ ops.Linear(channels, channels, bias=True),
+ )
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ ops.Linear(channels, 3 * channels, bias=True),
+ )
+
+ def forward(self, x, y):
+ shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1)
+ h = self.in_ln(x) * (1 + scale) + shift
+ return x + gate * self.mlp(h)
+
+
+class SimpleMLPAdaLN(nn.Module):
+ """Final small MLP that maps NerfEmbedder features to per-patch RGB."""
+
+ def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size):
+ super().__init__()
+ self.in_channels = in_channels
+ self.model_channels = model_channels
+ self.out_channels = out_channels
+ self.num_res_blocks = num_res_blocks
+ self.patch_size = patch_size
+
+ self.cond_embed = ops.Linear(z_channels, patch_size ** 2 * model_channels)
+ self.input_proj = ops.Linear(in_channels, model_channels)
+
+ self.res_blocks = nn.ModuleList(MLPResBlock(model_channels) for _ in range(num_res_blocks))
+
+ def forward(self, x, c):
+ x = self.input_proj(x)
+ c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1)
+ for block in self.res_blocks:
+ x = block(x, c)
+ return x
+
+
+class ResnetBlock(nn.Module):
+ """GroupNorm + Conv ResBlock used by the CoD Decoder."""
+
+ def __init__(self, *, in_channels, out_channels=None):
+ super().__init__()
+ out_channels = out_channels or in_channels
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+
+ self.norm1 = Normalize(in_channels)
+ self.conv1 = ops.Conv2d(in_channels, out_channels, 3, padding=1)
+ self.norm2 = Normalize(out_channels)
+ self.conv2 = ops.Conv2d(out_channels, out_channels, 3, padding=1)
+ if in_channels != out_channels:
+ self.nin_shortcut = ops.Conv2d(in_channels, out_channels, 1)
+
+ def forward(self, x):
+ h = self.conv1(nonlinearity(self.norm1(x)))
+ h = self.conv2(nonlinearity(self.norm2(h)))
+ if self.in_channels != self.out_channels:
+ x = self.nin_shortcut(x)
+ return x + h
+
+
+class AttnBlock(nn.Module):
+ """Patched (windowed) self-attention used by the CoD Decoder."""
+
+ def __init__(self, in_channels, patch_size=32):
+ super().__init__()
+ self.in_channels = in_channels
+ self.patch_size = patch_size
+ self.norm = Normalize(in_channels)
+ self.q = ops.Conv2d(in_channels, in_channels, 1)
+ self.k = ops.Conv2d(in_channels, in_channels, 1)
+ self.v = ops.Conv2d(in_channels, in_channels, 1)
+ self.proj_out = ops.Conv2d(in_channels, in_channels, 1)
+ # VAE attention selection: full-precision backends only (no sage/quantized attention)
+ self.optimized_attention = vae_attention()
+
+ def forward(self, x):
+ h_ = self.norm(x)
+ Q = self.q(h_)
+ K = self.k(h_)
+ V = self.v(h_)
+
+ d = self.patch_size
+ b, c, H, W = Q.shape
+ pad_h = (d - H % d) % d
+ pad_w = (d - W % d) % d
+ if pad_h or pad_w:
+ Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate")
+ K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate")
+ V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate")
+ _, _, H_pad, W_pad = Q.shape
+ nph, npw = H_pad // d, W_pad // d
+ np_ = nph * npw
+
+ def to_patches(t):
+ return (t.reshape(b, c, nph, d, npw, d)
+ .permute(0, 2, 4, 1, 3, 5)
+ .reshape(b * np_, c, d * d))
+
+ # [b*np, c, d*d]: attention over the d*d spatial positions of each window
+ Q = to_patches(Q)
+ K = to_patches(K)
+ V = to_patches(V)
+
+ h_ = self.optimized_attention(Q, K, V)
+ h_ = h_.reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad)
+ if pad_h or pad_w:
+ h_ = h_[:, :, :H, :W]
+ return x + self.proj_out(h_)
+
+
+class CoDDecoder(nn.Module):
+ """CoD Decoder: latent -> conditioning features for the denoiser (ds=16, light)."""
+
+ def __init__(self, out_ch=384, z_ch=128):
+ super().__init__()
+ self.conv_in = ops.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1)
+ self.block = nn.Sequential(
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
+ AttnBlock(out_ch, patch_size=32),
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
+ AttnBlock(out_ch, patch_size=32),
+ ResnetBlock(in_channels=out_ch, out_channels=out_ch),
+ )
+ self.norm_out = Normalize(out_ch)
+ self.conv_out = ops.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1)
+ self.ada = nn.Identity()
+
+ def forward(self, z):
+ h = self.block(self.conv_in(z))
+ h = self.conv_out(nonlinearity(self.norm_out(h)))
+ return self.ada(h)
+
+
+class DConvEncoder(nn.Module):
+ """DConvEncoder: image -> packed (mean, logvar) latent."""
+
+ def __init__(
+ self,
+ z_ch=128,
+ hidden_size=384,
+ num_blocks=21,
+ patch_size=16,
+ mlp_ratio=4.0,
+ head_size=768,
+ num_head_blocks=2,
+ out_ch_mult=2,
+ ):
+ super().__init__()
+ self.z_ch = z_ch
+ self.patch_size = patch_size
+ self.patch_cond_embed = ops.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True)
+ self.head_blocks = nn.ModuleList([
+ EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks)
+ ])
+ self.proj_down = ops.Conv2d(head_size, hidden_size, kernel_size=1, bias=True)
+ self.z_proj = ops.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True)
+ self.fuse_proj = ops.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True)
+ self.t_embedder = TimestepEmbedder(hidden_size)
+ self.blocks = nn.ModuleList([
+ DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks)
+ ])
+ self.norm_out = LayerNorm2d(hidden_size)
+ self.proj_out = ops.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True)
+
+ def forward_pred(self, z_t, t, y):
+ cond = self.patch_cond_embed(y)
+ for block in self.head_blocks:
+ cond = block(cond)
+ cond = self.proj_down(cond)
+
+ s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1))
+ c = self.t_embedder(t.view(-1), y.dtype)
+ for block in self.blocks:
+ s = block(s, c)
+ return self.proj_out(self.norm_out(s))
+
+
+class YEmbedder(nn.Module):
+ """Holds only the CoD decoder (the original Flux2-VAE encoder side is dropped at load)."""
+
+ def __init__(self, ch=384, z_ch=128):
+ super().__init__()
+ self.decoder = CoDDecoder(out_ch=ch, z_ch=z_ch)
+
+
+class DConvDenoiser(nn.Module):
+ """One-step DConv denoiser: latent (via cond) + zero noise -> reconstructed image."""
+
+ def __init__(
+ self,
+ patch_size=16,
+ in_channels=3,
+ hidden_size=384,
+ hidden_size_x=32,
+ mlp_ratio=4.0,
+ num_blocks=24,
+ num_cond_blocks=21,
+ bottleneck_dim=128,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.patch_size = patch_size
+ self.hidden_size = hidden_size
+ self.num_cond_blocks = num_cond_blocks
+
+ self.t_embedder = TimestepEmbedder(hidden_size)
+ self.y_embedder_x = ops.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0)
+ self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8)
+ self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True)
+ self.blocks = nn.ModuleList([
+ DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks)
+ ])
+ self.dec_net = SimpleMLPAdaLN(
+ in_channels=hidden_size_x,
+ model_channels=hidden_size_x,
+ out_channels=in_channels,
+ z_channels=hidden_size,
+ num_res_blocks=num_blocks - num_cond_blocks,
+ patch_size=patch_size,
+ )
+ self.final_layer = NerfFinalLayer(hidden_size_x, in_channels)
+ self.y_embedder = YEmbedder(ch=hidden_size, z_ch=bottleneck_dim)
+
+ def forward(self, x, t, cond):
+ b, _, h, w = x.shape
+ c = self.t_embedder(t.view(-1), x.dtype)
+
+ s = self.s_embedder(x, cond)
+ for block in self.blocks:
+ s = block(s, c)
+
+ length = s.shape[-2] * s.shape[-1]
+ s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size)
+
+ x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size)
+ x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1)
+ x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1)
+ x = self.x_embedder(x)
+
+ x = self.dec_net(x, s)
+ x = self.final_layer(x)
+ x = x.transpose(1, 2).reshape(b, length, -1)
+ return torch.nn.functional.fold(
+ x.transpose(1, 2).contiguous(), (h, w),
+ kernel_size=self.patch_size, stride=self.patch_size,
+ )
+
+
+class MageVAE(nn.Module):
+ """
+ Encode: DConvEncoder (one-step at t=0) -> posterior mean [B, 128, H/16, W/16]
+ Decode: DConvDenoiser + CoD Decoder -> image [B, 3, H, W] in [-1, 1]
+ """
+
+ latent_channels = 128
+ downsample_factor = 16
+
+ def __init__(self):
+ super().__init__()
+ self.dconv_encoder = DConvEncoder()
+ self.decoder_model = DConvDenoiser()
+
+ def encode(self, x):
+ B, _, H, W = x.shape
+ ps = self.dconv_encoder.patch_size
+ z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype)
+ t = torch.zeros(B, device=x.device, dtype=x.dtype)
+ out = self.dconv_encoder.forward_pred(z_t, t, x)
+ return out[:, : self.latent_channels] # posterior mean (sample_posterior=False)
+
+ def decode(self, z):
+ cond = self.decoder_model.y_embedder.decoder(z)
+ B = z.shape[0]
+ H = z.shape[2] * self.downsample_factor
+ W = z.shape[3] * self.downsample_factor
+ noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype)
+ t = torch.zeros(B, device=z.device, dtype=z.dtype)
+ return self.decoder_model.forward(noise, t, cond)
diff --git a/comfy/ldm/wan/model.py b/comfy/ldm/wan/model.py
index 1c9782a38..c042e93c4 100644
--- a/comfy/ldm/wan/model.py
+++ b/comfy/ldm/wan/model.py
@@ -552,6 +552,7 @@ class WanModel(torch.nn.Module):
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
# embeddings
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@@ -564,11 +565,13 @@ class WanModel(torch.nn.Module):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
+ img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
+ img_offset = full_ref.shape[1]
# In-context reference (Bernini)
context_latents = kwargs.get("context_latents", None)
@@ -589,6 +592,7 @@ class WanModel(torch.nn.Module):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -604,6 +608,11 @@ class WanModel(torch.nn.Module):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
+ x = out["img"]
+
# head
x = self.head(x, e)
@@ -777,6 +786,7 @@ class VaceWanModel(WanModel):
**kwargs,
):
# embeddings
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
transformer_options["grid_sizes"] = grid_sizes
@@ -807,6 +817,7 @@ class VaceWanModel(WanModel):
x_orig = x
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -822,6 +833,11 @@ class VaceWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
+ x = out["img"]
+
ii = self.vace_layers_mapping.get(i, None)
if ii is not None:
for iii in range(len(c)):
@@ -887,6 +903,7 @@ class CameraWanModel(WanModel):
**kwargs,
):
# embeddings
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if self.control_adapter is not None and camera_conditions is not None:
x = x + self.control_adapter(camera_conditions).to(x.dtype)
@@ -909,6 +926,7 @@ class CameraWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -924,6 +942,11 @@ class CameraWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
+ x = out["img"]
+
# head
x = self.head(x, e)
@@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel):
# embeddings
bs, _, time, height, width = x.shape
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
if control_video is not None:
x = x + self.cond_encoder(control_video)
@@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel):
context = self.text_embedding(context)
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
+
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
+ x = out["img"]
+
if audio_emb is not None:
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
# head
@@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel):
bs, _, time, height, width = x.shape
# embeddings
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
@@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel):
audio = None
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
+ x = out["img"]
+
# head
x = self.head(x, e)
@@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
+ x_input = x
+
+ img_offset = 0
if reference_latent is not None:
x = torch.cat((reference_latent, x), dim=2)
+ img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
+ (reference_latent.shape[3] // self.patch_size[1]) * \
+ (reference_latent.shape[4] // self.patch_size[2])
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
@@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
+ x = out["img"]
+
# head
x = self.head(x, e)
diff --git a/comfy/ldm/wan/model_animate.py b/comfy/ldm/wan/model_animate.py
index 84d7adec4..9ebe5694b 100644
--- a/comfy/ldm/wan/model_animate.py
+++ b/comfy/ldm/wan/model_animate.py
@@ -493,6 +493,7 @@ class AnimateWanModel(WanModel):
**kwargs,
):
# embeddings
+ x_input = x
x = self.patch_embedding(x.float()).to(x.dtype)
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
grid_sizes = x.shape[2:]
@@ -505,11 +506,13 @@ class AnimateWanModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
+ img_offset = 0
if self.ref_conv is not None:
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
+ img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@@ -522,6 +525,7 @@ class AnimateWanModel(WanModel):
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -537,6 +541,11 @@ class AnimateWanModel(WanModel):
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
+ x = out["img"]
+
if i % 5 == 0 and motion_vec is not None:
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)
diff --git a/comfy/ldm/wan/model_wandancer.py b/comfy/ldm/wan/model_wandancer.py
index 3caef6dc5..aeec1d725 100644
--- a/comfy/ldm/wan/model_wandancer.py
+++ b/comfy/ldm/wan/model_wandancer.py
@@ -111,6 +111,7 @@ class WanDancerModel(WanModel):
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
# embeddings
+ x_input = x
if int(fps + 0.5) != 30:
x = self.patch_embedding_global(x.float()).to(x.dtype)
else:
@@ -128,11 +129,13 @@ class WanDancerModel(WanModel):
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
full_ref = None
+ img_offset = 0
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
full_ref = kwargs.get("reference_latent", None)
if full_ref is not None:
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
x = torch.concat((full_ref, x), dim=1)
+ img_offset = full_ref.shape[1]
# context
context = self.text_embedding(context)
@@ -163,6 +166,7 @@ class WanDancerModel(WanModel):
context_img_len += clip_fea_ref.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
+ patches = transformer_options.get("patches", {})
blocks_replace = patches_replace.get("dit", {})
transformer_options["total_blocks"] = len(self.blocks)
transformer_options["block_type"] = "double"
@@ -177,6 +181,12 @@ class WanDancerModel(WanModel):
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
+
+ if "double_block" in patches:
+ for p in patches["double_block"]:
+ out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
+ x = out["img"]
+
if audio_emb is not None:
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)
diff --git a/comfy/ldm/wan/uni3c.py b/comfy/ldm/wan/uni3c.py
new file mode 100644
index 000000000..827ad2339
--- /dev/null
+++ b/comfy/ldm/wan/uni3c.py
@@ -0,0 +1,149 @@
+# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
+# Converted from the original diffusers based implementation.
+import torch
+import torch.nn as nn
+
+from comfy.ldm.flux.layers import EmbedND
+from .model import WanSelfAttention
+
+
+class Uni3CLayerNormZero(nn.Module):
+ def __init__(
+ self,
+ conditioning_dim,
+ embedding_dim,
+ eps=1e-5,
+ device=None, dtype=None, operations=None
+ ):
+ super().__init__()
+ self.silu = nn.SiLU()
+ self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
+ self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
+
+ def forward(self, x, temb):
+ shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
+ x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
+ return x, gate[:, None, :]
+
+
+class Uni3CAttentionBlock(nn.Module):
+ def __init__(
+ self,
+ dim,
+ ffn_dim,
+ num_heads,
+ time_embed_dim=5120,
+ eps=1e-6,
+ device=None, dtype=None, operations=None
+ ):
+ super().__init__()
+ operation_settings = {"operations": operations, "device": device, "dtype": dtype}
+ self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
+ self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
+ self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
+ self.ffn = nn.Sequential(
+ operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
+ operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
+
+ def forward(self, x, temb, freqs):
+ norm_x, gate_msa = self.norm1(x, temb)
+ x = x + gate_msa * self.self_attn(norm_x, freqs)
+ norm_x, gate_ff = self.norm2(x, temb)
+ x = x + gate_ff * self.ffn(norm_x)
+ return x
+
+
+class MaskCamEmbed(nn.Module):
+ def __init__(
+ self,
+ add_channels=7,
+ mid_channels=256,
+ conv_out_dim=5120,
+ device=None, dtype=None, operations=None
+ ):
+ super().__init__()
+ self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
+ self.mask_proj = nn.Sequential(
+ operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
+ operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
+ nn.SiLU())
+ self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
+
+ def forward(self, add_inputs):
+ add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
+ add_embeds = self.mask_proj(add_padded)
+ add_embeds = self.mask_zero_proj(add_embeds)
+ add_embeds = add_embeds.flatten(2).transpose(1, 2)
+ return add_embeds
+
+
+class WanUni3CControlnet(nn.Module):
+ def __init__(
+ self,
+ in_channels=36,
+ conv_out_dim=5120,
+ dim=1024,
+ ffn_dim=8192,
+ num_heads=16,
+ num_layers=20,
+ time_embed_dim=5120,
+ out_proj_dim=5120,
+ add_channels=7,
+ mid_channels=256,
+ device=None, dtype=None, operations=None
+ ):
+ super().__init__()
+ patch_size = (1, 2, 2)
+ self.num_layers = num_layers
+
+ self.controlnet_patch_embedding = operations.Conv3d(
+ in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
+ self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
+
+ if conv_out_dim != dim:
+ self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
+ else:
+ self.proj_in = nn.Identity()
+
+ self.controlnet_blocks = nn.ModuleList([
+ Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
+ for _ in range(num_layers)])
+ self.proj_out = nn.ModuleList([
+ operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
+ for _ in range(num_layers)])
+
+ head_dim = dim // num_heads
+ self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
+
+ def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
+ img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
+ img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
+ img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
+ img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
+ img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
+ freqs = self.rope_embedder(img_ids).movedim(1, 2)
+ return freqs
+
+ def process_input(self, control_input, render_mask=None, camera_embedding=None):
+ # render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
+ hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
+ t_len, h_len, w_len = hidden.shape[2:]
+ freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
+ hidden = hidden.flatten(2).transpose(1, 2)
+
+ add_inputs = None
+ if camera_embedding is not None and render_mask is not None:
+ add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
+ elif render_mask is not None:
+ add_inputs = render_mask
+
+ if add_inputs is not None:
+ hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
+
+ hidden = self.proj_in(hidden)
+ return hidden, freqs
+
+ def forward_block(self, block_index, hidden, temb, freqs):
+ hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
+ residual = self.proj_out[block_index](hidden)
+ return hidden, residual
diff --git a/comfy/logging.py b/comfy/logging.py
new file mode 100644
index 000000000..cc785296d
--- /dev/null
+++ b/comfy/logging.py
@@ -0,0 +1,10 @@
+import logging
+
+
+DETAIL = 15
+logging.addLevelName(DETAIL, "DETAIL")
+
+
+def detail(message, *args, **kwargs):
+ kwargs.setdefault("stacklevel", 2)
+ logging.log(DETAIL, message, *args, **kwargs)
diff --git a/comfy/model_base.py b/comfy/model_base.py
index 1b9fa7132..393eaf1c2 100644
--- a/comfy/model_base.py
+++ b/comfy/model_base.py
@@ -58,6 +58,7 @@ import comfy.ldm.omnigen.omnigen2
import comfy.ldm.seedvr.model
import comfy.ldm.boogu.model
import comfy.ldm.qwen_image.model
+import comfy.ldm.mage_flow.model
import comfy.ldm.joyimage.model
import comfy.ldm.ideogram4.model
import comfy.ldm.krea2.model
@@ -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)
diff --git a/comfy/model_detection.py b/comfy/model_detection.py
index 7dda87880..f670ff6f1 100644
--- a/comfy/model_detection.py
+++ b/comfy/model_detection.py
@@ -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"
diff --git a/comfy/model_management.py b/comfy/model_management.py
index 222005b6f..f7351224d 100644
--- a/comfy/model_management.py
+++ b/comfy/model_management.py
@@ -34,6 +34,7 @@ import comfy.utils
import comfy.quant_ops
import comfy_aimdo.host_buffer
import comfy_aimdo.vram_buffer
+from comfy.logging import detail
from typing import TYPE_CHECKING
if TYPE_CHECKING:
@@ -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()
diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py
index d70b42bf8..e44322e72 100644
--- a/comfy/model_patcher.py
+++ b/comfy/model_patcher.py
@@ -22,6 +22,7 @@ import collections
import inspect
import logging
import math
+import time
import uuid
from typing import Callable, Optional
@@ -37,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
diff --git a/comfy/ops.py b/comfy/ops.py
index 13c2604fb..9d692dcc7 100644
--- a/comfy/ops.py
+++ b/comfy/ops.py
@@ -41,7 +41,7 @@ def scaled_dot_product_attention(q, k, v, *args, **kwargs):
try:
- if torch.cuda.is_available() and comfy.model_management.WINDOWS:
+ if torch.cuda.is_available():
from torch.nn.attention import SDPBackend, sdpa_kernel
import inspect
if "set_priority" in inspect.signature(sdpa_kernel).parameters:
@@ -51,7 +51,10 @@ try:
SDPBackend.MATH,
]
- SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
+ if comfy.model_management.WINDOWS:
+ SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION)
+ else:
+ SDPA_BACKEND_PRIORITY.insert(1, SDPBackend.CUDNN_ATTENTION)
def scaled_dot_product_attention(q, k, v, *args, **kwargs):
if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower
@@ -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)
diff --git a/comfy/pinned_memory.py b/comfy/pinned_memory.py
index cb77c517a..d78ab3c76 100644
--- a/comfy/pinned_memory.py
+++ b/comfy/pinned_memory.py
@@ -9,14 +9,14 @@ import torch
from comfy.cli_args import args
-def _add_to_bucket(module, buckets, size, priority):
+def _add_to_bucket(module, module_pin, buckets, size, priority):
bucket = buckets.setdefault(size, [])
entry = [-priority, 0, module]
entry[1] = id(entry)
bisect.insort(bucket, entry)
- module._pin_balancer_entry = entry
+ module_pin["balancer_entry"] = entry
-def _steal_pin(module, stack, buckets, size, priority):
+def _steal_pin(module, stack, buckets, size, priority, subset):
bucket = buckets.get(size)
if bucket is None:
return False
@@ -31,34 +31,39 @@ def _steal_pin(module, stack, buckets, size, priority):
return False
*_, victim = bucket.pop()
- module._pin = victim._pin
- module._pin_registered = victim._pin_registered
- module._pin_stack_index = victim._pin_stack_index
- stack[module._pin_stack_index] = (module, stack[module._pin_stack_index][1])
+ module_pin = module._pins[subset]
+ victim_pin = victim._pins[subset]
+ module_pin["pin"] = victim_pin["pin"]
+ module_pin["registered"] = victim_pin["registered"]
+ module_pin["stack_index"] = victim_pin["stack_index"]
+ stack_index = module_pin["stack_index"]
+ stack[stack_index] = (module, stack[stack_index][1])
- victim._pin_registered = False
- del victim._pin
- del victim._pin_stack_index
- del victim._pin_balancer_entry
+ victim_pin["registered"] = False
+ del victim_pin["pin"]
+ del victim_pin["stack_index"]
+ del victim_pin["balancer_entry"]
- _add_to_bucket(module, buckets, size, priority)
+ _add_to_bucket(module, module_pin, buckets, size, priority)
return True
def get_pin(module, subset="weights"):
- pin = getattr(module, "_pin", None)
- if pin is None or module._pin_registered or args.disable_pinned_memory:
+ pins = module.__dict__.get("_pins")
+ module_pin = None if pins is None else pins.get(subset)
+ pin = None if module_pin is None else module_pin.get("pin")
+ if pin is None or module_pin["registered"] or args.disable_pinned_memory:
return pin
_, _, stack_split, pinned_size, *_ = module._pin_state[subset]
size = pin.nbytes
- comfy.model_management.ensure_pin_registerable(size)
+ comfy.model_management.ensure_pin_registerable(size, loaded=subset.endswith("-loaded"))
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
return pin
- module._pin_registered = True
- stack_split[0] = max(stack_split[0], module._pin_stack_index)
+ module_pin["registered"] = True
+ stack_split[0] = max(stack_split[0], module_pin["stack_index"])
comfy.model_management.TOTAL_PINNED_MEMORY += size
pinned_size[0] += size
return pin
@@ -72,23 +77,26 @@ def pin_memory(module, subset="weights", size=None):
if pin is not None:
return
+ pins = module.__dict__.setdefault("_pins", {})
+ module_pin = pins.setdefault(subset, {})
hostbuf, stack, stack_split, pinned_size, counter, buckets = pin_state[subset]
if size is None:
size = comfy.memory_management.vram_aligned_size([ module.weight, module.bias ])
- offset = hostbuf.size
registerable_size = size
- priority = getattr(module, "_pin_balancer_priority", None)
+ loaded = subset.endswith("-loaded")
+ priority = module_pin.get("balancer_priority")
if priority is None:
priority = comfy.utils.bit_reverse_range(counter[0], 16)
counter[0] += 1
- module._pin_balancer_priority = priority
+ module_pin["balancer_priority"] = priority
comfy.memory_management.extra_ram_release(comfy.memory_management.RAM_CACHE_HEADROOM)
- if (not comfy.model_management.ensure_pin_budget(size) or
- not comfy.model_management.ensure_pin_registerable(registerable_size)):
- return _steal_pin(module, stack, buckets, size, priority)
+ if (not comfy.model_management.ensure_pin_budget(size, loaded=loaded) or
+ not comfy.model_management.ensure_pin_registerable(registerable_size, loaded=loaded)):
+ return _steal_pin(module, stack, buckets, size, priority, subset)
+ offset = hostbuf.size
extended = False
try:
hostbuf.extend(size=size, register=False)
@@ -97,23 +105,23 @@ def pin_memory(module, subset="weights", size=None):
pin.untyped_storage()._comfy_hostbuf = hostbuf
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
- comfy.model_management.free_registrations(size)
+ comfy.model_management.free_registrations(size, loaded=loaded)
if torch.cuda.cudart().cudaHostRegister(pin.data_ptr(), size, 1) != 0:
comfy.model_management.discard_cuda_async_error()
del pin
hostbuf.truncate(offset, do_unregister=False)
- return _steal_pin(module, stack, buckets, size, priority)
+ return _steal_pin(module, stack, buckets, size, priority, subset)
except RuntimeError:
if extended:
hostbuf.truncate(offset, do_unregister=False)
- return _steal_pin(module, stack, buckets, size, priority)
+ return _steal_pin(module, stack, buckets, size, priority, subset)
- module._pin = pin
+ module_pin["pin"] = pin
stack.append((module, offset))
- module._pin_registered = True
- module._pin_stack_index = len(stack) - 1
- stack_split[0] = max(stack_split[0], module._pin_stack_index)
+ module_pin["registered"] = True
+ module_pin["stack_index"] = len(stack) - 1
+ stack_split[0] = max(stack_split[0], module_pin["stack_index"])
comfy.model_management.TOTAL_PINNED_MEMORY += size
pinned_size[0] += size
- _add_to_bucket(module, buckets, size, priority)
+ _add_to_bucket(module, module_pin, buckets, size, priority)
return True
diff --git a/comfy/samplers.py b/comfy/samplers.py
index 25c5a855f..9f571ece9 100755
--- a/comfy/samplers.py
+++ b/comfy/samplers.py
@@ -20,6 +20,7 @@ import comfy.hooks
import comfy.context_windows
import comfy.multigpu
import comfy.utils
+from comfy.logging import detail
import scipy.stats
import numpy
@@ -991,10 +992,15 @@ class KSAMPLER(Sampler):
noise = model_wrap.inner_model.model_sampling.noise_scaling(sigmas[0], noise, latent_image, self.max_denoise(model_wrap, sigmas))
- k_callback = None
total_steps = len(sigmas) - 1
- if callback is not None:
- k_callback = lambda x: callback(x["i"], x["denoised"], x["x"], total_steps)
+ first_step = True
+ def k_callback(x):
+ nonlocal first_step
+ if first_step:
+ detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, self.sampler_function.__name__, x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape))
+ first_step = False
+ if callback is not None:
+ callback(x["i"], x["denoised"], x["x"], total_steps)
samples = self.sampler_function(model_k, noise, sigmas, extra_args=extra_args, callback=k_callback, disable=disable_pbar, **self.extra_options)
samples = model_wrap.inner_model.model_sampling.inverse_noise_scaling(sigmas[-1], samples)
@@ -1270,10 +1276,13 @@ class CFGGuider:
return latent_image
if latent_image.is_nested:
+ sampler_shapes = [tuple(x.shape) for x in latent_image.unbind()]
latent_image, latent_shapes = comfy.utils.pack_latents(latent_image.unbind())
noise, _ = comfy.utils.pack_latents(noise.unbind())
else:
latent_shapes = [latent_image.shape]
+ sampler_shapes = [tuple(latent_image.shape)]
+ detail("Sampler: model=%s latent_shapes=%s", self.model_patcher.model.__class__.__name__, sampler_shapes)
if denoise_mask is not None:
if denoise_mask.is_nested:
diff --git a/comfy/sd.py b/comfy/sd.py
index 7b7b348b0..f265507f5 100644
--- a/comfy/sd.py
+++ b/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))
diff --git a/comfy/supported_models.py b/comfy/supported_models.py
index 2381014fd..d423ea6e6 100644
--- a/comfy/supported_models.py
+++ b/comfy/supported_models.py
@@ -27,6 +27,7 @@ import comfy.text_encoders.z_image
import comfy.text_encoders.ideogram4
import comfy.text_encoders.boogu
import comfy.text_encoders.krea2
+import comfy.text_encoders.mage_flow
import comfy.text_encoders.joyimage
import comfy.text_encoders.anima
import comfy.text_encoders.ace15
@@ -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,
diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py
index 0bba8341b..5163c1676 100644
--- a/comfy/text_encoders/gemma4.py
+++ b/comfy/text_encoders/gemma4.py
@@ -1,11 +1,15 @@
import torch
import torch.nn as nn
+import torchaudio.functional as AF
+import torchvision.transforms.functional as TVF
import numpy as np
+from tokenizers import Tokenizer
from dataclasses import dataclass
import math
from comfy import sd1_clip
import comfy.model_management
+import comfy.ops
from comfy.ldm.modules.attention import optimized_attention_for_device
from comfy.rmsnorm import rms_norm
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
@@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
+# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
+GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
+GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
+
@dataclass
class Gemma4Config:
vocab_size: int = 262144
@@ -35,6 +43,9 @@ class Gemma4Config:
transformer_type: str = "gemma4"
head_dim = 256
global_head_dim = 512
+ num_global_key_value_heads = None
+ attention_k_eq_v = False
+ vision_bidirectional = False
rms_norm_add = False
mlp_activation = "gelu_pytorch_tanh"
qkv_bias = False
@@ -51,6 +62,7 @@ class Gemma4Config:
num_kv_shared_layers: int = 18
use_double_wide_mlp: bool = False
stop_tokens = [1, 50, 106]
+ suppress_tokens = []
vision_config = GEMMA4_VISION_CONFIG
audio_config = GEMMA4_AUDIO_CONFIG
mm_tokens_per_image = 280
@@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config):
num_hidden_layers: int = 60
num_attention_heads: int = 32
num_key_value_heads: int = 16
+ vision_bidirectional = True
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
hidden_size_per_layer_input: int = 0
num_kv_shared_layers: int = 0
audio_config = None
vision_config = GEMMA4_VISION_31B_CONFIG
+@dataclass
+class Gemma4_12B_Config(Gemma4Config):
+ hidden_size: int = 3840
+ intermediate_size: int = 15360
+ num_hidden_layers: int = 48
+ num_attention_heads: int = 16
+ num_key_value_heads: int = 8
+ num_global_key_value_heads = 1
+ attention_k_eq_v = True
+ vision_bidirectional = True
+ sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
+ hidden_size_per_layer_input: int = 0
+ num_kv_shared_layers: int = 0
+ audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
+ vision_config = GEMMA4_UNIFIED_VISION_CONFIG
+ suppress_tokens = [258883, 258882]
+
# unfused RoPE as addcmul_ RoPE diverges from reference code
def _apply_rotary_pos_emb(x, freqs_cis):
@@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis):
return out
class Gemma4Attention(nn.Module):
- def __init__(self, config, head_dim, device=None, dtype=None, ops=None):
+ def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
super().__init__()
self.num_heads = config.num_attention_heads
- self.num_kv_heads = config.num_key_value_heads
+ self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
self.hidden_size = config.hidden_size
self.head_dim = head_dim
self.inner_size = self.num_heads * head_dim
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
- self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
+ # k_eq_v: V reuses the K projection (no separate v_proj weight)
+ self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
self.q_norm = None
@@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module):
shareable_kv = None
else:
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
- xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
+ if self.v_proj is not None:
+ xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
+ else:
+ xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
if self.k_norm is not None:
xk = self.k_norm(xk)
xv = rms_norm(xv)
@@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module):
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
- self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops)
+ # k_eq_v only on global layers, which then use num_global_key_value_heads
+ k_eq_v = config.attention_k_eq_v and not self.sliding_attention
+ num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
+ self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
num_kv_shared = config.num_kv_shared_layers
first_kv_shared = config.num_hidden_layers - num_kv_shared
@@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module):
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
- self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
- else:
- self.layer_scalar = None
+
+ # layer_scalar exists on every gemma4 variant, independent of per-layer input
+ self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
sliding_window = None
@@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module):
x = self.post_per_layer_input_norm(x)
x = residual + x
- if self.layer_scalar is not None:
- x = x * self.layer_scalar
+ x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
return x, present_key_value, shareable_kv
@@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module):
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
mask = mask + causal_mask if mask is not None else causal_mask
+ # Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
+ if self.config.vision_bidirectional and past_len == 0 and embeds_info:
+ block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
+ group = 0
+ for info in embeds_info:
+ if info.get("type") == "image":
+ start = info["index"]
+ block_ids[start:start + info["size"]] = group
+ group += 1
+ if group > 0:
+ same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
+ mask = mask.masked_fill(same_block, 0.0)
+
# Per-layer inputs
per_layer_inputs = None
if self.hidden_size_per_layer_input:
@@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module):
shared_global_kv = None # KV from last non-shared global layer
intermediate = None
+ all_intermediate = None
+ only_layers = None
+ if intermediate_output is not None:
+ if isinstance(intermediate_output, list):
+ all_intermediate = []
+ only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
+ elif intermediate_output == "all":
+ all_intermediate = []
+ intermediate_output = None
+ elif intermediate_output < 0:
+ intermediate_output = len(self.layers) + intermediate_output
+
next_key_values = []
for i, layer in enumerate(self.layers):
+ if all_intermediate is not None:
+ if only_layers is None or (i in only_layers):
+ all_intermediate.append(x.unsqueeze(1).clone())
+
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
layer_kwargs = {}
@@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module):
if self.norm is not None:
x = self.norm(x)
- if len(next_key_values) > 0:
+ if all_intermediate is not None:
+ if only_layers is None or (len(self.layers) in only_layers):
+ all_intermediate.append(x.unsqueeze(1).clone())
+ if len(all_intermediate) > 0:
+ intermediate = torch.cat(all_intermediate, dim=1)
+
+ if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
+ intermediate = self.norm(intermediate)
+
+ # Only hand back the KV cache when caching was actually requested; SDClipModel reads
+ # outputs[2] as the pooled output.
+ if past_key_values is not None and len(next_key_values) > 0:
return x, intermediate, next_key_values
return x, intermediate
@@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
cap = self.model.config.final_logit_softcapping
if cap:
logits = cap * torch.tanh(logits / cap)
+ if self.model.config.suppress_tokens:
+ logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
return logits
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
@@ -441,6 +519,28 @@ class Gemma4AudioMixin:
return None, None
+class Gemma4UnifiedBase(Gemma4Base):
+ """Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
+ def _init_model(self, config, dtype, device, operations):
+ self.num_layers = config.num_hidden_layers
+ self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
+ self.dtype = dtype
+ self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
+ self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
+ self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
+
+ def preprocess_embed(self, embed, device):
+ if embed["type"] == "image":
+ pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
+ patches, positions = self.vision_model.patchify(pixels)
+ vision_out = self.vision_model(patches, positions)
+ return self.multi_modal_projector(vision_out), None
+ if embed["type"] == "audio":
+ audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
+ return self.audio_projector(audio), None
+ return None, None
+
+
# Vision Encoder
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
@@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
+# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
+
+def _patches_merge(patches, positions_xy, length):
+ patch_size = math.isqrt(patches.shape[-1] // 3)
+ k = math.isqrt(patches.shape[-2] // length)
+ batch = patches.shape[:-2]
+
+ max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
+ kidx = torch.div(positions_xy, k, rounding_mode="floor")
+ rem = torch.remainder(positions_xy, k)
+ order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
+ perm = order.long().argsort(dim=-1)
+
+ merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
+ merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
+ merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
+
+ pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
+ pad = (positions_xy == -1).all(dim=-1, keepdim=True)
+ pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
+ pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
+ return merged, pos
+
+
+class Gemma4UnifiedVisionEmbedder(nn.Module):
+ """Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
+ def __init__(self, config, device=None, dtype=None, ops=None):
+ super().__init__()
+ self.patch_size = config["patch_size"]
+ self.pooling_kernel_size = config["pooling_kernel_size"]
+ patch_dim = config["model_patch_size"] ** 2 * 3
+ mm_embed_dim = config["mm_embed_dim"]
+ self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
+ self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
+ self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
+ self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
+ self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
+
+ def patchify(self, pixels):
+ """pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
+ ps, k = self.patch_size, self.pooling_kernel_size
+ out_patches, out_positions = [], []
+ for img in pixels:
+ ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
+ teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
+ grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
+ tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
+ n_model = teacher.shape[0] // (k * k)
+ mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
+ out_patches.append(mp.squeeze(0))
+ out_positions.append(mpos.squeeze(0))
+ return torch.stack(out_patches), torch.stack(out_positions)
+
+ def forward(self, pixel_values, image_position_ids):
+ x = self.patch_ln1(pixel_values)
+ x = self.patch_dense(x)
+ x = self.patch_ln2(x)
+
+ clamped = image_position_ids.clamp(min=0).long()
+ valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
+ axes = torch.arange(2, device=image_position_ids.device)
+ pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
+ pos_embs = (pos[clamped, axes] * valid).sum(-2)
+ x = x + pos_embs
+ return self.pos_norm(x)
+
+
# Audio Encoder
class Gemma4AudioConvSubsampler(nn.Module):
@@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
# Tokenizer and Wrappers
+def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
+ target_px = max_patches * patch_size ** 2
+ factor = math.sqrt(target_px / (height * width))
+ side_mult = pooling_kernel_size * patch_size
+ target_height = math.floor(factor * height / side_mult) * side_mult
+ target_width = math.floor(factor * width / side_mult) * side_mult
+
+ if target_height == 0 and target_width == 0:
+ raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
+
+ max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
+ if target_height == 0:
+ target_height = side_mult
+ target_width = min(math.floor(width / height) * side_mult, max_side_length)
+ elif target_width == 0:
+ target_width = side_mult
+ target_height = min(math.floor(height / width) * side_mult, max_side_length)
+
+ if target_height * target_width > target_px:
+ raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
+
+ return target_height, target_width
+
+
class Gemma4_Tokenizer():
tokenizer_json_data = None
@@ -998,25 +1189,35 @@ class Gemma4_Tokenizer():
return {"tokenizer_json": self.tokenizer_json_data}
return {}
- def _extract_mel_spectrogram(self, waveform, sample_rate):
- """Extract 128-bin log mel spectrogram.
- Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
- """
- # Mix to mono first, then resample to 16kHz
+ def _audio_token_count(self, num_samples):
+ # Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
+ _fl = 320 # int(round(16000 * 20.0 / 1000.0))
+ _hl = 160 # int(round(16000 * 10.0 / 1000.0))
+ _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
+ _t = _nmel
+ for _ in range(2):
+ _t = (_t + 2 - 3) // 2 + 1
+ return min(_t, 750)
+
+ @staticmethod
+ def _resample_16k(waveform, sample_rate):
+ """Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
+ load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
if waveform.dim() > 1 and waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
- audio = waveform.squeeze(0).float().numpy()
+ audio = waveform.float()
if sample_rate != 16000:
- # Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match)
- from scipy.signal import resample_poly, firwin
- from math import gcd
- g = gcd(sample_rate, 16000)
- up, down = 16000 // g, sample_rate // g
- L = max(up, down)
- h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5))
- audio = resample_poly(audio, up, down, window=h).astype(np.float32)
+ audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
+ lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
+ return audio.squeeze(0).contiguous()
+
+ def _extract_audio_features(self, waveform, sample_rate):
+ """Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
+ Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
+ """
+ audio = self._resample_16k(waveform, sample_rate).numpy()
n = len(audio)
# Pad to multiple of 128, build sample-level mask
@@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer():
if audio is not None:
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
- mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate)
- audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T])
+ feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
+ audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
# Process image/video frames
is_video = video is not None
@@ -1090,13 +1291,8 @@ class Gemma4_Tokenizer():
pooling_k = 3
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
max_patches = max_soft_tokens * pooling_k * pooling_k
- target_px = max_patches * patch_size * patch_size
- factor = (target_px / (h * w)) ** 0.5
- side_mult = pooling_k * patch_size
- target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
- target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
+ target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
- import torchvision.transforms.functional as TVF
for i in range(num_frames):
# rescaling to match reference code
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
@@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer():
llama_text = llama_template.format(text)
else:
# Build template from modalities present
- system = "<|turn>system\n<|think|>\n" if thinking else ""
+ system = "<|turn>system\n<|think|>\n\n" if thinking else ""
media = ""
if len(images) > 0:
if is_video:
@@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer():
if len(audio_features) > 0:
# Compute audio token count (always at 16kHz)
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
- _fl = 320 # int(round(16000 * 20.0 / 1000.0))
- _hl = 160 # int(round(16000 * 10.0 / 1000.0))
- _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
- _t = _nmel
- for _ in range(2):
- _t = (_t + 2 - 3) // 2 + 1
- n_audio_tokens = min(_t, 750)
+ n_audio_tokens = self._audio_token_count(num_samples)
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "