From 1aa5a19b2aa7c100a8282a9b833590cbedeaf9cb Mon Sep 17 00:00:00 2001 From: aszc-dev Date: Tue, 31 Oct 2023 19:48:46 +0100 Subject: [PATCH] Simplify LCM Sampler --- __init__.py | 6 +++--- coreml_suite/lcm/__init__.py | 4 ++-- coreml_suite/lcm/lcm_sampler.py | 38 ++++++++++++--------------------- 3 files changed, 19 insertions(+), 29 deletions(-) diff --git a/__init__.py b/__init__.py index 68f4c48..8410c43 100644 --- a/__init__.py +++ b/__init__.py @@ -4,19 +4,19 @@ import sys sys.path.append(os.path.dirname(__file__)) from coreml_suite.nodes import CoreMLLoaderUNet, CoreMLSampler, CoreMLModelAdapter -from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM +from coreml_suite.lcm import CoreMLConverterLCM, CoreMLSamplerLCM_Simple NODE_CLASS_MAPPINGS = { "CoreMLUNetLoader": CoreMLLoaderUNet, "CoreMLSampler": CoreMLSampler, "CoreMLModelAdapter": CoreMLModelAdapter, - "CoreMLSamplerLCM": CoreMLSamplerLCM, + "Core ML LCM Sampler (Simple)": CoreMLSamplerLCM_Simple, "CoreMLConverterLCM": CoreMLConverterLCM, } NODE_DISPLAY_NAME_MAPPINGS = { "CoreMLUNetLoader": "Load Core ML UNet", "CoreMLSampler": "Core ML Sampler", "CoreMLModelAdapter": "Core ML Adapter (Experimental)", - "CoreMLSamplerLCM": "Core ML LCM Sampler", + "Core ML LCM Sampler (Simple)": "Core ML LCM Sampler (Simple)", "CoreMLConverterLCM": "Convert LCM to Core ML", } diff --git a/coreml_suite/lcm/__init__.py b/coreml_suite/lcm/__init__.py index 8ae122a..2df93e0 100644 --- a/coreml_suite/lcm/__init__.py +++ b/coreml_suite/lcm/__init__.py @@ -1,4 +1,4 @@ -from .lcm_sampler import CoreMLSamplerLCM +from .lcm_sampler import CoreMLSamplerLCM_Simple from .nodes import CoreMLConverterLCM -__all__ = ["CoreMLSamplerLCM", "CoreMLConverterLCM"] +__all__ = ["CoreMLSamplerLCM_Simple", "CoreMLConverterLCM"] diff --git a/coreml_suite/lcm/lcm_sampler.py b/coreml_suite/lcm/lcm_sampler.py index 1c9c8ca..08d23a7 100644 --- a/coreml_suite/lcm/lcm_sampler.py +++ b/coreml_suite/lcm/lcm_sampler.py @@ -1,5 +1,4 @@ import os -import time import torch @@ -9,7 +8,7 @@ from coreml_suite.lcm.lcm_scheduler import LCMScheduler from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM -class CoreMLSamplerLCM: +class CoreMLSamplerLCM_Simple: def __init__(self): self.scheduler = LCMScheduler.from_pretrained( os.path.join(os.path.dirname(__file__), "scheduler_config.json") @@ -33,10 +32,7 @@ class CoreMLSamplerLCM: "round": 0.01, }, ), - "height": ("INT", {"default": 512, "min": 512, "max": 768}), - "width": ("INT", {"default": 512, "min": 512, "max": 768}), "num_images": ("INT", {"default": 1, "min": 1, "max": 64}), - "use_fp16": ("BOOLEAN", {"default": True}), "positive_prompt": ("STRING", {"multiline": True}), } } @@ -46,17 +42,16 @@ class CoreMLSamplerLCM: CATEGORY = "sampling" def sample( - self, - coreml_model, - seed, - steps, - cfg, - positive_prompt, - height, - width, - num_images, - use_fp16, + self, + coreml_model, + seed, + steps, + cfg, + positive_prompt, + num_images, ): + height = coreml_model.expected_inputs["sample"]["shape"][2] * 8 + width = coreml_model.expected_inputs["sample"]["shape"][3] * 8 model_config = get_model_config() wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model) @@ -68,18 +63,14 @@ class CoreMLSamplerLCM: safety_checker=None, ) - if use_fp16: - self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) - else: - self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float32) + self.pipe.to(torch_device=get_torch_device(), torch_dtype=torch.float16) - coreml_unet = wrapped_model - coreml_unet.config = self.pipe.unet.config + coreml_unet = wrapped_model + coreml_unet.config = self.pipe.unet.config - self.pipe.unet = coreml_unet + self.pipe.unet = coreml_unet torch.manual_seed(seed) - start_time = time.time() result = self.pipe( prompt=positive_prompt, @@ -92,7 +83,6 @@ class CoreMLSamplerLCM: output_type="np", ).images - print("LCM inference time: ", time.time() - start_time, "seconds") images_tensor = torch.from_numpy(result) return (images_tensor,)