Enable choosing attention implementation during conversion
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user