diff --git a/__init__.py b/__init__.py index 619a4f3..6c7499c 100644 --- a/__init__.py +++ b/__init__.py @@ -30,7 +30,7 @@ if "PYTEST_CURRENT_TEST" not in os.environ: h = logging.StreamHandler(sys.stdout) h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s")) log.addHandler(h) - for node in ["base", "hooks", "tools", "lazy"]: + for node in ["base", "hooks", "tools", "lazy", "anima"]: mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__) v3_modules.append(mod) diff --git a/doc/attention_couple.md b/doc/attention_couple.md index 5ef29ad..aff1060 100644 --- a/doc/attention_couple.md +++ b/doc/attention_couple.md @@ -1,7 +1,5 @@ # Attention Couple -NOTE: This is still considered an experimental feature, so the syntax may change. - Attention Couple is an attention-based implementation of regional prompting. it is faster and often more flexible than latent-based masking. The implementation is based on the one by [pamparamm](https://github.com/pamparamm/ComfyUI-ppm.git), modified to use ComfyUI's hook system. This enables it to work with prompt scheduling. @@ -12,6 +10,11 @@ As a consequence of this, however, you can also use `COUPLE` in your negative pr To enable batching negative prompts, run your positive and negative prompt through the `PPCAttentionCoupleBatchNegative` node. This will make the outputs identical to pamparamm's implementation and will also improve performance. It will fall back to the default behaviour in cases where batching can't be done, so it should always be safe to use. +## Anima + +There is a **very experimental** port of pamparamm's Anima support for Attention Couple in Prompt Control. Because ComfyUI lacks the built-in schedulable hooks required, you must first patch your model with `PC: Anima Attention Couple Model Patch` in addition to using `COUPLE` as usual. + +The code was hacked together with minimal thought, so expect bugs and misbehaviour. The port is also currently *not* compatible with NegPIP. ## Syntax diff --git a/prompt_control/anima_couple.py b/prompt_control/anima_couple.py new file mode 100644 index 0000000..18b2e8b --- /dev/null +++ b/prompt_control/anima_couple.py @@ -0,0 +1,138 @@ + +# Adapted from https://github.com/pamparamm/ComfyUI-ppm +import itertools +from collections.abc import Callable +from math import lcm + +import torch +import torch.nn.functional as F +from comfy.ldm.anima.model import Anima as AnimaDIT +from comfy.patcher_extension import WrapperExecutor +from comfy.sampler_helpers import convert_cond +from comfy.samplers import process_conds + +COND = 0 +UNCOND = 1 + + +def reshape_mask(mask: torch.Tensor, size: tuple[int, int], bs: int, num_tokens: int) -> torch.Tensor: + num_conds = mask.shape[0] + + mask_downsample = F.interpolate(mask, size=size, mode="nearest") + mask_downsample_reshaped = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(bs, dim=0) + + return mask_downsample_reshaped + + +def anima_sample_wrapper(executor, *args, **kwargs): + guider, _, extra_options, _, noise, latent_image, denoise_mask, *_ = args + seed = extra_options["seed"] + device = "cuda" # TODO: fix + + def pc_process_conds(pc_conds): + conds = [convert_cond([c])[0] for c in pc_conds] + conds = process_conds( + guider.inner_model, + noise, + {"positive": conds}, + device, + latent_image, + denoise_mask, + seed, + latent_shapes=[latent_image.shape], + ) + return [c["model_conds"]["c_crossattn"].cond for c in conds["positive"]] + + extra_options["model_options"]["transformer_options"]["pc_process_conds"] = pc_process_conds + return executor(*args, **kwargs) + + +def anima_forward_wrapper(executor: WrapperExecutor, *args, **kwargs): + """Model wrapper does something with activation shapes?""" + anima_model: AnimaDIT = executor.class_obj # type: ignore + + x: torch.Tensor = args[0] + 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"]) + patch_spatial = anima_model.patch_spatial + + 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 + kwargs["transformer_options"] = transformer_options + + return executor(*args, **kwargs) + + +def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, transformer_options): + """attention block wrapper""" + if "pc_couple" not in transformer_options: + return _forward(x, context, rope_emb, transformer_options) + c: torch.Tensor = context + + args = transformer_options["pc_couple"] + + 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"] + cond_or_uncond_couple = [] + + num_chunks = len(cond_or_uncond) + bs = x.shape[0] // num_chunks + + 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) + conds_c_tensor = torch.cat( + [cond.repeat(bs, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)], + dim=0, + ) + + xs, cs = [], [] + 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) + if cond_type == UNCOND: + xs.append(x_target) + cs.append(c_target) + cond_or_uncond_couple.append(UNCOND) + else: + xs.append(x_target.repeat(num_conds, 1, 1)) + cs.append(torch.cat([c_target, conds_c_tensor], dim=0)) + cond_or_uncond_couple.extend(itertools.repeat(COND, num_conds)) + + xs = torch.cat(xs, dim=0) + cs = torch.cat(cs, dim=0) + + out = _forward(xs, cs, rope_emb, transformer_options) + + size = tuple(transformer_options["activations_shape"][-2:]) + num_tokens = out.shape[1] + mask_downsample = reshape_mask(mask, size, bs, num_tokens) + + outputs = [] + cond_outputs = [] + i_cond = 0 + + for i, cond_type in enumerate(cond_or_uncond_couple): + pos, next_pos = i * bs, (i + 1) * bs + + if cond_type == UNCOND: + outputs.append(out[pos:next_pos]) + else: + pos_cond, next_pos_cond = i_cond * bs, (i_cond + 1) * bs + masked_output = out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond] + cond_outputs.append(masked_output) + i_cond += 1 + + if len(cond_outputs) > 0: + cond_output = torch.stack(cond_outputs).sum(0) + outputs.append(cond_output) + + return torch.cat(outputs, dim=0) diff --git a/prompt_control/attention_couple_ppm.py b/prompt_control/attention_couple_ppm.py index 19d015a..c3aaeeb 100644 --- a/prompt_control/attention_couple_ppm.py +++ b/prompt_control/attention_couple_ppm.py @@ -63,12 +63,14 @@ class AttentionCoupleHook(TransformerOptionsHook): def __init__(self): super().__init__(hook_scope=EnumHookScope.HookedOnly) - self.transformers_dict = { + self.transformers_dict: dict[str, Any] = { "patches": { "attn2_output_patch": [Proxy(self.attn2_output_patch)], "attn2_patch": [Proxy(self.attn2_patch)], - } + }, + "pc_couple": {}, } + self.has_negpip = False # The list will be calculated later. All clones must refer to the same kv dict self.kv: dict[str, list] = {"k": None, "v": None} # type: ignore @@ -77,6 +79,7 @@ class AttentionCoupleHook(TransformerOptionsHook): self.num_conds = len(conds) + 1 self.base_strength = base_cond[1].get("strength", 1.0) self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds] + self.comfy_conds = [base_cond] + conds self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds] base_mask = base_cond[1].get("mask", None) masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds] @@ -116,6 +119,11 @@ class AttentionCoupleHook(TransformerOptionsHook): self.mask = mask / mask.sum(dim=0, keepdim=True) def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]): + self.transformers_dict["pc_couple"] = { + "conds": self.comfy_conds, + "num_conds": self.num_conds, + "mask": self.mask, + } if self.kv["k"] is None: self.has_negpip = model.model_options.get("ppm_negpip", False) log.debug("AttentionCouple has_negpip=%s", self.has_negpip) diff --git a/prompt_control/nodes_anima.py b/prompt_control/nodes_anima.py new file mode 100644 index 0000000..e0b0a2d --- /dev/null +++ b/prompt_control/nodes_anima.py @@ -0,0 +1,64 @@ +# Adapted from ComfyUI-ppm into hook form + +from functools import partial + +import comfy.model_management +import comfy.patcher_extension +from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention +from comfy.model_base import Anima +from comfy.model_patcher import ModelPatcher +from comfy_api.latest import io + +from .anima_couple import ( + anima_forward_wrapper, + anima_sample_wrapper, + cosmos_attention_forward_couple, +) + + +class PCAnimaAttnCouplePatch(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="PCAnimaAttnCouplePatch", + display_name="PC: Anima attention Couple Model Patch", + category="promptcontrol/experimental", + inputs=[ + io.Model.Input("model"), + ], + outputs=[ + io.Model.Output(), + ], + ) + + @classmethod + def execute(cls, model: ModelPatcher) -> io.NodeOutput: + model_type = type(model.model) + m = model + + 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__, + anima_forward_wrapper, + ) + m.add_wrapper_with_key( + comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, + cls.__name__, + anima_sample_wrapper, + ) + + for block_name, _ 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", partial(cosmos_attention_forward_couple, attn_forward_prev) + ) + + return io.NodeOutput(m) + + +NODES = [PCAnimaAttnCouplePatch]