diff --git a/prompt_control/anima_couple.py b/prompt_control/anima_couple.py index ac3382b..650c153 100644 --- a/prompt_control/anima_couple.py +++ b/prompt_control/anima_couple.py @@ -1,11 +1,13 @@ # Adapted from https://github.com/pamparamm/ComfyUI-ppm import itertools from collections.abc import Callable +from functools import partial 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 @@ -23,6 +25,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 = partial(cosmos_attention_forward_couple, 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"] @@ -67,7 +86,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): diff --git a/prompt_control/nodes_anima.py b/prompt_control/nodes_anima.py index 7b3cfb3..70f994e 100644 --- a/prompt_control/nodes_anima.py +++ b/prompt_control/nodes_anima.py @@ -3,7 +3,6 @@ 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 @@ -11,20 +10,9 @@ from comfy_api.latest import io from .anima_couple import ( anima_forward_wrapper, anima_sample_wrapper, - cosmos_attention_forward_couple, ) -class CoupleForward: - def __init__(self, fn, block): - self.fn = fn - self.block = block - - def __call__(self, *args, **kwargs): - self.block.to("cuda") - return cosmos_attention_forward_couple(self.fn, *args, **kwargs) - - class PCAnimaAttnCouplePatch(io.ComfyNode): @classmethod def define_schema(cls) -> io.Schema: @@ -47,7 +35,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 +46,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)