From 3cd435081437f8c137203521da6cd5b9a37ec021 Mon Sep 17 00:00:00 2001 From: hnmr293 Date: Mon, 3 Apr 2023 23:49:59 +0900 Subject: [PATCH] improve type inference --- model/loader.py | 5 ++++- model/merge.py | 2 +- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/model/loader.py b/model/loader.py index ccbd3c5..b2c2f91 100644 --- a/model/loader.py +++ b/model/loader.py @@ -51,7 +51,10 @@ class Dict2Model: setattr(utils, 'load_torch_file', load_torch_file_hook) try: - return sd.load_checkpoint(config_path, None, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) + model, clip, vae = sd.load_checkpoint(config_path, None, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) + assert clip is not None + assert vae is not None + return (model, clip, vae) finally: setattr(sd, 'load_torch_file', load_torch_file_org) diff --git a/model/merge.py b/model/merge.py index 5e6bef8..e244d37 100644 --- a/model/merge.py +++ b/model/merge.py @@ -11,7 +11,7 @@ def merge( half: str, ignore_keys_only_in_B: bool = False, ): - result = dict() + result: Dict[str,torch.Tensor] = dict() for key in tqdm.tqdm(model_A.keys()): if key not in model_B: print(f' key {key} is found in model_A but not model_B')