1 Commits
Author SHA1 Message Date
Acly 00ccb7dcf0 Add LoadImageCached and SaveImageCached (renamed from SendImageHTTP)
* short-lived in-memory cache for image transfers
* upload images via HTTP to cache and load/reference them in workflows
* save/store images in workflows to cache and download them via HTTP
2025-10-20 13:43:25 +02:00
11 changed files with 85 additions and 1390 deletions
+8 -34
View File
@@ -68,15 +68,15 @@ This is typically faster than WebSocket, especially for large images.
This node will send a JSON message over WebSocket when an image is ready:
```json
{
"type": "executed",
"data": {
"node": "<node ID>",
"output": {
"images": [
{"source": "http", "id": "<image ID>", "content-type": "image/png", "type": "output"}
'type': 'executed',
'data': {
'node': '<node ID>',
'output': {
'images': [
{'source': 'http', 'id': '<image ID>', 'content-type': 'image/png', 'type': 'output'}
]
},
"prompt_id": "prompt ID"
'prompt_id': 'prompt ID'
}
}
```
@@ -205,8 +205,6 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
#### Paramters
* `folder_name`: sub-directory in ComfyUI's models folder.
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
* `limit=n`: (query parameter, optional) inspect at `n` models
* `offset=i`: (query parameter, optional) start with the `i`th model
#### Output
Lists available models with additional classification info:
@@ -220,7 +218,7 @@ Lists available models with additional classification info:
...
}
```
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, flux2, lumina2, z-image, chroma, qwen-image`
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, lumina2, chroma, qwen-image`
If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm`
@@ -230,24 +228,6 @@ Detection supports quantized models:
Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models.
#### Pagination
The query parameters limit and offset allow inspecting a subset of models per request.
Usually inspection is quite fast (it only looks at model headers), but it can be slow
in some cases due to anti-virus or slow harddrives.
```
GET /api/etn/model_info/checkpoints?limit=10&offset=20
```
This will return at most 10 models, starting with the 20th model in the list.
It also returns a special `_meta` entry in the output JSON:
```json
{
"checkpoint_20.safetensors": { ... },
"_meta": { "offset": 20, "count": 1, "total": 21 }
}
```
### GET /api/etn/languages
Returns a list of available languages for translation.
@@ -293,9 +273,3 @@ git clone https://github.com/Acly/comfyui-tooling-nodes.git
```
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)
+3 -18
View File
@@ -1,12 +1,10 @@
from comfy_api.latest import ComfyExtension, io
from . import api as api
from . import control, krita, nodes, region, tile, translation
from . import api as api, nodes, tile, region, nsfw, translation, krita
class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]:
node_list = [
return [
nodes.LoadImageCache,
nodes.SaveImageCache,
nodes.LoadImageBase64,
@@ -24,6 +22,7 @@ class ExternalToolingNodes(ComfyExtension):
region.DefineRegion,
region.ListRegionMasks,
region.AttentionMask,
nsfw.NSFWFilter,
translation.Translate,
krita.KritaOutput,
krita.KritaSendText,
@@ -33,21 +32,7 @@ class ExternalToolingNodes(ComfyExtension):
krita.KritaMaskLayer,
krita.Parameter,
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():
+15 -75
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from aiohttp import web
from typing import Any, NamedTuple
from typing import NamedTuple
from pathlib import Path
import json
import traceback
@@ -44,8 +44,6 @@ model_names = {
"CosmosI2V": "cosmos",
"CosmosT2IPredict2": "cosmos-predict2",
"CosmosI2VPredict2": "cosmos-predict2",
"ZImage": "z-image",
"Lumina2": "lumina2",
"WAN21_T2V": "wan21",
"WAN21_I2V": "wan21",
"WAN21_FunControl2V": "wan21-fun",
@@ -56,11 +54,6 @@ model_names = {
"ACEStep": "ace-step",
"Omnigen2": "omnigen2",
"QwenImage": "qwen-image",
"QwenImage21": "qwen-image21",
"ErnieImage": "ernie-image",
"Flux2": "flux2",
"Anima": "anima",
"Krea2": "krea2",
}
gguf_architectures = {
@@ -125,15 +118,12 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
raw_name = base_model.__class__.__name__
if raw_name == "SDXL":
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:
return {"base_model": "unknown"}
base_model_name = model_names.get(raw_name, "unknown")
result: dict[str, Any] = {"base_model": base_model_name}
result = {"base_model": base_model_name}
result["is_inpaint"] = (
base_model_name in ["sd15", "sdxl"] and input_count > 4
) or raw_name == "FluxInpaint"
@@ -152,7 +142,6 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
return result
return {"base_model": "unknown"}
except Exception as e:
print("[comfyui-tooling-nodes] Error inspecting file", filename)
traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
@@ -162,16 +151,12 @@ def detect_svdq(cfg: dict) -> str | None:
if comfy_config := md.get("comfy_config"):
if isinstance(comfy_config, str):
comfy_config = json.loads(comfy_config)
if model_class := comfy_config.get("model_class"):
return model_class
match md.get("model_class"):
case "NunchakuFluxTransformer2dModel":
return "Flux"
case "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
case "NunchakuZImageTransformer2DModel":
return "ZImage"
return comfy_config.get("model_class")
model_class = md.get("model_class")
if model_class == "NunchakuFluxTransformer2dModel":
return "Flux"
if model_class == "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
return None
@@ -183,9 +168,6 @@ def inspect_gguf(filename: str, model_type: str):
try:
path = folder_paths.get_full_path(model_type, filename)
if path is None:
raise Exception(f"Could not find full path for {model_type}/{filename}")
reader = gguf.GGUFReader(path)
arch_field = reader.get_field("general.architecture")
if arch_field is not None:
@@ -196,45 +178,19 @@ def inspect_gguf(filename: str, model_type: str):
arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8")
else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False}
if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5)
):
arch_str = "chroma"
# Detect Z-Image (modified Lumina2)
if arch_str == "lumina2":
for t in reader.tensors:
if t.name == "cap_embedder.1.bias" and t.shape[0] == 3840:
arch_str = "z-image"
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 = {
"base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False,
}
if result_type is not None:
result["type"] = result_type
try:
if file_type := reader.get_field("general.file_type"):
result["quant"] = file_type.contents().lower()
except Exception:
result["quant"] = reader.get_field("general.file_type").lower()
except Exception as e:
result["quant"] = "gguf"
return result
@@ -249,22 +205,17 @@ def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool)
return inspect_safetensors(filename, model_type, is_checkpoint)
def inspect_models(model_type: str, params: dict[str, str]):
def inspect_models(model_type: str):
try:
try:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
limit = int(params.get("limit", "1000"))
offset = int(params.get("offset", "0"))
files_range = files[offset : offset + limit]
is_checkpoint = model_type == "checkpoints"
info = {
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in files_range
for filename in files
}
if "limit" in params:
info["_meta"] = dict(offset=offset, count=len(files_range), total=len(files))
return web.json_response(info)
except Exception as e:
traceback.print_exc()
@@ -310,11 +261,11 @@ if _server is not None:
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
return inspect_models(folder_name, request.rel_url.query)
return inspect_models(folder_name)
@_server.routes.get("/api/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints", request.rel_url.query)
return inspect_models("checkpoints")
@_server.routes.get("/api/etn/languages")
async def languages(request):
@@ -350,11 +301,11 @@ if _server is not None:
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/image/{id}")
async def put_image(request: web.Request):
try:
id = request.match_info.get("id", "")
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)
content_type = request.headers.get("Content-Type", "application/octet-stream")
@@ -367,17 +318,6 @@ if _server is not None:
except Exception as e:
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}")
async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "")
-859
View File
@@ -1,859 +0,0 @@
"""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,)
+5 -3
View File
@@ -32,6 +32,7 @@ function loadImage(base64) {
}
const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=")
const outputIcon = loadImage("data:image/webp;base64,UklGRrIHAABXRUJQVlA4WAoAAAAQAAAAjwAAOwAAQUxQSKoCAAARkMbsnyFJ/2TVySSdzPJs27bxzbbtuznbtm3btm3bt+ikkkonlfxPU1X9n57zh4iYAPifZKQ2T7LIkKDkQ1FPrezCF/hjvq+dx7mwiqM3HnJ058yGXkorEfGCKZV6I6o+quOM4bMsgU5bfE1KOh4LEbGnv2RnUSer50DGFwxJ2qwWGTATERFfZPzBZNR9P1ZXTksgVdaODHgt/H5hCJjP0ME6eiLfCqTLypKBSPYdsmY2OjpWy0KOlF+EkIFk/DvnZ2tIzZA0a0kHUtok0KfWkxieJQQZBQksq3QWiXODEGSnYRsq76mxjJQK0cDdKjY1XpLSWyKYVYFTSyxLqCZSvexKppZnZDCZTNpD6BKTi2iIRbqzZc6gW8x5m1ptOCEuYaJ74B1T6RkhjPQXkugiuC9EBSk38wftRACZsdKLEHGRQjJSCyUgnwichSgdj4jYW64KqQsywN1E1JJqReqtWyEvLtOTFHMtfG8GXMYT6C5zQLIjqXiJm+guNw2ZOqTu+0uJ7sKyg2w+Urv9GdxdGoN0JKnh/qBnIE1LlH6LiHNAMZ5SWQkoJAJHcQ7iTUNlNyE7Vga4a7CsoFqP0AmQjqdm5dTWGJRTvqfTSu4ONe7VNRs0LiTzLK3cNEHsGWhuZ+gom0hlNMgXZ7T4WF16Q6YRuZZPAS7TYrGUoPgFEnZPUC3EKDEf0O6WSGFmMiVoxejw3UDcHC6c21oFNLZjVNhGgxrkHC+cOtQKtBa5aQkCLL4VBGDZsbYzu7uF6AGouPQ1t7mTIv5AMw8EZLmhbx0QizqHgJMer5MQwHkG7NP2YmgzCMqRPYff0cIWDSgOwbrcgOFnhcqLhQNaeSB4h1XxDZh29LX4kXVt5YABrZJBkM/acIBsx5IG/AyGpS1UrkrF4tm98K8uVlA4IOIEAABwHACdASqQADwAPm0uk0ckIiGhLjUJmIANiWgOuBpEsADI+tE9j+ifkl+QHyd1n+77tGXf04+Cpx/549gD9X+lF5t/2d9cn0T+gB/d/8R1gHoAeWx+3nwZfuZ6T+aq/0Dtm/zaCN5VfOk9UfrX8A/Ry9E1elWp+7DqL5PqqU3p7QQgEkexRJ8VwS4d4Xk5cyjjrDbvzKxXEwCzN96k0RfNIpiQ2YsSbcyoRxX5fFzllf+dOy8uCMy/ebkjS0ONwuzkRR55zP3zjA4e+C969ch1Ab9LgcrpUqQ8MOvEXnQmqQkxXM4x4nA1f+jJAAD+8tXq96F285yEhMGbOWPp352/yRnzPmWRyRibmd800tluUOW4IyIZz2Hw1xYA9/xsSgKy0yQS//7BbNldSPJ+MCX/mxBqrttfeQX/mf/+2AcZ7Z1wDrdNoOnt8ISIu33p34GUAqqEyPwtdrhMf55SjsQmwUtm/I/qPiQ0ZOPv5ci6kCP0Ddb2jRr38UOXvi54DlMMmkxTs7j/J3jUQYepc7xEgTVVJFcf+8P//18L/9gA//9fHuK11sTUs+RbYDzZKn0uM2PnTEUpAJAT4wETKSU4KDn3wR3rUGGycaAPVo40AjzO5g7VkynMJuo3M2vclcmfmS3ygBGDqjGHQybd03tnGdkGOCGLLlTEABV+UgJxYxH2YxRX9zUULXFnfnBpPxQcOdq+zs1zi9uI2iAqSKwh8Dhch2Ytz8iaZLW39S3+3pmGdITR49+nlHjcG4xNVSYRLFLRmEj/H/I+7qd90N6AF9aRDuUFH1O7ONRGjEQGvPMEF0Fj5atb5w9tjc1pcKTsaWvT2GbF9NQ31HmpaLAgs8szVbuEC8GHKCESxKmx+Hrh5ZwjrNihG3KL0H1n3/g/WetlSYEFYsYTXQmgyUGCVIILkJYrRIdLB5iVAPrseYWKCT8HgJuCUAhaqRO+6jn1fkplsC0yYCveVI+yyDsVr98kmO5arhQ3u+aqKVUvJ8xZL9as4008lN9DkKcRhvC4BwWdhupsqUYwLQmaQhLxP15875P/r43c8r4NI4sLDiCi7Rzww1dWNTyThiA07x8b/zTaFC9Sz+jtZDpRPoSf3LS+TmvHZQ+yv/N9nSK/CGpimH/qjTJOQRStf5ppvzT0FzGMX2tqNndJbZD8idLxJFXZekFF16KC/6scsX/lTNL+XFfTqsreVXu7bL/wjNVTPeGkJJE7aWcXP2+3qQTMv+LaO9INAsG3cyp5Co/F06O8XoVtYZXBjH3f3r9Y8Wp89/fqq2OfQSD2/Ujo1t0fNnMA14gpYdtm6+/RcRgNQGIPPGxgAaFjsfC4+63CcHr1nczuKyXiQjmoIH7n/0NCmJv3O+v/Lp30d3n/060TaO5ffQGrrx0O7TYUAC6pdQxfOeuX4/EsKJgMKTW18feF5m1SX4ODnH1SWutwnm5T/k0/l4YXbLUi8QRbdtx74QL9DJRtKP8bDT+yyf//8nAEApoAJH7jMHoQv7XKzIUdH1TDS7Phokc3PP5m68+eUTHU17v50avNmnEHCfybI4FC35LTpSaGqRsgNJJliiV37VIfbUlfDfgIqZmmxEHmnCQTSg2zcf5+9LWPblYTxLShx/2U34N3Rf/3Zvie6j8SS/8X+Yo+dXhKIyg1WX040AAAAA==")
function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded
@@ -76,8 +77,7 @@ function defaultParameterType(widgetType, connectedNode, connectedWidget) {
if (connectedNode.comfyClass === "CLIPTextEncode") {
paramType = "prompt (positive)"
}
const round = connectedWidget.options?.round
if ((paramType == "number" && round === undefined) || round === 1) {
if (connectedWidget.options?.round === 1) {
paramType = "number (integer)"
}
return paramType
@@ -96,7 +96,7 @@ function valueMatchesType(value, type, options) {
function optionalWidgetValue(widgets, index, fallback) {
const result = widgets.length > index ? widgets[index].value : null
return result === null || result === -1e10 || result === 1e10 ? fallback : result
return result === null || result === 0 ? fallback : result
}
function changeWidgets(node, type, connectedNode, connectedWidget) {
@@ -206,6 +206,8 @@ app.registerExtension({
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_KritaOutput") {
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
} else if (nodeData.name === "ETN_Parameter") {
setupParameterNode(nodeType)
} else if (nodeData.name === "ETN_SendText") {
+21 -120
View File
@@ -1,16 +1,14 @@
import sys
from enum import Enum
import torch
import numpy as np
from pathlib import Path
from typing import Any, NamedTuple
import comfy.samplers
import numpy as np
import server
import torch
from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io
from PIL import Image
import server
import comfy.samplers
from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io
from .nodes import SendImageWebSocket
@@ -75,13 +73,6 @@ class _BasicTypes(str):
BasicTypes = _BasicTypes("BASIC")
class OutputBatchMode(Enum):
default = "default"
images = "images"
animation = "animation"
layers = "layers"
class KritaOutput(io.ComfyNode):
@classmethod
def define_schema(cls):
@@ -89,41 +80,13 @@ class KritaOutput(io.ComfyNode):
node_id="ETN_KritaOutput",
display_name="Krita Output",
category="krita",
inputs=[
io.Image.Input("images"),
io.Int.Input("x", "offset x", default=0),
io.Int.Input("y", "offset y", default=0),
io.String.Input("name", default=""),
io.Combo.Input(
"batch_mode", OutputBatchMode, "batch mode", default=OutputBatchMode.default
),
io.Boolean.Input("resize_canvas", "resize canvas", default=False),
],
inputs=[io.Image.Input("images")],
is_output_node=True,
)
@classmethod
def execute( # type: ignore
cls,
images: torch.Tensor,
x: int = 0,
y: int = 0,
name="",
batch_mode: OutputBatchMode | str = OutputBatchMode.default,
resize_canvas=False,
):
batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
info = {
"name": name,
"offset_x": x,
"offset_y": y,
"batch_mode": batch_mode,
"resize_canvas": resize_canvas,
}
output = SendImageWebSocket.execute(images, "PNG")
assert isinstance(output.ui, dict)
output.ui["info"] = [info]
return output
def execute(cls, images: torch.Tensor):
return SendImageWebSocket.execute(images, "PNG")
class KritaSendText(io.ComfyNode):
@@ -142,7 +105,7 @@ class KritaSendText(io.ComfyNode):
)
@classmethod
def execute(cls, value: Any, name: str, type: str): # type: ignore
def execute(cls, value: Any, name: str, type: str):
mime = {
"text": "text/plain",
"markdown": "text/markdown",
@@ -170,28 +133,12 @@ class KritaCanvas(io.ComfyNode):
io.Int.Output(display_name="width"),
io.Int.Output(display_name="height"),
io.Int.Output(display_name="seed"),
io.Mask.Output(display_name="mask"),
],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
class SelectionContext(Enum):
automatic = "automatic"
entire_image = "entire image"
mask_bounds = "mask bounds"
_selection_context_help = """
Determines the section (crop bounding box) of the image and mask to transmit:
- automatic: area around the selection determined by Krita settings
- entire image: always use the entire canvas area
- mask bounds: tight bounding box of the current selection
This affects the Selection and Canvas nodes. The offset x/y outputs indicate the top-left corner of the context area relative to the full canvas."""
def execute(cls):
return io.NodeOutput(_placeholder_image(), 512, 512, 0)
class KritaSelection(io.ComfyNode):
@@ -201,26 +148,12 @@ class KritaSelection(io.ComfyNode):
node_id="ETN_KritaSelection",
display_name="Krita Selection",
category="krita",
inputs=[
io.Combo.Input(
"context",
options=SelectionContext,
default=SelectionContext.entire_image,
tooltip=_selection_context_help,
),
io.Int.Input("padding", "padding", default=0, min=0),
],
outputs=[
io.Mask.Output("mask", "mask"),
io.Boolean.Output("active", "active"),
io.Int.Output("x", "offset x"),
io.Int.Output("y", "offset y"),
],
outputs=[io.Mask.Output(display_name="mask"), io.Boolean.Output(display_name="active")],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(torch.ones(1, 512, 512), False, 0, 0)
def execute(cls):
return io.NodeOutput(torch.ones(1, 512, 512), False)
class KritaImageLayer(io.ComfyNode):
@@ -238,7 +171,7 @@ class KritaImageLayer(io.ComfyNode):
)
@classmethod
def execute(cls, name: str): # type: ignore
def execute(cls, name: str):
return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512))
@@ -256,7 +189,7 @@ class KritaMaskLayer(io.ComfyNode):
)
@classmethod
def execute(cls, name: str): # type: ignore
def execute(cls, name: str):
return io.NodeOutput(torch.ones(1, 512, 512))
@@ -284,14 +217,14 @@ class Parameter(io.ComfyNode):
io.String.Input("name", default="Parameter"),
io.Combo.Input("type", options=_param_types, default="auto"),
io.String.Input("default", default=""),
io.Float.Input("min", default=-1e10, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("max", default=1e10, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("min", default=0.0, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("max", default=1.0, min=-_fmax, max=_fmax, optional=True),
],
outputs=[io.AnyType.Output(display_name="value")],
)
@classmethod
def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
def execute(cls, name: str, type: str, default, min=0.0, max=1.0):
if type == "number":
return io.NodeOutput(float(default))
elif type == "number (integer)":
@@ -328,37 +261,5 @@ class KritaStyle(io.ComfyNode):
)
@classmethod
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
def execute(cls, name: str, sampler_preset: str):
raise NotImplementedError("This workflow must be started from Krita!")
+23 -50
View File
@@ -1,21 +1,20 @@
from __future__ import annotations
import base64
import time
from copy import copy
from dataclasses import dataclass
from io import BytesIO
import time
from typing import NamedTuple
from uuid import uuid4
from PIL import Image
import numpy as np
import base64
import torch
import torch.nn.functional as F
from io import BytesIO
from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
from comfy_api.latest import io
from PIL import Image
from server import BinaryEventTypes, PromptServer
class LoadImageBase64(io.ComfyNode):
@@ -30,14 +29,14 @@ class LoadImageBase64(io.ComfyNode):
)
@classmethod
def execute(cls, image: str): # type: ignore
def execute(cls, image: str):
_strip_prefix(image, "data:image/png;base64,")
imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata))
if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = torch.from_numpy(mask)[None,]
mask = torch.from_numpy(mask)
else:
mask = None
@@ -60,7 +59,7 @@ class LoadMaskBase64(io.ComfyNode):
)
@classmethod
def execute(cls, mask: str): # type: ignore
def execute(cls, mask: str):
_strip_prefix(mask, "data:image/png;base64,")
imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata))
@@ -86,7 +85,7 @@ class SendImageWebSocket(io.ComfyNode):
)
@classmethod
def execute(cls, images: torch.Tensor, format: str): # type: ignore
def execute(cls, images: torch.Tensor, format: str):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
@@ -108,9 +107,6 @@ class SendImageWebSocket(io.ComfyNode):
class ImageCache:
timeout = 600 # 10 minutes
max_size = 100 * 1024 * 1024 # 100 MB
@dataclass
class Entry:
data: bytes
@@ -118,15 +114,8 @@ class ImageCache:
timestamp: float
retrieved: int
class OldEntry(NamedTuple):
last_used: float
deleted: float
size: int
retrieved: int
def __init__(self):
self.images: dict[str, ImageCache.Entry] = {}
self.old: dict[str, ImageCache.OldEntry] = {}
def add(self, image: Image.Image, format: str):
key = uuid4().hex
@@ -148,13 +137,6 @@ class ImageCache:
def get(self, key: str, extend: bool = False):
entry = self.images.get(key)
if entry is None:
if old := self.old.get(key):
now = time.time()
print(
f"[comfyui-tooling-nodes] requested image {key} has been deleted ",
f"(last used {now - old.last_used:.0f}s ago, deleted {now - old.deleted:.0f}s ago, "
f"size {old.size / 1024**2:.1f}MB, retrieved {old.retrieved} times)",
)
return None, None
entry.retrieved += 1
if extend:
@@ -163,22 +145,14 @@ class ImageCache:
return entry.data, entry.content_type
def prune(self):
total_size = sum(len(entry.data) for entry in self.images.values())
if total_size <= self.max_size:
return
# Remove least recently used entries until under max size
sorted_entries = sorted(self.images.items(), key=lambda item: item[1].timestamp)
now = time.time()
for key, entry in sorted_entries:
age = now - entry.timestamp
if age > self.timeout or (age > 60 and entry.retrieved > 0):
self.old[key] = ImageCache.OldEntry(
entry.timestamp, now, len(entry.data), entry.retrieved
)
del self.images[key]
total_size -= len(entry.data)
if total_size <= self.max_size:
break
keys_to_delete = []
for key, entry in self.images.items():
d = now - entry.timestamp
if (d > 60 and entry.retrieved > 1) or d > 600:
keys_to_delete.append(key)
for key in keys_to_delete:
del self.images[key]
def __contains__(self, key: str):
return key in self.images
@@ -199,13 +173,12 @@ class LoadImageCache(io.ComfyNode):
)
@classmethod
def execute(cls, id: str): # type: ignore
def execute(cls, id: str):
image_data, content_type = image_cache.get(id, extend=True)
if image_data is None:
raise ValueError(f"Image with ID {id} not found in cache.")
img = Image.open(BytesIO(image_data))
w, h = img.size
c = len(img.getbands())
normalized = np.array(img).astype(np.float32) / 255.0
@@ -239,7 +212,7 @@ class SaveImageCache(io.ComfyNode):
)
@classmethod
def execute(cls, images: torch.Tensor, format: str): # type: ignore
def execute(cls, images: torch.Tensor, format: str):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
@@ -286,7 +259,7 @@ class ApplyMaskToImage(io.ComfyNode):
)
@classmethod
def execute(cls, image: torch.Tensor, mask: torch.Tensor): # type: ignore
def execute(cls, image: torch.Tensor, mask: torch.Tensor):
out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
@@ -302,7 +275,7 @@ class ApplyMaskToImage(io.ComfyNode):
# Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]):
alpha = mask[i] if is_mask_batch else mask[0]
out[i, 3, :, :] *= alpha
out[i, 3, :, :] = alpha
return (to_bhwc(out),)
@@ -331,7 +304,7 @@ class ReferenceImage(io.ComfyNode):
)
@classmethod
def execute( # type: ignore
def execute(
cls,
image: torch.Tensor,
weight: float,
@@ -361,7 +334,7 @@ class ApplyReferenceImages(io.ComfyNode):
)
@classmethod
def execute( # type: ignore
def execute(
cls,
conditioning: list[list],
clip_vision: ClipVisionModel,
+1 -4
View File
@@ -1,4 +1,5 @@
from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path
from tqdm import tqdm
import torch
@@ -38,10 +39,6 @@ class CLIPSafetyChecker(PreTrainedModel):
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)
# 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):
with torch.no_grad():
image_batch = self.vision_model(clip_input)[1]
+2 -2
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "3.4.0"
version = "3.0.0"
license = { file = "LICENSE" }
[project.urls]
@@ -13,7 +13,7 @@ line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741", "BLE001"]
ignore = ["E741"]
[tool.black]
line-length = 100
+1 -214
View File
@@ -1,34 +1,16 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
# Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py
# by @laksjdjf
from __future__ import annotations
from functools import partial
from typing import Any, NamedTuple
from typing import NamedTuple
import torch
import torch.nn.functional as F
import math
from torch import Tensor, Size
import comfy.model_management
import comfy.patcher_extension
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
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:
h, w = original_shape[2], original_shape[3]
hm, wm = mask.shape[2], mask.shape[3]
@@ -52,12 +34,6 @@ def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape:
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):
return a * b // math.gcd(a, b)
@@ -169,7 +145,6 @@ class AttentionMaskPatch:
mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
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.num_tokens = [cond.shape[1] for cond in self.conds]
self.num_conds = len(region_list)
@@ -178,8 +153,6 @@ class AttentionMaskPatch:
@staticmethod
def apply(model: ModelPatcher, regions: Region):
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):
assert k.mean() == v.mean(), "k and v must be the same."
@@ -248,189 +221,3 @@ class AttentionMaskPatch:
new_model.set_model_attn2_output_patch(attn2_output_patch)
new_model.set_attachments("etn_attention_mask", patch)
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
+6 -11
View File
@@ -9,13 +9,9 @@ IntArray = npt.NDArray[np.int_]
class TileLayout:
def __init__(
self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
):
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"
def __init__(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8"
assert min_tile_size % 8 == 0, "Tile size must be divisible by 8"
assert blending <= padding, "Blending must be smaller than padding"
self.image_size: IntArray = np.array(image.shape[-3:-1])
@@ -25,7 +21,7 @@ class TileLayout:
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
tile_size = np.ceil(image_size_with_overlap / self.tile_count)
self.tile_size: IntArray = (np.ceil(tile_size / multiple) * multiple).astype(int)
self.tile_size: IntArray = (np.ceil(tile_size / 8) * 8).astype(int)
def size(self, coord: IntArray):
return self.end(coord) - self.start(coord)
@@ -88,14 +84,13 @@ class CreateTileLayout(io.ComfyNode):
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("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")],
)
@classmethod
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, multiple))
def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int):
return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending))
class ExtractImageTile(io.ComfyNode):