mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
update extensions
This commit is contained in:
+3
-3
@@ -132,14 +132,14 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_
|
||||
if theta_func2:
|
||||
shared.state.textinfo = f"Loading B"
|
||||
print(f"Loading {secondary_model_info.filename}...")
|
||||
theta_1 = sd_models.read_state_dict(secondary_model_info.filename, map_location='cpu')
|
||||
theta_1 = sd_models.read_state_dict(secondary_model_info.filename)
|
||||
else:
|
||||
theta_1 = None
|
||||
|
||||
if theta_func1:
|
||||
shared.state.textinfo = f"Loading C"
|
||||
print(f"Loading {tertiary_model_info.filename}...")
|
||||
theta_2 = sd_models.read_state_dict(tertiary_model_info.filename, map_location='cpu')
|
||||
theta_2 = sd_models.read_state_dict(tertiary_model_info.filename)
|
||||
|
||||
shared.state.textinfo = 'Merging B and C'
|
||||
shared.state.sampling_steps = len(theta_1.keys())
|
||||
@@ -161,7 +161,7 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_
|
||||
|
||||
shared.state.textinfo = f"Loading {primary_model_info.filename}..."
|
||||
print(f"Loading {primary_model_info.filename}...")
|
||||
theta_0 = sd_models.read_state_dict(primary_model_info.filename, map_location='cpu')
|
||||
theta_0 = sd_models.read_state_dict(primary_model_info.filename)
|
||||
|
||||
print("Merging...")
|
||||
shared.state.textinfo = 'Merging A and B'
|
||||
|
||||
@@ -243,13 +243,10 @@ def read_state_dict(checkpoint_file):
|
||||
with rich.progress.open(checkpoint_file, 'rb') as f:
|
||||
if extension.lower() == ".safetensors":
|
||||
buffer = f.read()
|
||||
# load function always loads to cpu anyhow and model gets moved as needed
|
||||
pl_sd = safetensors.torch.load(buffer)
|
||||
elif extension.lower() == ".ckpt":
|
||||
buffer = io.BytesIO(f.read())
|
||||
pl_sd = torch.load(buffer)
|
||||
else:
|
||||
raise Exception(f"Unknown model type: {extension}")
|
||||
buffer = io.BytesIO(f.read())
|
||||
pl_sd = torch.load(buffer, map_location='cpu')
|
||||
|
||||
sd = get_state_dict_from_checkpoint(pl_sd)
|
||||
return sd
|
||||
|
||||
Reference in New Issue
Block a user