Refactor model imports and fix import paths
This commit is contained in:
+2
-1
@@ -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):
|
||||
|
||||
|
||||
+4
-4
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user