Extract lcm sampler from lcm sampling node
This commit is contained in:
+2
-2
@@ -5,7 +5,7 @@ sys.path.append(os.path.dirname(__file__))
|
||||
|
||||
from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter
|
||||
from coreml_suite.lcm import (
|
||||
CoreMLSamplerLCM,
|
||||
COREML_SAMPLER_LCM,
|
||||
CoreMLConverterLCM,
|
||||
)
|
||||
|
||||
@@ -13,7 +13,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"CoreMLUNetLoader": CoreMLLoaderUNet,
|
||||
"CoreMLSampler": CoreMLSampler,
|
||||
"CoreMLModelAdapter": CoreMLModelAdapter,
|
||||
"Core ML LCM Sampler": CoreMLSamplerLCM,
|
||||
"Core ML LCM Sampler": COREML_SAMPLER_LCM,
|
||||
"Core ML LCM Converter": CoreMLConverterLCM,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from .lcm_sampler import CoreMLSamplerLCM
|
||||
from .nodes import CoreMLConverterLCM
|
||||
from .nodes import CoreMLConverterLCM, COREML_SAMPLER_LCM
|
||||
|
||||
__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"]
|
||||
__all__ = ["COREML_SAMPLER_LCM", "CoreMLConverterLCM"]
|
||||
|
||||
@@ -1,9 +1,21 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
from coremltools import ComputeUnit
|
||||
from diffusers import LCMScheduler
|
||||
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
|
||||
|
||||
import comfy.utils
|
||||
import latent_preview
|
||||
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.lcm import lcm_converter
|
||||
from coreml_suite.lcm.sampler import CoreMLSamplerLCM
|
||||
from coreml_suite.logger import logger
|
||||
from coreml_suite.models import CoreMLModelWrapper
|
||||
from coreml_suite.nodes import CoreMLSampler
|
||||
|
||||
|
||||
class CoreMLConverterLCM:
|
||||
@@ -67,3 +79,74 @@ class CoreMLConverterLCM:
|
||||
target_path = lcm_converter.compile_model(out_path=out_path, out_name=out_name)
|
||||
|
||||
return (CoreMLModel(target_path, compute_unit, "compiled"),)
|
||||
|
||||
|
||||
class COREML_SAMPLER_LCM(CoreMLSampler):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
old_required = CoreMLSampler.INPUT_TYPES()["required"].copy()
|
||||
old_required["steps"][1]["default"] = 4
|
||||
old_required.pop("negative")
|
||||
old_required.pop("sampler_name")
|
||||
old_required.pop("scheduler")
|
||||
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,
|
||||
positive,
|
||||
latent_image=None,
|
||||
denoise=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
scheduler = LCMScheduler.from_pretrained(
|
||||
"SimianLuo/LCM_Dreamshaper_v7", subfolder="scheduler"
|
||||
)
|
||||
|
||||
model_config = get_model_config()
|
||||
wrapped_model = CoreMLModelWrapper(coreml_model)
|
||||
model = model_base.BaseModel(model_config, device=get_torch_device())
|
||||
model.diffusion_model = wrapped_model
|
||||
model_patcher = ModelPatcher(model, get_torch_device(), None)
|
||||
|
||||
if latent_image is None:
|
||||
logger.warning("No latent image provided, using empty tensor.")
|
||||
expected = coreml_model.expected_inputs["sample"]["shape"]
|
||||
latent_image = {"samples": torch.zeros(*expected).to(get_torch_device())}
|
||||
|
||||
x0_output = {}
|
||||
callback = latent_preview.prepare_callback(model_patcher, steps, x0_output)
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
|
||||
sampler = CoreMLSamplerLCM(scheduler)
|
||||
samples = sampler.sample(
|
||||
model_patcher,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
positive,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
)
|
||||
|
||||
out = latent_image.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent_image.copy()
|
||||
out_denoised["samples"] = model.process_latent_out(
|
||||
x0_output["x0"].to(get_torch_device())
|
||||
)
|
||||
else:
|
||||
out_denoised = out
|
||||
return out, out_denoised
|
||||
|
||||
@@ -1,69 +1,31 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import latent_preview
|
||||
from comfy import model_base, samplers
|
||||
from comfy import samplers
|
||||
from comfy.model_management import get_torch_device
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.sample import sample_custom, prepare_noise
|
||||
|
||||
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
|
||||
from coreml_suite.logger import logger
|
||||
from coreml_suite.models import CoreMLModelWrapper
|
||||
from coreml_suite.nodes import CoreMLSampler
|
||||
from coreml_suite.config import get_model_config
|
||||
|
||||
|
||||
class CoreMLSamplerLCM(CoreMLSampler):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
old_required = CoreMLSampler.INPUT_TYPES()["required"].copy()
|
||||
old_required["steps"][1]["default"] = 4
|
||||
old_required.pop("negative")
|
||||
old_required.pop("sampler_name")
|
||||
old_required.pop("scheduler")
|
||||
new_required = {"coreml_model": ("COREML_UNET",)}
|
||||
return {
|
||||
"required": new_required | old_required,
|
||||
"optional": {"latent_image": ("LATENT",)},
|
||||
}
|
||||
|
||||
CATEGORY = "Core ML Suite"
|
||||
|
||||
def __init__(self):
|
||||
self.scheduler = LCMScheduler.from_pretrained(
|
||||
"SimianLuo/LCM_Dreamshaper_v7", subfolder="scheduler"
|
||||
)
|
||||
class CoreMLSamplerLCM:
|
||||
def __init__(self, scheduler):
|
||||
self.scheduler = scheduler
|
||||
|
||||
def sample(
|
||||
self,
|
||||
coreml_model,
|
||||
model_patcher,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
positive,
|
||||
latent_image=None,
|
||||
latent_image,
|
||||
denoise=1.0,
|
||||
**kwargs,
|
||||
callback=None,
|
||||
disable_pbar=False,
|
||||
):
|
||||
model_config = get_model_config()
|
||||
wrapped_model = CoreMLModelWrapper(coreml_model)
|
||||
model = model_base.BaseModel(model_config, device=get_torch_device())
|
||||
model.diffusion_model = wrapped_model
|
||||
model_patcher = ModelPatcher(model, get_torch_device(), None)
|
||||
|
||||
positive[0][1]["control_apply_to_uncond"] = False
|
||||
|
||||
if latent_image is None:
|
||||
logger.warning("No latent image provided, using empty tensor.")
|
||||
expected = coreml_model.expected_inputs["sample"]["shape"]
|
||||
latent_image = {"samples": torch.zeros(*expected).to(get_torch_device())}
|
||||
|
||||
latent = latent_image["samples"].to(get_torch_device())
|
||||
|
||||
x0_output = {}
|
||||
callback = latent_preview.prepare_callback(model_patcher, steps, x0_output)
|
||||
|
||||
batch_size = latent.shape[0]
|
||||
dtype = latent.dtype
|
||||
device = get_torch_device()
|
||||
@@ -79,13 +41,9 @@ class CoreMLSamplerLCM(CoreMLSampler):
|
||||
}
|
||||
model_patcher.model_options |= model_options
|
||||
|
||||
batch_inds = latent_image.get("batch_index")
|
||||
noise = prepare_noise(latent, seed, batch_inds)
|
||||
|
||||
self.prepare_timesteps(denoise, device, steps)
|
||||
# self.scheduler.set_timesteps(steps, 50, device)
|
||||
|
||||
all_sigmas, sigmas = self.get_sigmas(steps, denoise)
|
||||
|
||||
model_patcher.model.model_sampling.set_sigmas(all_sigmas)
|
||||
sigma_to_timestep = {
|
||||
s.item(): t for s, t in zip(sigmas, self.scheduler.timesteps)
|
||||
@@ -95,6 +53,8 @@ class CoreMLSamplerLCM(CoreMLSampler):
|
||||
].expand(1)
|
||||
|
||||
noise_mask = latent_image.get("noise_mask")
|
||||
batch_inds = latent_image.get("batch_index")
|
||||
noise = prepare_noise(latent, seed, batch_inds)
|
||||
|
||||
sampler = samplers.ksampler("ddpm")()
|
||||
samples = sample_custom(
|
||||
@@ -108,26 +68,18 @@ class CoreMLSamplerLCM(CoreMLSampler):
|
||||
latent,
|
||||
noise_mask,
|
||||
callback,
|
||||
disable_pbar,
|
||||
seed,
|
||||
)
|
||||
return samples
|
||||
|
||||
out = latent_image.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent_image.copy()
|
||||
out_denoised["samples"] = model.process_latent_out(
|
||||
x0_output["x0"].to(device)
|
||||
)
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
def get_sigmas(self, steps):
|
||||
def get_sigmas(self, steps, denoise):
|
||||
alphas_cumprod = self.scheduler.alphas_cumprod
|
||||
sigmas = np.asarray(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5)
|
||||
skipping_step = len(sigmas) // steps
|
||||
s = sigmas[::-skipping_step][:steps]
|
||||
if len(s) == steps:
|
||||
s = np.append(s, sigmas[0])
|
||||
s = np.append(s, 0.0).astype(np.float32)
|
||||
sigmas = torch.from_numpy(sigmas).to(get_torch_device())
|
||||
return (sigmas, torch.from_numpy(s.copy()).to(get_torch_device()))
|
||||
|
||||
Reference in New Issue
Block a user