Add Advanced Sampler node

This commit is contained in:
aszc-dev
2024-06-28 15:52:54 +02:00
parent acecd10aee
commit 3f666ac0ea
3 changed files with 125 additions and 25 deletions
+3
View File
@@ -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",
+54 -4
View File
@@ -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
+68 -21
View File
@@ -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):