Move CoreMLModelWrapper to model.py

This commit is contained in:
Adrian Szczepański
2023-10-23 15:10:35 +02:00
parent 771e83bf33
commit 40a7228e8f
2 changed files with 43 additions and 39 deletions
+1 -39
View File
@@ -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
+42
View File
@@ -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