diff --git a/objectclear/pipelines/pipeline_objectclear.py b/objectclear/pipelines/pipeline_objectclear.py index d64ab25..f77c4a8 100644 --- a/objectclear/pipelines/pipeline_objectclear.py +++ b/objectclear/pipelines/pipeline_objectclear.py @@ -14,7 +14,7 @@ import inspect from typing import Any, Callable, Dict, List, Optional, Tuple, Union - +import os import numpy as np import PIL.Image import torch @@ -467,9 +467,10 @@ class ObjectClearPipeline( if self.config.apply_attention_guided_fusion: self.cross_attention_scores = {} - self.unet = self.unet_store_cross_attention_scores( - self.unet, self.cross_attention_scores - ) + # self.unet = self.unet_store_cross_attention_scores( + # self.unet, self.cross_attention_scores + # ) + self.original_state = None @classmethod @@ -550,7 +551,7 @@ class ObjectClearPipeline( print(image_embeds.shape,uncond_image_embeds.shape,123) return image_embeds, uncond_image_embeds - def unet_store_cross_attention_scores(self, unet, attention_scores): + def unet_store_cross_attention_scores(self, unet, attention_scores,applicable_layers=None): from diffusers.models.attention_processor import ( Attention, AttnProcessor, @@ -558,34 +559,42 @@ class ObjectClearPipeline( ) import types - UNET_LAYER_NAMES = [ - "down_blocks.0", - "down_blocks.1", - "down_blocks.2", - "mid_block", - "up_blocks.1", - "up_blocks.2", - "up_blocks.3", - ] - - start_layer = 0 - end_layer = 2 - applicable_layers = UNET_LAYER_NAMES[start_layer:end_layer] + # UNET_LAYER_NAMES = [ + # "down_blocks.0", + # "down_blocks.1", + # "down_blocks.2", + # "mid_block", + # "up_blocks.1", + # "up_blocks.2", + # "up_blocks.3", + # ] + # start_layer = 0 + # end_layer = 2 + # applicable_layers = UNET_LAYER_NAMES[start_layer:end_layer] + TARGET_LAYER = "down_blocks.1.attentions.0.transformer_blocks.0.attn2" + original_state = {} def make_new_get_attention_scores_fn(name): def new_get_attention_scores(module, query, key, attention_mask=None): attention_probs = module.old_get_attention_scores( query, key, attention_mask ) - attention_scores[name] = attention_probs + #attention_scores[name] = attention_probs + if name == TARGET_LAYER: + attention_scores[name] = attention_probs return attention_probs return new_get_attention_scores for name, module in unet.named_modules(): - if isinstance(module, Attention) and "attn2" in name: - if not any(layer in name for layer in applicable_layers): - continue + # if isinstance(module, Attention) and "attn2" in name: + # if not any(layer in name for layer in applicable_layers): + # continue + if isinstance(module, Attention) and name == TARGET_LAYER and "attn2" in name: + original_state[name] = { + "processor": module.processor, + "get_attention_scores": module.get_attention_scores + } if isinstance(module.processor, AttnProcessor2_0): module.set_processor(AttnProcessor()) module.old_get_attention_scores = module.get_attention_scores @@ -593,8 +602,20 @@ class ObjectClearPipeline( make_new_get_attention_scores_fn(name), module ) module.get_attention_scores = module.new_get_attention_scores + return unet, original_state + #return unet - return unet + def unet_restore_attention_processor(self, unet, original_state): + from diffusers.models.attention_processor import Attention + + for name, module in unet.named_modules(): + if isinstance(module, Attention) and "attn2" in name and name in original_state: + module.get_attention_scores = original_state[name]["get_attention_scores"] + module.set_processor(original_state[name]["processor"]) + if hasattr(module, "old_get_attention_scores"): + delattr(module, "old_get_attention_scores") + if hasattr(module, "new_get_attention_scores"): + delattr(module, "new_get_attention_scores") def resize_attn_map_divide2(self, attn_map, mask, fuse_index): bxh, num_noise_latents, num_text_tokens = attn_map.shape @@ -1887,6 +1908,11 @@ class ObjectClearPipeline( for i, t in enumerate(timesteps): if self.interrupt: continue + if i == len(timesteps) - 1 and self.config.apply_attention_guided_fusion: + self.unet, self.original_state = self.unet_store_cross_attention_scores( + self.unet, + self.cross_attention_scores + ) # expand the latents if we are doing classifier free guidance latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents @@ -1930,7 +1956,8 @@ class ObjectClearPipeline( # progressive attention mask blending fuse_index = 5 if self.config.apply_attention_guided_fusion: - if i == len(timesteps) - 1: + #if i == len(timesteps) - 1: + if i == len(timesteps) - 1 and self.config.apply_attention_guided_fusion: attn_key, attn_map = next(iter(self.cross_attention_scores.items())) attn_map = self.resize_attn_map_divide2(attn_map, mask, fuse_index) init_latents_proper = image_latents @@ -1939,7 +1966,13 @@ class ObjectClearPipeline( else: init_mask = attn_map attn_map = init_mask - self.clear_cross_attention_scores(self.cross_attention_scores) + #self.clear_cross_attention_scores(self.cross_attention_scores) + self.unet = self.unet_restore_attention_processor( + self.unet, + self.original_state + ) + + self.clear_cross_attention_scores(self.cross_attention_scores) if num_channels_unet == 4: init_latents_proper = image_latents