Merge: Typo and error handling (#2527)

* Typo and error handling

* typo typo

* error logging

* fix alpha vs beta

* Don't log VAE keys in unpruning
This commit is contained in:
AI-Casanova
2023-11-19 13:47:02 -06:00
committed by GitHub
parent 2b5adeeb92
commit 564b67bcd8
3 changed files with 13 additions and 7 deletions
+8 -2
View File
@@ -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:
+3 -3
View File
@@ -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))
+2 -2
View File
@@ -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):