Rearrange stuff

This commit is contained in:
aszc-dev
2024-06-28 15:52:53 +02:00
parent b326b3d3b9
commit c3038501eb
8 changed files with 79 additions and 110 deletions
+1 -1
View File
@@ -3,7 +3,7 @@ import sys
sys.path.append(os.path.dirname(__file__))
from coreml_suite import CoreMLLoaderUNet, CoreMLSampler
from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler
NODE_CLASS_MAPPINGS = {
"CoreMLUNetLoader": CoreMLLoaderUNet,
-4
View File
@@ -1,4 +0,0 @@
from coreml_suite.loaders import CoreMLLoaderUNet
from coreml_suite.samplers import CoreMLSampler
__all__ = ["CoreMLLoaderUNet", "CoreMLSampler"]
+20
View File
@@ -0,0 +1,20 @@
import torch
from torchvision.transforms.functional import resize
from coreml_suite.logger import logger
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}
+1 -1
View File
@@ -6,7 +6,7 @@ from comfy import supported_models_base
from comfy.latent_formats import SD15
from comfy.model_base import BaseModel
from coreml_suite.utils import expand_inputs, extract_residual_kwargs
from coreml_suite.controlnet import expand_inputs, extract_residual_kwargs
def get_model_config():
@@ -1,11 +1,65 @@
import os.path
import os
from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
import folder_paths
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite.logger import logger
from coreml_suite.latents import reshape_latent_image
from nodes import KSampler
from coreml_suite.models import CoreMLModelWrapper, get_model_config
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,
)
class CoreMLLoader:
@@ -51,34 +105,7 @@ class CoreMLLoader:
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
-74
View File
@@ -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 -1
View File
@@ -2,7 +2,7 @@ import pytest
import torch
from coreml_suite.samplers import reshape_latent_image
from coreml_suite.latents import reshape_latent_image
def test_fix_latents_no_latent_image():