diff --git a/modules/extras.py b/modules/extras.py index a0c048b4e..7bf9f8ad6 100644 --- a/modules/extras.py +++ b/modules/extras.py @@ -87,7 +87,10 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument "alpha_out_blocks"].split(",")] assert len(alpha) == 26 or len(alpha) == 20, "Alpha Block Weights are wrong length (26 or 20 for SDXL) falling back" kwargs["alpha"] = alpha - except Exception as e: + except KeyError as ke: + shared.log.warn(f"Merge: Malformed manual block weight at {ke} falling back") + kwargs["alpha"] = kwargs.get("alpha_preset", kwargs["alpha"]) + except AssertionError as e: shared.log.warn(f"Merge: {e}") kwargs["alpha"] = kwargs.get("alpha_preset", kwargs["alpha"]) finally: @@ -103,7 +106,10 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument [kwargs["beta_base"]] + kwargs["beta_in_blocks"].split(",") + [kwargs["beta_mid_block"]] + kwargs["beta_out_blocks"].split(",")] assert len(beta) == 26 or len(beta) == 20, "Beta Block Weights are wrong length (26 or 20 for SDXL) falling back" kwargs["beta"] = beta - except Exception as e: + except KeyError as ke: + shared.log.warn(f"Merge: Malformed manual block weight at {ke} falling back") + kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"]) + except AssertionError as e: shared.log.warn(f"Merge: {e}") kwargs["beta"] = kwargs.get("beta_preset", kwargs["beta"]) finally: diff --git a/modules/merging/merge.py b/modules/merging/merge.py index 09d03685b..c7118b11a 100644 --- a/modules/merging/merge.py +++ b/modules/merging/merge.py @@ -162,9 +162,9 @@ def un_prune_model( unpruned += 1 if precision == "fp16": merged.update({key: merged[key].half()}) - if unpruned != 0: - log.info(f"Merge: {unpruned} unmerged keys restored from Primary Model") - unpruned = 0 + if unpruned > 248: # VAE has 248 keys, and we are purposely restoring it here + log.info(f"Merge: {unpruned - 248} unmerged keys restored from Primary Model") + unpruned = 0 del original_a devices.torch_gc(force=True) original_b = TensorDict.from_dict(read_state_dict(models["model_b"], device)) diff --git a/modules/ui_models.py b/modules/ui_models.py index 9141dc9bb..ae28825f1 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -244,9 +244,9 @@ def create_ui(): def show_unload(device): if device == "gpu": - return gr.update(visble=True) + return gr.update(visible=True) else: - return gr.update(visble=False) + return gr.update(visible=False) def preset_visiblility(x):