Add support for CN models to LCM
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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),)
|
||||
|
||||
Reference in New Issue
Block a user