mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
Add VAE bake
This commit is contained in:
@@ -317,6 +317,17 @@ def run_MEHmodelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
except Exception as e:
|
||||
return fail(f"{e}")
|
||||
|
||||
bake_in_vae_filename = sd_vae.vae_dict.get(kwargs.get("bake_in_vae", None), None)
|
||||
if bake_in_vae_filename is not None:
|
||||
shared.log.info(f"Model merge: baking in VAE: {bake_in_vae_filename}")
|
||||
shared.state.textinfo = 'Baking in VAE'
|
||||
vae_dict = sd_vae.load_vae_dict(bake_in_vae_filename)
|
||||
for key in vae_dict.keys():
|
||||
theta_0_key = 'first_stage_model.' + key
|
||||
if theta_0_key in theta_0:
|
||||
theta_0[theta_0_key] = to_half(vae_dict[key], kwargs.get("precision", "fp16") == "fp16")
|
||||
del vae_dict
|
||||
|
||||
ckpt_dir = shared.opts.ckpt_dir or sd_models.model_path
|
||||
filename = kwargs.get("custom_name", "Unnamed_Merge")
|
||||
filename += "." + kwargs.get("checkpoint_format", None)
|
||||
|
||||
@@ -4,6 +4,7 @@ from random import shuffle
|
||||
from typing import NamedTuple
|
||||
import torch
|
||||
from scipy.optimize import linear_sum_assignment
|
||||
from modules.shared import log
|
||||
|
||||
SPECIAL_KEYS = [
|
||||
"first_stage_model.decoder.norm_out.weight",
|
||||
@@ -2227,7 +2228,7 @@ def inner_matching(
|
||||
if newL - oldL != 0:
|
||||
linear_sum += abs((newL - oldL).item())
|
||||
number += 1
|
||||
logging.info(f" permutation {p}: {newL - oldL}")
|
||||
log.info(f" permutation {p}: {newL - oldL}")
|
||||
|
||||
progress = progress or newL > oldL + 1e-12
|
||||
|
||||
|
||||
Reference in New Issue
Block a user