Wrapped Core ML Model is now diffusion_model attribute of BaseModel
This commit is contained in:
+8
-15
@@ -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
@@ -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,)
|
||||
|
||||
Reference in New Issue
Block a user