Generating SDXL with Core ML Sampler works
This commit is contained in:
+107
-3
@@ -46,6 +46,8 @@ class CoreMLInputs:
|
||||
self.t = t
|
||||
self.context = context
|
||||
self.control = control
|
||||
self.time_ids = kwargs.get("time_ids")
|
||||
self.text_embeds = kwargs.get("text_embeds")
|
||||
self.ts_cond = kwargs.get("timestep_cond")
|
||||
|
||||
def coreml_kwargs(self, expected_inputs):
|
||||
@@ -64,11 +66,22 @@ class CoreMLInputs:
|
||||
residual_kwargs = extract_residual_kwargs(expected_inputs, self.control)
|
||||
model_input_kwargs |= residual_kwargs
|
||||
|
||||
# LCM
|
||||
if self.ts_cond is not None:
|
||||
model_input_kwargs["timestep_cond"] = (
|
||||
self.ts_cond.cpu().numpy().astype(np.float16)
|
||||
)
|
||||
|
||||
# SDXL
|
||||
if "text_embeds" in expected_inputs:
|
||||
model_input_kwargs["text_embeds"] = (
|
||||
self.text_embeds.cpu().numpy().astype(np.float16)
|
||||
)
|
||||
if "time_ids" in expected_inputs:
|
||||
model_input_kwargs["time_ids"] = (
|
||||
self.time_ids.cpu().numpy().astype(np.float16)
|
||||
)
|
||||
|
||||
return model_input_kwargs
|
||||
|
||||
def chunks(self, expected_inputs):
|
||||
@@ -90,9 +103,100 @@ class CoreMLInputs:
|
||||
ts_cond_shape = expected_inputs["timestep_cond"]["shape"]
|
||||
chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape)
|
||||
|
||||
chunked_time_ids = [None] * len(chunked_x)
|
||||
if expected_inputs["time_ids"] is not None:
|
||||
time_ids_shape = expected_inputs["time_ids"]["shape"]
|
||||
if self.time_ids is None:
|
||||
self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to(
|
||||
self.x.device
|
||||
)
|
||||
chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape)
|
||||
|
||||
chunked_text_embeds = [None] * len(chunked_x)
|
||||
if expected_inputs["text_embeds"] is not None:
|
||||
text_embeds_shape = expected_inputs["text_embeds"]["shape"]
|
||||
if self.text_embeds is None:
|
||||
self.text_embeds = torch.zeros(
|
||||
len(chunked_x), *text_embeds_shape[1:]
|
||||
).to(self.x.device)
|
||||
chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape)
|
||||
|
||||
return [
|
||||
CoreMLInputs(x, t, context, control, timestep_cond=ts_cond)
|
||||
for x, t, context, control, ts_cond in zip(
|
||||
chunked_x, ts, chunked_context, chunked_control, chunked_ts_cond
|
||||
CoreMLInputs(
|
||||
x,
|
||||
t,
|
||||
context,
|
||||
control,
|
||||
timestep_cond=ts_cond,
|
||||
time_ids=time_ids,
|
||||
text_embeds=text_embeds,
|
||||
)
|
||||
for x, t, context, control, ts_cond, time_ids, text_embeds in zip(
|
||||
chunked_x,
|
||||
ts,
|
||||
chunked_context,
|
||||
chunked_control,
|
||||
chunked_ts_cond,
|
||||
chunked_time_ids,
|
||||
chunked_text_embeds,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def is_sdxl(coreml_model):
|
||||
return (
|
||||
"time_ids" in coreml_model.expected_inputs
|
||||
and "text_embeds" in coreml_model.expected_inputs
|
||||
)
|
||||
|
||||
|
||||
def sdxl_model_function_wrapper(time_ids, text_embeds):
|
||||
def wrapper(model_function, params):
|
||||
x = params["input"]
|
||||
t = params["timestep"]
|
||||
c = params["c"]
|
||||
|
||||
context = c.get("c_crossattn")
|
||||
|
||||
if context is None:
|
||||
return torch.zeros_like(x)
|
||||
|
||||
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def add_sdxl_model_options(model_patcher, positive, negative):
|
||||
mp = model_patcher.clone()
|
||||
|
||||
pos_dict = positive[0][1]
|
||||
neg_dict = negative[0][1]
|
||||
|
||||
pos_time_ids = [
|
||||
pos_dict.get("height", 768),
|
||||
pos_dict.get("width", 768),
|
||||
pos_dict.get("crop_h", 0),
|
||||
pos_dict.get("crop_w", 0),
|
||||
pos_dict.get("target_height", 768),
|
||||
pos_dict.get("target_width", 768),
|
||||
]
|
||||
|
||||
neg_time_ids = [
|
||||
neg_dict.get("height", 768),
|
||||
neg_dict.get("width", 768),
|
||||
neg_dict.get("crop_h", 0),
|
||||
neg_dict.get("crop_w", 0),
|
||||
neg_dict.get("target_height", 768),
|
||||
neg_dict.get("target_width", 768),
|
||||
]
|
||||
|
||||
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
|
||||
|
||||
text_embeds = torch.cat((pos_dict["pooled_output"], neg_dict["pooled_output"]))
|
||||
|
||||
model_options = {
|
||||
"model_function_wrapper": sdxl_model_function_wrapper(time_ids, text_embeds),
|
||||
}
|
||||
mp.model_options |= model_options
|
||||
|
||||
return mp
|
||||
|
||||
+10
-2
@@ -15,7 +15,7 @@ from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm
|
||||
from coreml_suite.logger import logger
|
||||
from nodes import KSampler, LoraLoader
|
||||
|
||||
from coreml_suite.models import CoreMLModelWrapper
|
||||
from coreml_suite.models import CoreMLModelWrapper, add_sdxl_model_options, is_sdxl
|
||||
from coreml_suite.config import get_model_config
|
||||
|
||||
|
||||
@@ -58,6 +58,9 @@ class CoreMLSampler(COREML_NODE, KSampler):
|
||||
negative is not None
|
||||
), "Negative conditioning is optional only for LCM models."
|
||||
|
||||
if is_sdxl(coreml_model):
|
||||
model_patcher = add_sdxl_model_options(model_patcher, positive, negative)
|
||||
|
||||
return super().sample(
|
||||
model_patcher,
|
||||
seed,
|
||||
@@ -84,7 +87,12 @@ class CoreMLSampler(COREML_NODE, KSampler):
|
||||
def get_model_patcher(self, coreml_model):
|
||||
model_config = get_model_config()
|
||||
wrapped_model = CoreMLModelWrapper(coreml_model)
|
||||
model = model_base.BaseModel(model_config, device=get_torch_device())
|
||||
|
||||
if is_sdxl(coreml_model):
|
||||
model = model_base.SDXL(model_config, device=get_torch_device())
|
||||
else:
|
||||
model = model_base.BaseModel(model_config, device=get_torch_device())
|
||||
|
||||
model.diffusion_model = wrapped_model
|
||||
model_patcher = ModelPatcher(model, get_torch_device(), None)
|
||||
return model_patcher
|
||||
|
||||
Reference in New Issue
Block a user