Converting refiner works
This commit is contained in:
Binary file not shown.
|
After Width: | Height: | Size: 1.5 MiB |
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user