diff --git a/coreml_suite/lcm/lcm_converter.py b/coreml_suite/lcm/lcm_converter.py index be3b4cc..5259584 100644 --- a/coreml_suite/lcm/lcm_converter.py +++ b/coreml_suite/lcm/lcm_converter.py @@ -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) diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 08d23a7..d9b18d7 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -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 diff --git a/coreml_suite/lcm/nodes.py b/coreml_suite/lcm/nodes.py index 296d8a9..fa9a527 100644 --- a/coreml_suite/lcm/nodes.py +++ b/coreml_suite/lcm/nodes.py @@ -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) diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 115a4fc..a7d4b32 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -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),)