diff --git a/region.py b/region.py index b2fd8c4..3a18005 100644 --- a/region.py +++ b/region.py @@ -1,16 +1,34 @@ +# Adapted from https://github.com/pamparamm/ComfyUI-ppm # Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py # by @laksjdjf from __future__ import annotations -from typing import NamedTuple +from functools import partial +from typing import Any, NamedTuple import torch import torch.nn.functional as F import math from torch import Tensor, Size +import comfy.model_management +import comfy.patcher_extension from comfy.model_patcher import ModelPatcher +from comfy.model_base import Anima, CosmosPredict2 +from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention +from comfy.sampler_helpers import convert_cond +from comfy.samplers import process_conds from comfy_api.latest import io +COND = 0 +UNCOND = 1 +ANIMA_COUPLE_WRAPPER_KEY = "etn_attention_mask_anima" +ANIMA_COUPLE_PATCH_KEY = "etn_attention_mask_patch" +CONDS_COUPLE_KEY = "etn_couple_conds" +COND_UNCOND_COUPLE_KEY = "etn_couple_cond_or_uncond" +COUPLE_ACTIVE_KEY = "etn_couple_active" +NUM_TOKENS_COUPLE_KEY = "etn_couple_num_tokens" + + def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: h, w = original_shape[2], original_shape[3] hm, wm = mask.shape[2], mask.shape[3] @@ -34,6 +52,12 @@ def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: return result +def reshape_mask(mask: Tensor, size: tuple[int, int], batch: int, target_size: int) -> Tensor: + result = F.interpolate(mask, size=size, mode="nearest") + result = result.view(mask.shape[0], target_size, 1) + return result.repeat_interleave(batch, dim=0) + + def lcm(a: int, b: int): return a * b // math.gcd(a, b) @@ -145,6 +169,7 @@ class AttentionMaskPatch: mask_sum = mask.sum(dim=0, keepdim=True) assert mask_sum.sum() > 0, "There are areas that are zero in all masks." self.mask = mask / mask_sum + self.region_conds = [r.conditioning for r in region_list] self.conds = [r.conditioning[0][0] for r in region_list] self.num_tokens = [cond.shape[1] for cond in self.conds] self.num_conds = len(region_list) @@ -153,6 +178,8 @@ class AttentionMaskPatch: @staticmethod def apply(model: ModelPatcher, regions: Region): patch = AttentionMaskPatch(regions.preprocess()) + if _is_anima_couple_model(model): + return patch.apply_anima(model) def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict): assert k.mean() == v.mean(), "k and v must be the same." @@ -221,3 +248,189 @@ class AttentionMaskPatch: new_model.set_model_attn2_output_patch(attn2_output_patch) new_model.set_attachments("etn_attention_mask", patch) return new_model + + def apply_anima(self, model: ModelPatcher): + new_model = model.clone() + _patch_cosmos_attention(new_model) + + device = comfy.model_management.get_torch_device() + conds_converted = [convert_cond(cond)[0] for cond in self.region_conds] + new_model.add_wrapper_with_key( + comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, + ANIMA_COUPLE_WRAPPER_KEY, + _anima_couple_sample_wrapper(conds_converted, device), + ) + new_model.add_wrapper_with_key( + comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, + ANIMA_COUPLE_WRAPPER_KEY, + _anima_couple_diffusion_wrapper(self), + ) + new_model.set_attachments("etn_attention_mask", self) + return new_model + + +def _is_anima_couple_model(model: ModelPatcher) -> bool: + model_type = type(model.model) + return issubclass(model_type, (Anima, CosmosPredict2)) + + +def _anima_couple_sample_wrapper(conds_converted: list, device): + def sample_wrapper(executor, *args, **kwargs): + if len(conds_converted) > 0: + guider = args[0] + extra_options: dict[str, Any] = args[2] + seed: int = extra_options["seed"] + noise: Tensor = args[4] + latent_image: Tensor = args[5] + denoise_mask: Tensor | None = args[6] + + conds_processed = process_conds( + guider.inner_model, + noise, + {"positive": conds_converted}, + device, + latent_image, + denoise_mask, + seed, + latent_shapes=[latent_image.shape], + )["positive"] + + conds_couple = [cond["model_conds"]["c_crossattn"].cond for cond in conds_processed] + + model_options: dict[str, Any] = extra_options["model_options"] + transformer_options: dict[str, Any] = model_options.get("transformer_options", {}).copy() + transformer_options[CONDS_COUPLE_KEY] = conds_couple + transformer_options[NUM_TOKENS_COUPLE_KEY] = [cond.shape[1] for cond in conds_couple] + model_options["transformer_options"] = transformer_options + + return executor(*args, **kwargs) + + return sample_wrapper + + +def _anima_couple_diffusion_wrapper(patch: AttentionMaskPatch): + def diffusion_wrapper(executor, *args, **kwargs): + anima_model = executor.class_obj + x: Tensor = args[0] + transformer_options: dict[str, Any] = kwargs.get("transformer_options", {}).copy() + patch_spatial = getattr(anima_model, "patch_spatial", 1) + + 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 + transformer_options[ANIMA_COUPLE_PATCH_KEY] = patch + kwargs["transformer_options"] = transformer_options + + return executor(*args, **kwargs) + + return diffusion_wrapper + + +def pre_cross_attention( + patch: AttentionMaskPatch, + transformer_options: dict, + x: Tensor, + context: Tensor, + rope_emb: Tensor | None, +) -> tuple[Tensor, Tensor, Tensor | None, dict]: + transformer_options = transformer_options.copy() + if CONDS_COUPLE_KEY not in transformer_options: + transformer_options[COND_UNCOND_COUPLE_KEY] = list(transformer_options["cond_or_uncond"]) + transformer_options[COUPLE_ACTIVE_KEY] = False + return x, context, rope_emb, transformer_options + + conds: list[Tensor] = transformer_options[CONDS_COUPLE_KEY] + num_tokens_c: list[int] = transformer_options[NUM_TOKENS_COUPLE_KEY] + cond_or_uncond = transformer_options["cond_or_uncond"] + + num_chunks = len(cond_or_uncond) + batch = x.shape[0] // num_chunks + x_chunks = x.chunk(num_chunks, dim=0) + c_chunks = context.chunk(num_chunks, dim=0) + lcm_tokens_c = lcm_for_list(num_tokens_c + [context.shape[1]]) + conds_c_tensor = torch.cat( + [cond.repeat(batch, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)], + dim=0, + ) + + xs, cs = [], [] + cond_or_uncond_couple = [] + for i, cond_type in enumerate(cond_or_uncond): + x_target = x_chunks[i] + c_target = c_chunks[i].repeat(1, lcm_tokens_c // context.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(patch.num_conds, 1, 1)) + cs.append(conds_c_tensor) + cond_or_uncond_couple.extend([COND] * patch.num_conds) + + transformer_options[COND_UNCOND_COUPLE_KEY] = cond_or_uncond_couple + transformer_options[COUPLE_ACTIVE_KEY] = True + + return torch.cat(xs, dim=0), torch.cat(cs, dim=0), rope_emb, transformer_options + + +def cross_attention_output(patch: AttentionMaskPatch, transformer_options: dict, out: Tensor): + cond_or_uncond = transformer_options[COND_UNCOND_COUPLE_KEY] + size = tuple(transformer_options["activations_shape"][-2:]) + batch = out.shape[0] // len(cond_or_uncond) + mask = patch.mask.to(out.device, dtype=out.dtype) + mask_downsample = reshape_mask(mask, size, batch, out.shape[1]) + + outputs = [] + cond_outputs = [] + i_cond = 0 + for i, cond_type in enumerate(cond_or_uncond): + pos, next_pos = i * batch, (i + 1) * batch + if cond_type == UNCOND: + outputs.append(out[pos:next_pos]) + else: + pos_cond, next_pos_cond = i_cond * batch, (i_cond + 1) * batch + cond_outputs.append(out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond]) + i_cond += 1 + + if len(cond_outputs) > 0: + outputs.append(torch.stack(cond_outputs).sum(0)) + + return torch.cat(outputs, dim=0) + + +def _patch_cosmos_attention(model_patcher: ModelPatcher): + cosmos_model = model_patcher.get_model_object("diffusion_model") + for block_name, block in ( + (n, b) + for n, b in cosmos_model.named_modules() + if ("cross_attn" in n or "self_attn" in n) and isinstance(b, CosmosAttention) + ): + patch_name = f"diffusion_model.{block_name}.forward" + if patch_name not in model_patcher.object_patches: + model_patcher.add_object_patch(patch_name, partial(_cosmos_attention_forward_patched, block)) + + +def _cosmos_attention_forward_patched( + self, + x: Tensor, + context: Tensor | None = None, + rope_emb: Tensor | None = None, + transformer_options: dict | None = None, +) -> Tensor: + transformer_options = transformer_options if transformer_options is not None else {} + patch: AttentionMaskPatch | None = transformer_options.get(ANIMA_COUPLE_PATCH_KEY) + + if context is not None and patch is not None: + x, context, rope_emb, transformer_options = pre_cross_attention( + patch, transformer_options, x, context, rope_emb + ) + + q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb) + output = self.compute_attention(q, k, v, transformer_options=transformer_options) + + if context is not None and patch is not None and transformer_options.get(COUPLE_ACTIVE_KEY, False): + output = cross_attention_output(patch, transformer_options, output) + + return output