diff --git a/modules/merging/merge.py b/modules/merging/merge.py index 92446b2a8..a057ed015 100644 --- a/modules/merging/merge.py +++ b/modules/merging/merge.py @@ -1,26 +1,27 @@ import os from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager -# from pathlib import Path from typing import Dict, Optional, Tuple - import safetensors.torch import torch from tqdm import tqdm - import modules.memstats import modules.devices as devices from modules.shared import log from modules.sd_models import read_state_dict from modules.merging import merge_methods from modules.merging.merge_utils import WeightClass -# from modules.merging.merge_model import SDModel from modules.merging.merge_rebasin import ( apply_permutation, sdunet_permutation_spec, update_model_a, weight_matching, ) +########################################################## +# Files in modules.merging are heavily modified +# versions of sd-meh by @s1dxl used with his blessing +# orginal code can be found @ https://github.com/s1dlx/meh +########################################################## MAX_TOKENS = 77 @@ -36,12 +37,6 @@ KEY_POSITION_IDS = ".".join( ) -NAI_KEYS = { - "cond_stage_model.transformer.embeddings.": "cond_stage_model.transformer.text_model.embeddings.", - "cond_stage_model.transformer.encoder.": "cond_stage_model.transformer.text_model.encoder.", - "cond_stage_model.transformer.final_layer_norm.": "cond_stage_model.transformer.text_model.final_layer_norm.", -} - def fix_clip(model: Dict) -> Dict: if KEY_POSITION_IDS in model.keys(): @@ -54,29 +49,6 @@ def fix_clip(model: Dict) -> Dict: return model -def fix_key(model: Dict, key: str) -> Dict: - for nk in NAI_KEYS: - if key.startswith(nk): - model[key.replace(nk, NAI_KEYS[nk])] = model[key] - del model[key] - - return model - - -# https://github.com/j4ded/sdweb-merge-block-weighted-gui/blob/master/scripts/mbw/merge_block_weighted.py#L115 -def fix_model(model: Dict) -> Dict: - for k in model.keys(): - model = fix_key(model, k) - return fix_clip(model) - - -# def load_sd_model(model: os.PathLike | str, device: torch.device = None) -> Dict: -# if isinstance(model, str): -# model = Path(model) -# -# return SDModel(model, device).load_model() - - def prune_sd_model(model: Dict) -> Dict: keys = list(model.keys()) for k in keys: @@ -201,7 +173,7 @@ def un_prune_model( merged.update({key: merged[key].half()}) del original_b - return fix_model(merged) + return fix_clip(merged) def simple_merge( @@ -247,7 +219,7 @@ def simple_merge( log_vram("after stage 2") - return fix_model(thetas["model_a"]) + return fix_clip(thetas["model_a"]) def rebasin_merge( diff --git a/modules/merging/merge_model.py b/modules/merging/merge_model.py deleted file mode 100644 index 37427c44b..000000000 --- a/modules/merging/merge_model.py +++ /dev/null @@ -1,54 +0,0 @@ -import os -from dataclasses import dataclass - -import safetensors -import torch -from tensordict import TensorDict - - - -@dataclass -class SDModel: - model_path: os.PathLike - device: torch.device - - def load_model(self): - # logging.info(f"Loading: {self.model_path}") - if os.path.splitext(self.model_path)[1] == ".safetensors": - ckpt = safetensors.torch.load_file( - self.model_path, - device=self.device, - ) - else: - ckpt = torch.load(self.model_path, map_location=self.device) - - return TensorDict.from_dict(get_state_dict_from_checkpoint(ckpt)) - - -# TODO: tidy up -# from: stable-diffusion-webui/modules/sd_models.py -def get_state_dict_from_checkpoint(pl_sd): - pl_sd = pl_sd.pop("state_dict", pl_sd) - pl_sd.pop("state_dict", None) - sd = {} - for k, v in pl_sd.items(): - if new_key := transform_checkpoint_dict_key(k): - sd[new_key] = v - - pl_sd.clear() - pl_sd.update(sd) - return pl_sd - - -chckpoint_dict_replacements = { - "cond_stage_model.transformer.embeddings.": "cond_stage_model.transformer.text_model.embeddings.", - "cond_stage_model.transformer.encoder.": "cond_stage_model.transformer.text_model.encoder.", - "cond_stage_model.transformer.final_layer_norm.": "cond_stage_model.transformer.text_model.final_layer_norm.", -} - - -def transform_checkpoint_dict_key(k): - for text, replacement in chckpoint_dict_replacements.items(): - if k.startswith(text): - k = replacement + k[len(text):] - return k diff --git a/modules/merging/merge_utils.py b/modules/merging/merge_utils.py index 41b6f7442..da0059967 100644 --- a/modules/merging/merge_utils.py +++ b/modules/merging/merge_utils.py @@ -1,5 +1,4 @@ import inspect -# import logging import re from modules.merging import merge_methods from modules.merging.merge_presets import BLOCK_WEIGHTS_PRESETS, SDXL_BLOCK_WEIGHTS_PRESETS