Move UNet logic to UNet Loader
This commit is contained in:
+21
-18
@@ -1,4 +1,5 @@
|
||||
from coremltools import ComputeUnit
|
||||
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||
|
||||
import folder_paths
|
||||
from comfy import supported_models_base, model_management
|
||||
@@ -46,19 +47,10 @@ class CoreMLLoader:
|
||||
sources = "compiled" if coreml_name.endswith(
|
||||
".mlmodelc") else "packages"
|
||||
|
||||
# TODO: This is a dummy model config, but it should be enough to
|
||||
# get the model to load - implement a proper model config
|
||||
model_config = supported_models_base.BASE({})
|
||||
model_config.latent_format = SD15()
|
||||
model_config.unet_config = {
|
||||
"disable_unet_model_creation": True,
|
||||
"num_res_blocks": 2,
|
||||
"attention_resolutions": [1, 2, 4],
|
||||
"channel_mult": [1, 2, 4, 4],
|
||||
"transformer_depth": [1, 1, 1, 0],
|
||||
}
|
||||
return (CoreMLModelWrapper(model_config, coreml_path, compute_unit,
|
||||
sources),)
|
||||
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):
|
||||
@@ -83,12 +75,23 @@ class CoreMLLoaderUNet(CoreMLLoader):
|
||||
PACKAGE_DIRNAME = "unet"
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
|
||||
def load(self, coreml_name, compute_unit):
|
||||
coreml_model = super().load(coreml_name, compute_unit)[0]
|
||||
def _load(self, coreml_path, compute_unit, sources):
|
||||
# TODO: This is a dummy model config, but it should be enough to
|
||||
# get the model to load - implement a proper model config
|
||||
model_config = supported_models_base.BASE({})
|
||||
model_config.latent_format = SD15()
|
||||
model_config.unet_config = {
|
||||
"disable_unet_model_creation": True,
|
||||
"num_res_blocks": 2,
|
||||
"attention_resolutions": [1, 2, 4],
|
||||
"channel_mult": [1, 2, 4, 4],
|
||||
"transformer_depth": [1, 1, 1, 0],
|
||||
}
|
||||
coreml_model = CoreMLModelWrapper(model_config, coreml_path,
|
||||
compute_unit, sources)
|
||||
|
||||
return (
|
||||
ModelPatcher(coreml_model, model_management.get_torch_device(),
|
||||
None),)
|
||||
return (ModelPatcher(coreml_model, model_management.get_torch_device(),
|
||||
None),)
|
||||
|
||||
|
||||
class CoreMLLoaderVAE(CoreMLLoader):
|
||||
|
||||
Reference in New Issue
Block a user