mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
Cleanup and Credits
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,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
|
||||
|
||||
Reference in New Issue
Block a user