Cleanup and Credits

This commit is contained in:
AI-Casanova
2023-11-18 11:04:43 -06:00
parent cea7464553
commit d2d54af7ac
3 changed files with 7 additions and 90 deletions
+7 -35
View File
@@ -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(
-54
View File
@@ -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
-1
View File
@@ -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