Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d4009e7b69 | ||
|
|
b20507e9af | ||
|
|
7ea420aef1 | ||
|
|
d4bda0e740 | ||
|
|
21ab05fd4f | ||
|
|
73b13e0d23 | ||
|
|
991ab6a40a | ||
|
|
1320e5cd9c | ||
|
|
ebdc700177 | ||
|
|
d752f41b66 | ||
|
|
ff9a5c91d2 | ||
|
|
d17a93323f | ||
|
|
9fb310700e | ||
|
|
f95a439d62 | ||
|
|
b044efe201 | ||
|
|
3f666ac0ea | ||
|
|
acecd10aee | ||
|
|
2f6597b8c0 | ||
|
|
a03a56a58e | ||
|
|
9b7e5a6cd8 | ||
|
|
600023382c | ||
|
|
7caf1cea1d | ||
|
|
c6229e5c5f | ||
|
|
5de5722474 | ||
|
|
a3cf825d79 | ||
|
|
75057e4ed2 | ||
|
|
a1d81faf68 | ||
|
|
1ff260fc36 | ||
|
|
83ad02748f | ||
|
|
4c9195bbc0 | ||
|
|
7643211d8d | ||
|
|
6be88fee2e | ||
|
|
1a82e4b48f | ||
|
|
7a5d040b61 | ||
|
|
b4313d731e | ||
|
|
9489503cbe | ||
|
|
edc8e39c83 | ||
|
|
de8915eb6c | ||
|
|
93ebaf4d5d | ||
|
|
1bc728d0ea | ||
|
|
33829c292f | ||
|
|
7b1c3c7ba7 | ||
|
|
8c9fbacb45 | ||
|
|
997c6a78ff | ||
|
|
d9be9c13e2 | ||
|
|
31a6ac6d2f | ||
|
|
11772e4e69 | ||
|
|
d4f3ed6fa9 | ||
|
|
1d450cca3c | ||
|
|
b2102592cd | ||
|
|
3d7473903b | ||
|
|
b12cd83041 | ||
|
|
ab567e48af | ||
|
|
12f667190f | ||
|
|
d612d1ffef | ||
|
|
d4666d3615 | ||
|
|
42c6a66a7c | ||
|
|
f4a1eb974b | ||
|
|
c09bbeabe2 | ||
|
|
43f8d330a0 | ||
|
|
4e32ca8dbc | ||
|
|
8f639eb2a0 | ||
|
|
bfe22d8d06 | ||
|
|
92080ae196 | ||
|
|
a638a79f81 | ||
|
|
d6f7188f7e | ||
|
|
ca715599c1 | ||
|
|
ef78f8596f | ||
|
|
3d6f8b7dcd | ||
|
|
83f49f0937 | ||
|
|
183d0b2707 | ||
|
|
ddbeb36d52 | ||
|
|
2ee81bd41d | ||
|
|
a1c66249e2 | ||
|
|
01fafd70f3 | ||
|
|
447b25c774 | ||
|
|
c3038501eb | ||
|
|
b326b3d3b9 |
@@ -2,8 +2,8 @@
|
|||||||
|
|
||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
Welcome! I've developed a set of custom nodes for ComfyUI that allows you to use Core ML models in your ComfyUI
|
Welcome! In this repository you'll find a set of custom nodes for [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
|
||||||
workflows.
|
that allows you to use Core ML models in your ComfyUI workflows.
|
||||||
These models are designed to leverage the Apple Neural Engine (ANE) on Apple Silicon (M1/M2) machines,
|
These models are designed to leverage the Apple Neural Engine (ANE) on Apple Silicon (M1/M2) machines,
|
||||||
thereby enhancing your workflows and improving performance.
|
thereby enhancing your workflows and improving performance.
|
||||||
|
|
||||||
@@ -48,6 +48,8 @@ That's it! You're now ready to start enhancing your ComfyUI workflows with Core
|
|||||||
- **VAE**: Variational Autoencoder. A model that learns a latent representation of images. It's used as a prior in
|
- **VAE**: Variational Autoencoder. A model that learns a latent representation of images. It's used as a prior in
|
||||||
Stable Diffusion.
|
Stable Diffusion.
|
||||||
- **Checkpoint**: A file that contains the weights of a model. It's used to load models in Stable Diffusion.
|
- **Checkpoint**: A file that contains the weights of a model. It's used to load models in Stable Diffusion.
|
||||||
|
- **LCM**: [Latent Consistency Model](https://latent-consistency-models.github.io/). A type of model designed to
|
||||||
|
generate images with as few steps as possible.
|
||||||
|
|
||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
> Note on Compute Units:
|
> Note on Compute Units:
|
||||||
@@ -64,6 +66,12 @@ These custom nodes come with a host of features, including:
|
|||||||
- Support for ANE (Apple Neural Engine)
|
- Support for ANE (Apple Neural Engine)
|
||||||
- Support for CPU and GPU
|
- Support for CPU and GPU
|
||||||
- Support for `mlmodelc` and `mlpackage` files
|
- Support for `mlmodelc` and `mlpackage` files
|
||||||
|
- Support for SDXL models
|
||||||
|
- Support for LCM models
|
||||||
|
- Support for LoRAs
|
||||||
|
- SD1.5 -> Core ML conversion
|
||||||
|
- SDXL -> Core ML conversion
|
||||||
|
- LCM -> Core ML conversion
|
||||||
|
|
||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
> Please note that using Core ML models can take a bit longer to load initially.
|
> Please note that using Core ML models can take a bit longer to load initially.
|
||||||
@@ -75,7 +83,18 @@ These custom nodes come with a host of features, including:
|
|||||||
|
|
||||||
## Installation
|
## Installation
|
||||||
|
|
||||||
The installation process is simple!
|
### Using ComfyUI-Manager
|
||||||
|
|
||||||
|
The easiest way to install the custom nodes is to use the ComfyUI-Manager. You can find the installation instructions
|
||||||
|
[here](https://github.com/ltdrdata/ComfyUI-Manager#installation). Once you've installed the ComfyUI-Manager, you can
|
||||||
|
install the custom nodes by following these steps:
|
||||||
|
|
||||||
|
- Open the ComfyUI-Manager by clicking the `Manager` button in the ComfyUI toolbar.
|
||||||
|
- Click the `Install Custom Nodes` button.
|
||||||
|
- Search for `Core ML` and click the `Install` button.
|
||||||
|
- Restart ComfyUI.
|
||||||
|
|
||||||
|
### Manual Installation
|
||||||
|
|
||||||
1. Clone this repository into the custom_nodes directory of your ComfyUI. If you're not sure how to do this, you can
|
1. Clone this repository into the custom_nodes directory of your ComfyUI. If you're not sure how to do this, you can
|
||||||
download the repository as a zip file and extract it into the same directory.
|
download the repository as a zip file and extract it into the same directory.
|
||||||
@@ -121,10 +140,6 @@ node is a `coreml_model` object that can be used with the Core ML Sampler.
|
|||||||
- **Outputs**:
|
- **Outputs**:
|
||||||
- **coreml_model**: A Core ML model that can be used with the Core ML Sampler.
|
- **coreml_model**: A Core ML model that can be used with the Core ML Sampler.
|
||||||
|
|
||||||
> [!NOTE]
|
|
||||||
> Some models are designed to support ControlNet. If you're using such a model,
|
|
||||||
> make sure to provide a ControlNet input; otherwise, the model will use random noise as ControlNet input.
|
|
||||||
|
|
||||||
#### Core ML Sampler (`CoreMLSampler`)
|
#### Core ML Sampler (`CoreMLSampler`)
|
||||||
|
|
||||||

|

|
||||||
@@ -143,6 +158,118 @@ resulting latent as you normally would in your workflow.
|
|||||||
- **LATENT**: The latent image output by the Core ML model. This can be decoded using a VAE Decoder or used as input
|
- **LATENT**: The latent image output by the Core ML model. This can be decoded using a VAE Decoder or used as input
|
||||||
to the next node in your workflow.
|
to the next node in your workflow.
|
||||||
|
|
||||||
|
#### Checkpoint Converter
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
You can use this node to convert any **SD1.5** based checkpoint to a Core ML model. The converted model is stored in the
|
||||||
|
`models/unet` directory and can be used with the `Core ML UNet Loader`. The conversion parameters are encoded in
|
||||||
|
the node name, so if the model already exists, the node will not convert it again.
|
||||||
|
|
||||||
|
- **Inputs**:
|
||||||
|
- **ckpt_name**: The name of the checkpoint to convert. This should be the name of the checkpoint file stored in the
|
||||||
|
`models/checkpoints` directory.
|
||||||
|
- **model_version**: Whether the model is based on SD1.5 or SDXL.
|
||||||
|
- **height**: The desired height of the image generated by the model. The default is 512. Must be a multiple of 8.
|
||||||
|
- **width**: The desired width of the image generated by the model. The default is 512. Must be a multiple of 8.
|
||||||
|
- **batch_size**: The batch size of generated images. If you're planning to generate batches of images, you can try
|
||||||
|
increasing this value to speed up the generation process. The default is 1.
|
||||||
|
- **attention_implementation**: The attention implementation used when converting the model. Choose SPLIT_EINSUM or
|
||||||
|
SPLIT_EINSUM_V2 for better ANE support. Choose ORIGINAL for better GPU support.
|
||||||
|
- **compute_unit**: The hardware on which the model should run. This is used only when loading the model and doesn't
|
||||||
|
affect the conversion process.
|
||||||
|
- **controlnet_support**: For the model to support ControlNet, it must be converted with this option set to True.
|
||||||
|
The
|
||||||
|
default is False.
|
||||||
|
- **lora_params** [optional]: Optional LoRA names and weights. If provided, the model will be converted with LoRA(s)
|
||||||
|
baked in. More on loading LoRAs below.
|
||||||
|
- **Outputs**:
|
||||||
|
- **coreml_model**: The converted Core ML model that can be used with Core ML Sampler.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> Some models use a custom config .yaml file. If you're using such a model, you'll need to place the config file in the
|
||||||
|
> `models/configs` directory. The config file should be named the same as the checkpoint file. For example, if the
|
||||||
|
> checkpoint file is named `juggernaut_aftermath.safetensors`, the config file should be
|
||||||
|
> named `juggernaut_aftermath.yaml`.
|
||||||
|
> The config file will be automatically loaded during conversion.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> For now, the converter relies heavilty on the model name to determine the conversion parameters. This means that if
|
||||||
|
> you change the model name, the node will convert the model again. Other than that, if you find the name too long or
|
||||||
|
> confusing, you can change it to anything you want.
|
||||||
|
|
||||||
|
#### LoRA Loader
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
This node allows you to load LoRAs and bake them into a model. Since this is a workaround (as model weights can't be
|
||||||
|
modified
|
||||||
|
after conversion), there are a few caveats to keep in mind:
|
||||||
|
|
||||||
|
- The LoRA weights and _strength_model_ parameter are baked into the model. This means that you can't change them
|
||||||
|
after conversion. This also means that you need to convert the model again if you want to change the LoRA weights.
|
||||||
|
- Loading LoRA affects CLIP, which is not a part of Core ML workflow, so you'll need to load CLIP separately,
|
||||||
|
either using `CLIPLoader` or `CheckpointLoaderSimple`. (See [example workflows](#example-workflows) for more details.)
|
||||||
|
- After conversion, if you want to load the model using `CoreMLUnetLoader`, you'll need to apply the same LoRAs to
|
||||||
|
CLIP manually. (See [example workflows](#example-workflows) for more details.)
|
||||||
|
- The LoRA names are encoded in the model name. This means that if you change the name of the LoRA file,
|
||||||
|
you'll need to change the model name as well, or the node will convert the model again. (Model strength is not
|
||||||
|
encoded, so if you want to change it, you'll need to delete the converted model manually)
|
||||||
|
- _strength_clip_ parameter only affects the CLIP model and is not baked into the converted model. This means that
|
||||||
|
you can change it after conversion.
|
||||||
|
|
||||||
|
- **Inputs**:
|
||||||
|
- **lora_name**: The name of the LoRA to load.
|
||||||
|
- **strength_model**: The strength of the LoRA model.
|
||||||
|
- **strength_clip**: The strength of the LoRA CLIP.
|
||||||
|
- **lora_params** [optional]: Optional output from other LoRA Loaders.
|
||||||
|
- **clip**: The CLIP model to use with the LoRA. This can be either output of the
|
||||||
|
`CLIPLoader`/`CheckpointLoaderSimple` or other LoRA Loaders.
|
||||||
|
- **Outputs**:
|
||||||
|
- **lora_params**: The LoRA parameters that can be passed to the Core ML Converter or other LoRA Loaders.
|
||||||
|
- **CLIP**: The CLIP model with LoRA applied.
|
||||||
|
|
||||||
|
#### LCM Converter
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
This node converts [SimianLuo/LCM_Dreamshaper_v7](https://huggingface.co/SimianLuo/LCM_Dreamshaper_v7) model to Core
|
||||||
|
ML. The converted model is stored in the `models/unet` directory and can be used with the Core ML UNet Loader. The
|
||||||
|
conversion parameteres are encoded in the node name, so if the model already exists, the node will not convert it again.
|
||||||
|
|
||||||
|
- **Inputs**:
|
||||||
|
- **height**: The desired height of the image generated by the model. The default is 512. Must be a multiple of 8.
|
||||||
|
- **width**: The desired width of the image generated by the model. The default is 512. Must be a multiple of 8.
|
||||||
|
- **batch_size**: The batch size of generated images. If you're planning to generate batches of images, you can try
|
||||||
|
increasing this value to speed up the generation process. The default is 1.
|
||||||
|
- **compute_unit**: The hardware on which the model should run. This is used only when loading the model and
|
||||||
|
doesn't affect the conversion process.
|
||||||
|
- **controlnet_support**: For the model to support ControlNet, it must be converted with this option set to True.
|
||||||
|
The default is False.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> The conversion process can take a while, so please be patient.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> When using the LCM model with Core ML Sampler, please set _sampler_name_ to `lcm` and _scheduler_ to `sgm_uniform`.
|
||||||
|
|
||||||
|
#### Core ML Adapter (Experimental) (`CoreMLModelAdapter`)
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
This node allows you to use a Core ML as a standard ComfyUI model. This is an experimental node and may not work with
|
||||||
|
all models and nodes. Please use with caution and pay attention to the expected inputs of the model.
|
||||||
|
|
||||||
|
- **Input**:
|
||||||
|
- **coreml_model**: The Core ML model to use as a ComfyUI model.
|
||||||
|
- **Output**:
|
||||||
|
- **MODEL**: The Core ML model wrapped in a ComfyUI model.
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> While this approach allows you to use Core ML models with many ComfyUI nodes (both standard and custom), the
|
||||||
|
> expected inputs of the model will not be checked, which may cause errors. Please make sure to use a model compatible
|
||||||
|
> with the expected parameters.
|
||||||
|
|
||||||
### Example Workflows
|
### Example Workflows
|
||||||
|
|
||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
@@ -180,10 +307,69 @@ being loaded using the standard ComfyUI nodes. Please refer to
|
|||||||
the [basic txt2img workflow](#basic-txt2img-with-core-ml-unet-loader) for more details on how to load the CLIP and VAE
|
the [basic txt2img workflow](#basic-txt2img-with-core-ml-unet-loader) for more details on how to load the CLIP and VAE
|
||||||
models.
|
models.
|
||||||
The ControlNet model used in this workflow is available
|
The ControlNet model used in this workflow is available
|
||||||
[here](https://huggingface.co/lllyasviel/ControlNet-v1-1/blob/main/control_v11p_sd15_lineart.pth).
|
[here](https://huggingface.co/lllyasviel/control_v11p_sd15_scribble/blob/main/diffusion_pytorch_model.fp16.safetensors).
|
||||||
Once downloaded, place the model in the `models/controlnet` directory.
|
Once downloaded, place the model in the `models/controlnet` directory.
|
||||||

|

|
||||||
|
|
||||||
|
#### Checkpoint conversion
|
||||||
|
|
||||||
|
This workflow uses the Checkpoint Converter to convert the checkpoint file. See
|
||||||
|
[Checkpoint Converter](#checkpoint-converter) description for more details.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
#### Checkpoint conversion with LoRA
|
||||||
|
|
||||||
|
This workflow uses the Checkpoint Converter to convert the checkpoint file with LoRA. See
|
||||||
|
[LoRA Loader](#lora-loader) description to read more about the caveats of using LoRA.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
#### LCM LoRA conversion
|
||||||
|
|
||||||
|
Please note that you can use multiple LoRAs with the same model. To do this, you'll need to use multiple LoRA Loaders.
|
||||||
|
> [!IMPORTANT]
|
||||||
|
> In this example, the model is passed through the adapter and `ModelSamplingDiscrete` nodes to a standard ComfyUI's
|
||||||
|
> KSampler (not Core ML Sampler). ModelSamplingDiscrete needs to be used to sample models with LCM LoRAs properly.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
#### Loader with LoRAs
|
||||||
|
|
||||||
|
This workflow uses the Core ML UNet Loader to load a model with LoRAs. The CLIP must be loaded separately and passed
|
||||||
|
through the same LoRA nodes as during conversion. See [LoRA Loader](#lora-loader) description to read more about the
|
||||||
|
caveats of using LoRA. Since _lora_name_ and _strength_model_ are baked into the model, it is not necessary to pass
|
||||||
|
them as inputs to the loader.
|
||||||
|
> [!IMPORTANT]
|
||||||
|
> In this example, the model is passed through the adapter and `ModelSamplingDiscrete` nodes to a standard ComfyUI's
|
||||||
|
> KSampler (not Core ML Sampler). ModelSamplingDiscrete needs to be used to sample models with LCM LoRAs properly.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
#### LCM conversion with ControlNet
|
||||||
|
|
||||||
|
This workflow uses LCM converter to
|
||||||
|
convert [SimianLuo/LCM_Dreamshaper_v7](https://huggingface.co/SimianLuo/LCM_Dreamshaper_v7)
|
||||||
|
model to Core ML. The converted model can then be used with or without ControlNet to generate images.
|
||||||
|

|
||||||
|
|
||||||
|
#### SDXL Base + Refiner conversion
|
||||||
|
|
||||||
|
This is a basic workflow for SDXL. You add LoRAs and ControlNets the same way as in the previous examples.
|
||||||
|
You can also skip the refiner step.
|
||||||
|
|
||||||
|
The models used in this workflow are available at the following links:
|
||||||
|
|
||||||
|
- [Base model + text_encoder (clip) + text_encoder_2 (clip2)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
|
||||||
|
- [Refiner model](https://huggingface.co/stabilityai/stable-diffusion-xl-refiner-1.0)
|
||||||
|
- [VAE](https://huggingface.co/stabilityai/sdxl-vae)
|
||||||
|
|
||||||
|
> [!IMPORTANT]
|
||||||
|
> SDXL on ANE is not supported. If loading of the model gets stuck, please try using CPU_AND_GPU or CPU_ONLY.
|
||||||
|
> For best results, use ORIGINAL attention implementation.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
## Limitations
|
## Limitations
|
||||||
|
|
||||||
- Core ML models are fixed in terms of their inputs and outputs.
|
- Core ML models are fixed in terms of their inputs and outputs.
|
||||||
@@ -191,13 +377,55 @@ Once downloaded, place the model in the `models/controlnet` directory.
|
|||||||
SD1.5).
|
SD1.5).
|
||||||
However, you can convert the model to a different input size using tools available
|
However, you can convert the model to a different input size using tools available
|
||||||
in the [apple/ml-stable-diffusion](https://github.com/apple/ml-stable-diffusion) repository.
|
in the [apple/ml-stable-diffusion](https://github.com/apple/ml-stable-diffusion) repository.
|
||||||
- For now, only Stable Diffusion v1.5 is supported.
|
- SD2.1 models are not supported.
|
||||||
- LoRA is not supported yet.
|
|
||||||
|
|
||||||
[^1]:
|
[^1]:
|
||||||
Unless [EnumeratedShapes](https://apple.github.io/coremltools/docs-guides/source/flexible-inputs.html#select-from-predetermined-shapes)
|
Unless [EnumeratedShapes](https://apple.github.io/coremltools/docs-guides/source/flexible-inputs.html#select-from-predetermined-shapes)
|
||||||
is used during conversion. Needs more testing.
|
is used during conversion. Needs more testing.
|
||||||
|
|
||||||
|
## FAQ
|
||||||
|
|
||||||
|
### Hardware and Performance
|
||||||
|
|
||||||
|
#### What's the difference between MPS, GPU, and ANE?
|
||||||
|
- **MPS (Metal Performance Shaders)**: Apple's framework for GPU acceleration. It's what PyTorch uses by default on Apple Silicon.
|
||||||
|
- **GPU**: The graphics processing unit on your Apple Silicon chip.
|
||||||
|
- **ANE (Apple Neural Engine)**: A specialized hardware accelerator for machine learning tasks.
|
||||||
|
|
||||||
|
#### Which compute unit should I choose?
|
||||||
|
- **CPU_AND_ANE**: Best for models converted with `--attention-implementation SPLIT_EINSUM`. This is the default and recommended option for most users.
|
||||||
|
- **CPU_AND_GPU**: Best for models converted with `--attention-implementation ORIGINAL`. Use this if you experience issues with ANE.
|
||||||
|
- **CPU_ONLY**: Use this as a fallback if you experience issues with both ANE and GPU.
|
||||||
|
|
||||||
|
#### Do I need `PYTORCH_ENABLE_MPS_FALLBACK=1`?
|
||||||
|
While our Core ML nodes don't use this environment variable directly, it may still be relevant for other parts of ComfyUI that use PyTorch with MPS backend. The setting of this variable is a user preference and depends on your specific needs and workflow requirements.
|
||||||
|
|
||||||
|
### Model Conversion and Compatibility
|
||||||
|
|
||||||
|
#### Is there a performance penalty when using the Core ML Adapter?
|
||||||
|
Yes, there might be a slight performance penalty compared to using directly converted models. However, the adapter provides more flexibility and compatibility with standard ComfyUI nodes.
|
||||||
|
|
||||||
|
#### Does the Core ML Adapter support SDXL?
|
||||||
|
Currently, SDXL support in the Core ML Adapter is limited. While it may work with some models, it's not officially supported and may cause issues.
|
||||||
|
|
||||||
|
#### Are `mlmodelc` and `mlpackage` formats safe?
|
||||||
|
Yes, both formats are safe to use. However, we recommend:
|
||||||
|
1. Always downloading original `.safetensors` files from trusted sources
|
||||||
|
2. Converting them yourself using our tools
|
||||||
|
3. Using the converted `.mlmodelc` files for better performance
|
||||||
|
|
||||||
|
#### Do Core ML models produce identical results to their safetensors counterparts?
|
||||||
|
While the results should be very similar, there might be slight differences due to:
|
||||||
|
- Different numerical precision
|
||||||
|
- Hardware-specific optimizations
|
||||||
|
- Different attention implementations
|
||||||
|
|
||||||
|
#### Should I convert models every time I queue a generation?
|
||||||
|
No! The conversion only happens once when you first use the converter node. After that, you should use the `CoreMLUnetLoader` to load the already converted model.
|
||||||
|
|
||||||
|
#### Will SDXL ever be supported on ANE?
|
||||||
|
Currently, there are technical limitations preventing SDXL from running efficiently on ANE. We recommend using `CPU_AND_GPU` or `CPU_ONLY` for SDXL models.
|
||||||
|
|
||||||
## Support
|
## Support
|
||||||
|
|
||||||
I'm here to help! If you have any questions or suggestions, don't hesitate to open an issue and I'll do my best
|
I'm here to help! If you have any questions or suggestions, don't hesitate to open an issue and I'll do my best
|
||||||
|
|||||||
@@ -3,13 +3,33 @@ import sys
|
|||||||
|
|
||||||
sys.path.append(os.path.dirname(__file__))
|
sys.path.append(os.path.dirname(__file__))
|
||||||
|
|
||||||
from coreml_suite import CoreMLLoaderUNet, CoreMLSampler
|
from coreml_suite.nodes import (
|
||||||
|
CoreMLLoaderUNet,
|
||||||
|
CoreMLSampler,
|
||||||
|
CoreMLSamplerAdvanced,
|
||||||
|
CoreMLModelAdapter,
|
||||||
|
CoreMLConverter,
|
||||||
|
COREML_LOAD_LORA,
|
||||||
|
)
|
||||||
|
from coreml_suite.lcm import (
|
||||||
|
COREML_CONVERT_LCM,
|
||||||
|
)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"CoreMLUNetLoader": CoreMLLoaderUNet,
|
"CoreMLUNetLoader": CoreMLLoaderUNet,
|
||||||
"CoreMLSampler": CoreMLSampler,
|
"CoreMLSampler": CoreMLSampler,
|
||||||
|
"CoreMLSamplerAdvanced": CoreMLSamplerAdvanced,
|
||||||
|
"CoreMLModelAdapter": CoreMLModelAdapter,
|
||||||
|
"Core ML LoRA Loader": COREML_LOAD_LORA,
|
||||||
|
"Core ML Converter": CoreMLConverter,
|
||||||
|
"Core ML LCM Converter": COREML_CONVERT_LCM,
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"CoreMLUNetLoader": "Load Core ML UNet",
|
"CoreMLUNetLoader": "Load Core ML UNet",
|
||||||
"CoreMLSampler": "Core ML Sampler",
|
"CoreMLSampler": "Core ML Sampler",
|
||||||
|
"CoreMLSamplerAdvanced": "Core ML Sampler (Advanced)",
|
||||||
|
"CoreMLModelAdapter": "Core ML Adapter (Experimental)",
|
||||||
|
"Core ML LoRA Loader": "Load LoRA to use with Core ML",
|
||||||
|
"Core ML Converter": "Convert Checkpoint to Core ML",
|
||||||
|
"Core ML LCM Converter": "Convert LCM to Core ML",
|
||||||
}
|
}
|
||||||
|
|||||||
|
After Width: | Height: | Size: 34 KiB |
|
After Width: | Height: | Size: 387 KiB |
|
After Width: | Height: | Size: 94 KiB |
|
After Width: | Height: | Size: 416 KiB |
|
After Width: | Height: | Size: 462 KiB |
|
After Width: | Height: | Size: 476 KiB |
|
After Width: | Height: | Size: 51 KiB |
|
After Width: | Height: | Size: 474 KiB |
|
After Width: | Height: | Size: 54 KiB |
|
After Width: | Height: | Size: 1.5 MiB |
|
Before Width: | Height: | Size: 469 KiB After Width: | Height: | Size: 508 KiB |
@@ -1,4 +1,2 @@
|
|||||||
from coreml_suite.loaders import CoreMLLoaderUNet
|
class COREML_NODE:
|
||||||
from coreml_suite.samplers import CoreMLSampler
|
CATEGORY = "Core ML Suite"
|
||||||
|
|
||||||
__all__ = ["CoreMLLoaderUNet", "CoreMLSampler"]
|
|
||||||
|
|||||||
@@ -0,0 +1,122 @@
|
|||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from comfy import supported_models_base
|
||||||
|
from comfy import latent_formats
|
||||||
|
from comfy.model_detection import convert_config
|
||||||
|
|
||||||
|
|
||||||
|
class ModelVersion(Enum):
|
||||||
|
SD15 = "sd15"
|
||||||
|
SDXL = "sdxl"
|
||||||
|
SDXL_REFINER = "sdxl_refiner"
|
||||||
|
LCM = "lcm"
|
||||||
|
|
||||||
|
|
||||||
|
config_map = {
|
||||||
|
ModelVersion.SD15: {
|
||||||
|
"use_checkpoint": False,
|
||||||
|
"image_size": 32,
|
||||||
|
"out_channels": 4,
|
||||||
|
"use_spatial_transformer": True,
|
||||||
|
"legacy": False,
|
||||||
|
"adm_in_channels": None,
|
||||||
|
"dtype": torch.float16,
|
||||||
|
"in_channels": 4,
|
||||||
|
"model_channels": 320,
|
||||||
|
"num_res_blocks": 2,
|
||||||
|
"attention_resolutions": [1, 2, 4],
|
||||||
|
"transformer_depth": [1, 1, 1, 0],
|
||||||
|
"channel_mult": [1, 2, 4, 4],
|
||||||
|
"transformer_depth_middle": 1,
|
||||||
|
"use_linear_in_transformer": False,
|
||||||
|
"context_dim": 768,
|
||||||
|
"num_heads": 8,
|
||||||
|
"disable_unet_model_creation": True,
|
||||||
|
},
|
||||||
|
ModelVersion.SDXL: {
|
||||||
|
"use_checkpoint": False,
|
||||||
|
"image_size": 32,
|
||||||
|
"out_channels": 4,
|
||||||
|
"use_spatial_transformer": True,
|
||||||
|
"legacy": False,
|
||||||
|
"num_classes": "sequential",
|
||||||
|
"adm_in_channels": 2816,
|
||||||
|
"dtype": torch.float16,
|
||||||
|
"in_channels": 4,
|
||||||
|
"model_channels": 320,
|
||||||
|
"num_res_blocks": 2,
|
||||||
|
"attention_resolutions": [2, 4],
|
||||||
|
"transformer_depth": [0, 2, 10],
|
||||||
|
"channel_mult": [1, 2, 4],
|
||||||
|
"transformer_depth_middle": 10,
|
||||||
|
"use_linear_in_transformer": True,
|
||||||
|
"context_dim": 2048,
|
||||||
|
"num_head_channels": 64,
|
||||||
|
"disable_unet_model_creation": True,
|
||||||
|
},
|
||||||
|
ModelVersion.SDXL_REFINER: {
|
||||||
|
"use_checkpoint": False,
|
||||||
|
"image_size": 32,
|
||||||
|
"out_channels": 4,
|
||||||
|
"use_spatial_transformer": True,
|
||||||
|
"legacy": False,
|
||||||
|
"num_classes": "sequential",
|
||||||
|
"adm_in_channels": 2560,
|
||||||
|
"dtype": torch.float16,
|
||||||
|
"in_channels": 4,
|
||||||
|
"model_channels": 384,
|
||||||
|
"num_res_blocks": 2,
|
||||||
|
"attention_resolutions": [2, 4],
|
||||||
|
"transformer_depth": [0, 4, 4, 0],
|
||||||
|
"channel_mult": [1, 2, 4, 4],
|
||||||
|
"transformer_depth_middle": 4,
|
||||||
|
"use_linear_in_transformer": True,
|
||||||
|
"context_dim": 1280,
|
||||||
|
"num_head_channels": 64,
|
||||||
|
"disable_unet_model_creation": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
latent_format_map = {
|
||||||
|
ModelVersion.SD15: latent_formats.SD15,
|
||||||
|
ModelVersion.SDXL: latent_formats.SDXL,
|
||||||
|
ModelVersion.SDXL_REFINER: latent_formats.SDXL,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_model_config(model_version: ModelVersion):
|
||||||
|
unet_config = convert_config(config_map[model_version])
|
||||||
|
config = supported_models_base.BASE(unet_config)
|
||||||
|
config.latent_format = latent_format_map[model_version]()
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def unet_config_from_diffusers_unet(state_dict):
|
||||||
|
match = {}
|
||||||
|
attention_resolutions = []
|
||||||
|
|
||||||
|
attn_res = 1
|
||||||
|
for i in range(5):
|
||||||
|
k = "down_blocks.{}.attentions.1.transformer_blocks.0.attn2.to_k.weight".format(
|
||||||
|
i
|
||||||
|
)
|
||||||
|
if k in state_dict:
|
||||||
|
match["context_dim"] = state_dict[k].shape[1]
|
||||||
|
attention_resolutions.append(attn_res)
|
||||||
|
attn_res *= 2
|
||||||
|
|
||||||
|
match["attention_resolutions"] = attention_resolutions
|
||||||
|
|
||||||
|
match["model_channels"] = state_dict["conv_in.weight"].shape[0]
|
||||||
|
match["in_channels"] = state_dict["conv_in.weight"].shape[1]
|
||||||
|
match["adm_in_channels"] = None
|
||||||
|
if "class_embedding.linear_1.weight" in state_dict:
|
||||||
|
match["adm_in_channels"] = state_dict["class_embedding.linear_1.weight"].shape[
|
||||||
|
1
|
||||||
|
]
|
||||||
|
elif "add_embedding.linear_1.weight" in state_dict:
|
||||||
|
match["adm_in_channels"] = state_dict["add_embedding.linear_1.weight"].shape[1]
|
||||||
|
|
||||||
|
print(match)
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
from itertools import chain
|
||||||
|
from math import ceil
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from coreml_suite.latents import chunk_batch
|
||||||
|
|
||||||
|
|
||||||
|
def expand_inputs(inputs):
|
||||||
|
expanded = inputs.copy()
|
||||||
|
for k, v in inputs.items():
|
||||||
|
if isinstance(v, np.ndarray):
|
||||||
|
expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v
|
||||||
|
elif isinstance(v, torch.Tensor):
|
||||||
|
expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v
|
||||||
|
elif isinstance(v, list):
|
||||||
|
expanded[k] = v * 2 if len(v) == 1 else v
|
||||||
|
elif isinstance(v, dict):
|
||||||
|
expand_inputs(v)
|
||||||
|
return expanded
|
||||||
|
|
||||||
|
|
||||||
|
def extract_residual_kwargs(expected_inputs, control):
|
||||||
|
if "additional_residual_0" not in expected_inputs.keys():
|
||||||
|
return {}
|
||||||
|
if control is None:
|
||||||
|
return no_control(expected_inputs)
|
||||||
|
|
||||||
|
residual_kwargs = {
|
||||||
|
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
|
||||||
|
for i, r in enumerate(chain(control["output"], control["middle"]))
|
||||||
|
}
|
||||||
|
return residual_kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def no_control(expected_inputs):
|
||||||
|
shapes_dict = {
|
||||||
|
k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional")
|
||||||
|
}
|
||||||
|
residual_kwargs = {
|
||||||
|
k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16)
|
||||||
|
for k, shape in shapes_dict.items()
|
||||||
|
}
|
||||||
|
return residual_kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_control(cn, target_size):
|
||||||
|
if cn is None:
|
||||||
|
return [None] * target_size
|
||||||
|
|
||||||
|
num_chunks = ceil(cn["output"][0].shape[0] / target_size)
|
||||||
|
|
||||||
|
out = [{"output": [], "middle": []} for _ in range(num_chunks)]
|
||||||
|
|
||||||
|
for k, v in cn.items():
|
||||||
|
for i, x in enumerate(v):
|
||||||
|
chunks = chunk_batch(x, (target_size, *x.shape[1:]))
|
||||||
|
for j, chunk in enumerate(chunks):
|
||||||
|
out[j][k].append(chunk)
|
||||||
|
|
||||||
|
return out
|
||||||
@@ -0,0 +1,362 @@
|
|||||||
|
import gc
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import time
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
|
import coremltools as ct
|
||||||
|
import numpy as np
|
||||||
|
import python_coreml_stable_diffusion.unet
|
||||||
|
import torch
|
||||||
|
from diffusers import (
|
||||||
|
StableDiffusionPipeline,
|
||||||
|
LatentConsistencyModelPipeline,
|
||||||
|
StableDiffusionXLPipeline,
|
||||||
|
)
|
||||||
|
from python_coreml_stable_diffusion.unet import (
|
||||||
|
UNet2DConditionModel,
|
||||||
|
UNet2DConditionModelXL,
|
||||||
|
AttentionImplementations,
|
||||||
|
)
|
||||||
|
|
||||||
|
from coreml_suite.config import ModelVersion
|
||||||
|
from coreml_suite.lcm.unet import UNet2DConditionModelLCM
|
||||||
|
from coreml_suite.logger import logger
|
||||||
|
from folder_paths import get_folder_paths
|
||||||
|
|
||||||
|
|
||||||
|
class StableDiffusionLCMPipeline(LatentConsistencyModelPipeline):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
MODEL_TYPE_TO_UNET_CLS = {
|
||||||
|
ModelVersion.SD15: UNet2DConditionModel,
|
||||||
|
ModelVersion.SDXL: UNet2DConditionModelXL,
|
||||||
|
ModelVersion.LCM: UNet2DConditionModelLCM,
|
||||||
|
}
|
||||||
|
|
||||||
|
MODEL_TYPE_TO_PIPE_CLS = {
|
||||||
|
ModelVersion.SD15: StableDiffusionPipeline,
|
||||||
|
ModelVersion.SDXL: StableDiffusionXLPipeline,
|
||||||
|
ModelVersion.LCM: StableDiffusionLCMPipeline,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_unet(model_type: ModelVersion, ref_pipe):
|
||||||
|
ref_unet = ref_pipe.unet
|
||||||
|
|
||||||
|
unet_cls = MODEL_TYPE_TO_UNET_CLS[model_type]
|
||||||
|
cml_unet = unet_cls.from_config(ref_unet.config).eval()
|
||||||
|
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
|
||||||
|
|
||||||
|
return cml_unet
|
||||||
|
|
||||||
|
|
||||||
|
def get_encoder_hidden_states_shape(ref_pipe, batch_size):
|
||||||
|
text_encoder = (
|
||||||
|
ref_pipe.text_encoder_2
|
||||||
|
if hasattr(ref_pipe, "text_encoder_2")
|
||||||
|
else ref_pipe.text_encoder
|
||||||
|
)
|
||||||
|
|
||||||
|
text_token_sequence_length = text_encoder.config.max_position_embeddings
|
||||||
|
hidden_size = (text_encoder.config.hidden_size,)
|
||||||
|
|
||||||
|
encoder_hidden_states_shape = (
|
||||||
|
batch_size,
|
||||||
|
ref_pipe.unet.config.cross_attention_dim or hidden_size,
|
||||||
|
1,
|
||||||
|
text_token_sequence_length,
|
||||||
|
)
|
||||||
|
|
||||||
|
return encoder_hidden_states_shape
|
||||||
|
|
||||||
|
|
||||||
|
def get_coreml_inputs(sample_inputs):
|
||||||
|
coreml_sample_unet_inputs = {
|
||||||
|
k: v.numpy().astype(np.float16) for k, v in sample_inputs.items()
|
||||||
|
}
|
||||||
|
return [
|
||||||
|
ct.TensorType(
|
||||||
|
name=k,
|
||||||
|
shape=v.shape,
|
||||||
|
dtype=v.numpy().dtype if isinstance(v, torch.Tensor) else v.dtype,
|
||||||
|
)
|
||||||
|
for k, v in coreml_sample_unet_inputs.items()
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def load_coreml_model(out_path):
|
||||||
|
logger.info(f"Loading model from {out_path}")
|
||||||
|
|
||||||
|
start = time.time()
|
||||||
|
coreml_model = ct.models.MLModel(out_path)
|
||||||
|
logger.info(f"Loading {out_path} took {time.time() - start:.1f} seconds")
|
||||||
|
|
||||||
|
return coreml_model
|
||||||
|
|
||||||
|
|
||||||
|
def convert_to_coreml(
|
||||||
|
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")
|
||||||
|
coreml_model = load_coreml_model(out_path)
|
||||||
|
else:
|
||||||
|
logger.info(f"Converting {submodule_name} to CoreML..")
|
||||||
|
coreml_model = ct.convert(
|
||||||
|
torchscript_module,
|
||||||
|
convert_to="mlprogram",
|
||||||
|
minimum_deployment_target=ct.target.macOS13,
|
||||||
|
inputs=sample_inputs,
|
||||||
|
outputs=[
|
||||||
|
ct.TensorType(name=name, dtype=np.float32) for name in output_names
|
||||||
|
],
|
||||||
|
skip_model_load=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
del torchscript_module
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
return coreml_model
|
||||||
|
|
||||||
|
|
||||||
|
def get_out_path(submodule_name, model_name):
|
||||||
|
fname = f"{model_name}_{submodule_name}.mlpackage"
|
||||||
|
unet_path = get_folder_paths(submodule_name)[0]
|
||||||
|
out_path = os.path.join(unet_path, fname)
|
||||||
|
return out_path
|
||||||
|
|
||||||
|
|
||||||
|
def compile_coreml_model(source_model_path, output_dir, final_name):
|
||||||
|
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
|
||||||
|
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
|
||||||
|
if os.path.exists(target_path):
|
||||||
|
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
|
||||||
|
return target_path
|
||||||
|
|
||||||
|
logger.info(f"Compiling {source_model_path}")
|
||||||
|
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
|
||||||
|
|
||||||
|
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
|
||||||
|
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
|
||||||
|
shutil.move(compiled_output, target_path)
|
||||||
|
|
||||||
|
return target_path
|
||||||
|
|
||||||
|
|
||||||
|
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
|
||||||
|
sample_unet_inputs = dict(
|
||||||
|
[
|
||||||
|
("sample", torch.rand(*sample_shape)),
|
||||||
|
(
|
||||||
|
"timestep",
|
||||||
|
torch.tensor([scheduler.timesteps[0].item()] * batch_size).to(
|
||||||
|
torch.float32
|
||||||
|
),
|
||||||
|
),
|
||||||
|
("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
return sample_unet_inputs
|
||||||
|
|
||||||
|
|
||||||
|
def lcm_inputs(sample_unet_inputs):
|
||||||
|
batch_size = sample_unet_inputs["sample"].shape[0]
|
||||||
|
return {"timestep_cond": torch.randn(batch_size, 256).to(torch.float32)}
|
||||||
|
|
||||||
|
|
||||||
|
def sdxl_inputs(sample_unet_inputs, ref_pipe):
|
||||||
|
sample_shape = sample_unet_inputs["sample"].shape
|
||||||
|
batch_size = sample_shape[0]
|
||||||
|
h = sample_shape[2] * 8
|
||||||
|
w = sample_shape[3] * 8
|
||||||
|
original_size = (h, w)
|
||||||
|
crops_coords_top_left = (0, 0)
|
||||||
|
|
||||||
|
is_refiner = (
|
||||||
|
hasattr(ref_pipe.config, "requires_aesthetics_score")
|
||||||
|
and ref_pipe.config.requires_aesthetics_score
|
||||||
|
)
|
||||||
|
|
||||||
|
if is_refiner:
|
||||||
|
aesthetic_score = (6.0,)
|
||||||
|
time_ids_list = list(original_size + crops_coords_top_left + aesthetic_score)
|
||||||
|
else:
|
||||||
|
target_size = (h, w)
|
||||||
|
time_ids_list = list(original_size + crops_coords_top_left + target_size)
|
||||||
|
|
||||||
|
time_ids = torch.tensor(time_ids_list).repeat(batch_size, 1).to(torch.int64)
|
||||||
|
text_embeds_shape = (batch_size, ref_pipe.text_encoder_2.config.hidden_size)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"time_ids": time_ids,
|
||||||
|
"text_embeds": torch.randn(*text_embeds_shape).to(torch.float32),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_inputs_spec(inputs):
|
||||||
|
inputs_spec = {k: (v.shape, v.dtype) for k, v in inputs.items()}
|
||||||
|
return inputs_spec
|
||||||
|
|
||||||
|
|
||||||
|
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]
|
||||||
|
h, w = sample_shape[2:]
|
||||||
|
|
||||||
|
# conv_in
|
||||||
|
out_h, out_w = calculate_conv2d_output_shape(
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
reference_unet.conv_in,
|
||||||
|
)
|
||||||
|
additional_residuals_shapes.append(
|
||||||
|
(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
|
||||||
|
]
|
||||||
|
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
|
||||||
|
)
|
||||||
|
additional_residuals_shapes.append(
|
||||||
|
(
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
|
||||||
|
additional_inputs = {}
|
||||||
|
for i, shape in enumerate(additional_residuals_shapes):
|
||||||
|
sample_residual_input = torch.rand(*shape)
|
||||||
|
additional_inputs[f"additional_residual_{i}"] = sample_residual_input
|
||||||
|
|
||||||
|
return additional_inputs
|
||||||
|
|
||||||
|
|
||||||
|
def convert_unet(
|
||||||
|
ref_pipe,
|
||||||
|
model_version: ModelVersion,
|
||||||
|
unet_out_path: str,
|
||||||
|
batch_size: int = 1,
|
||||||
|
sample_size: tuple[int, int] = (64, 64),
|
||||||
|
controlnet_support: bool = False,
|
||||||
|
):
|
||||||
|
coreml_unet = get_unet(model_version, ref_pipe)
|
||||||
|
ref_unet = ref_pipe.unet
|
||||||
|
|
||||||
|
sample_shape = (
|
||||||
|
batch_size, # B
|
||||||
|
ref_unet.config.in_channels, # C
|
||||||
|
sample_size[0], # H
|
||||||
|
sample_size[1], # W
|
||||||
|
)
|
||||||
|
|
||||||
|
encoder_hidden_states_shape = get_encoder_hidden_states_shape(ref_pipe, batch_size)
|
||||||
|
|
||||||
|
scheduler = ref_pipe.scheduler
|
||||||
|
scheduler.set_timesteps(50)
|
||||||
|
|
||||||
|
sample_inputs = get_sample_input(
|
||||||
|
batch_size, encoder_hidden_states_shape, sample_shape, scheduler
|
||||||
|
)
|
||||||
|
|
||||||
|
if model_version == ModelVersion.LCM:
|
||||||
|
sample_inputs |= lcm_inputs(sample_inputs)
|
||||||
|
|
||||||
|
if model_version == ModelVersion.SDXL:
|
||||||
|
sample_inputs |= sdxl_inputs(sample_inputs, ref_pipe)
|
||||||
|
|
||||||
|
if controlnet_support:
|
||||||
|
sample_inputs |= add_cnet_support(sample_shape, ref_unet)
|
||||||
|
|
||||||
|
sample_inputs_spec = get_inputs_spec(sample_inputs)
|
||||||
|
|
||||||
|
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())
|
||||||
|
)
|
||||||
|
logger.info("Done.")
|
||||||
|
|
||||||
|
coreml_sample_inputs = get_coreml_inputs(sample_inputs)
|
||||||
|
|
||||||
|
coreml_unet = convert_to_coreml(
|
||||||
|
"unet", traced_unet, coreml_sample_inputs, ["noise_pred"], unet_out_path
|
||||||
|
)
|
||||||
|
|
||||||
|
del traced_unet
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
coreml_unet.save(unet_out_path)
|
||||||
|
logger.info(f"Saved unet into {unet_out_path}")
|
||||||
|
|
||||||
|
|
||||||
|
def convert(
|
||||||
|
ckpt_path: str,
|
||||||
|
model_version: ModelVersion,
|
||||||
|
unet_out_path: str,
|
||||||
|
batch_size: int = 1,
|
||||||
|
sample_size: tuple[int, int] = (64, 64),
|
||||||
|
controlnet_support: bool = False,
|
||||||
|
lora_weights: list[tuple[Union[str, os.PathLike], float]] = None,
|
||||||
|
attn_impl: str = AttentionImplementations.SPLIT_EINSUM.name,
|
||||||
|
config_path: str = None,
|
||||||
|
):
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
|
||||||
|
ref_pipe = get_pipeline(ckpt_path, config_path, model_version)
|
||||||
|
|
||||||
|
for i, lora_weight in enumerate(lora_weights or []):
|
||||||
|
lora_path, strength = lora_weight
|
||||||
|
adapter_name = f"lora_{i}"
|
||||||
|
ref_pipe.load_lora_weights(lora_path, adapter_name=adapter_name)
|
||||||
|
ref_pipe.set_adapters([adapter_name], adapter_weights=[strength])
|
||||||
|
ref_pipe.fuse_lora()
|
||||||
|
|
||||||
|
convert_unet(
|
||||||
|
ref_pipe,
|
||||||
|
model_version,
|
||||||
|
unet_out_path,
|
||||||
|
batch_size,
|
||||||
|
sample_size,
|
||||||
|
controlnet_support,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_pipeline(ckpt_path, config_path, model_version):
|
||||||
|
pipe_cls = MODEL_TYPE_TO_PIPE_CLS[model_version]
|
||||||
|
ref_pipe = pipe_cls.from_single_file(ckpt_path, original_config_file=config_path)
|
||||||
|
return ref_pipe
|
||||||
|
|
||||||
|
|
||||||
|
def compile_model(out_path, out_name, submodule_name):
|
||||||
|
# Compile the model
|
||||||
|
target_path = compile_coreml_model(
|
||||||
|
out_path, get_folder_paths(submodule_name)[0], f"{out_name}_{submodule_name}"
|
||||||
|
)
|
||||||
|
logger.info(f"Compiled {out_path} to {target_path}")
|
||||||
|
return target_path
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_batch(input_tensor, target_shape):
|
||||||
|
if input_tensor.shape == target_shape:
|
||||||
|
return [input_tensor]
|
||||||
|
|
||||||
|
batch_size = input_tensor.shape[0]
|
||||||
|
target_batch_size = target_shape[0]
|
||||||
|
|
||||||
|
num_chunks = batch_size // target_batch_size
|
||||||
|
if num_chunks == 0:
|
||||||
|
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to(
|
||||||
|
input_tensor.device
|
||||||
|
)
|
||||||
|
return [torch.cat((input_tensor, padding), dim=0)]
|
||||||
|
|
||||||
|
mod = batch_size % target_batch_size
|
||||||
|
if mod != 0:
|
||||||
|
chunks = list(torch.chunk(input_tensor[:-mod], num_chunks))
|
||||||
|
padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to(
|
||||||
|
input_tensor.device
|
||||||
|
)
|
||||||
|
padded = torch.cat((input_tensor[-mod:], padding), dim=0)
|
||||||
|
chunks.append(padded)
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
chunks = list(torch.chunk(input_tensor, num_chunks))
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
|
def merge_chunks(chunks, orig_shape):
|
||||||
|
merged = torch.cat(chunks, dim=0)
|
||||||
|
if merged.shape == orig_shape:
|
||||||
|
return merged
|
||||||
|
return merged[: orig_shape[0]]
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .nodes import COREML_CONVERT_LCM
|
||||||
|
|
||||||
|
__all__ = ["COREML_CONVERT_LCM"]
|
||||||
@@ -0,0 +1,297 @@
|
|||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
import gc
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from diffusers import UNet2DConditionModel, LCMScheduler
|
||||||
|
from diffusers.loaders import LoraLoaderMixin
|
||||||
|
|
||||||
|
from comfy.model_management import get_torch_device
|
||||||
|
from coreml_suite.lcm.unet import UNet2DConditionModelLCM
|
||||||
|
|
||||||
|
from transformers import CLIPTextModel
|
||||||
|
import coremltools as ct
|
||||||
|
|
||||||
|
from folder_paths import get_folder_paths
|
||||||
|
|
||||||
|
logging.basicConfig()
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
logger.setLevel(logging.DEBUG)
|
||||||
|
|
||||||
|
MODEL_VERSION = "SimianLuo/LCM_Dreamshaper_v7"
|
||||||
|
MODEL_NAME = MODEL_VERSION.split("/")[-1] + "_4k"
|
||||||
|
|
||||||
|
import python_coreml_stable_diffusion.unet as unet
|
||||||
|
|
||||||
|
unet.ATTENTION_IMPLEMENTATION_IN_EFFECT = unet.AttentionImplementations.SPLIT_EINSUM
|
||||||
|
|
||||||
|
|
||||||
|
def get_unets():
|
||||||
|
ref_unet = UNet2DConditionModel.from_pretrained(
|
||||||
|
MODEL_VERSION,
|
||||||
|
subfolder="unet",
|
||||||
|
device_map=None,
|
||||||
|
low_cpu_mem_usage=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
cml_unet = UNet2DConditionModelLCM.from_config(ref_unet.config).eval()
|
||||||
|
cml_unet.load_state_dict(ref_unet.state_dict(), strict=False)
|
||||||
|
|
||||||
|
return cml_unet, ref_unet
|
||||||
|
|
||||||
|
|
||||||
|
def get_encoder_hidden_states_shape(unet_config, batch_size):
|
||||||
|
text_encoder = CLIPTextModel.from_pretrained(
|
||||||
|
MODEL_VERSION, subfolder="text_encoder"
|
||||||
|
)
|
||||||
|
|
||||||
|
text_token_sequence_length = text_encoder.config.max_position_embeddings
|
||||||
|
hidden_size = (text_encoder.config.hidden_size,)
|
||||||
|
|
||||||
|
encoder_hidden_states_shape = (
|
||||||
|
batch_size,
|
||||||
|
unet_config.cross_attention_dim or hidden_size,
|
||||||
|
1,
|
||||||
|
text_token_sequence_length,
|
||||||
|
)
|
||||||
|
|
||||||
|
return encoder_hidden_states_shape
|
||||||
|
|
||||||
|
|
||||||
|
def get_scheduler():
|
||||||
|
scheduler = LCMScheduler.from_pretrained(MODEL_VERSION, subfolder="scheduler")
|
||||||
|
scheduler.set_timesteps(50, get_torch_device(), 50)
|
||||||
|
return scheduler
|
||||||
|
|
||||||
|
|
||||||
|
def get_coreml_inputs(sample_inputs):
|
||||||
|
coreml_sample_unet_inputs = {
|
||||||
|
k: v.numpy().astype(np.float16) for k, v in sample_inputs.items()
|
||||||
|
}
|
||||||
|
return [
|
||||||
|
ct.TensorType(
|
||||||
|
name=k,
|
||||||
|
shape=v.shape,
|
||||||
|
dtype=v.numpy().dtype if isinstance(v, torch.Tensor) else v.dtype,
|
||||||
|
)
|
||||||
|
for k, v in coreml_sample_unet_inputs.items()
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def load_coreml_model(out_path):
|
||||||
|
logger.info(f"Loading model from {out_path}")
|
||||||
|
|
||||||
|
start = time.time()
|
||||||
|
coreml_model = ct.models.MLModel(out_path)
|
||||||
|
logger.info(f"Loading {out_path} took {time.time() - start:.1f} seconds")
|
||||||
|
|
||||||
|
return coreml_model
|
||||||
|
|
||||||
|
|
||||||
|
def convert_to_coreml(
|
||||||
|
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")
|
||||||
|
coreml_model = load_coreml_model(out_path)
|
||||||
|
else:
|
||||||
|
logger.info(f"Converting {submodule_name} to CoreML..")
|
||||||
|
coreml_model = ct.convert(
|
||||||
|
torchscript_module,
|
||||||
|
convert_to="mlprogram",
|
||||||
|
minimum_deployment_target=ct.target.macOS13,
|
||||||
|
inputs=sample_inputs,
|
||||||
|
outputs=[
|
||||||
|
ct.TensorType(name=name, dtype=np.float32) for name in output_names
|
||||||
|
],
|
||||||
|
skip_model_load=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
del torchscript_module
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
return coreml_model
|
||||||
|
|
||||||
|
|
||||||
|
def get_out_path(submodule_name, model_name):
|
||||||
|
fname = f"{model_name}_{submodule_name}.mlpackage"
|
||||||
|
unet_path = get_folder_paths(submodule_name)[0]
|
||||||
|
out_path = os.path.join(unet_path, fname)
|
||||||
|
return out_path
|
||||||
|
|
||||||
|
|
||||||
|
def compile_coreml_model(source_model_path, output_dir, final_name):
|
||||||
|
"""Compiles Core ML models using the coremlcompiler utility from Xcode toolchain"""
|
||||||
|
target_path = os.path.join(output_dir, f"{final_name}.mlmodelc")
|
||||||
|
if os.path.exists(target_path):
|
||||||
|
logger.warning(f"Found existing compiled model at {target_path}! Skipping..")
|
||||||
|
return target_path
|
||||||
|
|
||||||
|
logger.info(f"Compiling {source_model_path}")
|
||||||
|
source_model_name = os.path.basename(os.path.splitext(source_model_path)[0])
|
||||||
|
|
||||||
|
os.system(f"xcrun coremlcompiler compile {source_model_path} {output_dir}")
|
||||||
|
compiled_output = os.path.join(output_dir, f"{source_model_name}.mlmodelc")
|
||||||
|
shutil.move(compiled_output, target_path)
|
||||||
|
|
||||||
|
return target_path
|
||||||
|
|
||||||
|
|
||||||
|
def get_sample_input(batch_size, encoder_hidden_states_shape, sample_shape, scheduler):
|
||||||
|
sample_unet_inputs = dict(
|
||||||
|
[
|
||||||
|
("sample", torch.rand(*sample_shape)),
|
||||||
|
(
|
||||||
|
"timestep",
|
||||||
|
torch.tensor([scheduler.timesteps[0].item()] * batch_size).to(
|
||||||
|
torch.float32
|
||||||
|
),
|
||||||
|
),
|
||||||
|
("encoder_hidden_states", torch.rand(*encoder_hidden_states_shape)),
|
||||||
|
("timestep_cond", torch.randn(batch_size, 256).to(torch.float32)),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
return sample_unet_inputs
|
||||||
|
|
||||||
|
|
||||||
|
def get_unet_inputs_spec(sample_unet_inputs):
|
||||||
|
sample_unet_inputs_spec = {
|
||||||
|
k: (v.shape, v.dtype) for k, v in sample_unet_inputs.items()
|
||||||
|
}
|
||||||
|
return sample_unet_inputs_spec
|
||||||
|
|
||||||
|
|
||||||
|
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]
|
||||||
|
h, w = sample_shape[2:]
|
||||||
|
|
||||||
|
# conv_in
|
||||||
|
out_h, out_w = calculate_conv2d_output_shape(
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
reference_unet.conv_in,
|
||||||
|
)
|
||||||
|
additional_residuals_shapes.append(
|
||||||
|
(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
|
||||||
|
]
|
||||||
|
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
|
||||||
|
)
|
||||||
|
additional_residuals_shapes.append(
|
||||||
|
(
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
|
||||||
|
additional_inputs = {}
|
||||||
|
for i, shape in enumerate(additional_residuals_shapes):
|
||||||
|
sample_residual_input = torch.rand(*shape)
|
||||||
|
additional_inputs[f"additional_residual_{i}"] = sample_residual_input
|
||||||
|
|
||||||
|
return additional_inputs
|
||||||
|
|
||||||
|
|
||||||
|
def convert(
|
||||||
|
out_path: str,
|
||||||
|
batch_size: int = 1,
|
||||||
|
sample_size: tuple[int, int] = (64, 64),
|
||||||
|
controlnet_support: bool = False,
|
||||||
|
lora_paths: list[str] = None,
|
||||||
|
):
|
||||||
|
lora_paths = lora_paths or []
|
||||||
|
coreml_unet, ref_unet = get_unets()
|
||||||
|
|
||||||
|
for lora_path in lora_paths:
|
||||||
|
lora_sd, network_alphas = LoraLoaderMixin.lora_state_dict(lora_path)
|
||||||
|
LoraLoaderMixin.load_lora_into_unet(lora_sd, network_alphas, ref_unet)
|
||||||
|
ref_unet.fuse_lora()
|
||||||
|
|
||||||
|
sample_shape = (
|
||||||
|
batch_size, # B
|
||||||
|
ref_unet.config.in_channels, # C
|
||||||
|
sample_size[0], # H
|
||||||
|
sample_size[1], # W
|
||||||
|
)
|
||||||
|
|
||||||
|
encoder_hidden_states_shape = get_encoder_hidden_states_shape(
|
||||||
|
ref_unet.config, batch_size
|
||||||
|
)
|
||||||
|
|
||||||
|
scheduler = get_scheduler()
|
||||||
|
|
||||||
|
sample_inputs = get_sample_input(
|
||||||
|
batch_size, encoder_hidden_states_shape, sample_shape, scheduler
|
||||||
|
)
|
||||||
|
|
||||||
|
if controlnet_support:
|
||||||
|
sample_inputs |= add_cnet_support(sample_shape, ref_unet)
|
||||||
|
|
||||||
|
sample_inputs_spec = get_unet_inputs_spec(sample_inputs)
|
||||||
|
|
||||||
|
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())
|
||||||
|
)
|
||||||
|
logger.info("Done.")
|
||||||
|
|
||||||
|
coreml_sample_inputs = get_coreml_inputs(sample_inputs)
|
||||||
|
|
||||||
|
coreml_unet = convert_to_coreml(
|
||||||
|
"unet", traced_unet, coreml_sample_inputs, ["noise_pred"], out_path
|
||||||
|
)
|
||||||
|
|
||||||
|
del traced_unet
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
coreml_unet.save(out_path)
|
||||||
|
logger.info(f"Saved unet into {out_path}")
|
||||||
|
|
||||||
|
|
||||||
|
def compile_model(out_path, out_name):
|
||||||
|
# Compile the model
|
||||||
|
target_path = compile_coreml_model(
|
||||||
|
out_path, get_folder_paths("unet")[0], f"{out_name}_unet"
|
||||||
|
)
|
||||||
|
logger.info(f"Compiled {out_path} to {target_path}")
|
||||||
|
return target_path
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
h = 512
|
||||||
|
w = 512
|
||||||
|
sample_size = (h // 8, w // 8)
|
||||||
|
batch_size = 4
|
||||||
|
|
||||||
|
cn_support_str = "_cn" if True else ""
|
||||||
|
|
||||||
|
out_name = f"{MODEL_NAME}_{batch_size}x{w}x{h}{cn_support_str}"
|
||||||
|
|
||||||
|
out_path = get_out_path("unet", f"{out_name}")
|
||||||
|
if not os.path.exists(out_path):
|
||||||
|
convert(out_path=out_path, sample_size=sample_size, batch_size=batch_size)
|
||||||
|
compile_model(out_path=out_path, out_name=out_name)
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
from coremltools import ComputeUnit
|
||||||
|
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||||
|
|
||||||
|
from coreml_suite import COREML_NODE
|
||||||
|
from coreml_suite.lcm import converter as lcm_converter
|
||||||
|
|
||||||
|
|
||||||
|
class COREML_CONVERT_LCM(COREML_NODE):
|
||||||
|
"""Converts a LCM model to Core ML."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"height": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
|
||||||
|
"width": ("INT", {"default": 512, "min": 512, "max": 768, "step": 8}),
|
||||||
|
"batch_size": ("INT", {"default": 1, "min": 1, "max": 64}),
|
||||||
|
"compute_unit": (
|
||||||
|
[
|
||||||
|
ComputeUnit.CPU_AND_NE.name,
|
||||||
|
ComputeUnit.CPU_AND_GPU.name,
|
||||||
|
ComputeUnit.ALL.name,
|
||||||
|
ComputeUnit.CPU_ONLY.name,
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"controlnet_support": ("BOOLEAN", {"default": False}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("COREML_UNET",)
|
||||||
|
RETURN_NAMES = ("coreml_model",)
|
||||||
|
FUNCTION = "convert"
|
||||||
|
|
||||||
|
def convert(self, height, width, batch_size, compute_unit, controlnet_support):
|
||||||
|
"""Converts a LCM model to Core ML.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
height (int): Height of the target image.
|
||||||
|
width (int): Width of the target image.
|
||||||
|
batch_size (int): Batch size.
|
||||||
|
compute_unit (str): Compute unit to use when loading the model.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
coreml_model: The converted Core ML model.
|
||||||
|
|
||||||
|
The converted model is also saved to "models/unet" directory and
|
||||||
|
can be loaded with the "LCMCoreMLLoaderUNet" node.
|
||||||
|
"""
|
||||||
|
h = height
|
||||||
|
w = width
|
||||||
|
sample_size = (h // 8, w // 8)
|
||||||
|
batch_size = batch_size
|
||||||
|
cn_support_str = "_cn" if controlnet_support else ""
|
||||||
|
|
||||||
|
out_name = f"{lcm_converter.MODEL_NAME}_{batch_size}x{w}x{h}{cn_support_str}"
|
||||||
|
|
||||||
|
out_path = lcm_converter.get_out_path("unet", f"{out_name}")
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name)
|
||||||
|
|
||||||
|
return (CoreMLModel(target_path, compute_unit, "compiled"),)
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
from overrides import overrides
|
||||||
|
from python_coreml_stable_diffusion.unet import UNet2DConditionModel, TimestepEmbedding
|
||||||
|
|
||||||
|
|
||||||
|
class UNet2DConditionModelLCM(UNet2DConditionModel):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
time_cond_proj_dim=None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
timestep_input_dim = self.config.block_out_channels[0]
|
||||||
|
time_embed_dim = self.config.block_out_channels[0] * 4
|
||||||
|
|
||||||
|
time_embedding = TimestepEmbedding(
|
||||||
|
timestep_input_dim, time_embed_dim, cond_proj_dim=time_cond_proj_dim
|
||||||
|
)
|
||||||
|
self.time_embedding = time_embedding
|
||||||
|
|
||||||
|
@overrides(check_signature=False)
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
sample,
|
||||||
|
timestep,
|
||||||
|
encoder_hidden_states,
|
||||||
|
timestep_cond,
|
||||||
|
*additional_residuals,
|
||||||
|
):
|
||||||
|
# 0. Project (or look-up) time embeddings
|
||||||
|
t_emb = self.time_proj(timestep)
|
||||||
|
emb = self.time_embedding(t_emb, timestep_cond)
|
||||||
|
|
||||||
|
# 1. center input if necessary
|
||||||
|
if self.config.center_input_sample:
|
||||||
|
sample = 2 * sample - 1.0
|
||||||
|
|
||||||
|
# 2. pre-process
|
||||||
|
sample = self.conv_in(sample)
|
||||||
|
|
||||||
|
# 3. down
|
||||||
|
down_block_res_samples = (sample,)
|
||||||
|
for downsample_block in self.down_blocks:
|
||||||
|
if (
|
||||||
|
hasattr(downsample_block, "attentions")
|
||||||
|
and downsample_block.attentions is not None
|
||||||
|
):
|
||||||
|
sample, res_samples = downsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
|
||||||
|
|
||||||
|
down_block_res_samples += res_samples
|
||||||
|
|
||||||
|
if additional_residuals:
|
||||||
|
new_down_block_res_samples = ()
|
||||||
|
for i, down_block_res_sample in enumerate(down_block_res_samples):
|
||||||
|
down_block_res_sample = down_block_res_sample + additional_residuals[i]
|
||||||
|
new_down_block_res_samples += (down_block_res_sample,)
|
||||||
|
down_block_res_samples = new_down_block_res_samples
|
||||||
|
|
||||||
|
# 4. mid
|
||||||
|
sample = self.mid_block(
|
||||||
|
sample, emb, encoder_hidden_states=encoder_hidden_states
|
||||||
|
)
|
||||||
|
|
||||||
|
if additional_residuals:
|
||||||
|
sample = sample + additional_residuals[-1]
|
||||||
|
|
||||||
|
# 5. up
|
||||||
|
for upsample_block in self.up_blocks:
|
||||||
|
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||||
|
down_block_res_samples = down_block_res_samples[
|
||||||
|
: -len(upsample_block.resnets)
|
||||||
|
]
|
||||||
|
|
||||||
|
if (
|
||||||
|
hasattr(upsample_block, "attentions")
|
||||||
|
and upsample_block.attentions is not None
|
||||||
|
):
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
res_hidden_states_tuple=res_samples,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample, temb=emb, res_hidden_states_tuple=res_samples
|
||||||
|
)
|
||||||
|
|
||||||
|
# 6. post-process
|
||||||
|
sample = self.conv_norm_out(sample)
|
||||||
|
sample = self.conv_act(sample)
|
||||||
|
sample = self.conv_out(sample)
|
||||||
|
|
||||||
|
return (sample,)
|
||||||
@@ -0,0 +1,73 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
from comfy.model_management import get_torch_device
|
||||||
|
from comfy_extras.nodes_model_advanced import ModelSamplingDiscreteDistilled, LCM
|
||||||
|
|
||||||
|
|
||||||
|
def is_lcm(coreml_model):
|
||||||
|
return "timestep_cond" in coreml_model.expected_inputs
|
||||||
|
|
||||||
|
|
||||||
|
def get_w_embedding(w, embedding_dim=512, dtype=torch.float32):
|
||||||
|
assert len(w.shape) == 1
|
||||||
|
w = w * 1000.0
|
||||||
|
|
||||||
|
half_dim = embedding_dim // 2
|
||||||
|
emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1)
|
||||||
|
emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb)
|
||||||
|
emb = w.to(dtype)[:, None] * emb[None, :]
|
||||||
|
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
|
||||||
|
if embedding_dim % 2 == 1: # zero pad
|
||||||
|
emb = torch.nn.functional.pad(emb, (0, 1))
|
||||||
|
assert emb.shape == (w.shape[0], embedding_dim)
|
||||||
|
return emb
|
||||||
|
|
||||||
|
|
||||||
|
def model_function_wrapper(w_embedding):
|
||||||
|
def wrapper(model_function, params):
|
||||||
|
x = params["input"]
|
||||||
|
t = params["timestep"]
|
||||||
|
c = params["c"]
|
||||||
|
|
||||||
|
context = c.get("c_crossattn")
|
||||||
|
|
||||||
|
if context is None:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
|
||||||
|
return model_function(x, t, **c, timestep_cond=w_embedding)
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
|
||||||
|
def lcm_patch(model):
|
||||||
|
m = model.clone()
|
||||||
|
sampling_type = LCM
|
||||||
|
sampling_base = ModelSamplingDiscreteDistilled
|
||||||
|
|
||||||
|
class ModelSamplingAdvanced(sampling_base, sampling_type):
|
||||||
|
pass
|
||||||
|
|
||||||
|
model_sampling = ModelSamplingAdvanced()
|
||||||
|
m.add_object_patch("model_sampling", model_sampling)
|
||||||
|
|
||||||
|
return m
|
||||||
|
|
||||||
|
|
||||||
|
def add_lcm_model_options(model_patcher, cfg, latent_image):
|
||||||
|
mp = model_patcher.clone()
|
||||||
|
|
||||||
|
latent = latent_image["samples"].to(get_torch_device())
|
||||||
|
batch_size = latent.shape[0]
|
||||||
|
dtype = latent.dtype
|
||||||
|
device = get_torch_device()
|
||||||
|
|
||||||
|
w = torch.tensor(cfg).repeat(batch_size)
|
||||||
|
w_embedding = get_w_embedding(w, embedding_dim=256).to(device=device, dtype=dtype)
|
||||||
|
|
||||||
|
model_options = {
|
||||||
|
"model_function_wrapper": model_function_wrapper(w_embedding),
|
||||||
|
"sampler_cfg_function": lambda x: x["cond"].to(device),
|
||||||
|
}
|
||||||
|
mp.model_options |= model_options
|
||||||
|
|
||||||
|
return mp
|
||||||
@@ -1,84 +0,0 @@
|
|||||||
import os.path
|
|
||||||
|
|
||||||
from coremltools import ComputeUnit
|
|
||||||
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
|
||||||
|
|
||||||
import folder_paths
|
|
||||||
|
|
||||||
from coreml_suite.logger import logger
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLLoader:
|
|
||||||
PACKAGE_DIRNAME = ""
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"coreml_name": (list(s.coreml_filenames().keys()),),
|
|
||||||
"compute_unit": (
|
|
||||||
[
|
|
||||||
ComputeUnit.CPU_AND_NE.name,
|
|
||||||
ComputeUnit.CPU_AND_GPU.name,
|
|
||||||
ComputeUnit.ALL.name,
|
|
||||||
ComputeUnit.CPU_ONLY.name,
|
|
||||||
],
|
|
||||||
),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
FUNCTION = "load"
|
|
||||||
CATEGORY = "Core ML Suite"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def coreml_filenames(cls):
|
|
||||||
extensions = (".mlmodelc", ".mlpackage")
|
|
||||||
all_paths = folder_paths.get_filename_list_(cls.PACKAGE_DIRNAME)[1]
|
|
||||||
coreml_paths = folder_paths.filter_files_extensions(all_paths, extensions)
|
|
||||||
|
|
||||||
return {os.path.split(p)[-1]: p for p in coreml_paths}
|
|
||||||
|
|
||||||
def load(self, coreml_name, compute_unit):
|
|
||||||
logger.info(f"Loading {coreml_name} to {compute_unit}")
|
|
||||||
|
|
||||||
coreml_path = self.coreml_filenames()[coreml_name]
|
|
||||||
|
|
||||||
sources = "compiled" if coreml_name.endswith(".mlmodelc") else "packages"
|
|
||||||
|
|
||||||
return self._load(coreml_path, compute_unit, sources)
|
|
||||||
|
|
||||||
def _load(self, coreml_path, compute_unit, sources):
|
|
||||||
return (CoreMLModel(coreml_path, compute_unit, sources),)
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLLoaderCkpt(CoreMLLoader):
|
|
||||||
PACKAGE_DIRNAME = "checkpoints"
|
|
||||||
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
|
|
||||||
|
|
||||||
def load(self, coreml_name, compute_unit):
|
|
||||||
# TODO: Implement this
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLLoaderTextEncoder(CoreMLLoader):
|
|
||||||
PACKAGE_DIRNAME = "clip"
|
|
||||||
RETURN_TYPES = ("CLIP",)
|
|
||||||
|
|
||||||
def load(self, coreml_name, compute_unit):
|
|
||||||
# TODO: Implement this
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLLoaderUNet(CoreMLLoader):
|
|
||||||
PACKAGE_DIRNAME = "unet"
|
|
||||||
RETURN_TYPES = ("COREML_UNET",)
|
|
||||||
RETURN_NAMES = ("coreml_model",)
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLLoaderVAE(CoreMLLoader):
|
|
||||||
PACKAGE_DIRNAME = "vae"
|
|
||||||
RETURN_TYPES = ("VAE",)
|
|
||||||
|
|
||||||
def load(self, coreml_name, compute_unit):
|
|
||||||
# TODO: Implement this
|
|
||||||
pass
|
|
||||||
@@ -1,63 +1,295 @@
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from comfy import model_base
|
||||||
from comfy import supported_models_base
|
from comfy.model_management import get_torch_device
|
||||||
from comfy.latent_formats import SD15
|
from comfy.model_patcher import ModelPatcher
|
||||||
from comfy.model_base import BaseModel
|
from coreml_suite.config import get_model_config, ModelVersion
|
||||||
|
from coreml_suite.controlnet import extract_residual_kwargs, chunk_control
|
||||||
from coreml_suite.utils import expand_inputs, extract_residual_kwargs
|
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||||
|
from coreml_suite.lcm.utils import is_lcm
|
||||||
|
from coreml_suite.logger import logger
|
||||||
|
|
||||||
|
|
||||||
def get_model_config():
|
class CoreMLModelWrapper:
|
||||||
# TODO: This is a dummy model config, but it should be enough to
|
def __init__(self, coreml_model):
|
||||||
# get the model to load - implement a proper model config
|
self.coreml_model = coreml_model
|
||||||
model_config = supported_models_base.BASE({})
|
self.dtype = torch.float16
|
||||||
model_config.latent_format = SD15()
|
|
||||||
model_config.unet_config = {
|
def __call__(self, x, t, context, control, transformer_options=None, **kwargs):
|
||||||
"disable_unet_model_creation": True,
|
inputs = CoreMLInputs(x, t, context, control, **kwargs)
|
||||||
"num_res_blocks": 2,
|
input_list = inputs.chunks(self.expected_inputs)
|
||||||
"attention_resolutions": [1, 2, 4],
|
|
||||||
"channel_mult": [1, 2, 4, 4],
|
chunked_out = [
|
||||||
"transformer_depth": [1, 1, 1, 0],
|
self.get_torch_outputs(
|
||||||
}
|
self.coreml_model(**input_kwargs.coreml_kwargs(self.expected_inputs)),
|
||||||
return model_config
|
x.device,
|
||||||
|
)
|
||||||
|
for input_kwargs in input_list
|
||||||
|
]
|
||||||
|
merged_out = merge_chunks(chunked_out, x.shape)
|
||||||
|
|
||||||
|
return merged_out
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_torch_outputs(model_output, device):
|
||||||
|
return torch.from_numpy(model_output["noise_pred"]).to(device)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def expected_inputs(self):
|
||||||
|
return self.coreml_model.expected_inputs
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_lcm(self):
|
||||||
|
return is_lcm(self.coreml_model)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_sdxl_base(self):
|
||||||
|
return is_sdxl_base(self.coreml_model)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_sdxl_refiner(self):
|
||||||
|
return is_sdxl_refiner(self.coreml_model)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def config(self):
|
||||||
|
if self.is_sdxl_base:
|
||||||
|
return get_model_config(ModelVersion.SDXL)
|
||||||
|
|
||||||
|
if self.is_sdxl_refiner:
|
||||||
|
return get_model_config(ModelVersion.SDXL_REFINER)
|
||||||
|
|
||||||
|
return get_model_config(ModelVersion.SD15)
|
||||||
|
|
||||||
|
|
||||||
class CoreMLModelWrapper(BaseModel):
|
class CoreMLModelWrapperLCM(CoreMLModelWrapper):
|
||||||
def __init__(self, model_config, coreml_model):
|
def __init__(self, coreml_model):
|
||||||
super().__init__(model_config)
|
super().__init__(coreml_model)
|
||||||
self.diffusion_model = coreml_model
|
self.config = None
|
||||||
|
|
||||||
def apply_model(
|
|
||||||
self,
|
|
||||||
x,
|
|
||||||
t,
|
|
||||||
c_concat=None,
|
|
||||||
c_crossattn=None,
|
|
||||||
c_adm=None,
|
|
||||||
control=None,
|
|
||||||
transformer_options={},
|
|
||||||
):
|
|
||||||
sample = x.cpu().numpy().astype(np.float16)
|
|
||||||
|
|
||||||
context = c_crossattn.cpu().numpy().astype(np.float16)
|
class CoreMLInputs:
|
||||||
|
def __init__(self, x, t, context, control, **kwargs):
|
||||||
|
self.x = x
|
||||||
|
self.t = t
|
||||||
|
self.context = context
|
||||||
|
self.control = control
|
||||||
|
self.time_ids = kwargs.get("time_ids")
|
||||||
|
self.text_embeds = kwargs.get("text_embeds")
|
||||||
|
self.ts_cond = kwargs.get("timestep_cond")
|
||||||
|
|
||||||
|
def coreml_kwargs(self, expected_inputs):
|
||||||
|
sample = self.x.cpu().numpy().astype(np.float16)
|
||||||
|
|
||||||
|
context = self.context.cpu().numpy().astype(np.float16)
|
||||||
context = context.transpose(0, 2, 1)[:, :, None, :]
|
context = context.transpose(0, 2, 1)[:, :, None, :]
|
||||||
|
|
||||||
t = t.cpu().numpy().astype(np.float16)
|
t = self.t.cpu().numpy().astype(np.float16)
|
||||||
|
|
||||||
model_input_kwargs = {
|
model_input_kwargs = {
|
||||||
"sample": sample,
|
"sample": sample,
|
||||||
"encoder_hidden_states": context,
|
"encoder_hidden_states": context,
|
||||||
"timestep": t,
|
"timestep": t,
|
||||||
}
|
}
|
||||||
residual_kwargs = extract_residual_kwargs(self.diffusion_model, control)
|
residual_kwargs = extract_residual_kwargs(expected_inputs, self.control)
|
||||||
model_input_kwargs |= residual_kwargs
|
model_input_kwargs |= residual_kwargs
|
||||||
model_input_kwargs = expand_inputs(model_input_kwargs)
|
|
||||||
|
|
||||||
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
|
# LCM
|
||||||
return torch.from_numpy(np_out).to(x.device)
|
if self.ts_cond is not None:
|
||||||
|
model_input_kwargs["timestep_cond"] = (
|
||||||
|
self.ts_cond.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
|
||||||
def get_dtype(self):
|
# SDXL
|
||||||
# Hardcoding torch-compatible dtype (used for memory allocation)
|
if "text_embeds" in expected_inputs:
|
||||||
return torch.float16
|
model_input_kwargs["text_embeds"] = (
|
||||||
|
self.text_embeds.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
if "time_ids" in expected_inputs:
|
||||||
|
model_input_kwargs["time_ids"] = (
|
||||||
|
self.time_ids.cpu().numpy().astype(np.float16)
|
||||||
|
)
|
||||||
|
|
||||||
|
return model_input_kwargs
|
||||||
|
|
||||||
|
def chunks(self, expected_inputs):
|
||||||
|
sample_shape = expected_inputs["sample"]["shape"]
|
||||||
|
timestep_shape = expected_inputs["timestep"]["shape"]
|
||||||
|
hidden_shape = expected_inputs["encoder_hidden_states"]["shape"]
|
||||||
|
context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1])
|
||||||
|
|
||||||
|
chunked_x = chunk_batch(self.x, sample_shape)
|
||||||
|
ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0]))
|
||||||
|
chunked_context = chunk_batch(self.context, context_shape)
|
||||||
|
|
||||||
|
chunked_control = [None] * len(chunked_x)
|
||||||
|
if self.control is not None:
|
||||||
|
chunked_control = chunk_control(self.control, sample_shape[0])
|
||||||
|
|
||||||
|
chunked_ts_cond = [None] * len(chunked_x)
|
||||||
|
if self.ts_cond is not None:
|
||||||
|
ts_cond_shape = expected_inputs["timestep_cond"]["shape"]
|
||||||
|
chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape)
|
||||||
|
|
||||||
|
chunked_time_ids = [None] * len(chunked_x)
|
||||||
|
if expected_inputs.get("time_ids") is not None:
|
||||||
|
time_ids_shape = expected_inputs["time_ids"]["shape"]
|
||||||
|
if self.time_ids is None:
|
||||||
|
self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to(
|
||||||
|
self.x.device
|
||||||
|
)
|
||||||
|
chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape)
|
||||||
|
|
||||||
|
chunked_text_embeds = [None] * len(chunked_x)
|
||||||
|
if expected_inputs.get("text_embeds") is not None:
|
||||||
|
text_embeds_shape = expected_inputs["text_embeds"]["shape"]
|
||||||
|
if self.text_embeds is None:
|
||||||
|
self.text_embeds = torch.zeros(
|
||||||
|
len(chunked_x), *text_embeds_shape[1:]
|
||||||
|
).to(self.x.device)
|
||||||
|
chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape)
|
||||||
|
|
||||||
|
return [
|
||||||
|
CoreMLInputs(
|
||||||
|
x,
|
||||||
|
t,
|
||||||
|
context,
|
||||||
|
control,
|
||||||
|
timestep_cond=ts_cond,
|
||||||
|
time_ids=time_ids,
|
||||||
|
text_embeds=text_embeds,
|
||||||
|
)
|
||||||
|
for x, t, context, control, ts_cond, time_ids, text_embeds in zip(
|
||||||
|
chunked_x,
|
||||||
|
ts,
|
||||||
|
chunked_context,
|
||||||
|
chunked_control,
|
||||||
|
chunked_ts_cond,
|
||||||
|
chunked_time_ids,
|
||||||
|
chunked_text_embeds,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl(coreml_model):
|
||||||
|
return (
|
||||||
|
"time_ids" in coreml_model.expected_inputs
|
||||||
|
and "text_embeds" in coreml_model.expected_inputs
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl_base(coreml_model):
|
||||||
|
return (
|
||||||
|
is_sdxl(coreml_model)
|
||||||
|
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_sdxl_refiner(coreml_model):
|
||||||
|
return (
|
||||||
|
is_sdxl(coreml_model)
|
||||||
|
and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False):
|
||||||
|
def wrapper(model_function, params):
|
||||||
|
x = params["input"]
|
||||||
|
t = params["timestep"]
|
||||||
|
c = params["c"]
|
||||||
|
|
||||||
|
context = c.get("c_crossattn")
|
||||||
|
|
||||||
|
if context is None:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
|
||||||
|
if refiner and context is not None:
|
||||||
|
# converted refiner accepts only g clip
|
||||||
|
c["c_crossattn"] = context[:, :, 768:]
|
||||||
|
|
||||||
|
return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds)
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
|
||||||
|
def add_sdxl_model_options(model_patcher, positive, negative):
|
||||||
|
mp = model_patcher.clone()
|
||||||
|
|
||||||
|
pos_dict = positive[0][1]
|
||||||
|
neg_dict = negative[0][1]
|
||||||
|
|
||||||
|
pos_pooled = pos_dict["pooled_output"]
|
||||||
|
neg_pooled = neg_dict["pooled_output"]
|
||||||
|
|
||||||
|
pos_time_ids = [
|
||||||
|
pos_dict.get("height", 768),
|
||||||
|
pos_dict.get("width", 768),
|
||||||
|
pos_dict.get("crop_h", 0),
|
||||||
|
pos_dict.get("crop_w", 0),
|
||||||
|
]
|
||||||
|
|
||||||
|
neg_time_ids = [
|
||||||
|
neg_dict.get("height", 768),
|
||||||
|
neg_dict.get("width", 768),
|
||||||
|
neg_dict.get("crop_h", 0),
|
||||||
|
neg_dict.get("crop_w", 0),
|
||||||
|
]
|
||||||
|
|
||||||
|
if model_patcher.model.diffusion_model.is_sdxl_base:
|
||||||
|
pos_time_ids += [
|
||||||
|
pos_dict.get("target_height", 768),
|
||||||
|
pos_dict.get("target_width", 768),
|
||||||
|
]
|
||||||
|
|
||||||
|
neg_time_ids += [
|
||||||
|
neg_dict.get("target_height", 768),
|
||||||
|
neg_dict.get("target_width", 768),
|
||||||
|
]
|
||||||
|
|
||||||
|
is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner
|
||||||
|
if is_refiner:
|
||||||
|
pos_time_ids += [
|
||||||
|
pos_dict.get("aesthetic_score", 6),
|
||||||
|
]
|
||||||
|
|
||||||
|
neg_time_ids += [
|
||||||
|
neg_dict.get("aesthetic_score", 2.5),
|
||||||
|
]
|
||||||
|
|
||||||
|
time_ids = torch.tensor([pos_time_ids, neg_time_ids])
|
||||||
|
text_embeds = torch.cat((pos_pooled, neg_pooled))
|
||||||
|
|
||||||
|
model_options = {
|
||||||
|
"model_function_wrapper": sdxl_model_function_wrapper(
|
||||||
|
time_ids, text_embeds, is_refiner
|
||||||
|
),
|
||||||
|
}
|
||||||
|
mp.model_options |= model_options
|
||||||
|
|
||||||
|
return mp
|
||||||
|
|
||||||
|
|
||||||
|
def get_latent_image(coreml_model, latent_image):
|
||||||
|
if latent_image is not None:
|
||||||
|
return latent_image
|
||||||
|
|
||||||
|
logger.warning("No latent image provided, using empty tensor.")
|
||||||
|
expected = coreml_model.expected_inputs["sample"]["shape"]
|
||||||
|
batch_size = max(expected[0] // 2, 1)
|
||||||
|
latent_image = {"samples": torch.zeros(batch_size, *expected[1:])}
|
||||||
|
return latent_image
|
||||||
|
|
||||||
|
|
||||||
|
def get_model_patcher(coreml_model):
|
||||||
|
wrapped_model = CoreMLModelWrapper(coreml_model)
|
||||||
|
|
||||||
|
if wrapped_model.is_sdxl_base:
|
||||||
|
model = model_base.SDXL(wrapped_model.config, device=get_torch_device())
|
||||||
|
elif wrapped_model.is_sdxl_refiner:
|
||||||
|
model = model_base.SDXLRefiner(wrapped_model.config, device=get_torch_device())
|
||||||
|
else:
|
||||||
|
model = model_base.BaseModel(wrapped_model.config, device=get_torch_device())
|
||||||
|
|
||||||
|
model.diffusion_model = wrapped_model
|
||||||
|
model_patcher = ModelPatcher(model, get_torch_device(), None)
|
||||||
|
return model_patcher
|
||||||
|
|||||||
@@ -0,0 +1,373 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
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 coreml_suite import COREML_NODE
|
||||||
|
from coreml_suite import converter
|
||||||
|
from coreml_suite.config import ModelVersion
|
||||||
|
from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm
|
||||||
|
from coreml_suite.logger import logger
|
||||||
|
from nodes import KSampler, LoraLoader, KSamplerAdvanced
|
||||||
|
|
||||||
|
from coreml_suite.models import (
|
||||||
|
add_sdxl_model_options,
|
||||||
|
is_sdxl,
|
||||||
|
get_model_patcher,
|
||||||
|
get_latent_image,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLSampler(COREML_NODE, KSampler):
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
old_required = KSampler.INPUT_TYPES()["required"].copy()
|
||||||
|
old_required.pop("model")
|
||||||
|
old_required.pop("negative")
|
||||||
|
old_required.pop("latent_image")
|
||||||
|
new_required = {"coreml_model": ("COREML_UNET",)}
|
||||||
|
return {
|
||||||
|
"required": new_required | old_required,
|
||||||
|
"optional": {"negative": ("CONDITIONING",), "latent_image": ("LATENT",)},
|
||||||
|
}
|
||||||
|
|
||||||
|
def sample(
|
||||||
|
self,
|
||||||
|
coreml_model,
|
||||||
|
seed,
|
||||||
|
steps,
|
||||||
|
cfg,
|
||||||
|
sampler_name,
|
||||||
|
scheduler,
|
||||||
|
positive,
|
||||||
|
negative=None,
|
||||||
|
latent_image=None,
|
||||||
|
denoise=1.0,
|
||||||
|
):
|
||||||
|
model_patcher = get_model_patcher(coreml_model)
|
||||||
|
latent_image = get_latent_image(coreml_model, latent_image)
|
||||||
|
|
||||||
|
if is_lcm(coreml_model):
|
||||||
|
negative = [[None, {}]]
|
||||||
|
positive[0][1]["control_apply_to_uncond"] = False
|
||||||
|
model_patcher = add_lcm_model_options(model_patcher, cfg, latent_image)
|
||||||
|
model_patcher = lcm_patch(model_patcher)
|
||||||
|
else:
|
||||||
|
assert (
|
||||||
|
negative is not None
|
||||||
|
), "Negative conditioning is optional only for LCM models."
|
||||||
|
|
||||||
|
if is_sdxl(coreml_model):
|
||||||
|
model_patcher = add_sdxl_model_options(model_patcher, positive, negative)
|
||||||
|
|
||||||
|
return super().sample(
|
||||||
|
model_patcher,
|
||||||
|
seed,
|
||||||
|
steps,
|
||||||
|
cfg,
|
||||||
|
sampler_name,
|
||||||
|
scheduler,
|
||||||
|
positive,
|
||||||
|
negative,
|
||||||
|
latent_image,
|
||||||
|
denoise,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLSamplerAdvanced(COREML_NODE, KSamplerAdvanced):
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
old_required = KSamplerAdvanced.INPUT_TYPES()["required"].copy()
|
||||||
|
old_required.pop("model")
|
||||||
|
old_required.pop("negative")
|
||||||
|
old_required.pop("latent_image")
|
||||||
|
new_required = {"coreml_model": ("COREML_UNET",)}
|
||||||
|
return {
|
||||||
|
"required": new_required | old_required,
|
||||||
|
"optional": {"negative": ("CONDITIONING",), "latent_image": ("LATENT",)},
|
||||||
|
}
|
||||||
|
|
||||||
|
def sample(
|
||||||
|
self,
|
||||||
|
coreml_model,
|
||||||
|
add_noise,
|
||||||
|
noise_seed,
|
||||||
|
steps,
|
||||||
|
cfg,
|
||||||
|
sampler_name,
|
||||||
|
scheduler,
|
||||||
|
positive,
|
||||||
|
start_at_step,
|
||||||
|
end_at_step,
|
||||||
|
return_with_leftover_noise,
|
||||||
|
negative=None,
|
||||||
|
latent_image=None,
|
||||||
|
denoise=1.0,
|
||||||
|
):
|
||||||
|
model_patcher = get_model_patcher(coreml_model)
|
||||||
|
latent_image = get_latent_image(coreml_model, latent_image)
|
||||||
|
|
||||||
|
if is_lcm(coreml_model):
|
||||||
|
negative = [[None, {}]]
|
||||||
|
positive[0][1]["control_apply_to_uncond"] = False
|
||||||
|
model_patcher = add_lcm_model_options(model_patcher, cfg, latent_image)
|
||||||
|
model_patcher = lcm_patch(model_patcher)
|
||||||
|
else:
|
||||||
|
assert (
|
||||||
|
negative is not None
|
||||||
|
), "Negative conditioning is optional only for LCM models."
|
||||||
|
|
||||||
|
if is_sdxl(coreml_model):
|
||||||
|
model_patcher = add_sdxl_model_options(model_patcher, positive, negative)
|
||||||
|
|
||||||
|
return super().sample(
|
||||||
|
model_patcher,
|
||||||
|
add_noise,
|
||||||
|
noise_seed,
|
||||||
|
steps,
|
||||||
|
cfg,
|
||||||
|
sampler_name,
|
||||||
|
scheduler,
|
||||||
|
positive,
|
||||||
|
negative,
|
||||||
|
latent_image,
|
||||||
|
start_at_step,
|
||||||
|
end_at_step,
|
||||||
|
return_with_leftover_noise,
|
||||||
|
denoise,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLLoader(COREML_NODE):
|
||||||
|
PACKAGE_DIRNAME = ""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"coreml_name": (list(s.coreml_filenames().keys()),),
|
||||||
|
"compute_unit": (
|
||||||
|
[
|
||||||
|
ComputeUnit.CPU_AND_NE.name,
|
||||||
|
ComputeUnit.CPU_AND_GPU.name,
|
||||||
|
ComputeUnit.ALL.name,
|
||||||
|
ComputeUnit.CPU_ONLY.name,
|
||||||
|
],
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
FUNCTION = "load"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def coreml_filenames(cls):
|
||||||
|
extensions = (".mlmodelc", ".mlpackage")
|
||||||
|
all_paths = folder_paths.get_filename_list_(cls.PACKAGE_DIRNAME)[1]
|
||||||
|
coreml_paths = folder_paths.filter_files_extensions(all_paths, extensions)
|
||||||
|
|
||||||
|
return {os.path.split(p)[-1]: p for p in coreml_paths}
|
||||||
|
|
||||||
|
def load(self, coreml_name, compute_unit):
|
||||||
|
logger.info(f"Loading {coreml_name} to {compute_unit}")
|
||||||
|
|
||||||
|
coreml_path = self.coreml_filenames()[coreml_name]
|
||||||
|
|
||||||
|
sources = "compiled" if coreml_name.endswith(".mlmodelc") else "packages"
|
||||||
|
|
||||||
|
return (CoreMLModel(coreml_path, compute_unit, sources),)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLLoaderUNet(CoreMLLoader):
|
||||||
|
PACKAGE_DIRNAME = "unet"
|
||||||
|
RETURN_TYPES = ("COREML_UNET",)
|
||||||
|
RETURN_NAMES = ("coreml_model",)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLModelAdapter(COREML_NODE):
|
||||||
|
"""
|
||||||
|
Adapter Node to use CoreML models as Comfy models. This is an experimental
|
||||||
|
feature and may not work as expected.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"coreml_model": ("COREML_UNET",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("MODEL",)
|
||||||
|
|
||||||
|
FUNCTION = "wrap"
|
||||||
|
CATEGORY = "Core ML Suite"
|
||||||
|
|
||||||
|
def wrap(self, coreml_model):
|
||||||
|
model_patcher = get_model_patcher(coreml_model)
|
||||||
|
return (model_patcher,)
|
||||||
|
|
||||||
|
|
||||||
|
class CoreMLConverter(COREML_NODE):
|
||||||
|
"""Converts a LCM model to Core ML."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||||
|
"model_version": (
|
||||||
|
[
|
||||||
|
ModelVersion.SD15.name,
|
||||||
|
ModelVersion.SDXL.name,
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"height": ("INT", {"default": 512, "min": 256, "max": 2048, "step": 8}),
|
||||||
|
"width": ("INT", {"default": 512, "min": 256, "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,
|
||||||
|
ComputeUnit.CPU_AND_GPU.name,
|
||||||
|
ComputeUnit.ALL.name,
|
||||||
|
ComputeUnit.CPU_ONLY.name,
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"controlnet_support": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"lora_params": ("LORA_PARAMS",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("COREML_UNET",)
|
||||||
|
RETURN_NAMES = ("coreml_model",)
|
||||||
|
FUNCTION = "convert"
|
||||||
|
|
||||||
|
def convert(
|
||||||
|
self,
|
||||||
|
ckpt_name,
|
||||||
|
model_version,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
batch_size,
|
||||||
|
attention_implementation,
|
||||||
|
compute_unit,
|
||||||
|
controlnet_support,
|
||||||
|
lora_params=None,
|
||||||
|
):
|
||||||
|
"""Converts a LCM model to Core ML.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
height (int): Height of the target image.
|
||||||
|
width (int): Width of the target image.
|
||||||
|
batch_size (int): Batch size.
|
||||||
|
compute_unit (str): Compute unit to use when loading the model.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
coreml_model: The converted Core ML model.
|
||||||
|
|
||||||
|
The converted model is also saved to "models/unet" directory and
|
||||||
|
can be loaded with the "LCMCoreMLLoaderUNet" node.
|
||||||
|
"""
|
||||||
|
model_version = ModelVersion[model_version]
|
||||||
|
|
||||||
|
lora_params = lora_params or {}
|
||||||
|
lora_params = [(k, v[0]) for k, v in lora_params.items()]
|
||||||
|
lora_params = sorted(lora_params, key=lambda lora: lora[0])
|
||||||
|
lora_weights = [(self.lora_path(lora[0]), lora[1]) for lora in lora_params]
|
||||||
|
|
||||||
|
h = height
|
||||||
|
w = width
|
||||||
|
sample_size = (h // 8, w // 8)
|
||||||
|
batch_size = batch_size
|
||||||
|
cn_support_str = "_cn" if controlnet_support else ""
|
||||||
|
lora_str = (
|
||||||
|
"_" + "_".join(lora_param[0].split(".")[0] for lora_param in lora_params)
|
||||||
|
if lora_params
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
|
||||||
|
attn_str = (
|
||||||
|
"_"
|
||||||
|
+ {"SPLIT_EINSUM": "se", "SPLIT_EINSUM_V2": "se2", "ORIGINAL": "orig"}[
|
||||||
|
attention_implementation
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
out_name = f"{ckpt_name.split('.')[0]}{lora_str}_{batch_size}x{w}x{h}{cn_support_str}{attn_str}"
|
||||||
|
out_name = out_name.replace(" ", "_")
|
||||||
|
|
||||||
|
logger.info(f"Converting {ckpt_name} to {out_name}")
|
||||||
|
logger.info(f"Batch size: {batch_size}")
|
||||||
|
logger.info(f"Width: {w}, Height: {h}")
|
||||||
|
logger.info(f"ControlNet support: {controlnet_support}")
|
||||||
|
logger.info(f"Attention implementation: {attention_implementation}")
|
||||||
|
|
||||||
|
if lora_params:
|
||||||
|
logger.info(f"LoRAs used:")
|
||||||
|
for lora_param in lora_params:
|
||||||
|
logger.info(f" {lora_param[0]} - strength: {lora_param[1]}")
|
||||||
|
|
||||||
|
unet_out_path = converter.get_out_path("unet", f"{out_name}")
|
||||||
|
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||||
|
|
||||||
|
config_filename = ckpt_name.split(".")[0] + ".yaml"
|
||||||
|
config_path = folder_paths.get_full_path("configs", config_filename)
|
||||||
|
if config_path:
|
||||||
|
logger.info(f"Using config file {config_path}")
|
||||||
|
|
||||||
|
converter.convert(
|
||||||
|
ckpt_path=ckpt_path,
|
||||||
|
model_version=model_version,
|
||||||
|
unet_out_path=unet_out_path,
|
||||||
|
sample_size=sample_size,
|
||||||
|
batch_size=batch_size,
|
||||||
|
controlnet_support=controlnet_support,
|
||||||
|
lora_weights=lora_weights,
|
||||||
|
attn_impl=attention_implementation,
|
||||||
|
config_path=config_path,
|
||||||
|
)
|
||||||
|
unet_target_path = converter.compile_model(
|
||||||
|
out_path=unet_out_path, out_name=out_name, submodule_name="unet"
|
||||||
|
)
|
||||||
|
|
||||||
|
return (CoreMLModel(unet_target_path, compute_unit, "compiled"),)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def lora_path(lora_name):
|
||||||
|
return folder_paths.get_full_path("loras", lora_name)
|
||||||
|
|
||||||
|
|
||||||
|
class COREML_LOAD_LORA(COREML_NODE, LoraLoader):
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
required = LoraLoader.INPUT_TYPES()["required"].copy()
|
||||||
|
required.pop("model")
|
||||||
|
return {
|
||||||
|
"required": required,
|
||||||
|
"optional": {"lora_params": ("LORA_PARAMS",)},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("CLIP", "LORA_PARAMS")
|
||||||
|
RETURN_NAMES = ("CLIP", "lora_params")
|
||||||
|
|
||||||
|
def load_lora(
|
||||||
|
self, clip, lora_name, strength_model, strength_clip, lora_params=None
|
||||||
|
):
|
||||||
|
_, lora_clip = super().load_lora(
|
||||||
|
None, clip, lora_name, strength_model, strength_clip
|
||||||
|
)
|
||||||
|
|
||||||
|
lora_params = lora_params or {}
|
||||||
|
lora_params[lora_name] = (strength_model, strength_clip)
|
||||||
|
|
||||||
|
return lora_clip, lora_params
|
||||||
@@ -1,74 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torchvision.transforms.functional import resize
|
|
||||||
|
|
||||||
from comfy.model_management import get_torch_device
|
|
||||||
from comfy.model_patcher import ModelPatcher
|
|
||||||
from coreml_suite.logger import logger
|
|
||||||
from nodes import KSampler
|
|
||||||
|
|
||||||
from coreml_suite.models import CoreMLModelWrapper, get_model_config
|
|
||||||
|
|
||||||
|
|
||||||
def reshape_latent_image(latent_image, target_shape):
|
|
||||||
if latent_image is None:
|
|
||||||
logger.warning("No latent image provided, using zeros.")
|
|
||||||
return {"samples": torch.zeros(target_shape)}
|
|
||||||
|
|
||||||
if latent_image["samples"].shape == target_shape:
|
|
||||||
return latent_image
|
|
||||||
|
|
||||||
logger.warning(
|
|
||||||
"Latent image shape does not match model input shape,"
|
|
||||||
" resizing to match models expected input shape."
|
|
||||||
)
|
|
||||||
resized = resize(latent_image["samples"], target_shape[-2:])
|
|
||||||
return {"samples": resized}
|
|
||||||
|
|
||||||
|
|
||||||
class CoreMLSampler(KSampler):
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(s):
|
|
||||||
old_required = KSampler.INPUT_TYPES()["required"].copy()
|
|
||||||
old_required.pop("model")
|
|
||||||
old_required.pop("latent_image")
|
|
||||||
new_required = {"coreml_model": ("COREML_UNET",)}
|
|
||||||
return {
|
|
||||||
"required": new_required | old_required,
|
|
||||||
"optional": {"latent_image": ("LATENT",)},
|
|
||||||
}
|
|
||||||
|
|
||||||
CATEGORY = "Core ML Suite"
|
|
||||||
|
|
||||||
def sample(
|
|
||||||
self,
|
|
||||||
coreml_model,
|
|
||||||
seed,
|
|
||||||
steps,
|
|
||||||
cfg,
|
|
||||||
sampler_name,
|
|
||||||
scheduler,
|
|
||||||
positive,
|
|
||||||
negative,
|
|
||||||
latent_image=None,
|
|
||||||
denoise=1.0,
|
|
||||||
):
|
|
||||||
sample_shape = coreml_model.expected_inputs["sample"]["shape"]
|
|
||||||
latent_image = reshape_latent_image(latent_image, sample_shape)
|
|
||||||
latent_image["samples"] = latent_image["samples"][0:1]
|
|
||||||
|
|
||||||
model_config = get_model_config()
|
|
||||||
wrapped_model = CoreMLModelWrapper(model_config, coreml_model)
|
|
||||||
model = ModelPatcher(wrapped_model, get_torch_device(), None)
|
|
||||||
|
|
||||||
return super().sample(
|
|
||||||
model,
|
|
||||||
seed,
|
|
||||||
steps,
|
|
||||||
cfg,
|
|
||||||
sampler_name,
|
|
||||||
scheduler,
|
|
||||||
positive,
|
|
||||||
negative,
|
|
||||||
latent_image,
|
|
||||||
denoise,
|
|
||||||
)
|
|
||||||
@@ -1,62 +0,0 @@
|
|||||||
from itertools import chain
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from coreml_suite.logger import logger
|
|
||||||
|
|
||||||
|
|
||||||
def expand_inputs(inputs):
|
|
||||||
expanded = inputs.copy()
|
|
||||||
for k, v in inputs.items():
|
|
||||||
if isinstance(v, np.ndarray):
|
|
||||||
expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v
|
|
||||||
elif isinstance(v, torch.Tensor):
|
|
||||||
expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v
|
|
||||||
elif isinstance(v, list):
|
|
||||||
expanded[k] = v * 2 if len(v) == 1 else v
|
|
||||||
elif isinstance(v, dict):
|
|
||||||
expand_inputs(v)
|
|
||||||
return expanded
|
|
||||||
|
|
||||||
|
|
||||||
def extract_residual_kwargs(model, control):
|
|
||||||
if "additional_residual_0" not in model.expected_inputs.keys():
|
|
||||||
return {}
|
|
||||||
if control is None:
|
|
||||||
return no_control(model)
|
|
||||||
|
|
||||||
residual_kwargs = {
|
|
||||||
"additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16)
|
|
||||||
for i, r in enumerate(chain(control["output"], control["middle"]))
|
|
||||||
}
|
|
||||||
return residual_kwargs
|
|
||||||
|
|
||||||
|
|
||||||
def no_control(model):
|
|
||||||
# Dirty hack to get the expected input shape when doing partial ControlNet
|
|
||||||
# 0.18215 is the latent scale factor (IDK, it kinda works)
|
|
||||||
# TODO: Find a better way to do this or tweak the values
|
|
||||||
|
|
||||||
logger.warning(
|
|
||||||
"No ControlNet input, despite the model supports it. "
|
|
||||||
"Using random noise as ControlNet residuals. "
|
|
||||||
"For better results, please use a ControlNet or a model "
|
|
||||||
"that does not support ControlNet."
|
|
||||||
)
|
|
||||||
residuals_names = [
|
|
||||||
name
|
|
||||||
for name in model.expected_inputs.keys()
|
|
||||||
if name.startswith("additional_residual")
|
|
||||||
]
|
|
||||||
residual_kwargs = {
|
|
||||||
"additional_residual_{}".format(i): 0.18215
|
|
||||||
* torch.randn(
|
|
||||||
*model.expected_inputs["additional_residual_{}".format(i)]["shape"]
|
|
||||||
)
|
|
||||||
.cpu()
|
|
||||||
.numpy()
|
|
||||||
.astype(dtype=np.float16)
|
|
||||||
for i in range(len(residuals_names))
|
|
||||||
}
|
|
||||||
return residual_kwargs
|
|
||||||
@@ -1,2 +1,6 @@
|
|||||||
git+https://github.com/apple/ml-stable-diffusion.git
|
git+https://github.com/apple/ml-stable-diffusion.git
|
||||||
coremltools
|
coremltools>=7.1
|
||||||
|
overrides
|
||||||
|
diffusers>=0.22
|
||||||
|
peft>=0.6.2
|
||||||
|
omegaconf>=2.3
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import requests
|
||||||
|
from PIL import Image
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from folder_paths import get_save_image_path, get_output_directory
|
||||||
|
|
||||||
|
IMAGE_PREFIX = "E2E-1.5-CoreML"
|
||||||
|
|
||||||
|
|
||||||
|
class OutputImageRepository:
|
||||||
|
def __init__(self, name_prefix):
|
||||||
|
self.name_prefix = name_prefix
|
||||||
|
|
||||||
|
def list_images(self):
|
||||||
|
full_output_folder, _, _, _, _ = get_save_image_path(
|
||||||
|
self.name_prefix, get_output_directory(), 512, 512
|
||||||
|
)
|
||||||
|
return full_output_folder, os.listdir(full_output_folder)
|
||||||
|
|
||||||
|
def delete_images(self):
|
||||||
|
full_output_folder, images = self.list_images()
|
||||||
|
for image in images:
|
||||||
|
os.remove(os.path.join(full_output_folder, image))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def output_image_repository():
|
||||||
|
repo = OutputImageRepository(IMAGE_PREFIX)
|
||||||
|
yield repo
|
||||||
|
repo.delete_images()
|
||||||
|
|
||||||
|
|
||||||
|
def test_basic_conversion_1_5(output_image_repository):
|
||||||
|
with open("integration/workflows/e2e-1.5-basic-conversion.json") as f:
|
||||||
|
prompt = json.load(f)
|
||||||
|
queue_prompt(prompt)
|
||||||
|
|
||||||
|
full_output_folder, images = output_image_repository.list_images()
|
||||||
|
assert len(images) == 2
|
||||||
|
assert all(image.startswith(IMAGE_PREFIX) for image in images)
|
||||||
|
assert all(image.endswith(".png") for image in images)
|
||||||
|
assert all(
|
||||||
|
os.path.isfile(os.path.join(full_output_folder, image)) for image in images
|
||||||
|
)
|
||||||
|
|
||||||
|
image1 = Image.open(os.path.join(full_output_folder, images[0]))
|
||||||
|
image2 = Image.open(os.path.join(full_output_folder, images[1]))
|
||||||
|
assert psnr(np.array(image1), np.array(image2)) > 30
|
||||||
|
assert psnr(np.array(image2), np.array(image1)) > 30
|
||||||
|
|
||||||
|
|
||||||
|
def psnr(img1, img2):
|
||||||
|
mse = np.mean((img1 - img2) ** 2)
|
||||||
|
if mse == 0:
|
||||||
|
return 100
|
||||||
|
PIXEL_MAX = 255.0
|
||||||
|
return 20 * np.log10(PIXEL_MAX / np.sqrt(mse))
|
||||||
|
|
||||||
|
|
||||||
|
def queue_prompt(prompt: dict):
|
||||||
|
p = {"prompt": prompt}
|
||||||
|
data = json.dumps(p).encode("utf-8")
|
||||||
|
req = requests.post("http://localhost:8188/prompt", data=data)
|
||||||
|
assert req.status_code == 200
|
||||||
|
while True:
|
||||||
|
req = requests.get("http://localhost:8188/prompt")
|
||||||
|
if req.json()["exec_info"]["queue_remaining"] == 0:
|
||||||
|
break
|
||||||
@@ -0,0 +1,182 @@
|
|||||||
|
{
|
||||||
|
"3": {
|
||||||
|
"inputs": {
|
||||||
|
"seed": 0,
|
||||||
|
"steps": 20,
|
||||||
|
"cfg": 8,
|
||||||
|
"sampler_name": "dpmpp_2m",
|
||||||
|
"scheduler": "karras",
|
||||||
|
"denoise": 1,
|
||||||
|
"model": [
|
||||||
|
"4",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"positive": [
|
||||||
|
"6",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"negative": [
|
||||||
|
"7",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"latent_image": [
|
||||||
|
"5",
|
||||||
|
0
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "KSampler",
|
||||||
|
"_meta": {
|
||||||
|
"title": "KSampler"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"4": {
|
||||||
|
"inputs": {
|
||||||
|
"ckpt_name": "dreamshaper_8.safetensors"
|
||||||
|
},
|
||||||
|
"class_type": "CheckpointLoaderSimple",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Load Checkpoint"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"5": {
|
||||||
|
"inputs": {
|
||||||
|
"width": 512,
|
||||||
|
"height": 512,
|
||||||
|
"batch_size": 1
|
||||||
|
},
|
||||||
|
"class_type": "EmptyLatentImage",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Empty Latent Image"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"6": {
|
||||||
|
"inputs": {
|
||||||
|
"text": "beautiful scenery nature glass bottle landscape, purple galaxy bottle",
|
||||||
|
"clip": [
|
||||||
|
"4",
|
||||||
|
1
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"_meta": {
|
||||||
|
"title": "CLIP Text Encode (Prompt)"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"7": {
|
||||||
|
"inputs": {
|
||||||
|
"text": "text, watermark",
|
||||||
|
"clip": [
|
||||||
|
"4",
|
||||||
|
1
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "CLIPTextEncode",
|
||||||
|
"_meta": {
|
||||||
|
"title": "CLIP Text Encode (Prompt)"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"8": {
|
||||||
|
"inputs": {
|
||||||
|
"samples": [
|
||||||
|
"3",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"vae": [
|
||||||
|
"4",
|
||||||
|
2
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "VAEDecode",
|
||||||
|
"_meta": {
|
||||||
|
"title": "VAE Decode"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"9": {
|
||||||
|
"inputs": {
|
||||||
|
"filename_prefix": "E2E-1.5-MPS",
|
||||||
|
"images": [
|
||||||
|
"8",
|
||||||
|
0
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "SaveImage",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Save Image"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"10": {
|
||||||
|
"inputs": {
|
||||||
|
"ckpt_name": "dreamshaper_8.safetensors",
|
||||||
|
"model_version": "SD15",
|
||||||
|
"height": 512,
|
||||||
|
"width": 512,
|
||||||
|
"batch_size": 1,
|
||||||
|
"attention_implementation": "SPLIT_EINSUM",
|
||||||
|
"compute_unit": "CPU_AND_NE",
|
||||||
|
"controlnet_support": false
|
||||||
|
},
|
||||||
|
"class_type": "Core ML Converter",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Convert Checkpoint to Core ML"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"11": {
|
||||||
|
"inputs": {
|
||||||
|
"seed": 0,
|
||||||
|
"steps": 20,
|
||||||
|
"cfg": 8,
|
||||||
|
"sampler_name": "dpmpp_2m",
|
||||||
|
"scheduler": "karras",
|
||||||
|
"denoise": 1,
|
||||||
|
"coreml_model": [
|
||||||
|
"10",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"positive": [
|
||||||
|
"6",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"negative": [
|
||||||
|
"7",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"latent_image": [
|
||||||
|
"5",
|
||||||
|
0
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "CoreMLSampler",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Core ML Sampler"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"13": {
|
||||||
|
"inputs": {
|
||||||
|
"samples": [
|
||||||
|
"11",
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"vae": [
|
||||||
|
"4",
|
||||||
|
2
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "VAEDecode",
|
||||||
|
"_meta": {
|
||||||
|
"title": "VAE Decode"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"14": {
|
||||||
|
"inputs": {
|
||||||
|
"filename_prefix": "E2E-1.5-CoreML",
|
||||||
|
"images": [
|
||||||
|
"13",
|
||||||
|
0
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"class_type": "SaveImage",
|
||||||
|
"_meta": {
|
||||||
|
"title": "Save Image"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,19 +0,0 @@
|
|||||||
import pytest
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from coreml_suite.samplers import reshape_latent_image
|
|
||||||
|
|
||||||
|
|
||||||
def test_fix_latents_no_latent_image():
|
|
||||||
reshaped = reshape_latent_image(None, (2, 4, 64, 64))
|
|
||||||
assert reshaped["samples"].shape == (2, 4, 64, 64)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"latent_shape", [(2, 4, 64, 64), (2, 4, 128, 128), (2, 4, 32, 32), (2, 4, 128, 64)]
|
|
||||||
)
|
|
||||||
def test_reshape_latents(latent_shape):
|
|
||||||
latent_image = {"samples": torch.zeros(latent_shape)}
|
|
||||||
reshaped = reshape_latent_image(latent_image, (2, 4, 64, 64))
|
|
||||||
assert reshaped["samples"].shape == (2, 4, 64, 64)
|
|
||||||
@@ -0,0 +1,125 @@
|
|||||||
|
import pytest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from comfy.model_management import get_torch_device
|
||||||
|
from coreml_suite.latents import chunk_batch, merge_chunks
|
||||||
|
from coreml_suite.controlnet import chunk_control
|
||||||
|
from coreml_suite.models import (
|
||||||
|
CoreMLInputs,
|
||||||
|
)
|
||||||
|
from coreml_suite.config import get_model_config
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def expected_inputs():
|
||||||
|
expected = {
|
||||||
|
"sample": {"shape": (2, 4, 64, 64)},
|
||||||
|
"timestep": {"shape": (2,)},
|
||||||
|
"timestep_cond": {"shape": (2, 256)},
|
||||||
|
"encoder_hidden_states": {"shape": (2, 768, 1, 77)},
|
||||||
|
"additional_residual_0": {"shape": (2, 320, 64, 64)},
|
||||||
|
"additional_residual_1": {"shape": (2, 640, 32, 32)},
|
||||||
|
}
|
||||||
|
return expected
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def model_config():
|
||||||
|
return get_model_config()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||||
|
def test_batch_chunking(batch_size):
|
||||||
|
latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||||
|
target_shape = (4, 4, 64, 64)
|
||||||
|
|
||||||
|
chunked = chunk_batch(latent_image, target_shape)
|
||||||
|
|
||||||
|
for chunk in chunked:
|
||||||
|
assert chunk.shape == target_shape
|
||||||
|
|
||||||
|
if batch_size % target_shape[0] != 0:
|
||||||
|
assert chunked[-1][batch_size % target_shape[0] :].sum() == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9])
|
||||||
|
def test_merge_chunks(batch_size):
|
||||||
|
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
|
||||||
|
target_shape = (4, 4, 64, 64)
|
||||||
|
chunked = chunk_batch(input_tensor, target_shape)
|
||||||
|
|
||||||
|
merged = merge_chunks(chunked, input_tensor.shape)
|
||||||
|
|
||||||
|
assert merged.shape == input_tensor.shape
|
||||||
|
assert torch.equal(input_tensor, merged)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def inputs():
|
||||||
|
x = torch.randn(1, 4, 64, 64).to(get_torch_device())
|
||||||
|
t = torch.randn([1]).to(get_torch_device())
|
||||||
|
c_crossattn = torch.randn(1, 77, 768).to(get_torch_device())
|
||||||
|
control = {
|
||||||
|
"output": [
|
||||||
|
torch.randn(1, 320, 64, 64).to(get_torch_device()),
|
||||||
|
torch.randn(1, 640, 32, 32).to(get_torch_device()),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
timestep_cond = torch.randn(1, 256).to(get_torch_device())
|
||||||
|
|
||||||
|
return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"b, target_size, num_chunks",
|
||||||
|
[
|
||||||
|
(1, 2, 1),
|
||||||
|
(1, 1, 1),
|
||||||
|
(2, 2, 1),
|
||||||
|
(3, 2, 2),
|
||||||
|
(4, 2, 2),
|
||||||
|
(5, 3, 2),
|
||||||
|
(9, 4, 3),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_chunking_controlnet(b, target_size, num_chunks):
|
||||||
|
cn = {
|
||||||
|
"output": [
|
||||||
|
torch.randn(b, 320, 64, 64).to(get_torch_device()),
|
||||||
|
torch.randn(b, 640, 32, 32).to(get_torch_device()),
|
||||||
|
],
|
||||||
|
"middle": [
|
||||||
|
torch.randn(b, 1280, 8, 8).to(get_torch_device()),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
chunked = chunk_control(cn, target_size)
|
||||||
|
|
||||||
|
assert len(chunked) == num_chunks
|
||||||
|
for chunk in chunked:
|
||||||
|
assert chunk["output"][0].shape == (target_size, 320, 64, 64)
|
||||||
|
assert chunk["output"][1].shape == (target_size, 640, 32, 32)
|
||||||
|
assert chunk["middle"][0].shape == (target_size, 1280, 8, 8)
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunking_no_control():
|
||||||
|
cn = None
|
||||||
|
target_size = 2
|
||||||
|
|
||||||
|
chunked = chunk_control(cn, target_size)
|
||||||
|
|
||||||
|
assert chunked == [None, None]
|
||||||
|
|
||||||
|
|
||||||
|
def test_chunking_inputs(expected_inputs, inputs):
|
||||||
|
chunked = inputs.chunks(expected_inputs)
|
||||||
|
|
||||||
|
assert len(chunked) == 1
|
||||||
|
|
||||||
|
assert chunked[0].x.shape == (2, 4, 64, 64)
|
||||||
|
assert chunked[0].t.shape == (2,)
|
||||||
|
assert chunked[0].context.shape == (2, 77, 768)
|
||||||
|
assert chunked[0].control["output"][0].shape == (2, 320, 64, 64)
|
||||||
|
assert chunked[0].control["output"][1].shape == (2, 640, 32, 32)
|
||||||
|
assert chunked[0].ts_cond.shape == (2, 256)
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
from coreml_suite.controlnet import no_control
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_control():
|
||||||
|
expected_inputs = {
|
||||||
|
"additional_residual_0": {"shape": (2, 2, 2)},
|
||||||
|
"additional_residual_1": {"shape": (2, 4, 4)},
|
||||||
|
"additional_residual_2": {"shape": (2, 8, 8)},
|
||||||
|
}
|
||||||
|
|
||||||
|
residual_kwargs = no_control(expected_inputs)
|
||||||
|
|
||||||
|
assert len(residual_kwargs) == 3
|
||||||
|
assert residual_kwargs["additional_residual_0"].shape == (2, 2, 2)
|
||||||
|
assert residual_kwargs["additional_residual_1"].shape == (2, 4, 4)
|
||||||
|
assert residual_kwargs["additional_residual_2"].shape == (2, 8, 8)
|
||||||