2022-10-13 23:00:38 -06:00
|
|
|
import collections
|
2022-09-17 03:05:04 -06:00
|
|
|
import os.path
|
|
|
|
import sys
|
|
|
|
from collections import namedtuple
|
|
|
|
import torch
|
|
|
|
from omegaconf import OmegaConf
|
|
|
|
|
|
|
|
from ldm.util import instantiate_from_config
|
|
|
|
|
2022-10-02 06:03:39 -06:00
|
|
|
from modules import shared, modelloader, devices
|
2022-09-27 10:01:13 -06:00
|
|
|
from modules.paths import models_path
|
|
|
|
|
|
|
|
model_dir = "Stable-diffusion"
|
2022-09-30 02:42:40 -06:00
|
|
|
model_path = os.path.abspath(os.path.join(models_path, model_dir))
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-08 14:26:48 -06:00
|
|
|
CheckpointInfo = namedtuple("CheckpointInfo", ['filename', 'title', 'hash', 'model_name', 'config'])
|
2022-09-17 03:05:04 -06:00
|
|
|
checkpoints_list = {}
|
2022-10-13 23:00:38 -06:00
|
|
|
checkpoints_loaded = collections.OrderedDict()
|
2022-09-17 03:05:04 -06:00
|
|
|
|
|
|
|
try:
|
|
|
|
# this silences the annoying "Some weights of the model checkpoint were not used when initializing..." message at start.
|
|
|
|
|
|
|
|
from transformers import logging
|
|
|
|
|
|
|
|
logging.set_verbosity_error()
|
|
|
|
except Exception:
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
2022-10-02 12:09:10 -06:00
|
|
|
def setup_model():
|
2022-09-27 10:01:13 -06:00
|
|
|
if not os.path.exists(model_path):
|
|
|
|
os.makedirs(model_path)
|
2022-10-02 12:09:10 -06:00
|
|
|
|
2022-09-29 18:59:36 -06:00
|
|
|
list_models()
|
|
|
|
|
|
|
|
|
2022-09-28 15:59:44 -06:00
|
|
|
def checkpoint_tiles():
|
|
|
|
return sorted([x.title for x in checkpoints_list.values()])
|
|
|
|
|
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
def list_models():
|
|
|
|
checkpoints_list.clear()
|
2022-10-02 12:09:10 -06:00
|
|
|
model_list = modelloader.load_models(model_path=model_path, command_path=shared.cmd_opts.ckpt_dir, ext_filter=[".ckpt"])
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-09-30 02:42:40 -06:00
|
|
|
def modeltitle(path, shorthash):
|
2022-09-17 03:05:04 -06:00
|
|
|
abspath = os.path.abspath(path)
|
|
|
|
|
2022-10-02 12:22:20 -06:00
|
|
|
if shared.cmd_opts.ckpt_dir is not None and abspath.startswith(shared.cmd_opts.ckpt_dir):
|
|
|
|
name = abspath.replace(shared.cmd_opts.ckpt_dir, '')
|
2022-09-30 02:42:40 -06:00
|
|
|
elif abspath.startswith(model_path):
|
|
|
|
name = abspath.replace(model_path, '')
|
2022-09-17 03:05:04 -06:00
|
|
|
else:
|
|
|
|
name = os.path.basename(path)
|
|
|
|
|
|
|
|
if name.startswith("\\") or name.startswith("/"):
|
|
|
|
name = name[1:]
|
|
|
|
|
2022-09-28 15:59:44 -06:00
|
|
|
shortname = os.path.splitext(name.replace("/", "_").replace("\\", "_"))[0]
|
|
|
|
|
2022-09-30 02:42:40 -06:00
|
|
|
return f'{name} [{shorthash}]', shortname
|
2022-09-17 03:05:04 -06:00
|
|
|
|
|
|
|
cmd_ckpt = shared.cmd_opts.ckpt
|
|
|
|
if os.path.exists(cmd_ckpt):
|
|
|
|
h = model_hash(cmd_ckpt)
|
2022-09-30 02:42:40 -06:00
|
|
|
title, short_model_name = modeltitle(cmd_ckpt, h)
|
2022-10-08 14:26:48 -06:00
|
|
|
checkpoints_list[title] = CheckpointInfo(cmd_ckpt, title, h, short_model_name, shared.cmd_opts.config)
|
2022-10-02 08:24:50 -06:00
|
|
|
shared.opts.data['sd_model_checkpoint'] = title
|
2022-09-17 03:05:04 -06:00
|
|
|
elif cmd_ckpt is not None and cmd_ckpt != shared.default_sd_model_file:
|
2022-09-27 10:01:13 -06:00
|
|
|
print(f"Checkpoint in --ckpt argument not found (Possible it was moved to {model_path}: {cmd_ckpt}", file=sys.stderr)
|
|
|
|
for filename in model_list:
|
|
|
|
h = model_hash(filename)
|
2022-09-30 02:42:40 -06:00
|
|
|
title, short_model_name = modeltitle(filename, h)
|
2022-10-08 14:26:48 -06:00
|
|
|
|
|
|
|
basename, _ = os.path.splitext(filename)
|
|
|
|
config = basename + ".yaml"
|
|
|
|
if not os.path.exists(config):
|
|
|
|
config = shared.cmd_opts.config
|
|
|
|
|
|
|
|
checkpoints_list[title] = CheckpointInfo(filename, title, h, short_model_name, config)
|
2022-09-30 02:42:40 -06:00
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-09-28 15:30:09 -06:00
|
|
|
def get_closet_checkpoint_match(searchString):
|
2022-09-29 12:08:03 -06:00
|
|
|
applicable = sorted([info for info in checkpoints_list.values() if searchString in info.title], key = lambda x:len(x.title))
|
2022-09-30 02:42:40 -06:00
|
|
|
if len(applicable) > 0:
|
2022-09-28 15:30:09 -06:00
|
|
|
return applicable[0]
|
|
|
|
return None
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-09-30 02:42:40 -06:00
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
def model_hash(filename):
|
|
|
|
try:
|
|
|
|
with open(filename, "rb") as file:
|
|
|
|
import hashlib
|
|
|
|
m = hashlib.sha256()
|
|
|
|
|
|
|
|
file.seek(0x100000)
|
|
|
|
m.update(file.read(0x10000))
|
|
|
|
return m.hexdigest()[0:8]
|
|
|
|
except FileNotFoundError:
|
|
|
|
return 'NOFILE'
|
|
|
|
|
|
|
|
|
|
|
|
def select_checkpoint():
|
|
|
|
model_checkpoint = shared.opts.sd_model_checkpoint
|
|
|
|
checkpoint_info = checkpoints_list.get(model_checkpoint, None)
|
|
|
|
if checkpoint_info is not None:
|
|
|
|
return checkpoint_info
|
|
|
|
|
|
|
|
if len(checkpoints_list) == 0:
|
2022-09-18 14:52:01 -06:00
|
|
|
print(f"No checkpoints found. When searching for checkpoints, looked at:", file=sys.stderr)
|
2022-10-02 12:09:10 -06:00
|
|
|
if shared.cmd_opts.ckpt is not None:
|
|
|
|
print(f" - file {os.path.abspath(shared.cmd_opts.ckpt)}", file=sys.stderr)
|
|
|
|
print(f" - directory {model_path}", file=sys.stderr)
|
|
|
|
if shared.cmd_opts.ckpt_dir is not None:
|
|
|
|
print(f" - directory {os.path.abspath(shared.cmd_opts.ckpt_dir)}", file=sys.stderr)
|
2022-09-18 14:52:01 -06:00
|
|
|
print(f"Can't run without a checkpoint. Find and place a .ckpt file into any of those locations. The program will exit.", file=sys.stderr)
|
|
|
|
exit(1)
|
2022-09-17 03:05:04 -06:00
|
|
|
|
|
|
|
checkpoint_info = next(iter(checkpoints_list.values()))
|
|
|
|
if model_checkpoint is not None:
|
|
|
|
print(f"Checkpoint {model_checkpoint} not found; loading fallback {checkpoint_info.title}", file=sys.stderr)
|
|
|
|
|
|
|
|
return checkpoint_info
|
|
|
|
|
|
|
|
|
2022-10-09 01:23:31 -06:00
|
|
|
def get_state_dict_from_checkpoint(pl_sd):
|
|
|
|
if "state_dict" in pl_sd:
|
|
|
|
return pl_sd["state_dict"]
|
|
|
|
|
|
|
|
return pl_sd
|
|
|
|
|
|
|
|
|
2022-10-08 14:26:48 -06:00
|
|
|
def load_model_weights(model, checkpoint_info):
|
|
|
|
checkpoint_file = checkpoint_info.filename
|
|
|
|
sd_model_hash = checkpoint_info.hash
|
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
if checkpoint_info not in checkpoints_loaded:
|
|
|
|
print(f"Loading weights [{sd_model_hash}] from {checkpoint_file}")
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-15 01:35:18 -06:00
|
|
|
pl_sd = torch.load(checkpoint_file, map_location=shared.weight_load_location)
|
2022-10-13 23:00:38 -06:00
|
|
|
if "global_step" in pl_sd:
|
|
|
|
print(f"Global Step: {pl_sd['global_step']}")
|
2022-10-09 01:23:31 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
sd = get_state_dict_from_checkpoint(pl_sd)
|
|
|
|
model.load_state_dict(sd, strict=False)
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
if shared.cmd_opts.opt_channelslast:
|
|
|
|
model.to(memory_format=torch.channels_last)
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
if not shared.cmd_opts.no_half:
|
|
|
|
model.half()
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
devices.dtype = torch.float32 if shared.cmd_opts.no_half else torch.float16
|
|
|
|
devices.dtype_vae = torch.float32 if shared.cmd_opts.no_half or shared.cmd_opts.no_half_vae else torch.float16
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
vae_file = os.path.splitext(checkpoint_file)[0] + ".vae.pt"
|
2022-10-02 06:03:39 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
if not os.path.exists(vae_file) and shared.cmd_opts.vae_path is not None:
|
|
|
|
vae_file = shared.cmd_opts.vae_path
|
2022-10-10 11:46:55 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
if os.path.exists(vae_file):
|
|
|
|
print(f"Loading VAE weights from: {vae_file}")
|
2022-10-15 01:35:18 -06:00
|
|
|
vae_ckpt = torch.load(vae_file, map_location=shared.weight_load_location)
|
2022-10-13 23:00:38 -06:00
|
|
|
vae_dict = {k: v for k, v in vae_ckpt["state_dict"].items() if k[0:4] != "loss"}
|
|
|
|
model.first_stage_model.load_state_dict(vae_dict)
|
2022-10-07 01:40:22 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
model.first_stage_model.to(devices.dtype_vae)
|
2022-10-07 01:40:22 -06:00
|
|
|
|
2022-10-13 23:00:38 -06:00
|
|
|
checkpoints_loaded[checkpoint_info] = model.state_dict().copy()
|
|
|
|
while len(checkpoints_loaded) > shared.opts.sd_checkpoint_cache:
|
|
|
|
checkpoints_loaded.popitem(last=False) # LRU
|
|
|
|
else:
|
|
|
|
print(f"Loading weights [{sd_model_hash}] from cache")
|
|
|
|
checkpoints_loaded.move_to_end(checkpoint_info)
|
|
|
|
model.load_state_dict(checkpoints_loaded[checkpoint_info])
|
2022-10-10 07:11:14 -06:00
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
model.sd_model_hash = sd_model_hash
|
2022-10-08 13:12:24 -06:00
|
|
|
model.sd_model_checkpoint = checkpoint_file
|
2022-10-08 14:26:48 -06:00
|
|
|
model.sd_checkpoint_info = checkpoint_info
|
2022-09-17 03:05:04 -06:00
|
|
|
|
|
|
|
|
|
|
|
def load_model():
|
|
|
|
from modules import lowvram, sd_hijack
|
|
|
|
checkpoint_info = select_checkpoint()
|
|
|
|
|
2022-10-08 14:26:48 -06:00
|
|
|
if checkpoint_info.config != shared.cmd_opts.config:
|
2022-10-09 01:31:47 -06:00
|
|
|
print(f"Loading config from: {checkpoint_info.config}")
|
2022-10-08 14:26:48 -06:00
|
|
|
|
|
|
|
sd_config = OmegaConf.load(checkpoint_info.config)
|
2022-09-17 03:05:04 -06:00
|
|
|
sd_model = instantiate_from_config(sd_config.model)
|
2022-10-08 14:26:48 -06:00
|
|
|
load_model_weights(sd_model, checkpoint_info)
|
2022-09-17 03:05:04 -06:00
|
|
|
|
|
|
|
if shared.cmd_opts.lowvram or shared.cmd_opts.medvram:
|
|
|
|
lowvram.setup_for_low_vram(sd_model, shared.cmd_opts.medvram)
|
|
|
|
else:
|
|
|
|
sd_model.to(shared.device)
|
|
|
|
|
|
|
|
sd_hijack.model_hijack.hijack(sd_model)
|
|
|
|
|
|
|
|
sd_model.eval()
|
|
|
|
|
|
|
|
print(f"Model loaded.")
|
|
|
|
return sd_model
|
|
|
|
|
|
|
|
|
2022-09-17 04:49:36 -06:00
|
|
|
def reload_model_weights(sd_model, info=None):
|
2022-09-29 06:40:28 -06:00
|
|
|
from modules import lowvram, devices, sd_hijack
|
2022-09-17 04:49:36 -06:00
|
|
|
checkpoint_info = info or select_checkpoint()
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-10-08 13:12:24 -06:00
|
|
|
if sd_model.sd_model_checkpoint == checkpoint_info.filename:
|
2022-09-17 03:05:04 -06:00
|
|
|
return
|
|
|
|
|
2022-10-08 14:26:48 -06:00
|
|
|
if sd_model.sd_checkpoint_info.config != checkpoint_info.config:
|
2022-10-13 23:00:38 -06:00
|
|
|
checkpoints_loaded.clear()
|
2022-10-09 04:23:30 -06:00
|
|
|
shared.sd_model = load_model()
|
|
|
|
return shared.sd_model
|
2022-10-08 14:26:48 -06:00
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
if shared.cmd_opts.lowvram or shared.cmd_opts.medvram:
|
|
|
|
lowvram.send_everything_to_cpu()
|
|
|
|
else:
|
|
|
|
sd_model.to(devices.cpu)
|
|
|
|
|
2022-09-29 06:40:28 -06:00
|
|
|
sd_hijack.model_hijack.undo_hijack(sd_model)
|
|
|
|
|
2022-10-08 14:26:48 -06:00
|
|
|
load_model_weights(sd_model, checkpoint_info)
|
2022-09-17 03:05:04 -06:00
|
|
|
|
2022-09-29 06:40:28 -06:00
|
|
|
sd_hijack.model_hijack.hijack(sd_model)
|
|
|
|
|
2022-09-17 03:05:04 -06:00
|
|
|
if not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram:
|
|
|
|
sd_model.to(devices.device)
|
|
|
|
|
|
|
|
print(f"Weights loaded.")
|
|
|
|
return sd_model
|