hugging-apps/echo-memory
0
1import os, torch, json, importlib, logging2from typing import List3 4logger = logging.getLogger(__name__)5 6from .downloader import download_models, download_customized_models, Preset_model_id, Preset_model_website7 8from .sd_text_encoder import SDTextEncoder9from .sd_unet import SDUNet10from .sd_vae_encoder import SDVAEEncoder11from .sd_vae_decoder import SDVAEDecoder12from .lora import get_lora_loaders13 14from .sdxl_text_encoder import SDXLTextEncoder, SDXLTextEncoder215from .sdxl_unet import SDXLUNet16from .sdxl_vae_decoder import SDXLVAEDecoder17from .sdxl_vae_encoder import SDXLVAEEncoder18 19from .sd3_text_encoder import SD3TextEncoder1, SD3TextEncoder2, SD3TextEncoder320from .sd3_dit import SD3DiT21from .sd3_vae_decoder import SD3VAEDecoder22from .sd3_vae_encoder import SD3VAEEncoder23 24from .sd_controlnet import SDControlNet25from .sdxl_controlnet import SDXLControlNetUnion26 27from .sd_motion import SDMotionModel28from .sdxl_motion import SDXLMotionModel29 30from .svd_image_encoder import SVDImageEncoder31from .svd_unet import SVDUNet32from .svd_vae_decoder import SVDVAEDecoder33from .svd_vae_encoder import SVDVAEEncoder34 35from .sd_ipadapter import SDIpAdapter, IpAdapterCLIPImageEmbedder36from .sdxl_ipadapter import SDXLIpAdapter, IpAdapterXLCLIPImageEmbedder37 38from .hunyuan_dit_text_encoder import HunyuanDiTCLIPTextEncoder, HunyuanDiTT5TextEncoder39from .hunyuan_dit import HunyuanDiT40from .hunyuan_video_vae_decoder import HunyuanVideoVAEDecoder41from .hunyuan_video_vae_encoder import HunyuanVideoVAEEncoder42 43from .flux_dit import FluxDiT44from .flux_text_encoder import FluxTextEncoder245from .flux_vae import FluxVAEEncoder, FluxVAEDecoder46from .flux_ipadapter import FluxIpAdapter47 48from .cog_vae import CogVAEEncoder, CogVAEDecoder49from .cog_dit import CogDiT50 51from ..extensions.RIFE import IFNet52from ..extensions.ESRGAN import RRDBNet53 54from ..configs.model_config import model_loader_configs, huggingface_model_loader_configs, patch_model_loader_configs55from .utils import load_state_dict, init_weights_on_device, hash_state_dict_keys, split_state_dict_with_prefix56 57 58def load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device):59 loaded_model_names, loaded_models = [], []60 for model_name, model_class in zip(model_names, model_classes):61 print(f" model_name: {model_name} model_class: {model_class.__name__}")62 state_dict_converter = model_class.state_dict_converter()63 if model_resource == "civitai":64 state_dict_results = state_dict_converter.from_civitai(state_dict)65 elif model_resource == "diffusers":66 state_dict_results = state_dict_converter.from_diffusers(state_dict)67 if isinstance(state_dict_results, tuple):68 model_state_dict, extra_kwargs = state_dict_results69 print(f" This model is initialized with extra kwargs: {extra_kwargs}")70 else:71 model_state_dict, extra_kwargs = state_dict_results, {}72 torch_dtype = torch.float32 if extra_kwargs.get("upcast_to_float32", False) else torch_dtype73 with init_weights_on_device():74 model = model_class(**extra_kwargs)75 if hasattr(model, "eval"):76 model = model.eval()77 model.load_state_dict(model_state_dict, assign=True)78 model = model.to(dtype=torch_dtype, device=device)79 loaded_model_names.append(model_name)80 loaded_models.append(model)81 return loaded_model_names, loaded_models82 83 84def load_model_from_huggingface_folder(file_path, model_names, model_classes, torch_dtype, device):85 loaded_model_names, loaded_models = [], []86 for model_name, model_class in zip(model_names, model_classes):87 if torch_dtype in [torch.float32, torch.float16, torch.bfloat16]:88 model = model_class.from_pretrained(file_path, torch_dtype=torch_dtype).eval()89 else:90 model = model_class.from_pretrained(file_path).eval().to(dtype=torch_dtype)91 if torch_dtype == torch.float16 and hasattr(model, "half"):92 model = model.half()93 try:94 model = model.to(device=device)95 except:96 pass97 loaded_model_names.append(model_name)98 loaded_models.append(model)99 return loaded_model_names, loaded_models100 101 102def load_single_patch_model_from_single_file(state_dict, model_name, model_class, base_model, extra_kwargs, torch_dtype, device):103 print(f" model_name: {model_name} model_class: {model_class.__name__} extra_kwargs: {extra_kwargs}")104 base_state_dict = base_model.state_dict()105 base_model.to("cpu")106 del base_model107 model = model_class(**extra_kwargs)108 model.load_state_dict(base_state_dict, strict=False)109 model.load_state_dict(state_dict, strict=False)110 model.to(dtype=torch_dtype, device=device)111 return model112 113 114def load_patch_model_from_single_file(state_dict, model_names, model_classes, extra_kwargs, model_manager, torch_dtype, device):115 loaded_model_names, loaded_models = [], []116 for model_name, model_class in zip(model_names, model_classes):117 while True:118 for model_id in range(len(model_manager.model)):119 base_model_name = model_manager.model_name[model_id]120 if base_model_name == model_name:121 base_model_path = model_manager.model_path[model_id]122 base_model = model_manager.model[model_id]123 print(f" Adding patch model to {base_model_name} ({base_model_path})")124 patched_model = load_single_patch_model_from_single_file(125 state_dict, model_name, model_class, base_model, extra_kwargs, torch_dtype, device)126 loaded_model_names.append(base_model_name)127 loaded_models.append(patched_model)128 model_manager.model.pop(model_id)129 model_manager.model_path.pop(model_id)130 model_manager.model_name.pop(model_id)131 break132 else:133 break134 return loaded_model_names, loaded_models135 136 137 138class ModelDetectorTemplate:139 def __init__(self):140 pass141 142 def match(self, file_path="", state_dict={}):143 return False144 145 def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):146 return [], []147 148 149 150class ModelDetectorFromSingleFile:151 def __init__(self, model_loader_configs=[]):152 self.keys_hash_with_shape_dict = {}153 self.keys_hash_dict = {}154 for metadata in model_loader_configs:155 self.add_model_metadata(*metadata)156 157 158 def add_model_metadata(self, keys_hash, keys_hash_with_shape, model_names, model_classes, model_resource):159 self.keys_hash_with_shape_dict[keys_hash_with_shape] = (model_names, model_classes, model_resource)160 if keys_hash is not None:161 self.keys_hash_dict[keys_hash] = (model_names, model_classes, model_resource)162 163 164 def match(self, file_path="", state_dict={}):165 if isinstance(file_path, str) and os.path.isdir(file_path):166 return False167 if state_dict is None or len(state_dict) == 0:168 # Handle list of file paths (for split model files)169 if isinstance(file_path, list):170 state_dict = {}171 for path in file_path:172 state_dict.update(load_state_dict(path))173 else:174 state_dict = load_state_dict(file_path)175 keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)176 if keys_hash_with_shape in self.keys_hash_with_shape_dict:177 return True178 keys_hash = hash_state_dict_keys(state_dict, with_shape=False)179 if keys_hash in self.keys_hash_dict:180 return True181 # Debug: log hash if it's a list of files (merged model)182 if isinstance(file_path, list) and len(state_dict) > 0:183 logger.info(f" Debug: ModelDetectorFromSingleFile - hash_with_shape={keys_hash_with_shape}, hash={keys_hash}, keys_count={len(state_dict)}")184 logger.info(f" Debug: Available hashes count: {len(self.keys_hash_with_shape_dict)}")185 if keys_hash_with_shape in self.keys_hash_with_shape_dict:186 logger.info(f" Debug: Hash FOUND in keys_hash_with_shape_dict!")187 else:188 logger.warning(f" Debug: Hash NOT FOUND in keys_hash_with_shape_dict")189 logger.info(f" Debug: Sample hashes in dict: {list(self.keys_hash_with_shape_dict.keys())[:5]}")190 return False191 192 193 def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):194 if state_dict is None or len(state_dict) == 0:195 # Handle list of file paths (for split model files)196 if isinstance(file_path, list):197 state_dict = {}198 for path in file_path:199 state_dict.update(load_state_dict(path))200 else:201 state_dict = load_state_dict(file_path)202 203 # Load models with strict matching204 keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)205 if keys_hash_with_shape in self.keys_hash_with_shape_dict:206 model_names, model_classes, model_resource = self.keys_hash_with_shape_dict[keys_hash_with_shape]207 loaded_model_names, loaded_models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device)208 return loaded_model_names, loaded_models209 210 # Load models without strict matching211 # (the shape of parameters may be inconsistent, and the state_dict_converter will modify the model architecture)212 keys_hash = hash_state_dict_keys(state_dict, with_shape=False)213 if keys_hash in self.keys_hash_dict:214 model_names, model_classes, model_resource = self.keys_hash_dict[keys_hash]215 loaded_model_names, loaded_models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device)216 return loaded_model_names, loaded_models217 218 return [], []219 220 221 222class ModelDetectorFromSplitedSingleFile(ModelDetectorFromSingleFile):223 def __init__(self, model_loader_configs=[]):224 super().__init__(model_loader_configs)225 226 227 def match(self, file_path="", state_dict={}):228 if isinstance(file_path, str) and os.path.isdir(file_path):229 return False230 if state_dict is None or len(state_dict) == 0:231 # Handle list of file paths (for split model files)232 if isinstance(file_path, list):233 state_dict = {}234 for path in file_path:235 state_dict.update(load_state_dict(path))236 else:237 state_dict = load_state_dict(file_path)238 # First try to match the complete state_dict (for merged models)239 if super().match(file_path, state_dict):240 return True241 # If complete match fails, try split matching242 splited_state_dict = split_state_dict_with_prefix(state_dict)243 for sub_state_dict in splited_state_dict:244 if super().match(file_path, sub_state_dict):245 return True246 return False247 248 249 def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):250 # Load state_dict if empty251 if state_dict is None or len(state_dict) == 0:252 # Handle list of file paths (for split model files)253 if isinstance(file_path, list):254 state_dict = {}255 for path in file_path:256 state_dict.update(load_state_dict(path))257 else:258 state_dict = load_state_dict(file_path)259 # First try to load the complete state_dict (for merged models)260 if super().match(file_path, state_dict):261 loaded_model_names, loaded_models = super().load(file_path, state_dict, device, torch_dtype, **kwargs)262 if loaded_model_names:263 return loaded_model_names, loaded_models264 # If complete load fails, try split loading265 splited_state_dict = split_state_dict_with_prefix(state_dict)266 valid_state_dict = {}267 for sub_state_dict in splited_state_dict:268 if super().match(file_path, sub_state_dict):269 valid_state_dict.update(sub_state_dict)270 if super().match(file_path, valid_state_dict):271 loaded_model_names, loaded_models = super().load(file_path, valid_state_dict, device, torch_dtype, **kwargs)272 else:273 loaded_model_names, loaded_models = [], []274 for sub_state_dict in splited_state_dict:275 if super().match(file_path, sub_state_dict):276 loaded_model_names_, loaded_models_ = super().load(file_path, valid_state_dict, device, torch_dtype, **kwargs)277 loaded_model_names += loaded_model_names_278 loaded_models += loaded_models_279 return loaded_model_names, loaded_models280 281 282 283class ModelDetectorFromHuggingfaceFolder:284 def __init__(self, model_loader_configs=[]):285 self.architecture_dict = {}286 for metadata in model_loader_configs:287 self.add_model_metadata(*metadata)288 289 290 def add_model_metadata(self, architecture, huggingface_lib, model_name, redirected_architecture):291 self.architecture_dict[architecture] = (huggingface_lib, model_name, redirected_architecture)292 293 294 def match(self, file_path="", state_dict={}):295 if not isinstance(file_path, str) or os.path.isfile(file_path):296 return False297 file_list = os.listdir(file_path)298 if "config.json" not in file_list:299 return False300 with open(os.path.join(file_path, "config.json"), "r") as f:301 config = json.load(f)302 if "architectures" not in config and "_class_name" not in config:303 return False304 return True305 306 307 def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):308 with open(os.path.join(file_path, "config.json"), "r") as f:309 config = json.load(f)310 loaded_model_names, loaded_models = [], []311 architectures = config["architectures"] if "architectures" in config else [config["_class_name"]]312 for architecture in architectures:313 huggingface_lib, model_name, redirected_architecture = self.architecture_dict[architecture]314 if redirected_architecture is not None:315 architecture = redirected_architecture316 model_class = importlib.import_module(huggingface_lib).__getattribute__(architecture)317 loaded_model_names_, loaded_models_ = load_model_from_huggingface_folder(file_path, [model_name], [model_class], torch_dtype, device)318 loaded_model_names += loaded_model_names_319 loaded_models += loaded_models_320 return loaded_model_names, loaded_models321 322 323 324class ModelDetectorFromPatchedSingleFile:325 def __init__(self, model_loader_configs=[]):326 self.keys_hash_with_shape_dict = {}327 for metadata in model_loader_configs:328 self.add_model_metadata(*metadata)329 330 331 def add_model_metadata(self, keys_hash_with_shape, model_name, model_class, extra_kwargs):332 self.keys_hash_with_shape_dict[keys_hash_with_shape] = (model_name, model_class, extra_kwargs)333 334 335 def match(self, file_path="", state_dict={}):336 if not isinstance(file_path, str) or os.path.isdir(file_path):337 return False338 if state_dict is None or len(state_dict) == 0:339 state_dict = load_state_dict(file_path)340 keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)341 if keys_hash_with_shape in self.keys_hash_with_shape_dict:342 return True343 return False344 345 346 def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, model_manager=None, **kwargs):347 if state_dict is None or len(state_dict) == 0:348 state_dict = load_state_dict(file_path)349 350 # Load models with strict matching351 loaded_model_names, loaded_models = [], []352 keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)353 if keys_hash_with_shape in self.keys_hash_with_shape_dict:354 model_names, model_classes, extra_kwargs = self.keys_hash_with_shape_dict[keys_hash_with_shape]355 loaded_model_names_, loaded_models_ = load_patch_model_from_single_file(356 state_dict, model_names, model_classes, extra_kwargs, model_manager, torch_dtype, device)357 loaded_model_names += loaded_model_names_358 loaded_models += loaded_models_359 return loaded_model_names, loaded_models360 361 362 363class ModelManager:364 def __init__(365 self,366 torch_dtype=torch.float16,367 device="cuda",368 model_id_list: List[Preset_model_id] = [],369 downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],370 file_path_list: List[str] = [],371 ):372 self.torch_dtype = torch_dtype373 self.device = device374 self.model = []375 self.model_path = []376 self.model_name = []377 downloaded_files = download_models(model_id_list, downloading_priority) if len(model_id_list) > 0 else []378 self.model_detector = [379 ModelDetectorFromSingleFile(model_loader_configs),380 ModelDetectorFromSplitedSingleFile(model_loader_configs),381 ModelDetectorFromHuggingfaceFolder(huggingface_model_loader_configs),382 ModelDetectorFromPatchedSingleFile(patch_model_loader_configs),383 ]384 self.load_models(downloaded_files + file_path_list)385 386 387 def load_model_from_single_file(self, file_path="", state_dict={}, model_names=[], model_classes=[], model_resource=None):388 print(f"Loading models from file: {file_path}")389 if state_dict is None or len(state_dict) == 0:390 state_dict = load_state_dict(file_path)391 model_names, models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, self.torch_dtype, self.device)392 for model_name, model in zip(model_names, models):393 self.model.append(model)394 self.model_path.append(file_path)395 self.model_name.append(model_name)396 print(f" The following models are loaded: {model_names}.")397 398 399 def load_model_from_huggingface_folder(self, file_path="", model_names=[], model_classes=[]):400 print(f"Loading models from folder: {file_path}")401 model_names, models = load_model_from_huggingface_folder(file_path, model_names, model_classes, self.torch_dtype, self.device)402 for model_name, model in zip(model_names, models):403 self.model.append(model)404 self.model_path.append(file_path)405 self.model_name.append(model_name)406 print(f" The following models are loaded: {model_names}.")407 408 409 def load_patch_model_from_single_file(self, file_path="", state_dict={}, model_names=[], model_classes=[], extra_kwargs={}):410 print(f"Loading patch models from file: {file_path}")411 model_names, models = load_patch_model_from_single_file(412 state_dict, model_names, model_classes, extra_kwargs, self, self.torch_dtype, self.device)413 for model_name, model in zip(model_names, models):414 self.model.append(model)415 self.model_path.append(file_path)416 self.model_name.append(model_name)417 print(f" The following patched models are loaded: {model_names}.")418 419 420 def load_lora(self, file_path="", state_dict={}, lora_alpha=1.0):421 if isinstance(file_path, list):422 for file_path_ in file_path:423 self.load_lora(file_path_, state_dict=state_dict, lora_alpha=lora_alpha)424 else:425 print(f"Loading LoRA models from file: {file_path}")426 is_loaded = False427 if state_dict is None or len(state_dict) == 0:428 state_dict = load_state_dict(file_path)429 for model_name, model, model_path in zip(self.model_name, self.model, self.model_path):430 for lora in get_lora_loaders():431 match_results = lora.match(model, state_dict)432 if match_results is not None:433 print(f" Adding LoRA to {model_name} ({model_path}).")434 lora_prefix, model_resource = match_results435 lora.load(model, state_dict, lora_prefix, alpha=lora_alpha, model_resource=model_resource)436 is_loaded = True437 break438 if not is_loaded:439 print(f" Cannot load LoRA: {file_path}")440 441 442 def load_model(self, file_path, model_names=None, device=None, torch_dtype=None):443 print(f"Loading models from: {file_path}")444 if device is None: device = self.device445 if torch_dtype is None: torch_dtype = self.torch_dtype446 if isinstance(file_path, list):447 state_dict = {}448 for path in file_path:449 state_dict.update(load_state_dict(path))450 logger.info(f" Merged state_dict from {len(file_path)} files, total keys: {len(state_dict)}")451 elif os.path.isfile(file_path):452 state_dict = load_state_dict(file_path)453 else:454 state_dict = None455 for i, model_detector in enumerate(self.model_detector):456 detector_name = model_detector.__class__.__name__457 if model_detector.match(file_path, state_dict):458 logger.info(f" Matched by {detector_name}")459 model_names, models = model_detector.load(460 file_path, state_dict,461 device=device, torch_dtype=torch_dtype,462 allowed_model_names=model_names, model_manager=self463 )464 for model_name, model in zip(model_names, models):465 self.model.append(model)466 self.model_path.append(file_path)467 self.model_name.append(model_name)468 print(f" The following models are loaded: {model_names}.")469 break470 else:471 if isinstance(file_path, list) and len(state_dict) > 0:472 logger.info(f" {detector_name} did not match")473 else:474 print(f" We cannot detect the model type. No models are loaded.")475 if isinstance(file_path, list) and len(state_dict) > 0:476 from .utils import hash_state_dict_keys477 actual_hash = hash_state_dict_keys(state_dict, with_shape=True)478 logger.warning(f" Debug: Actual hash = {actual_hash}")479 logger.warning(f" Debug: First detector has {len(self.model_detector[0].keys_hash_with_shape_dict)} configured hashes")480 # Check if hash exists in config481 if actual_hash in self.model_detector[0].keys_hash_with_shape_dict:482 logger.error(f" Debug: Hash EXISTS in detector but match() returned False!")483 else:484 logger.warning(f" Debug: Hash does NOT exist in detector config")485 logger.info(f" Debug: Sample configured hashes: {list(self.model_detector[0].keys_hash_with_shape_dict.keys())[:10]}")486 487 488 def load_models(self, file_path_list, model_names=None, device=None, torch_dtype=None):489 for file_path in file_path_list:490 self.load_model(file_path, model_names, device=device, torch_dtype=torch_dtype)491 492 493 def fetch_model(self, model_name, file_path=None, require_model_path=False):494 fetched_models = []495 fetched_model_paths = []496 for model, model_path, model_name_ in zip(self.model, self.model_path, self.model_name):497 if file_path is not None and file_path != model_path:498 continue499 if model_name == model_name_:500 fetched_models.append(model)501 fetched_model_paths.append(model_path)502 if len(fetched_models) == 0:503 print(f"No {model_name} models available.")504 return None505 if len(fetched_models) == 1:506 print(f"Using {model_name} from {fetched_model_paths[0]}.")507 else:508 print(f"More than one {model_name} models are loaded in model manager: {fetched_model_paths}. Using {model_name} from {fetched_model_paths[0]}.")509 if require_model_path:510 return fetched_models[0], fetched_model_paths[0]511 else:512 return fetched_models[0]513 514 515 def to(self, device):516 for model in self.model:517 model.to(device)518 519 