Rearrange stuff
This commit is contained in:
+1
-1
@@ -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,
|
||||
|
||||
@@ -1,4 +0,0 @@
|
||||
from coreml_suite.loaders import CoreMLLoaderUNet
|
||||
from coreml_suite.samplers import CoreMLSampler
|
||||
|
||||
__all__ = ["CoreMLLoaderUNet", "CoreMLSampler"]
|
||||
|
||||
@@ -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}
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user