diff --git a/coreml_suite/converter.py b/coreml_suite/converter.py index 1d196a2..52c8c7a 100644 --- a/coreml_suite/converter.py +++ b/coreml_suite/converter.py @@ -6,9 +6,13 @@ from enum import Enum, auto import coremltools as ct import numpy as np +import python_coreml_stable_diffusion.unet import torch from diffusers import StableDiffusionPipeline, LatentConsistencyModelPipeline -from python_coreml_stable_diffusion.unet import UNet2DConditionModel +from python_coreml_stable_diffusion.unet import ( + UNet2DConditionModel, + AttentionImplementations, +) from coreml_suite.lcm.unet import UNet2DConditionModelLCM from coreml_suite.logger import logger @@ -273,11 +277,16 @@ def convert( sample_size: tuple[int, int] = (64, 64), controlnet_support: bool = False, lora_paths: list[str | os.PathLike] = None, + attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name, ): if os.path.exists(unet_out_path): logger.info(f"Found existing model at {unet_out_path}! Skipping..") return + python_coreml_stable_diffusion.unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = ( + AttentionImplementations(attn_impl) + ) + model_type = ModelType.SD15 pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_type] diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index a842e67..2a703e3 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -3,6 +3,7 @@ import os import torch from coremltools import ComputeUnit from python_coreml_stable_diffusion.coreml_model import CoreMLModel +from python_coreml_stable_diffusion.unet import AttentionImplementations import folder_paths from comfy import model_base @@ -174,6 +175,13 @@ class COREML_CONVERT(COREML_NODE): "height": ("INT", {"default": 512, "min": 512, "max": 2048, "step": 8}), "width": ("INT", {"default": 512, "min": 512, "max": 2048, "step": 8}), "batch_size": ("INT", {"default": 1, "min": 1, "max": 64}), + "attention_implementation": ( + [ + AttentionImplementations.SPLIT_EINSUM.name, + AttentionImplementations.SPLIT_EINSUM_V2.name, + AttentionImplementations.ORIGINAL.name, + ], + ), "compute_unit": ( [ ComputeUnit.CPU_AND_NE.name, @@ -199,6 +207,7 @@ class COREML_CONVERT(COREML_NODE): height, width, batch_size, + attention_implementation, compute_unit, controlnet_support, lora_stack=None, @@ -249,6 +258,7 @@ class COREML_CONVERT(COREML_NODE): batch_size=batch_size, controlnet_support=controlnet_support, lora_paths=lora_paths, + attn_impl=attention_implementation, ) unet_target_path = converter.compile_model( out_path=unet_out_path, out_name=out_name, submodule_name="unet"