Wrapped Core ML Model is now diffusion_model attribute of BaseModel

This commit is contained in:
aszc-dev
2023-11-08 17:51:27 +01:00
parent 6d095a67a2
commit 701443f59e
2 changed files with 19 additions and 21 deletions
+8 -15
View File
@@ -24,18 +24,16 @@ def get_model_config():
return model_config
class CoreMLModelWrapper(BaseModel):
def __init__(self, model_config, coreml_model):
super().__init__(model_config)
class CoreMLModelWrapper:
def __init__(self, coreml_model):
self.diffusion_model = coreml_model
self.dtype = torch.float16
def apply_model(
self,
x,
t,
c_concat=None,
c_crossattn=None,
c_adm=None,
control=None,
transformer_options={},
**kwargs,
@@ -108,18 +106,13 @@ class CoreMLModelWrapper(BaseModel):
def expected_inputs(self):
return self.diffusion_model.expected_inputs
def __call__(self, latents, ts, encoder_hidden_states, **kwargs):
return (
self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs),
def __call__(self, latents, ts, context, control, transformer_options, **kwargs):
return self.apply_model(
latents, ts, context, control, transformer_options, **kwargs
)
class CoreMLModelWrapperLCM(CoreMLModelWrapper):
def __init__(self, model_config, coreml_model):
super().__init__(model_config, coreml_model)
def __init__(self, coreml_model):
super().__init__(coreml_model)
self.config = None
def __call__(self, latents, ts, encoder_hidden_states, **kwargs):
return (
self.apply_model(latents, ts, c_crossattn=encoder_hidden_states, **kwargs),
)
+11 -6
View File
@@ -5,6 +5,7 @@ from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
import folder_paths
from comfy import model_base
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite.logger import logger
@@ -41,8 +42,10 @@ class CoreMLSampler(KSampler):
denoise=1.0,
):
model_config = get_model_config()
wrapped_model = CoreMLModelWrapper(model_config, coreml_model)
model = ModelPatcher(wrapped_model, get_torch_device(), None)
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.")
@@ -50,7 +53,7 @@ class CoreMLSampler(KSampler):
latent_image = {"samples": torch.zeros(expected[0] // 2, *expected[1:])}
return super().sample(
model,
model_patcher,
seed,
steps,
cfg,
@@ -130,6 +133,8 @@ class CoreMLModelAdapter:
def wrap(self, coreml_model):
model_config = get_model_config()
wrapped_model = CoreMLModelWrapper(model_config, coreml_model)
patched_model = ModelPatcher(wrapped_model, get_torch_device(), None)
return (patched_model,)
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)
return (model_patcher,)