Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0abe742480 | ||
|
|
b3ae4aa2d9 | ||
|
|
1b8d81ce5a | ||
|
|
ca01116495 | ||
|
|
5d3194f4d4 | ||
|
|
d5812f900b | ||
|
|
7064288fbe | ||
|
|
d82675092e | ||
|
|
a1e51904de | ||
|
|
ffa130239b | ||
|
|
2fd51d0d47 | ||
|
|
d3c75155b4 | ||
|
|
cbaef8d9c5 | ||
|
|
b2783d82a6 | ||
|
|
09759222de | ||
|
|
7fc3df1174 | ||
|
|
ed99942f86 | ||
|
|
2d395424ea | ||
|
|
9b9ea62dd8 | ||
|
|
ad36f89af3 | ||
|
|
7130dcb2df | ||
|
|
77186eda87 | ||
|
|
24a7bd1a77 | ||
|
|
2d14a03ad8 |
@@ -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
@@ -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():
|
||||||
|
|||||||
@@ -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
@@ -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,)
|
||||||
@@ -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!")
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user