Add support for CN models to LCM

This commit is contained in:
aszc-dev
2023-11-03 01:27:48 +01:00
parent b90591dfd4
commit 1937f39cca
4 changed files with 45 additions and 39 deletions
+25 -15
View File
@@ -98,7 +98,7 @@ def load_coreml_model(out_path):
def convert_to_coreml(
submodule_name, torchscript_module, sample_inputs, output_names, out_path
submodule_name, torchscript_module, sample_inputs, output_names, out_path
):
if os.path.exists(out_path):
logger.info(f"Skipping export because {out_path} already exists")
@@ -171,6 +171,7 @@ def get_unet_inputs_spec(sample_unet_inputs):
def add_cnet_support(sample_shape, reference_unet):
from python_coreml_stable_diffusion.unet import calculate_conv2d_output_shape
additional_residuals_shapes = []
batch_size = sample_shape[0]
@@ -183,27 +184,32 @@ def add_cnet_support(sample_shape, reference_unet):
reference_unet.conv_in,
)
additional_residuals_shapes.append(
(batch_size, reference_unet.conv_in.out_channels, out_h, out_w))
(batch_size, reference_unet.conv_in.out_channels, out_h, out_w)
)
# down_blocks
for down_block in reference_unet.down_blocks:
additional_residuals_shapes += [
(batch_size, resnet.out_channels, out_h, out_w) for resnet in
down_block.resnets
(batch_size, resnet.out_channels, out_h, out_w)
for resnet in down_block.resnets
]
if hasattr(down_block,
"downsamplers") and down_block.downsamplers is not None:
if hasattr(down_block, "downsamplers") and down_block.downsamplers is not None:
for downsampler in down_block.downsamplers:
out_h, out_w = calculate_conv2d_output_shape(out_h, out_w,
downsampler.conv)
out_h, out_w = calculate_conv2d_output_shape(
out_h, out_w, downsampler.conv
)
additional_residuals_shapes.append(
(batch_size, down_block.downsamplers[-1].conv.out_channels, out_h,
out_w))
(
batch_size,
down_block.downsamplers[-1].conv.out_channels,
out_h,
out_w,
)
)
# mid_block
additional_residuals_shapes.append(
(
batch_size, reference_unet.mid_block.resnets[-1].out_channels, out_h, out_w)
(batch_size, reference_unet.mid_block.resnets[-1].out_channels, out_h, out_w)
)
additional_inputs = {}
@@ -215,8 +221,10 @@ def add_cnet_support(sample_shape, reference_unet):
def convert(
out_path: str, batch_size: int = 1, sample_size: tuple[int, int] = (64, 64),
controlnet_support: bool = False
out_path: str,
batch_size: int = 1,
sample_size: tuple[int, int] = (64, 64),
controlnet_support: bool = False,
):
coreml_unet, ref_unet = get_unets()
@@ -244,7 +252,9 @@ def convert(
logger.info(f"Sample UNet inputs spec: {sample_inputs_spec}")
logger.info("JIT tracing..")
traced_unet = torch.jit.trace(coreml_unet, example_inputs=list(sample_inputs.values()))
traced_unet = torch.jit.trace(
coreml_unet, example_inputs=list(sample_inputs.values())
)
logger.info("Done.")
coreml_sample_inputs = get_coreml_inputs(sample_inputs)
+7 -7
View File
@@ -42,13 +42,13 @@ class CoreMLSamplerLCM_Simple:
CATEGORY = "sampling"
def sample(
self,
coreml_model,
seed,
steps,
cfg,
positive_prompt,
num_images,
self,
coreml_model,
seed,
steps,
cfg,
positive_prompt,
num_images,
):
height = coreml_model.expected_inputs["sample"]["shape"][2] * 8
width = coreml_model.expected_inputs["sample"]["shape"][3] * 8
+12 -7
View File
@@ -16,12 +16,14 @@ class CoreMLConverterLCM:
"height": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
"width": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
"batch_size": ("INT", {"default": 4, "min": 1, "max": 64}),
"compute_unit": ([
ComputeUnit.CPU_AND_NE.name,
ComputeUnit.CPU_AND_GPU.name,
ComputeUnit.ALL.name,
ComputeUnit.CPU_ONLY.name,
],),
"compute_unit": (
[
ComputeUnit.CPU_AND_NE.name,
ComputeUnit.CPU_AND_GPU.name,
ComputeUnit.ALL.name,
ComputeUnit.CPU_ONLY.name,
],
),
"controlnet_support": ("BOOLEAN", {"default": False}),
}
}
@@ -57,7 +59,10 @@ class CoreMLConverterLCM:
if not os.path.exists(out_path):
lcm_converter.convert(
out_path=out_path, sample_size=sample_size, batch_size=batch_size, controlnet_support=controlnet_support
out_path=out_path,
sample_size=sample_size,
batch_size=batch_size,
controlnet_support=controlnet_support,
)
target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name)
+1 -10
View File
@@ -41,9 +41,7 @@ class CoreMLModelWrapper(BaseModel):
):
chunked_in = self.chunk_inputs(x, t, c_crossattn, control)
chunked_out = [
self._apply_model(
x, t, c_concat, c_crossattn, c_adm, control, transformer_options
)
self._apply_model(x, t, c_crossattn, control)
for x, t, c_crossattn, control in zip(*chunked_in)
]
@@ -105,12 +103,5 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper):
super().__init__(model_config, coreml_model)
self.config = None
def _apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None,
control=None, transformer_options={}):
model_input_kwargs = self.prepare_inputs(x, t, c_crossattn, control)
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
return torch.from_numpy(np_out).to(x.device)
def __call__(self, latents, ts, encoder_hidden_states, **kwargs):
return (self.apply_model(latents, ts, c_crossattn=encoder_hidden_states),)