Synchronize code to improve inference speed
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user