mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
device fixes
This commit is contained in:
+8
-1
@@ -11,7 +11,7 @@ import safetensors.torch
|
||||
from modules.merging.merge import merge_models
|
||||
from modules.merging.merge_utils import TRIPLE_METHODS
|
||||
|
||||
from modules import shared, images, sd_models, sd_vae, sd_models_config
|
||||
from modules import shared, images, sd_models, sd_vae, sd_models_config, devices
|
||||
|
||||
checkpoint_dict_skip_on_merge = ["cond_stage_model.transformer.text_model.embeddings.position_ids"]
|
||||
|
||||
@@ -313,7 +313,14 @@ def run_MEHmodelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
kwargs.pop("beta_preset", None)
|
||||
|
||||
if kwargs["device"] == "cuda":
|
||||
kwargs["device"] == devices.device
|
||||
sd_models.unload_model_weights()
|
||||
elif kwargs["device"] == "shuffle":
|
||||
kwargs["device"] = torch.device("cpu")
|
||||
kwargs["work_device"] = devices.device
|
||||
else:
|
||||
kwargs["device"] = torch.device("cpu")
|
||||
|
||||
|
||||
try:
|
||||
theta_0 = merge_models(**kwargs)
|
||||
|
||||
@@ -160,7 +160,6 @@ def create_ui():
|
||||
def sd_model_choices():
|
||||
return ['None'] + sd_models.checkpoint_tiles()
|
||||
|
||||
|
||||
with gr.Row(equal_height=False):
|
||||
with gr.Column(variant='compact'):
|
||||
with FormRow():
|
||||
@@ -291,9 +290,6 @@ def create_ui():
|
||||
for key in list(kwargs.keys()):
|
||||
if kwargs[key] in [None, "None", "", 0, []]:
|
||||
del kwargs[key]
|
||||
if kwargs["device"] == "shuffle":
|
||||
kwargs["device"] = "cpu"
|
||||
kwargs["work_device"] = "cuda"
|
||||
|
||||
try:
|
||||
results = extras.run_MEHmodelmerger(dummy_component, **kwargs)
|
||||
@@ -332,7 +328,6 @@ def create_ui():
|
||||
else:
|
||||
return gr.Slider.update(value=None, visible=False)
|
||||
|
||||
|
||||
def load_presets(presets, ratio):
|
||||
for i, p in enumerate(presets):
|
||||
presets[i] = BLOCK_WEIGHTS_PRESETS[p]
|
||||
|
||||
Reference in New Issue
Block a user