Simplify LCM Sampler

This commit is contained in:
aszc-dev
2023-11-03 01:27:48 +01:00
parent 8a814b7a56
commit 1aa5a19b2a
3 changed files with 19 additions and 29 deletions
+3 -3
View File
@@ -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",
}
+2 -2
View File
@@ -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"]
+14 -24
View File
@@ -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,)