24 Commits
Author SHA1 Message Date
Acly 0abe742480 Version 3.4.0 2026-10-04 17:59:27 +09:00
Acly b3ae4aa2d9 Fix missing maskb batch dim; Make ApplyMaskToImage respect existing alpha 2026-09-26 18:44:00 +09:00
fukc-gihtub 1b8d81ce5a Add QwenImage21 model type 2026-09-23 11:33:15 +09:00
Sen-sou ca01116495 fix: add compatibility with ComfyUI INT8 models 2026-08-19 10:27:17 +02:00
Acly 5d3194f4d4 Version 3.3.0 2026-06-28 11:22:20 +09:00
fukc-gihtub d5812f900b Add Krea2 model type 2026-06-27 13:22:23 +09:00
Mutive 7064288fbe Add Anima region attention support (#67)
* Add Anima region attention support

* Added adaptation acknowledgement for Anima Attention couple implementation

* Address Anima attention cleanup feedback
2026-06-27 13:21:43 +09:00
Acly d82675092e Version 3.2.0 2026-05-31 11:15:36 +02:00
Acly a1e51904de Fix some type errors 2026-05-30 13:50:52 +02:00
Acly ffa130239b Add Anima Control LLLite nodes
* from kohya-ss/ComfyUI-Anima-LLLite
* split nodes into load/apply to cache model loading
2026-05-30 13:50:34 +02:00
Acly 2fd51d0d47 Workaround for import failure in transformers with certain pytorch installs #66
* seems to affects pytorch compiled with DISTRIBUTED=0, eg. Windows ROCm
* should probably be fixed in transformers somehow?
2026-05-30 11:49:19 +02:00
Acly d3c75155b4 Fix already cached image PUT requests (#65)
* Support already cached image PUT requests without aborting the connection
* conditionally send 100 Continue instead
2026-05-11 09:14:53 +02:00
Acly cbaef8d9c5 Version 3.1.4, fix type checks 2026-05-03 18:15:02 +02:00
FeepingCreature b2783d82a6 Add Anima model type 2026-05-01 17:53:32 +02:00
VERIGEN 09759222de Add ERNIE Image model detection 2026-05-01 17:53:02 +02:00
Acly 7fc3df1174 Version 3.1.3 2026-02-21 20:33:28 +01:00
Acly ed99942f86 Add mask output to KritaCanvas 2026-02-20 13:08:25 +01:00
Acly 2d395424ea Fix nsfw filter with transformers>5 2026-02-05 16:28:32 +01:00
Acly 9b9ea62dd8 Version 3.1.2 2026-01-31 18:01:29 +01:00
Acly ad36f89af3 Tiles: make multiple for tile layout configurable
* can now ensure eg. multiple of 16 tiles to be compatible with flux2 latent downsample factor
* default is 8, which matches previous hardcoded value
2026-01-26 17:28:46 +01:00
Alex 7130dcb2df Add KritaStyleAndPrompt node for synced prompts across workspaces
New node ETN_KritaStyleAndPrompt that works like KritaStyle but:
- Prompts and style sync between Generate/Live/Animation/Graph workspaces
- Outputs fully prepared prompts (wildcards evaluated, style merged)
- Model output includes extracted LoRAs from prompts
2026-01-24 13:15:12 +01:00
Acly 77186eda87 Model inspection: detect Flux 2 klein GGUF variants 2026-01-20 17:11:14 +01:00
Acly 24a7bd1a77 Version 3.1.1 2026-01-18 21:02:30 +01:00
Acly 2d14a03ad8 Model inspection: detect variants of Flux 2 (Klein-4B, Klein-9B) 2026-01-16 19:28:58 +01:00
10 changed files with 1214 additions and 43 deletions
+6
View File
@@ -293,3 +293,9 @@ git clone https://github.com/Acly/comfyui-tooling-nodes.git
``` ```
Restart ComfyUI and the nodes are functional. Restart ComfyUI and the nodes are functional.
## Acknowledgements
* Region nodes adapted from [laksjdjf/cgem156-ComfyUI](https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py)
* Control nodes adapted from [kohya-ss/ComfyUI-Anima-LLLite](https://github.com/kohya-ss/ComfyUI-Anima-LLLite)
+18 -3
View File
@@ -1,10 +1,12 @@
from comfy_api.latest import ComfyExtension, io from comfy_api.latest import ComfyExtension, io
from . import api as api, nodes, tile, region, nsfw, translation, krita
from . import api as api
from . import control, krita, nodes, region, tile, translation
class ExternalToolingNodes(ComfyExtension): class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]: async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [ node_list = [
nodes.LoadImageCache, nodes.LoadImageCache,
nodes.SaveImageCache, nodes.SaveImageCache,
nodes.LoadImageBase64, nodes.LoadImageBase64,
@@ -22,7 +24,6 @@ class ExternalToolingNodes(ComfyExtension):
region.DefineRegion, region.DefineRegion,
region.ListRegionMasks, region.ListRegionMasks,
region.AttentionMask, region.AttentionMask,
nsfw.NSFWFilter,
translation.Translate, translation.Translate,
krita.KritaOutput, krita.KritaOutput,
krita.KritaSendText, krita.KritaSendText,
@@ -32,7 +33,21 @@ class ExternalToolingNodes(ComfyExtension):
krita.KritaMaskLayer, krita.KritaMaskLayer,
krita.Parameter, krita.Parameter,
krita.KritaStyle, krita.KritaStyle,
krita.KritaStyleAndPrompt,
control.ControlApply,
control.ControlLoad,
] ]
try: # see #66
from . import nsfw
node_list.append(nsfw.NSFWFilter)
except (ImportError, ModuleNotFoundError):
import traceback
print("[comfyui-tooling-nodes] WARNING: Could not import all nodes.")
traceback.print_exc()
return node_list
async def comfy_entrypoint(): async def comfy_entrypoint():
+36 -2
View File
@@ -56,7 +56,11 @@ model_names = {
"ACEStep": "ace-step", "ACEStep": "ace-step",
"Omnigen2": "omnigen2", "Omnigen2": "omnigen2",
"QwenImage": "qwen-image", "QwenImage": "qwen-image",
"QwenImage21": "qwen-image21",
"ErnieImage": "ernie-image",
"Flux2": "flux2", "Flux2": "flux2",
"Anima": "anima",
"Krea2": "krea2",
} }
gguf_architectures = { gguf_architectures = {
@@ -121,6 +125,9 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
raw_name = base_model.__class__.__name__ raw_name = base_model.__class__.__name__
if raw_name == "SDXL": if raw_name == "SDXL":
model_type = base_model.model_type(cfg).name.lower().replace("_", "-") model_type = base_model.model_type(cfg).name.lower().replace("_", "-")
if raw_name == "Flux2":
hidden_size = unet_config.get("hidden_size", 0)
model_type = {3072: "klein-4b", 4096: "klein-9b"}.get(hidden_size, "dev")
if not raw_name: if not raw_name:
return {"base_model": "unknown"} return {"base_model": "unknown"}
@@ -190,7 +197,6 @@ def inspect_gguf(filename: str, model_type: str):
else: # stable-diffusion.cpp, requires conversion. not handled for now else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False} return {"base_model": "flux", "is_inpaint": False}
# Detect Chroma (modified Flux)
if arch_str == "flux" and any( if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer") t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5) for t in itertools.islice(reader.tensors, 5)
@@ -204,10 +210,27 @@ def inspect_gguf(filename: str, model_type: str):
arch_str = "z-image" arch_str = "z-image"
break break
# Detect Flux variants
result_type = None
if arch_str == "flux":
for t in reader.tensors:
if t.name.startswith("distilled_guidance_layer"):
arch_str = "chroma"
break
elif t.name == "double_stream_modulation_img.lin.weight":
arch_str = "flux2"
if t.shape[0] == 3072:
result_type = "klein-4b"
elif t.shape[0] == 4096:
result_type = "klein-9b"
break
result = { result = {
"base_model": gguf_architectures.get(arch_str, arch_str), "base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False, "is_inpaint": False,
} }
if result_type is not None:
result["type"] = result_type
try: try:
if file_type := reader.get_field("general.file_type"): if file_type := reader.get_field("general.file_type"):
result["quant"] = file_type.contents().lower() result["quant"] = file_type.contents().lower()
@@ -327,11 +350,11 @@ if _server is not None:
except Exception as e: except Exception as e:
return web.json_response(dict(error=str(e)), status=500) return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/image/{id}")
async def put_image(request: web.Request): async def put_image(request: web.Request):
try: try:
id = request.match_info.get("id", "") id = request.match_info.get("id", "")
if id in image_cache: if id in image_cache:
await request.release() # Consume and discard the data to avoid connection abort
return web.json_response(dict(status="cached"), status=200) return web.json_response(dict(status="cached"), status=200)
content_type = request.headers.get("Content-Type", "application/octet-stream") content_type = request.headers.get("Content-Type", "application/octet-stream")
@@ -344,6 +367,17 @@ if _server is not None:
except Exception as e: except Exception as e:
return web.json_response(dict(error=str(e)), status=500) return web.json_response(dict(error=str(e)), status=500)
async def _put_image_expect_handler(request: web.Request):
if request.match_info.get("id", "") in image_cache:
# Skip "100 Continue" since we don't need the data, return 200 immediately.
return web.json_response(dict(status="cached"), status=200)
# otherwise run default aiohttp handler
return None
_server.app.router.add_route(
"PUT", "/api/etn/image/{id}", put_image, expect_handler=_put_image_expect_handler
)
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}") @_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request): async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "") folder_name = request.match_info.get("folder_name", "")
+859
View File
@@ -0,0 +1,859 @@
"""ControlNet-LLLite for Anima (DiT) — ComfyUI port (v2 architecture).
Adapted from kohya-ss/ComfyUI-Anima-LLLite
https://github.com/kohya-ss/ComfyUI-Anima-LLLite
Apache-2.0 license
Adapted from kohya-ss/sd-scripts. The on-disk weight format is the v2
named-key format (per-module key prefix = lllite_name, shared encoder under
``lllite_conditioning1.*``, depth embedding split per-module as
``{name}.depth_embed``); legacy ``lllite_modules.*`` files are rejected.
Differences vs. the sd-scripts reference (``networks/control_net_lllite_anima.py``):
* No dependency on ``library.utils`` — uses stdlib logging.
* Module discovery filters the LLM-Adapter sub-tree by class identity in
addition to the path-based check (ComfyUI ships two distinct ``Attention``
classes that share the bare class name).
* ``LLLiteModuleDiT`` keeps a ``restore()`` method (and an idempotent
``apply_to()``); ComfyUI patches/unpatches the original Linear around
every sampler call via ``set_model_unet_function_wrapper``.
* Forward pass casts ``x`` and ``cond_emb`` to the LLLite parameter dtype
so autocast / mixed-precision flows that hand us a different dtype than
the LLLite weights still work.
* CFG batch-size and sequence-length mismatches fall back to identity
instead of asserting, so a slightly-off cond image cannot abort sampling.
* The training-side ``AnimaControlNetLLLiteWrapper`` is omitted; ComfyUI
integrates via ``model_function_wrapper`` in nodes.py instead.
"""
from __future__ import annotations
from copy import copy
import logging
import os
from dataclasses import dataclass
from typing import Any
import folder_paths
import safetensors
import safetensors.torch
import torch
import torch.nn.functional as F
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
from torch import nn
logger = logging.getLogger("comfyui-tooling-nodes")
# Class names of the modules that LLLite injects into. The LLM-Adapter uses
# a different ``Attention`` class with the same bare name; we filter it by
# path (``llm_adapter`` in the qualified name) and by the ``is_selfattn``
# attribute presence.
TARGET_ATTENTION_CLASS = "Attention"
TARGET_MLP_CLASS = "GPT2FeedForward"
LLM_ADAPTER_NAME = "llm_adapter"
LLLITE_ARCH_VERSION = "2"
# ----------------------------------------------------------------------------
# target_layers: atomic specifiers and presets
# ----------------------------------------------------------------------------
ATOMIC_SPECIFIERS: tuple[str, ...] = (
"self_attn_q_pre",
"self_attn_kv_pre",
"cross_attn_q_pre",
"mlp_fc1_pre",
)
PRESETS: dict = {
"self_attn_q": ("self_attn_q_pre",),
"self_attn_qkv": ("self_attn_q_pre", "self_attn_kv_pre"),
"self_attn_qkv_cross_q": ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre"),
}
def parse_target_layers(spec: str) -> tuple[str, ...]:
"""Resolve a ``target_layers`` spec to a canonical atomic tuple.
Accepts a preset name (``"self_attn_qkv"``) or a comma-separated list of
atomic specifiers (``"self_attn_q_pre,mlp_fc1_pre"``). Returns the atomics
in ``ATOMIC_SPECIFIERS`` order with duplicates removed.
"""
if not isinstance(spec, str):
raise TypeError(f"target_layers must be str, got {type(spec).__name__}")
spec = spec.strip()
if not spec:
raise ValueError("target_layers spec is empty")
if spec in PRESETS:
parts = list(PRESETS[spec])
else:
parts = [p.strip() for p in spec.split(",") if p.strip()]
bad = [p for p in parts if p not in ATOMIC_SPECIFIERS]
if bad:
raise ValueError(
f"unknown target_layers atomic specifier(s): {bad}. "
f"valid atomic={list(ATOMIC_SPECIFIERS)}, presets={list(PRESETS)}"
)
return tuple(a for a in ATOMIC_SPECIFIERS if a in parts)
# ----------------------------------------------------------------------------
# Conditioning1 trunk (v2)
# ----------------------------------------------------------------------------
def _gn(channels: int) -> nn.GroupNorm:
g = 8
while g > 1 and channels % g != 0:
g //= 2
return nn.GroupNorm(g, channels)
class _ResBlock(nn.Module):
def __init__(self, ch: int):
super().__init__()
self.norm1 = _gn(ch)
self.conv1 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
self.norm2 = _gn(ch)
self.conv2 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = self.conv1(F.silu(self.norm1(x)))
h = self.conv2(F.silu(self.norm2(h)))
return x + h
ASPP_DEFAULT_DILATIONS: tuple[int, ...] = (1, 2, 4, 8)
class _ASPP(nn.Module):
def __init__(self, ch: int, dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS):
super().__init__()
assert len(dilations) >= 1, "ASPP needs at least one dilation"
branches = []
for d in dilations:
if d == 1:
conv = nn.Conv2d(ch, ch, kernel_size=1)
else:
conv = nn.Conv2d(ch, ch, kernel_size=3, padding=d, dilation=d)
branches.append(nn.Sequential(conv, _gn(ch), nn.SiLU()))
self.branches = nn.ModuleList(branches)
self.global_pool = nn.AdaptiveAvgPool2d(1)
self.global_conv = nn.Sequential(nn.Conv2d(ch, ch, kernel_size=1), _gn(ch), nn.SiLU())
n_branches = len(dilations) + 1
self.proj = nn.Sequential(nn.Conv2d(ch * n_branches, ch, kernel_size=1), _gn(ch), nn.SiLU())
def forward(self, x: torch.Tensor) -> torch.Tensor:
h, w = x.shape[-2:]
outs = [b(x) for b in self.branches]
g = self.global_conv(self.global_pool(x))
g = F.interpolate(g, size=(h, w), mode="bilinear", align_corners=False)
outs.append(g)
return self.proj(torch.cat(outs, dim=1))
class _Conditioning1(nn.Module):
def __init__(
self,
cond_dim: int,
cond_emb_dim: int,
n_resblocks: int,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
):
super().__init__()
assert cond_dim % 2 == 0, f"cond_dim must be even, got {cond_dim}"
assert cond_in_channels >= 1, f"cond_in_channels must be >= 1, got {cond_in_channels}"
ch_half = cond_dim // 2
self.cond_in_channels = cond_in_channels
self.conv1 = nn.Conv2d(cond_in_channels, ch_half, kernel_size=4, stride=4, padding=0)
self.norm1 = _gn(ch_half)
self.conv2 = nn.Conv2d(ch_half, ch_half, kernel_size=3, stride=1, padding=1)
self.norm2 = _gn(ch_half)
self.conv3 = nn.Conv2d(ch_half, cond_dim, kernel_size=4, stride=4, padding=0)
self.norm3 = _gn(cond_dim)
self.resblocks = nn.ModuleList([_ResBlock(cond_dim) for _ in range(n_resblocks)])
self.aspp = _ASPP(cond_dim, aspp_dilations) if use_aspp else None
self.proj = nn.Conv2d(cond_dim, cond_emb_dim, kernel_size=1)
self.out_norm = nn.LayerNorm(cond_emb_dim)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = F.silu(self.norm1(self.conv1(x)))
h = F.silu(self.norm2(self.conv2(h)))
h = F.silu(self.norm3(self.conv3(h)))
for rb in self.resblocks:
h = rb(h)
if self.aspp is not None:
h = self.aspp(h)
h = self.proj(h)
b, c, hh, ww = h.shape
h = h.view(b, c, hh * ww).permute(0, 2, 1).contiguous()
h = self.out_norm(h)
return h
# ----------------------------------------------------------------------------
# LLLite module (v2: FiLM + SiLU + 5D path + depth embedding)
# ----------------------------------------------------------------------------
class LLLiteModuleDiT(nn.Module):
def __init__(
self,
name: str,
org_module: nn.Linear,
cond_emb_dim: int,
mlp_dim: int,
dropout: float | None = None,
multiplier: float = 1.0,
):
super().__init__()
self.lllite_name = name
# Wrap in a list so the original Linear is not registered as a submodule
# and its weights stay out of state_dict.
self.org_module = [org_module]
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.dropout = dropout
self.multiplier = multiplier
in_dim = org_module.in_features
self.down = nn.Linear(in_dim, mlp_dim)
self.mid = nn.Linear(mlp_dim + cond_emb_dim, mlp_dim)
# FiLM: cond_local -> (gamma, beta), zero-init for identity at start.
self.cond_to_film = nn.Linear(cond_emb_dim, 2 * mlp_dim)
nn.init.zeros_(self.cond_to_film.weight)
nn.init.zeros_(self.cond_to_film.bias)
self.up = nn.Linear(mlp_dim, in_dim)
nn.init.zeros_(self.up.weight)
nn.init.zeros_(self.up.bias)
self.cond_emb: torch.Tensor | None = None
self.org_forward = None
# Set by the parent ControlNetLLLiteDiT after construction.
self.layer_idx: int = -1
self._depth_embeds_ref: list[nn.Parameter] = []
def apply_to(self):
if self.org_forward is None:
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def restore(self):
if self.org_forward is not None:
self.org_module[0].forward = self.org_forward
self.org_forward = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Input layouts:
# self/cross attention q/k/v: (B, S, D) — already flattened in the Anima block
# mlp.layer1: (B, T, H, W, D) — passed un-flattened
# Flatten the 5D case to 3D for the LLLite path and reshape on exit.
if self.multiplier == 0.0 or self.cond_emb is None:
return self.org_forward(x)
orig_shape = x.shape
is_5d = x.dim() == 5
if is_5d:
B, T, H, W, D = orig_shape
x = x.reshape(B, T * H * W, D)
cx = self.cond_emb # (B_c, S, cond_emb_dim)
# Broadcast cond_emb to the runtime batch (CFG cond+uncond, multi-cond).
if x.shape[0] != cx.shape[0]:
if x.shape[0] % cx.shape[0] != 0:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1)
if x.shape[1] != cx.shape[1]:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
# Run the LLLite mini-MLP in its own parameter dtype, then cast the
# correction back to ``x``'s dtype before adding. Robust to autocast
# flows where x and LLLite weights have different dtypes.
param_dtype = self.down.weight.dtype
x_proc = x if x.dtype == param_dtype else x.to(param_dtype)
if cx.dtype != param_dtype or cx.device != x.device:
cx = cx.to(device=x.device, dtype=param_dtype)
# Per-module depth embedding (zero-init so it's a no-op at train start).
if self._depth_embeds_ref:
depth_e = self._depth_embeds_ref[0][self.layer_idx]
if depth_e.dtype != param_dtype or depth_e.device != x.device:
depth_e = depth_e.to(device=x.device, dtype=param_dtype)
cond_local = cx + depth_e
else:
cond_local = cx
h = F.silu(self.down(x_proc))
gb = self.cond_to_film(cond_local)
gamma, beta = gb.chunk(2, dim=-1)
m = self.mid(torch.cat([cond_local, h], dim=-1))
m = m * (1 + gamma) + beta
m = F.silu(m)
if self.dropout is not None and self.training:
m = F.dropout(m, p=self.dropout)
out = self.up(m) * self.multiplier
if out.dtype != x.dtype:
out = out.to(x.dtype)
y = self.org_forward(x + out)
if is_5d:
# org Linear out_features may differ from in_features — recover with -1.
y = y.reshape(orig_shape[0], orig_shape[1], orig_shape[2], orig_shape[3], -1)
return y
# ----------------------------------------------------------------------------
# ControlNetLLLiteDiT
# ----------------------------------------------------------------------------
class ControlNetLLLiteDiT(nn.Module):
def __init__(
self,
dit: nn.Module,
cond_emb_dim: int = 32,
mlp_dim: int = 64,
target_layers: str = "self_attn_q",
dropout: float | None = None,
multiplier: float = 1.0,
cond_dim: int = 64,
cond_resblocks: int = 1,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
inpaint_masked_input: bool = False,
):
super().__init__()
atomics = parse_target_layers(target_layers)
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.target_layers = target_layers
self.target_atomics = atomics
self.dropout = dropout
self.multiplier = multiplier
self.cond_dim = cond_dim
self.cond_resblocks = cond_resblocks
self.use_aspp = use_aspp
self.aspp_dilations = tuple(aspp_dilations) if use_aspp else ()
# 4ch (RGB+mask) inpainting metadata. `inpaint_masked_input` records the training-time
# RGB-masking policy for cond_image preparation; it does not alter the forward pass here.
self.cond_in_channels = cond_in_channels
self.inpaint_masked_input = inpaint_masked_input
self.conditioning1 = _Conditioning1(
cond_dim,
cond_emb_dim,
cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
)
modules = self._create_modules(dit, cond_emb_dim, mlp_dim, atomics, dropout, multiplier)
self.lllite_modules = nn.ModuleList(modules)
n = len(self.lllite_modules)
self.depth_embeds = nn.Parameter(torch.zeros(n, cond_emb_dim))
for i, m in enumerate(self.lllite_modules):
m.layer_idx = i
m._depth_embeds_ref = [self.depth_embeds]
aspp_info = f"aspp={'on' + str(list(self.aspp_dilations)) if use_aspp else 'off'}"
inpaint_info = (
f", inpaint=on(masked_input={inpaint_masked_input})" if cond_in_channels != 3 else ""
)
logger.info(
"ControlNet-LLLite (Anima v%s): created %d modules for target=%r "
"(atomics=%s), cond_in_channels=%d, cond_dim=%d, cond_resblocks=%d, %s, "
"cond_emb_dim=%d, mlp_dim=%d%s",
LLLITE_ARCH_VERSION,
n,
target_layers,
list(atomics),
cond_in_channels,
cond_dim,
cond_resblocks,
aspp_info,
cond_emb_dim,
mlp_dim,
inpaint_info,
)
@staticmethod
def _attn_atomic_match(is_self_attn: bool, child_name: str, atomics: tuple[str, ...]) -> bool:
if "output_proj" in child_name:
return False
if is_self_attn:
if child_name == "q_proj":
return "self_attn_q_pre" in atomics
if child_name in ("k_proj", "v_proj"):
return "self_attn_kv_pre" in atomics
return False
else:
if child_name == "q_proj":
return "cross_attn_q_pre" in atomics
return False # cross_attn K,V live in text-embedding space
def _create_modules(
self,
dit: nn.Module,
cond_emb_dim: int,
mlp_dim: int,
atomics: tuple[str, ...],
dropout: float | None,
multiplier: float,
) -> list[LLLiteModuleDiT]:
modules: list[LLLiteModuleDiT] = []
want_mlp_fc1 = "mlp_fc1_pre" in atomics
any_attn = any(
a in atomics for a in ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre")
)
for name, module in dit.named_modules():
if LLM_ADAPTER_NAME in name:
continue
cls = module.__class__.__name__
def _is_linear_like(module):
return (
hasattr(module, "in_features")
and hasattr(module, "out_features")
and callable(getattr(module, "forward", None))
)
if any_attn and cls == TARGET_ATTENTION_CLASS:
# The Anima-block Attention exposes is_selfattn; the LLM-Adapter
# Attention does not — skip the latter even if path filter misses.
if not hasattr(module, "is_selfattn"):
continue
is_self_attn = bool(module.is_selfattn)
for child_name, child in module.named_children():
if not _is_linear_like(child):
continue
if not self._attn_atomic_match(is_self_attn, child_name, atomics):
continue
full_name = f"lllite_dit.{name}.{child_name}".replace(".", "_")
modules.append(
LLLiteModuleDiT(
full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier
)
)
elif want_mlp_fc1 and cls == TARGET_MLP_CLASS:
child = getattr(module, "layer1", None)
if not _is_linear_like(child):
continue
full_name = f"lllite_dit.{name}.layer1".replace(".", "_")
modules.append(
LLLiteModuleDiT(full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier)
)
return modules
def set_cond_image(self, cond_image: torch.Tensor | None):
"""cond_image: (B, 3, H*16, W*16) in [-1, 1]; ``None`` clears."""
if cond_image is None:
for m in self.lllite_modules:
m.cond_emb = None
return
cx = self.conditioning1(cond_image) # (B, S, cond_emb_dim)
for m in self.lllite_modules:
m.cond_emb = cx
def clear_cond_image(self):
self.set_cond_image(None)
def set_multiplier(self, multiplier: float):
self.multiplier = multiplier
for m in self.lllite_modules:
m.multiplier = multiplier
def apply_to(self):
for m in self.lllite_modules:
m.apply_to()
def restore(self):
for m in self.lllite_modules:
m.restore()
# ----------------------------------------------------------------------------
# Save / load (named-key format; legacy lllite_modules.* is rejected)
# ----------------------------------------------------------------------------
_INTERNAL_MODULES_PREFIX = "lllite_modules."
_INTERNAL_COND_PREFIX = "conditioning1."
_INTERNAL_DEPTH_KEY = "depth_embeds"
_SAVED_COND_PREFIX = "lllite_conditioning1."
_SAVED_DEPTH_SUFFIX = ".depth_embed"
def _from_saved_state_dict(lllite: ControlNetLLLiteDiT, weights_sd: dict) -> dict:
"""Rewrite a v2 named-key state dict back to the internal layout."""
name_to_idx = {m.lllite_name: i for i, m in enumerate(lllite.lllite_modules)}
n_modules = len(name_to_idx)
out: dict = {}
depth_slices: dict = {}
for k, v in weights_sd.items():
if k.startswith(_SAVED_COND_PREFIX):
out[_INTERNAL_COND_PREFIX + k[len(_SAVED_COND_PREFIX) :]] = v
continue
if k.endswith(_SAVED_DEPTH_SUFFIX):
name = k[: -len(_SAVED_DEPTH_SUFFIX)]
if name in name_to_idx:
depth_slices[name_to_idx[name]] = v
continue
head, dot, tail = k.partition(".")
if dot and head in name_to_idx:
out[f"{_INTERNAL_MODULES_PREFIX}{name_to_idx[head]}.{tail}"] = v
continue
out[k] = v
if depth_slices:
missing = [i for i in range(n_modules) if i not in depth_slices]
if missing:
raise RuntimeError(f"depth_embed slices missing for module idx(es) {missing}")
out[_INTERNAL_DEPTH_KEY] = torch.stack([depth_slices[i] for i in range(n_modules)], dim=0)
return out
def load_lllite_weights(lllite: ControlNetLLLiteDiT, file: str, strict: bool = False):
weights_sd = safetensors.torch.load_file(file)
if any(k.startswith(_INTERNAL_MODULES_PREFIX) for k in weights_sd):
raise RuntimeError(
f"weights at {file} appear to be in a legacy ControlNet-LLLite weight format "
f"(keys starting with '{_INTERNAL_MODULES_PREFIX}'). The current code uses a "
f"named-key format (per-module key prefix = lllite_name, e.g. "
f"'lllite_dit_blocks_0_self_attn_q_proj.down.weight'). Re-train with the current codebase."
)
converted = _from_saved_state_dict(lllite, weights_sd)
info = lllite.load_state_dict(converted, strict=strict)
logger.info("loaded LLLite weights from %s: %s", file, info)
return info
def read_lllite_metadata(file: str) -> dict:
if os.path.splitext(file)[1] != ".safetensors":
raise RuntimeError(f"Must use .safetensors files, got {file}")
with safetensors.safe_open(file, framework="pt") as f:
return f.metadata() or {}
# ----------------------------------------------------------------------------
# ComfyUI nodes for Anima ControlNet-LLLite
# ----------------------------------------------------------------------------
def _get_inner_dit(model) -> torch.nn.Module:
"""Reach the underlying Anima DiT (nn.Module) from a ComfyUI ModelPatcher."""
inner = getattr(model, "model", None)
if inner is None:
raise RuntimeError("Input MODEL has no .model attribute (not a ModelPatcher?)")
dit = getattr(inner, "diffusion_model", None)
if dit is None:
raise RuntimeError("MODEL.model has no .diffusion_model — not a UNet/DiT model?")
return dit
def _target_cond_hw(latent_h: int, latent_w: int, patch_spatial: int = 2) -> tuple[int, int]:
"""Return the (H, W) the cond image / mask must be resized to.
The LLLite ``conditioning1`` Conv has stride 16, so the cond image must be
sized to ``latent_HW * 8`` in input pixel space (= ``token_HW * 16`` after
DiT patchify with patch_spatial=2). The DiT internally pads the latent up
to a multiple of ``patch_spatial`` (see ``MiniTrainDIT.forward`` →
``pad_to_patch_size``), so we mirror that rounding here — otherwise odd
latent dims (e.g. 1032 px → 129 latent) yield a token-count mismatch that
silently bypasses every LLLite module.
"""
padded_h = ((latent_h + patch_spatial - 1) // patch_spatial) * patch_spatial
padded_w = ((latent_w + patch_spatial - 1) // patch_spatial) * patch_spatial
return padded_h * 8, padded_w * 8
def _prepare_cond_image(
image: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI IMAGE (B,H,W,3) in [0,1] → (1,3,H*8,W*8) in [-1,1]."""
if image.ndim == 4 and image.shape[-1] == 3:
# (B, H, W, 3) -> (B, 3, H, W)
img = image.permute(0, 3, 1, 2).contiguous()
else:
raise ValueError(f"Unexpected cond image shape: {tuple(image.shape)} (expected B,H,W,3)")
img = img[:1] # use first frame only
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if img.shape[-2] != target_h or img.shape[-1] != target_w:
img = F.interpolate(img, size=(target_h, target_w), mode="bicubic", align_corners=False)
img = img.clamp(0.0, 1.0)
img = img * 2.0 - 1.0
return img.to(device=device, dtype=dtype)
def _prepare_mask(
mask: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI MASK (B,H,W) in [0,1] → (1,1,H*8,W*8) binarized at 0.5.
Returns the mask in ``{0.0, 1.0}`` (1 = inpaint area, 0 = keep). The caller
is responsible for the ``*2-1`` rescale before concat with RGB.
"""
if mask.ndim == 3:
m = mask.unsqueeze(1) # (B, 1, H, W)
elif mask.ndim == 4 and mask.shape[1] == 1:
m = mask
else:
raise ValueError(f"Unexpected mask shape: {tuple(mask.shape)} (expected B,H,W or B,1,H,W)")
m = m[:1]
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if m.shape[-2] != target_h or m.shape[-1] != target_w:
m = F.interpolate(m.float(), size=(target_h, target_w), mode="nearest")
m = (m >= 0.5).to(dtype=dtype)
return m.to(device=device)
def _build_inpaint_cond_image(
rgb_pm1: torch.Tensor, mask01: torch.Tensor, masked_input: bool
) -> torch.Tensor:
"""rgb_pm1: (1,3,H,W) in [-1,1], mask01: (1,1,H,W) in {0,1}. Returns (1,4,H,W).
Mirrors ``_build_inpaint_cond_image`` in the sd-scripts training / inference
code: the mask channel is rescaled to ``[-1, +1]`` (matches the RGB range),
and if ``masked_input`` is set the RGB is zeroed where ``mask >= 0.5``.
"""
if masked_input:
keep = (mask01 < 0.5).to(rgb_pm1.dtype)
rgb_pm1 = rgb_pm1 * keep
mask_pm1 = mask01.to(rgb_pm1.dtype) * 2.0 - 1.0
return torch.cat([rgb_pm1, mask_pm1], dim=1)
ETNControlNet = io.Custom("ETN_CONTROL_NET")
class ControlLoad(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_load",
display_name="Load ControlNet (tooling-nodes)",
description="Loads ControlNet weights. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
io.Combo.Input("weights", folder_paths.get_filename_list("controlnet")),
],
outputs=[
io.Model.Output("out_model", "model"),
ETNControlNet.Output("control_net"),
],
)
@classmethod
def execute(cls, model: ModelPatcher, weights: str): # type: ignore[override]
weights_path = folder_paths.get_full_path("controlnet", weights)
if weights_path is None or not os.path.isfile(weights_path):
raise FileNotFoundError(f"LLLite weights not found: {weights}")
# Architecture is fully determined by the trained weights — read everything
# from metadata rather than exposing knobs that would just cause load errors.
meta = read_lllite_metadata(weights_path)
if "lllite.version" not in meta:
raise RuntimeError(
"Unrecognized model. This node currently only loads Anima LLLite weights."
)
ce_dim = int(meta.get("lllite.cond_emb_dim", 32))
m_dim = int(meta.get("lllite.mlp_dim", 64))
# v2 records the canonical atomic form under lllite.target_atomics; fall back
# to the legacy preset key, then to the v1 default.
tl = meta.get("lllite.target_atomics", meta.get("lllite.target_layers", "self_attn_q"))
cond_dim = int(meta.get("lllite.cond_dim", 64))
cond_resblocks = int(meta.get("lllite.cond_resblocks", 1))
use_aspp = str(meta.get("lllite.use_aspp", "false")).lower() == "true"
aspp_dilations_meta = meta.get("lllite.aspp_dilations")
if use_aspp and aspp_dilations_meta:
aspp_dilations = tuple(int(d) for d in aspp_dilations_meta.split(",") if d.strip())
else:
aspp_dilations = ASPP_DEFAULT_DILATIONS
cond_in_channels = int(meta.get("lllite.cond_in_channels", 3))
inpaint_masked_input = (
str(meta.get("lllite.inpaint_masked_input", "false")).lower() == "true"
)
lllite = ControlNetLLLiteDiT(
_get_inner_dit(model),
cond_emb_dim=ce_dim,
mlp_dim=m_dim,
target_layers=tl,
multiplier=1.0,
cond_dim=cond_dim,
cond_resblocks=cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
inpaint_masked_input=inpaint_masked_input,
)
load_lllite_weights(lllite, weights_path, strict=False)
lllite.eval().requires_grad_(False)
return io.NodeOutput(model, lllite)
class ControlApply(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_apply",
display_name="Apply ControlNet (tooling-nodes)",
description="Applies ControlNet conditioning. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
ETNControlNet.Input("control_net"),
io.Image.Input("image"),
io.Mask.Input("mask", optional=True),
io.Float.Input("strength", default=1.0, min=-10.0, max=10.0, step=0.01),
io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001),
io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001),
],
outputs=[io.Model.Output("model")],
)
@classmethod
def execute( # type: ignore[override]
cls,
model: ModelPatcher,
control_net: ControlNetLLLiteDiT,
image: torch.Tensor,
strength: float,
start_percent: float,
end_percent: float,
mask: torch.Tensor | None = None,
):
dit = _get_inner_dit(model)
patch_spatial = int(getattr(dit, "patch_spatial", 2))
lllite = control_net
lllite.set_multiplier(strength)
# Mask / cond_in_channels consistency: 4ch weights need a MASK, 3ch weights ignore it.
if lllite.cond_in_channels == 4 and mask is None:
raise ValueError("ControlNet weights require a mask input (inpaint mode)")
if lllite.cond_in_channels != 4 and mask is not None:
mask = None
# Convert percent range -> sigma range (start_percent=0 → sigma_max).
model_sampling = model.get_model_object("model_sampling")
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
# Capture image / mask tensors (cloned to detach from any upstream caching)
src_image = image.detach().clone()
src_mask = mask.detach().clone() if mask is not None else None
is_inpaint = lllite.cond_in_channels == 4
# Cache for the per-resolution preprocessed cond image (avoids repeat resize)
cache: dict[str, Any] = {"cond_image_pp": None, "key": None, "lllite_loaded_to": None}
# Capture any previously-installed wrapper BEFORE we clone — model_options
# has a single "model_function_wrapper" slot, so without delegation a second
# wrapper-installing node would silently no-op the first. Mirrors the
# ChromaRadianceOptions pattern in comfy_extras/nodes_chroma_radiance.py.
old_wrapper = model.model_options.get("model_function_wrapper")
def _call_next(apply_model, input_x, timestep, c):
if old_wrapper is not None:
return old_wrapper(apply_model, {"input": input_x, "timestep": timestep, "c": c})
return apply_model(input_x, timestep, **c)
def wrapper(apply_model, args):
input_x = args["input"]
timestep = args["timestep"]
c = args["c"]
# Step-range gate: skip LLLite entirely when current sigma is outside
# [sigma_end, sigma_start]. percent_to_sigma maps 0.0 → sigma_max,
# 1.0 → sigma_min, so the active window is sigma_end <= sigma <= sigma_start.
sigma = float(timestep.max().item())
if not (sigma_end <= sigma <= sigma_start):
return _call_next(apply_model, input_x, timestep, c)
# Anima latent shape: (B, C, T, H, W) — take spatial dims from the tail.
latent_h, latent_w = int(input_x.shape[-2]), int(input_x.shape[-1])
device = input_x.device
dtype = input_x.dtype
# Move LLLite to the runtime device/dtype lazily.
tag = (device, dtype)
if cache["lllite_loaded_to"] != tag:
lllite.to(device=device, dtype=dtype)
cache["lllite_loaded_to"] = tag
cache["cond_image_pp"] = None # invalidate
key = (latent_h, latent_w, device, dtype)
if cache["key"] != key or cache["cond_image_pp"] is None:
rgb = _prepare_cond_image(
src_image, latent_h, latent_w, device, dtype, patch_spatial
)
if is_inpaint:
assert src_mask is not None, "Cannot use inpaint control-net without a mask"
mk = _prepare_mask(src_mask, latent_h, latent_w, device, dtype, patch_spatial)
cache["cond_image_pp"] = _build_inpaint_cond_image(
rgb, mk, lllite.inpaint_masked_input
)
else:
cache["cond_image_pp"] = rgb
cache["key"] = key
lllite.set_multiplier(strength)
lllite.set_cond_image(cache["cond_image_pp"])
lllite.apply_to()
try:
return _call_next(apply_model, input_x, timestep, c)
finally:
lllite.restore()
lllite.clear_cond_image()
m = model.clone()
m.set_model_unet_function_wrapper(wrapper)
return (m,)
+46 -12
View File
@@ -1,15 +1,16 @@
import sys import sys
import torch
import numpy as np
from enum import Enum from enum import Enum
from pathlib import Path from pathlib import Path
from typing import Any, NamedTuple from typing import Any, NamedTuple
from PIL import Image
import server
import comfy.samplers import comfy.samplers
import numpy as np
import server
import torch
from comfy.comfy_types.node_typing import IO from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io from comfy_api.latest import io
from PIL import Image
from .nodes import SendImageWebSocket from .nodes import SendImageWebSocket
@@ -102,7 +103,7 @@ class KritaOutput(io.ComfyNode):
) )
@classmethod @classmethod
def execute( def execute( # type: ignore
cls, cls,
images: torch.Tensor, images: torch.Tensor,
x: int = 0, x: int = 0,
@@ -141,7 +142,7 @@ class KritaSendText(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, value: Any, name: str, type: str): def execute(cls, value: Any, name: str, type: str): # type: ignore
mime = { mime = {
"text": "text/plain", "text": "text/plain",
"markdown": "text/markdown", "markdown": "text/markdown",
@@ -169,12 +170,13 @@ class KritaCanvas(io.ComfyNode):
io.Int.Output(display_name="width"), io.Int.Output(display_name="width"),
io.Int.Output(display_name="height"), io.Int.Output(display_name="height"),
io.Int.Output(display_name="seed"), io.Int.Output(display_name="seed"),
io.Mask.Output(display_name="mask"),
], ],
) )
@classmethod @classmethod
def execute(cls): def execute(cls, **kwargs):
return io.NodeOutput(_placeholder_image(), 512, 512, 0) return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
class SelectionContext(Enum): class SelectionContext(Enum):
@@ -236,7 +238,7 @@ class KritaImageLayer(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, name: str): def execute(cls, name: str): # type: ignore
return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512)) return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512))
@@ -254,7 +256,7 @@ class KritaMaskLayer(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, name: str): def execute(cls, name: str): # type: ignore
return io.NodeOutput(torch.ones(1, 512, 512)) return io.NodeOutput(torch.ones(1, 512, 512))
@@ -289,7 +291,7 @@ class Parameter(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, name: str, type: str, default, min=0.0, max=1.0): def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
if type == "number": if type == "number":
return io.NodeOutput(float(default)) return io.NodeOutput(float(default))
elif type == "number (integer)": elif type == "number (integer)":
@@ -326,5 +328,37 @@ class KritaStyle(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, name: str, sampler_preset: str): def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!")
class KritaStyleAndPrompt(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaStyleAndPrompt",
display_name="Krita Style & Prompt",
category="krita",
inputs=[
io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]),
],
outputs=[
io.Model.Output(display_name="model (with loras)"),
io.Clip.Output(display_name="clip"),
io.Vae.Output(display_name="vae"),
io.String.Output(display_name="positive prompt (evaluated)"),
io.String.Output(display_name="negative prompt (evaluated)"),
io.Combo.Output(
display_name="sampler name", options=comfy.samplers.KSampler.SAMPLERS
),
io.Combo.Output(
display_name="scheduler", options=comfy.samplers.KSampler.SCHEDULERS
),
io.Int.Output(display_name="steps"),
io.Float.Output(display_name="guidance"),
],
)
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!") raise NotImplementedError("This workflow must be started from Krita!")
+18 -16
View File
@@ -1,20 +1,21 @@
from __future__ import annotations from __future__ import annotations
import base64
import time
from copy import copy from copy import copy
from dataclasses import dataclass from dataclasses import dataclass
import time from io import BytesIO
from typing import NamedTuple from typing import NamedTuple
from uuid import uuid4 from uuid import uuid4
from PIL import Image
import numpy as np import numpy as np
import base64
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from io import BytesIO
from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel from comfy.sd import StyleModel
from comfy_api.latest import io from comfy_api.latest import io
from PIL import Image
from server import BinaryEventTypes, PromptServer
class LoadImageBase64(io.ComfyNode): class LoadImageBase64(io.ComfyNode):
@@ -29,14 +30,14 @@ class LoadImageBase64(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, image: str): def execute(cls, image: str): # type: ignore
_strip_prefix(image, "data:image/png;base64,") _strip_prefix(image, "data:image/png;base64,")
imgdata = base64.b64decode(image) imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
if "A" in img.getbands(): if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0 mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = torch.from_numpy(mask) mask = torch.from_numpy(mask)[None,]
else: else:
mask = None mask = None
@@ -59,7 +60,7 @@ class LoadMaskBase64(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, mask: str): def execute(cls, mask: str): # type: ignore
_strip_prefix(mask, "data:image/png;base64,") _strip_prefix(mask, "data:image/png;base64,")
imgdata = base64.b64decode(mask) imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
@@ -85,7 +86,7 @@ class SendImageWebSocket(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, images: torch.Tensor, format: str): def execute(cls, images: torch.Tensor, format: str): # type: ignore
results = [] results = []
for tensor in images: for tensor in images:
array = 255.0 * tensor.cpu().numpy() array = 255.0 * tensor.cpu().numpy()
@@ -198,12 +199,13 @@ class LoadImageCache(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, id: str): def execute(cls, id: str): # type: ignore
image_data, content_type = image_cache.get(id, extend=True) image_data, content_type = image_cache.get(id, extend=True)
if image_data is None: if image_data is None:
raise ValueError(f"Image with ID {id} not found in cache.") raise ValueError(f"Image with ID {id} not found in cache.")
img = Image.open(BytesIO(image_data)) img = Image.open(BytesIO(image_data))
w, h = img.size w, h = img.size
c = len(img.getbands()) c = len(img.getbands())
normalized = np.array(img).astype(np.float32) / 255.0 normalized = np.array(img).astype(np.float32) / 255.0
@@ -237,7 +239,7 @@ class SaveImageCache(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, images: torch.Tensor, format: str): def execute(cls, images: torch.Tensor, format: str): # type: ignore
results = [] results = []
for tensor in images: for tensor in images:
array = 255.0 * tensor.cpu().numpy() array = 255.0 * tensor.cpu().numpy()
@@ -284,7 +286,7 @@ class ApplyMaskToImage(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, image: torch.Tensor, mask: torch.Tensor): def execute(cls, image: torch.Tensor, mask: torch.Tensor): # type: ignore
out = to_bchw(image) out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
@@ -300,7 +302,7 @@ class ApplyMaskToImage(io.ComfyNode):
# Apply each mask in the batch to its corresponding image's alpha channel # Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]): for i in range(out.shape[0]):
alpha = mask[i] if is_mask_batch else mask[0] alpha = mask[i] if is_mask_batch else mask[0]
out[i, 3, :, :] = alpha out[i, 3, :, :] *= alpha
return (to_bhwc(out),) return (to_bhwc(out),)
@@ -329,7 +331,7 @@ class ReferenceImage(io.ComfyNode):
) )
@classmethod @classmethod
def execute( def execute( # type: ignore
cls, cls,
image: torch.Tensor, image: torch.Tensor,
weight: float, weight: float,
@@ -359,7 +361,7 @@ class ApplyReferenceImages(io.ComfyNode):
) )
@classmethod @classmethod
def execute( def execute( # type: ignore
cls, cls,
conditioning: list[list], conditioning: list[list],
clip_vision: ClipVisionModel, clip_vision: ClipVisionModel,
+4 -1
View File
@@ -1,5 +1,4 @@
from __future__ import annotations from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path from pathlib import Path
from tqdm import tqdm from tqdm import tqdm
import torch import torch
@@ -39,6 +38,10 @@ class CLIPSafetyChecker(PreTrainedModel):
self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False) self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False) self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False)
# Model requires post_init after transformers v4.57.3
if hasattr(self, "post_init"):
self.post_init()
def forward(self, clip_input, images: Tensor, sensitivity: float): def forward(self, clip_input, images: Tensor, sensitivity: float):
with torch.no_grad(): with torch.no_grad():
image_batch = self.vision_model(clip_input)[1] image_batch = self.vision_model(clip_input)[1]
+2 -2
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui-tooling-nodes" name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "3.1.0" version = "3.4.0"
license = { file = "LICENSE" } license = { file = "LICENSE" }
[project.urls] [project.urls]
@@ -13,7 +13,7 @@ line-length = 100
preview = true preview = true
[tool.ruff.lint] [tool.ruff.lint]
ignore = ["E741"] ignore = ["E741", "BLE001"]
[tool.black] [tool.black]
line-length = 100 line-length = 100
+214 -1
View File
@@ -1,16 +1,34 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
# Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py # Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py
# by @laksjdjf # by @laksjdjf
from __future__ import annotations from __future__ import annotations
from typing import NamedTuple from functools import partial
from typing import Any, NamedTuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import math import math
from torch import Tensor, Size from torch import Tensor, Size
import comfy.model_management
import comfy.patcher_extension
from comfy.model_patcher import ModelPatcher from comfy.model_patcher import ModelPatcher
from comfy.model_base import Anima, CosmosPredict2
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
from comfy_api.latest import io from comfy_api.latest import io
COND = 0
UNCOND = 1
ANIMA_COUPLE_WRAPPER_KEY = "etn_attention_mask_anima"
ANIMA_COUPLE_PATCH_KEY = "etn_attention_mask_patch"
CONDS_COUPLE_KEY = "etn_couple_conds"
COND_UNCOND_COUPLE_KEY = "etn_couple_cond_or_uncond"
COUPLE_ACTIVE_KEY = "etn_couple_active"
NUM_TOKENS_COUPLE_KEY = "etn_couple_num_tokens"
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
h, w = original_shape[2], original_shape[3] h, w = original_shape[2], original_shape[3]
hm, wm = mask.shape[2], mask.shape[3] hm, wm = mask.shape[2], mask.shape[3]
@@ -34,6 +52,12 @@ def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape:
return result return result
def reshape_mask(mask: Tensor, size: tuple[int, int], batch: int, target_size: int) -> Tensor:
result = F.interpolate(mask, size=size, mode="nearest")
result = result.view(mask.shape[0], target_size, 1)
return result.repeat_interleave(batch, dim=0)
def lcm(a: int, b: int): def lcm(a: int, b: int):
return a * b // math.gcd(a, b) return a * b // math.gcd(a, b)
@@ -145,6 +169,7 @@ class AttentionMaskPatch:
mask_sum = mask.sum(dim=0, keepdim=True) mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks." assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum self.mask = mask / mask_sum
self.region_conds = [r.conditioning for r in region_list]
self.conds = [r.conditioning[0][0] for r in region_list] self.conds = [r.conditioning[0][0] for r in region_list]
self.num_tokens = [cond.shape[1] for cond in self.conds] self.num_tokens = [cond.shape[1] for cond in self.conds]
self.num_conds = len(region_list) self.num_conds = len(region_list)
@@ -153,6 +178,8 @@ class AttentionMaskPatch:
@staticmethod @staticmethod
def apply(model: ModelPatcher, regions: Region): def apply(model: ModelPatcher, regions: Region):
patch = AttentionMaskPatch(regions.preprocess()) patch = AttentionMaskPatch(regions.preprocess())
if _is_anima_couple_model(model):
return patch.apply_anima(model)
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict): def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
assert k.mean() == v.mean(), "k and v must be the same." assert k.mean() == v.mean(), "k and v must be the same."
@@ -221,3 +248,189 @@ class AttentionMaskPatch:
new_model.set_model_attn2_output_patch(attn2_output_patch) new_model.set_model_attn2_output_patch(attn2_output_patch)
new_model.set_attachments("etn_attention_mask", patch) new_model.set_attachments("etn_attention_mask", patch)
return new_model return new_model
def apply_anima(self, model: ModelPatcher):
new_model = model.clone()
_patch_cosmos_attention(new_model)
device = comfy.model_management.get_torch_device()
conds_converted = [convert_cond(cond)[0] for cond in self.region_conds]
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_sample_wrapper(conds_converted, device),
)
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_diffusion_wrapper(self),
)
new_model.set_attachments("etn_attention_mask", self)
return new_model
def _is_anima_couple_model(model: ModelPatcher) -> bool:
model_type = type(model.model)
return issubclass(model_type, (Anima, CosmosPredict2))
def _anima_couple_sample_wrapper(conds_converted: list, device):
def sample_wrapper(executor, *args, **kwargs):
if len(conds_converted) > 0:
guider = args[0]
extra_options: dict[str, Any] = args[2]
seed: int = extra_options["seed"]
noise: Tensor = args[4]
latent_image: Tensor = args[5]
denoise_mask: Tensor | None = args[6]
conds_processed = process_conds(
guider.inner_model,
noise,
{"positive": conds_converted},
device,
latent_image,
denoise_mask,
seed,
latent_shapes=[latent_image.shape],
)["positive"]
conds_couple = [cond["model_conds"]["c_crossattn"].cond for cond in conds_processed]
model_options: dict[str, Any] = extra_options["model_options"]
transformer_options: dict[str, Any] = model_options.get("transformer_options", {}).copy()
transformer_options[CONDS_COUPLE_KEY] = conds_couple
transformer_options[NUM_TOKENS_COUPLE_KEY] = [cond.shape[1] for cond in conds_couple]
model_options["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return sample_wrapper
def _anima_couple_diffusion_wrapper(patch: AttentionMaskPatch):
def diffusion_wrapper(executor, *args, **kwargs):
anima_model = executor.class_obj
x: Tensor = args[0]
transformer_options: dict[str, Any] = kwargs.get("transformer_options", {}).copy()
patch_spatial = getattr(anima_model, "patch_spatial", 1)
activations_shape = list(x.shape)
activations_shape[-2] = activations_shape[-2] // patch_spatial
activations_shape[-1] = activations_shape[-1] // patch_spatial
transformer_options["activations_shape"] = activations_shape
transformer_options[ANIMA_COUPLE_PATCH_KEY] = patch
kwargs["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return diffusion_wrapper
def pre_cross_attention(
patch: AttentionMaskPatch,
transformer_options: dict,
x: Tensor,
context: Tensor,
rope_emb: Tensor | None,
) -> tuple[Tensor, Tensor, Tensor | None, dict]:
transformer_options = transformer_options.copy()
if CONDS_COUPLE_KEY not in transformer_options:
transformer_options[COND_UNCOND_COUPLE_KEY] = list(transformer_options["cond_or_uncond"])
transformer_options[COUPLE_ACTIVE_KEY] = False
return x, context, rope_emb, transformer_options
conds: list[Tensor] = transformer_options[CONDS_COUPLE_KEY]
num_tokens_c: list[int] = transformer_options[NUM_TOKENS_COUPLE_KEY]
cond_or_uncond = transformer_options["cond_or_uncond"]
num_chunks = len(cond_or_uncond)
batch = x.shape[0] // num_chunks
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = context.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm_for_list(num_tokens_c + [context.shape[1]])
conds_c_tensor = torch.cat(
[cond.repeat(batch, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)],
dim=0,
)
xs, cs = [], []
cond_or_uncond_couple = []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // context.shape[1], 1)
if cond_type == UNCOND:
xs.append(x_target)
cs.append(c_target)
cond_or_uncond_couple.append(UNCOND)
else:
xs.append(x_target.repeat(patch.num_conds, 1, 1))
cs.append(conds_c_tensor)
cond_or_uncond_couple.extend([COND] * patch.num_conds)
transformer_options[COND_UNCOND_COUPLE_KEY] = cond_or_uncond_couple
transformer_options[COUPLE_ACTIVE_KEY] = True
return torch.cat(xs, dim=0), torch.cat(cs, dim=0), rope_emb, transformer_options
def cross_attention_output(patch: AttentionMaskPatch, transformer_options: dict, out: Tensor):
cond_or_uncond = transformer_options[COND_UNCOND_COUPLE_KEY]
size = tuple(transformer_options["activations_shape"][-2:])
batch = out.shape[0] // len(cond_or_uncond)
mask = patch.mask.to(out.device, dtype=out.dtype)
mask_downsample = reshape_mask(mask, size, batch, out.shape[1])
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond):
pos, next_pos = i * batch, (i + 1) * batch
if cond_type == UNCOND:
outputs.append(out[pos:next_pos])
else:
pos_cond, next_pos_cond = i_cond * batch, (i_cond + 1) * batch
cond_outputs.append(out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond])
i_cond += 1
if len(cond_outputs) > 0:
outputs.append(torch.stack(cond_outputs).sum(0))
return torch.cat(outputs, dim=0)
def _patch_cosmos_attention(model_patcher: ModelPatcher):
cosmos_model = model_patcher.get_model_object("diffusion_model")
for block_name, block in (
(n, b)
for n, b in cosmos_model.named_modules()
if ("cross_attn" in n or "self_attn" in n) and isinstance(b, CosmosAttention)
):
patch_name = f"diffusion_model.{block_name}.forward"
if patch_name not in model_patcher.object_patches:
model_patcher.add_object_patch(patch_name, partial(_cosmos_attention_forward_patched, block))
def _cosmos_attention_forward_patched(
self,
x: Tensor,
context: Tensor | None = None,
rope_emb: Tensor | None = None,
transformer_options: dict | None = None,
) -> Tensor:
transformer_options = transformer_options if transformer_options is not None else {}
patch: AttentionMaskPatch | None = transformer_options.get(ANIMA_COUPLE_PATCH_KEY)
if context is not None and patch is not None:
x, context, rope_emb, transformer_options = pre_cross_attention(
patch, transformer_options, x, context, rope_emb
)
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
output = self.compute_attention(q, k, v, transformer_options=transformer_options)
if context is not None and patch is not None and transformer_options.get(COUPLE_ACTIVE_KEY, False):
output = cross_attention_output(patch, transformer_options, output)
return output
+11 -6
View File
@@ -9,9 +9,13 @@ IntArray = npt.NDArray[np.int_]
class TileLayout: class TileLayout:
def __init__(self, image: Tensor, min_tile_size: int, padding: int, blending: int): def __init__(
assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8" self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
assert min_tile_size % 8 == 0, "Tile size must be divisible by 8" ):
assert all([x % multiple == 0 for x in image.shape[-3:-1]]), (
"Image size must be divisible by multiple"
)
assert min_tile_size % multiple == 0, "Tile size must be divisible by multiple"
assert blending <= padding, "Blending must be smaller than padding" assert blending <= padding, "Blending must be smaller than padding"
self.image_size: IntArray = np.array(image.shape[-3:-1]) self.image_size: IntArray = np.array(image.shape[-3:-1])
@@ -21,7 +25,7 @@ class TileLayout:
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
tile_size = np.ceil(image_size_with_overlap / self.tile_count) tile_size = np.ceil(image_size_with_overlap / self.tile_count)
self.tile_size: IntArray = (np.ceil(tile_size / 8) * 8).astype(int) self.tile_size: IntArray = (np.ceil(tile_size / multiple) * multiple).astype(int)
def size(self, coord: IntArray): def size(self, coord: IntArray):
return self.end(coord) - self.start(coord) return self.end(coord) - self.start(coord)
@@ -84,13 +88,14 @@ class CreateTileLayout(io.ComfyNode):
io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8), io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8),
io.Int.Input("padding", default=32, min=0, max=8192, step=8), io.Int.Input("padding", default=32, min=0, max=8192, step=8),
io.Int.Input("blending", default=8, min=0, max=256, step=8), io.Int.Input("blending", default=8, min=0, max=256, step=8),
io.Int.Input("multiple", default=8, min=1, max=1024, step=1),
], ],
outputs=[io.Custom("TileLayout").Output(display_name="layout")], outputs=[io.Custom("TileLayout").Output(display_name="layout")],
) )
@classmethod @classmethod
def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int): def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int):
return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending)) return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending, multiple))
class ExtractImageTile(io.ComfyNode): class ExtractImageTile(io.ComfyNode):