From 11541024483c2f271d734d6ca704b8ec02b6feac Mon Sep 17 00:00:00 2001 From: Ayberk Aksoy Date: Sat, 30 Dec 2023 23:29:12 -0500 Subject: [PATCH] Refactor model imports and fix import paths --- model/iter.py | 3 ++- model/merge2.py | 8 ++++---- sample.py | 2 +- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/model/iter.py b/model/iter.py index 351a406..cd9f09e 100644 --- a/model/iter.py +++ b/model/iter.py @@ -1,7 +1,8 @@ from typing import List, Callable, Any, Optional import torch import tqdm -from comfy.sd import ModelPatcher, CLIP, VAE +from comfy.sd import CLIP, VAE +from comfy.model_patcher import ModelPatcher class CondForModels(torch.Tensor): diff --git a/model/merge2.py b/model/merge2.py index 7d0aed8..4508632 100644 --- a/model/merge2.py +++ b/model/merge2.py @@ -5,7 +5,7 @@ from .loader import Dict2Model from .merge import block_index, weighted_sum_block, StateDictMergerBlockWeighted from .iter import iterize_model -from comfy.ldm.models.diffusion.ddpm import LatentDiffusion +# from comfy.ldm.models.diffusion.ddpm import LatentDiffusion class MergedModule(torch.nn.Module): @@ -41,15 +41,15 @@ class MergedModule(torch.nn.Module): ATTR_ALPHAS = 'mbw_alphas' ATTR_INDEX = 'mbw_index' -def get_current_alpha(model: LatentDiffusion) -> Optional[List[float]]: +def get_current_alpha(model) -> Optional[List[float]]: if hasattr(model, ATTR_ALPHAS): return getattr(model, ATTR_ALPHAS)[getattr(model, ATTR_INDEX)] else: return None def mbw_on_the_fly( - model_A: LatentDiffusion, - model_B: LatentDiffusion, + model_A, + model_B, alphas_list: List[List[float]], base_alpha: float, ): diff --git a/sample.py b/sample.py index be85b47..fa7532c 100644 --- a/sample.py +++ b/sample.py @@ -6,7 +6,7 @@ import comfy.sample import comfy.model_management import comfy.samplers from nodes import common_ksampler -from comfy.sd import ModelPatcher +from comfy.model_patcher import ModelPatcher from .model.iter import iterize_model, CondForModels from .model import merge2