Add torch.device

This commit is contained in:
AI-Casanova
2023-11-17 20:35:06 -06:00
parent abb4739590
commit cb8be0fc6e
3 changed files with 23 additions and 26 deletions
+1 -1
View File
@@ -311,7 +311,7 @@ def run_MEHmodelmerger(id_task, **kwargs): # pylint: disable=unused-argument
kwargs.pop("beta_mid_block", None)
kwargs.pop("beta_out_blocks", None)
kwargs.pop("beta_preset", None)
if kwargs["device"] == "cuda":
sd_models.unload_model_weights()
+20 -21
View File
@@ -1,4 +1,3 @@
import gc
import os
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
@@ -10,6 +9,7 @@ import torch
from tqdm import tqdm
import modules.memstats
import modules.devices as devices
from modules.shared import log
from modules.merging import merge_methods
from modules.merging.merge_utils import WeightClass
@@ -69,7 +69,7 @@ def fix_model(model: Dict) -> Dict:
return fix_clip(model)
def load_sd_model(model: os.PathLike | str, device: str = "cpu") -> Dict:
def load_sd_model(model: os.PathLike | str, device: torch.device = None) -> Dict:
if isinstance(model, str):
model = Path(model)
@@ -102,7 +102,7 @@ def log_vram(txt=""):
def load_thetas(
models: Dict[str, os.PathLike | str],
prune: bool,
device: str,
device: torch.device,
precision: str,
) -> Dict:
log_vram("before loading models")
@@ -111,13 +111,12 @@ def load_thetas(
else:
thetas = {k: load_sd_model(m, device) for k, m in models.items()}
if device == "cuda":
for model_key, model in thetas.items():
for key, block in model.items():
if precision == "fp16":
thetas[model_key].update({key: block.to(device).half()})
else:
thetas[model_key].update({key: block.to(device)})
for model_key, model in thetas.items():
for key, block in model.items():
if precision == "fp16":
thetas[model_key].update({key: block.to(device).half()})
else:
thetas[model_key].update({key: block.to(device)})
log_vram("models loaded")
return thetas
@@ -129,8 +128,8 @@ def merge_models(
precision: str = "fp16",
weights_clip: bool = False,
re_basin: bool = False,
device: str = "cpu",
work_device: Optional[str] = None,
device: torch.device = None,
work_device: torch.device = None,
prune: bool = False,
threads: int = 1,
**kwargs,
@@ -178,7 +177,7 @@ def un_prune_model(
if prune:
log.info("Un-pruning merged model")
del thetas
gc.collect()
devices.torch_gc(force=True)
log_vram("remove thetas")
original_a = load_sd_model(models["model_a"], device)
for key in tqdm(original_a.keys(), desc="un-prune model a"):
@@ -189,7 +188,7 @@ def un_prune_model(
if precision == "fp16":
merged.update({key: merged[key].half()})
del original_a
gc.collect()
devices.torch_gc(force=True)
# log_vram("remove original_a")
original_b = load_sd_model(models["model_b"], device)
for key in tqdm(original_b.keys(), desc="un-prune model b"):
@@ -210,8 +209,8 @@ def simple_merge(
merge_mode: str,
precision: str = "fp16",
weights_clip: bool = False,
device: str = "cpu",
work_device: Optional[str] = None,
device: torch.device = None,
work_device: torch.device = None,
threads: int = 1,
) -> Dict:
futures = []
@@ -257,8 +256,8 @@ def rebasin_merge(
precision: str = "fp16",
weights_clip: bool = False,
iterations: int = 1,
device="cpu",
work_device=None,
device: torch.device = None,
work_device: torch.device = None,
threads: int = 1,
):
# WARNING: not sure how this does when 3 models are involved...
@@ -345,8 +344,8 @@ def merge_key(
merge_mode: str,
precision: str = "fp16",
weights_clip: bool = False,
device: str = "cpu",
work_device: Optional[str] = None,
device: torch.device = None,
work_device: torch.device = None,
) -> Optional[Tuple[str, Dict]]:
if work_device is None:
work_device = device
@@ -413,7 +412,7 @@ def get_merge_method_args(
current_bases: Dict,
thetas: Dict,
key: str,
work_device: str,
work_device: torch.device,
) -> Dict:
merge_method_args = {
"a": thetas["model_a"][key].to(work_device),
+2 -4
View File
@@ -1,4 +1,3 @@
# import logging
import os
from dataclasses import dataclass
@@ -6,17 +5,16 @@ import safetensors
import torch
from tensordict import TensorDict
# logging.getLogger("sd_meh").addHandler(logging.NullHandler())
@dataclass
class SDModel:
model_path: os.PathLike
device: str
device: torch.device
def load_model(self):
# logging.info(f"Loading: {self.model_path}")
if self.model_path.suffix == ".safetensors":
if os.path.splitext(self.model_path) == ".safetensors":
ckpt = safetensors.torch.load_file(
self.model_path,
device=self.device,