Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
665d6d4f86 | ||
|
|
08a1e86fda | ||
|
|
89002e0ba1 | ||
|
|
31bdd7cb33 | ||
|
|
d0ace191ba |
@@ -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:])
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user