Apply Anima attention wrappers just-in-time before execution

This avoids using stale block references and seems to fix LoRAs
This commit is contained in:
asagi4
2026-07-01 19:00:15 +03:00
parent 6547749a0f
commit 0a698eb7ab
2 changed files with 26 additions and 20 deletions
+26 -1
View File
@@ -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):
-19
View File
@@ -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)