Converting refiner works

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent f95a439d62
commit 9fb310700e
3 changed files with 37 additions and 18 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.5 MiB

+19 -5
View File
@@ -52,7 +52,11 @@ def get_unet(model_type: ModelVersion, ref_pipe):
def get_encoder_hidden_states_shape(ref_pipe, batch_size):
text_encoder = ref_pipe.text_encoder
text_encoder = (
ref_pipe.text_encoder_2
if hasattr(ref_pipe, "text_encoder_2")
else ref_pipe.text_encoder
)
text_token_sequence_length = text_encoder.config.max_position_embeddings
hidden_size = (text_encoder.config.hidden_size,)
@@ -166,11 +170,21 @@ def sdxl_inputs(sample_unet_inputs, ref_pipe):
batch_size = sample_shape[0]
h = sample_shape[2] * 8
w = sample_shape[3] * 8
original_size = (h, w) # output_resolution
crops_coords_top_left = (0, 0) # topleft_crop_cond
target_size = (h, w) # resolution_cond
original_size = (h, w)
crops_coords_top_left = (0, 0)
is_refiner = (
hasattr(ref_pipe.config, "requires_aesthetics_score")
and ref_pipe.config.requires_aesthetics_score
)
if is_refiner:
aesthetic_score = (6.0,)
time_ids_list = list(original_size + crops_coords_top_left + aesthetic_score)
else:
target_size = (h, w)
time_ids_list = list(original_size + crops_coords_top_left + target_size)
time_ids_list = list(original_size + crops_coords_top_left + target_size)
time_ids = torch.tensor(time_ids_list).repeat(batch_size, 1).to(torch.int64)
text_embeds_shape = (batch_size, ref_pipe.text_encoder_2.config.hidden_size)
+18 -13
View File
@@ -192,7 +192,7 @@ def is_sdxl_refiner(coreml_model):
)
def sdxl_model_function_wrapper(time_ids, text_embeds):
def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False):
def wrapper(model_function, params):
x = params["input"]
t = params["timestep"]
@@ -203,6 +203,10 @@ def sdxl_model_function_wrapper(time_ids, text_embeds):
if context is None:
return torch.zeros_like(x)
if refiner and context is not None:
# converted refiner accepts only g clip
c["c_crossattn"] = context[:, :, 768:]
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
return wrapper
@@ -214,6 +218,9 @@ def add_sdxl_model_options(model_patcher, positive, negative):
pos_dict = positive[0][1]
neg_dict = negative[0][1]
pos_pooled = pos_dict["pooled_output"]
neg_pooled = neg_dict["pooled_output"]
pos_time_ids = [
pos_dict.get("height", 768),
pos_dict.get("width", 768),
@@ -229,35 +236,33 @@ def add_sdxl_model_options(model_patcher, positive, negative):
]
if model_patcher.model.diffusion_model.is_sdxl_base:
base_pos_time_ids = [
pos_time_ids += [
pos_dict.get("target_height", 768),
pos_dict.get("target_width", 768),
]
pos_time_ids += base_pos_time_ids
base_neg_time_ids = [
neg_time_ids += [
neg_dict.get("target_height", 768),
neg_dict.get("target_width", 768),
]
neg_time_ids += base_neg_time_ids
if model_patcher.model.diffusion_model.is_sdxl_refiner:
refiner_pos_time_ids = [
is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner
if is_refiner:
pos_time_ids += [
pos_dict.get("aesthetic_score", 6),
]
pos_time_ids += refiner_pos_time_ids
refiner_neg_time_ids = [
neg_time_ids += [
neg_dict.get("aesthetic_score", 2.5),
]
neg_time_ids += refiner_neg_time_ids
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
text_embeds = torch.cat((pos_dict["pooled_output"], neg_dict["pooled_output"]))
text_embeds = torch.cat((pos_pooled, neg_pooled))
model_options = {
"model_function_wrapper": sdxl_model_function_wrapper(time_ids, text_embeds),
"model_function_wrapper": sdxl_model_function_wrapper(
time_ids, text_embeds, is_refiner
),
}
mp.model_options |= model_options