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:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user