Compare commits

...
Author SHA1 Message Date
asagi4 665d6d4f86 Change spot where wrappers are applied 2026-07-01 13:53:46 +03:00
asagi4 08a1e86fda WIP: negpip hack 2026-07-01 13:33:42 +03:00
asagi4 89002e0ba1 Timers for debugging 2026-07-01 13:29:46 +03:00
asagi4 31bdd7cb33 Don't spam DEF expansions at INFO log level 2026-07-01 13:29:46 +03:00
asagi4 d0ace191ba Only force block to GPU if COUPLE is actually in use
This seems to interfere with LoRAs somehow, so minimize impact.
2026-07-01 13:24:46 +03:00
4 changed files with 95 additions and 14 deletions
+73 -4
View File
@@ -6,6 +6,7 @@ from math import lcm
import torch
import torch.nn.functional as F
from comfy.ldm.anima.model import Anima as AnimaDIT
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.patcher_extension import WrapperExecutor
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
@@ -13,6 +14,18 @@ from comfy.samplers import process_conds
COND = 0
UNCOND = 1
COND_NEGPIP_MASK_KEY = "c_ppm_negpip_mask"
NEGPIP_MASKS_COUPLE_KEY = "ppm_couple_negpip_masks"
NEGPIP_MASK_KEY = "ppm_negpip_mask"
class CoupleForward:
def __init__(self, fn):
self.fn = fn
def __call__(self, *args, **kwargs):
return cosmos_attention_forward_couple(self.fn, *args, **kwargs)
def reshape_mask(mask: torch.Tensor, size: tuple[int, int], bs: int, num_tokens: int) -> torch.Tensor:
num_conds = mask.shape[0]
@@ -23,6 +36,23 @@ def reshape_mask(mask: torch.Tensor, size: tuple[int, int], bs: int, num_tokens:
return mask_downsample_reshaped
def wrap_forwards(anima_model):
backups = {}
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
backups[block_name] = b.forward
b.forward = CoupleForward(b.forward)
return backups
def unwrap_forwards(anima_model, backups):
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
b.forward = backups[block_name]
def anima_sample_wrapper(executor, *args, **kwargs):
guider, _, extra_options, _, noise, latent_image, denoise_mask, *_ = args
seed = extra_options["seed"]
@@ -40,10 +70,17 @@ def anima_sample_wrapper(executor, *args, **kwargs):
seed,
latent_shapes=[latent_image.shape],
)
return [
conds_p = [
c["model_conds"]["c_crossattn"].cond * pc_conds[i][1].get("strength", 1.0)
for i, c in enumerate(conds["positive"])
]
# TODO: weight?
negpip_masks_couple = []
if all(COND_NEGPIP_MASK_KEY in cond["model_conds"] for cond in conds["positive"]):
negpip_masks_couple = [
cond["model_conds"][COND_NEGPIP_MASK_KEY].cond.to(device) for cond in conds["positive"]
]
return conds_p, negpip_masks_couple
extra_options["model_options"]["transformer_options"]["pc_process_conds"] = pc_process_conds
return executor(*args, **kwargs)
@@ -57,7 +94,9 @@ def anima_forward_wrapper(executor: WrapperExecutor, *args, **kwargs):
transformer_options: dict = kwargs.get("transformer_options", {}).copy()
pc = transformer_options.get("pc_couple")
if pc and "processed_conds" not in pc:
pc["processed_conds"] = transformer_options["pc_process_conds"](pc["conds"])
conds, negpip_masks = transformer_options["pc_process_conds"](pc["conds"])
pc["processed_conds"] = conds
transformer_options[COND_NEGPIP_MASK_KEY] = negpip_masks if negpip_masks else None
patch_spatial = anima_model.patch_spatial
activations_shape = list(x.shape)
@@ -67,7 +106,13 @@ def anima_forward_wrapper(executor: WrapperExecutor, *args, **kwargs):
transformer_options["activations_shape"] = activations_shape
kwargs["transformer_options"] = transformer_options
return executor(*args, **kwargs)
b = {}
if pc:
b = wrap_forwards(anima_model)
r = executor(*args, **kwargs)
if pc:
unwrap_forwards(anima_model, b)
return r
def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, transformer_options):
@@ -82,6 +127,7 @@ def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, tr
mask = args["mask"]
conds = args["processed_conds"][1:]
num_conds = len(conds) + 1
num_tokens_c: list[int] = [c.shape[1] for c in conds]
cond_or_uncond = transformer_options["cond_or_uncond"]
@@ -90,6 +136,9 @@ def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, tr
num_chunks = len(cond_or_uncond)
bs = x.shape[0] // num_chunks
n = transformer_options.get(NEGPIP_MASK_KEY)
negpip_masks = transformer_options.get(NEGPIP_MASKS_COUPLE_KEY)
has_negpip = n is not None and negpip_masks is not None
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = c.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm(c.shape[1], *num_tokens_c)
@@ -98,7 +147,16 @@ def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, tr
dim=0,
)
xs, cs = [], []
if has_negpip:
n_chunks = n.chunk(num_chunks, dim=0)
num_tokens_n: list[int] = [mask.shape[1] for mask in negpip_masks]
lcm_tokens_n = lcm(*(num_tokens_n + [n.shape[1]]))
conds_n_tensor = torch.cat(
[mask.repeat(bs, lcm_tokens_n // num_tokens_n[i], 1) for i, mask in enumerate(negpip_masks)],
dim=0,
)
xs, cs, ns = [], [], []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // c.shape[1], 1)
@@ -111,9 +169,20 @@ def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, tr
cs.append(torch.cat([c_target, conds_c_tensor], dim=0))
cond_or_uncond_couple.extend(itertools.repeat(COND, num_conds))
if has_negpip:
n_target = n_chunks[i].repeat(1, lcm_tokens_n // n.shape[1], 1)
if cond_type == UNCOND:
ns.append(x_target)
else:
ns.append(torch.cat([n_target, conds_n_tensor], dim=0))
xs = torch.cat(xs, dim=0)
cs = torch.cat(cs, dim=0)
if has_negpip:
ns = torch.cat(ns, dim=0)
transformer_options[NEGPIP_MASK_KEY] = ns
out = _forward(xs, cs, rope_emb, transformer_options)
size = tuple(transformer_options["activations_shape"][-2:])
+1 -2
View File
@@ -88,7 +88,6 @@ def expand_macros(text):
iterations += 1
if iterations > 10:
raise ValueError("Unable to resolve DEFs, make sure there are no cycles!")
return text
for search, replace in replacements:
res = substitute_defcall(res, search, replace)
if res == prevres:
@@ -96,7 +95,7 @@ def expand_macros(text):
prevres = res
if res.strip() != text.strip():
res = res.strip()
log.info("DEFs expanded to: %s", res)
log.debug("DEFs expanded to: %s", res)
return res
+4 -8
View File
@@ -21,7 +21,10 @@ class CoupleForward:
self.block = block
def __call__(self, *args, **kwargs):
self.block.to("cuda")
transformer_options = kwargs["transformer_options"]
pc = transformer_options.get("pc_couple")
if pc:
self.block.to("cuda")
return cosmos_attention_forward_couple(self.fn, *args, **kwargs)
@@ -47,7 +50,6 @@ class PCAnimaAttnCouplePatch(io.ComfyNode):
if issubclass(model_type, Anima):
m = model.clone()
anima_model = model.get_model_object("diffusion_model")
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
cls.__name__,
@@ -59,12 +61,6 @@ class PCAnimaAttnCouplePatch(io.ComfyNode):
anima_sample_wrapper,
)
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
attn_forward_prev = m.get_model_object(f"diffusion_model.{block_name}.forward")
m.add_object_patch(f"diffusion_model.{block_name}.forward", CoupleForward(attn_forward_prev, b))
return io.NodeOutput(m)
+17
View File
@@ -6,6 +6,7 @@ import re
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from timeit import default_timer as timer
from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
if TYPE_CHECKING:
@@ -34,6 +35,22 @@ except ImportError:
log = logging.getLogger("comfyui-prompt-control")
TIME_CONTEXT = []
def push_timer(name):
TIME_CONTEXT.append((name, timer()))
name = ":".join(n[0] for n in TIME_CONTEXT)
log.info("START: %s", name)
def pop_timer():
global TIME_CONTEXT
name = ":".join(n[0] for n in TIME_CONTEXT)
_, start = TIME_CONTEXT.pop()
end = timer()
log.info("%s: took %s seconds", name, round(end - start, 3))
def flatten(x):
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x: