From 40a7228e8f0957a51f6474056d884138b4a3dd83 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrian=20Szczepa=C5=84ski?= Date: Wed, 18 Oct 2023 15:44:10 +0200 Subject: [PATCH] Move CoreMLModelWrapper to model.py --- loaders.py | 40 +--------------------------------------- model.py | 42 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 43 insertions(+), 39 deletions(-) create mode 100644 model.py diff --git a/loaders.py b/loaders.py index 13f62fe..55e54fc 100644 --- a/loaders.py +++ b/loaders.py @@ -1,15 +1,11 @@ -import numpy as np -import torch from coremltools import ComputeUnit -from python_coreml_stable_diffusion.coreml_model import CoreMLModel import folder_paths from comfy import supported_models_base, model_management from comfy.latent_formats import SD15 -from comfy.model_base import BaseModel from comfy.model_patcher import ModelPatcher from .logger import logger -from .utils import expand_inputs, extract_residual_kwargs +from .model import CoreMLModelWrapper class CoreMLLoader: @@ -91,37 +87,3 @@ class CoreMLLoaderVAE(CoreMLLoader): def load(self, coreml_name, compute_unit): # TODO: Implement this pass - - -class CoreMLModelWrapper(BaseModel): - def __init__(self, model_config, mlpackage_path, compute_unit, - sources="packages"): - super().__init__(model_config) - self.diffusion_model = CoreMLModel(mlpackage_path, compute_unit, - sources) - - def apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, - control=None, transformer_options={}): - sample = x.cpu().numpy().astype(np.float16) - - context = c_crossattn.cpu().numpy().astype(np.float16) - context = context.transpose(0, 2, 1)[:, :, None, :] - - t = t.cpu().numpy().astype(np.float16) - - model_input_kwargs = { - "sample": sample, - "encoder_hidden_states": context, - "timestep": t, - } - residual_kwargs = extract_residual_kwargs(self.diffusion_model, - control) - model_input_kwargs |= residual_kwargs - model_input_kwargs = expand_inputs(model_input_kwargs) - - np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] - return torch.from_numpy(np_out).to(x.device) - - def get_dtype(self): - # Hardcoding torch-compatible dtype (used for memory allocation) - return torch.float16 diff --git a/model.py b/model.py new file mode 100644 index 0000000..b3af7ff --- /dev/null +++ b/model.py @@ -0,0 +1,42 @@ +import numpy as np +import torch + +from python_coreml_stable_diffusion.coreml_model import CoreMLModel + +from comfy.model_base import BaseModel + +from .utils import expand_inputs, extract_residual_kwargs + + +class CoreMLModelWrapper(BaseModel): + def __init__(self, model_config, mlpackage_path, compute_unit, + sources="packages"): + super().__init__(model_config) + self.diffusion_model = CoreMLModel(mlpackage_path, compute_unit, + sources) + + def apply_model(self, x, t, c_concat=None, c_crossattn=None, c_adm=None, + control=None, transformer_options={}): + sample = x.cpu().numpy().astype(np.float16) + + context = c_crossattn.cpu().numpy().astype(np.float16) + context = context.transpose(0, 2, 1)[:, :, None, :] + + t = t.cpu().numpy().astype(np.float16) + + model_input_kwargs = { + "sample": sample, + "encoder_hidden_states": context, + "timestep": t, + } + residual_kwargs = extract_residual_kwargs(self.diffusion_model, + control) + model_input_kwargs |= residual_kwargs + model_input_kwargs = expand_inputs(model_input_kwargs) + + np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"] + return torch.from_numpy(np_out).to(x.device) + + def get_dtype(self): + # Hardcoding torch-compatible dtype (used for memory allocation) + return torch.float16