Refactor model config

This commit is contained in:
aszc-dev
2023-11-09 00:04:28 +01:00
parent c26099b334
commit fa0735746c
5 changed files with 48 additions and 29 deletions
+36
View File
@@ -0,0 +1,36 @@
import torch
from comfy import supported_models_base
from comfy import latent_formats
from comfy.model_detection import convert_config
SD15 = {
"use_checkpoint": False,
"image_size": 32,
"out_channels": 4,
"use_spatial_transformer": True,
"legacy": False,
"adm_in_channels": None,
"dtype": torch.float16,
"in_channels": 4,
"model_channels": 320,
"num_res_blocks": 2,
"attention_resolutions": [1, 2, 4],
"transformer_depth": [1, 1, 1, 0],
"channel_mult": [1, 2, 4, 4],
"transformer_depth_middle": 1,
"use_linear_in_transformer": False,
"context_dim": 768,
"num_heads": 8,
"disable_unet_model_creation": True,
}
def get_unet_config():
return convert_config(SD15)
def get_model_config():
config = supported_models_base.BASE(get_unet_config())
config.latent_format = latent_formats.SD15()
return config
+9 -5
View File
@@ -6,12 +6,14 @@ from diffusers.utils.torch_utils import randn_tensor
from tqdm import tqdm
import latent_preview
from comfy import model_base
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite.lcm.lcm_scheduler import LCMScheduler
from coreml_suite.logger import logger
from coreml_suite.models import get_model_config, CoreMLModelWrapperLCM
from coreml_suite.models import CoreMLModelWrapper
from coreml_suite.nodes import CoreMLSampler
from coreml_suite.config import get_model_config
class CoreMLSamplerLCM(CoreMLSampler):
@@ -47,8 +49,10 @@ class CoreMLSamplerLCM(CoreMLSampler):
**kwargs,
):
model_config = get_model_config()
wrapped_model = CoreMLModelWrapperLCM(model_config, coreml_model)
patched_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.")
@@ -57,11 +61,11 @@ class CoreMLSamplerLCM(CoreMLSampler):
positive = positive[0][0]
callback = latent_preview.prepare_callback(patched_model, steps, None)
callback = latent_preview.prepare_callback(model_patcher, steps, None)
torch.manual_seed(seed)
return self._sample(
patched_model, steps, cfg, positive, latent_image, denoise, callback
model_patcher, steps, cfg, positive, latent_image, denoise, callback
)
def _sample(
-18
View File
@@ -1,28 +1,10 @@
import numpy as np
import torch
from comfy import supported_models_base
from comfy.latent_formats import SD15
from coreml_suite.controlnet import extract_residual_kwargs, chunk_control
from coreml_suite.latents import chunk_batch, merge_chunks
def get_model_config():
# TODO: This is a dummy model config, but it should be enough to
# get the model to load - implement a proper model config
model_config = supported_models_base.BASE({})
model_config.latent_format = SD15()
model_config.unet_config = {
"disable_unet_model_creation": True,
"num_res_blocks": 2,
"attention_resolutions": [1, 2, 4],
"channel_mult": [1, 2, 4, 4],
"transformer_depth": [1, 1, 1, 0],
}
return model_config
class CoreMLModelWrapper:
def __init__(self, coreml_model):
self.coreml_model = coreml_model
+2 -1
View File
@@ -11,7 +11,8 @@ from comfy.model_patcher import ModelPatcher
from coreml_suite.logger import logger
from nodes import KSampler
from coreml_suite.models import CoreMLModelWrapper, get_model_config
from coreml_suite.models import CoreMLModelWrapper
from coreml_suite.config import get_model_config
class CoreMLSampler(KSampler):
+1 -5
View File
@@ -1,5 +1,3 @@
from unittest import mock
import pytest
import torch
@@ -8,11 +6,9 @@ from comfy.model_management import get_torch_device
from coreml_suite.latents import chunk_batch, merge_chunks
from coreml_suite.controlnet import chunk_control
from coreml_suite.models import (
CoreMLModelWrapper,
get_model_config,
CoreMLModelWrapperLCM,
CoreMLInputs,
)
from coreml_suite.config import get_model_config
@pytest.fixture