Enable choosing attention implementation during conversion

This commit is contained in:
aszc-dev
2023-11-17 22:55:13 +01:00
committed by aszc
parent 5477e3d71a
commit 8092a19173
2 changed files with 20 additions and 1 deletions
+10 -1
View File
@@ -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]
+10
View File
@@ -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"