From 3f666ac0eaa4e992b1055522905e18ece20cc43c Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Sun, 19 Nov 2023 23:40:44 +0100 Subject: [PATCH] Add Advanced Sampler node --- __init__.py | 3 ++ coreml_suite/models.py | 58 +++++++++++++++++++++++++-- coreml_suite/nodes.py | 89 ++++++++++++++++++++++++++++++++---------- 3 files changed, 125 insertions(+), 25 deletions(-) diff --git a/__init__.py b/__init__.py index f102d00..440fc76 100644 --- a/__init__.py +++ b/__init__.py @@ -6,6 +6,7 @@ sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import ( CoreMLLoaderUNet, CoreMLSampler, + CoreMLSamplerAdvanced, CoreMLModelAdapter, COREML_CONVERT, COREML_LOAD_LORA, @@ -17,6 +18,7 @@ from coreml_suite.lcm import ( NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, + "CoreMLSamplerAdvanced": CoreMLSamplerAdvanced, "CoreMLModelAdapter": CoreMLModelAdapter, "Core ML LoRA Loader": COREML_LOAD_LORA, "Core ML Converter": COREML_CONVERT, @@ -25,6 +27,7 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "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", diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 79e1e48..f03518f 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -1,8 +1,13 @@ import numpy as np import torch +from comfy import model_base +from comfy.model_management import get_torch_device +from comfy.model_patcher import ModelPatcher +from coreml_suite.config import get_model_config from coreml_suite.controlnet import extract_residual_kwargs, chunk_control from coreml_suite.latents import chunk_batch, merge_chunks +from coreml_suite.logger import logger class CoreMLModelWrapper: @@ -177,8 +182,6 @@ def add_sdxl_model_options(model_patcher, positive, negative): pos_dict.get("width", 768), pos_dict.get("crop_h", 0), pos_dict.get("crop_w", 0), - pos_dict.get("target_height", 768), - pos_dict.get("target_width", 768), ] neg_time_ids = [ @@ -186,10 +189,32 @@ def add_sdxl_model_options(model_patcher, positive, negative): neg_dict.get("width", 768), neg_dict.get("crop_h", 0), neg_dict.get("crop_w", 0), - neg_dict.get("target_height", 768), - neg_dict.get("target_width", 768), ] + if model_patcher.model.diffusion_model.expected_inputs["time_ids"]["shape"][1] == 6: + base_pos_time_ids = [ + pos_dict.get("target_height", 768), + pos_dict.get("target_width", 768), + ] + pos_time_ids += base_pos_time_ids + + base_neg_time_ids = [ + neg_dict.get("target_height", 768), + neg_dict.get("target_width", 768), + ] + neg_time_ids += base_neg_time_ids + + else: + refiner_pos_time_ids = [ + pos_dict.get("aesthetic_score", 6), + ] + pos_time_ids += refiner_pos_time_ids + + refiner_neg_time_ids = [ + neg_dict.get("aesthetic_score", 2.5), + ] + neg_time_ids += refiner_neg_time_ids + time_ids = torch.tensor([pos_time_ids, neg_time_ids]) text_embeds = torch.cat((pos_dict["pooled_output"], neg_dict["pooled_output"])) @@ -200,3 +225,28 @@ def add_sdxl_model_options(model_patcher, positive, negative): 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): + model_config = get_model_config() + wrapped_model = CoreMLModelWrapper(coreml_model) + + if is_sdxl(coreml_model): + model = model_base.SDXL(model_config, device=get_torch_device()) + else: + model = model_base.BaseModel(model_config, device=get_torch_device()) + + model.diffusion_model = wrapped_model + model_patcher = ModelPatcher(model, get_torch_device(), None) + return model_patcher diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index a60794a..fbdab04 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -13,9 +13,15 @@ from coreml_suite import COREML_NODE from coreml_suite import converter 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 +from nodes import KSampler, LoraLoader, KSamplerAdvanced -from coreml_suite.models import CoreMLModelWrapper, add_sdxl_model_options, is_sdxl +from coreml_suite.models import ( + CoreMLModelWrapper, + add_sdxl_model_options, + is_sdxl, + get_model_patcher, + get_latent_image, +) from coreml_suite.config import get_model_config @@ -45,8 +51,8 @@ class CoreMLSampler(COREML_NODE, KSampler): latent_image=None, denoise=1.0, ): - model_patcher = self.get_model_patcher(coreml_model) - latent_image = self.get_latent_image(coreml_model, latent_image) + model_patcher = get_model_patcher(coreml_model) + latent_image = get_latent_image(coreml_model, latent_image) if is_lcm(coreml_model): negative = [[None, {}]] @@ -74,28 +80,69 @@ class CoreMLSampler(COREML_NODE, KSampler): denoise, ) - def get_latent_image(self, 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 +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 get_model_patcher(self, coreml_model): - model_config = get_model_config() - wrapped_model = CoreMLModelWrapper(coreml_model) + 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 = model_base.SDXL(model_config, device=get_torch_device()) - else: - model = model_base.BaseModel(model_config, device=get_torch_device()) + model_patcher = add_sdxl_model_options(model_patcher, positive, negative) - model.diffusion_model = wrapped_model - model_patcher = ModelPatcher(model, get_torch_device(), None) - return model_patcher + 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):