Refactor model config
This commit is contained in:
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 +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
|
||||
|
||||
Reference in New Issue
Block a user