Synchronize code to improve inference speed

This commit is contained in:
smthemex
2025-11-24 18:01:01 +08:00
committed by GitHub
parent df683256cc
commit 258e7124ff
+58 -25
View File
@@ -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