Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
model_manager.py519 linesDownload Raw Back to root
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