From a8ccf028756bd95f28e8fb1eace381e337a1ee04 Mon Sep 17 00:00:00 2001 From: Lawrence Ling <5894666+Lawrr@users.noreply.github.com> Date: Thu, 30 May 2024 21:07:38 +1000 Subject: [PATCH] Improve image loader validation to only consider model file --- py/better_combos.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/py/better_combos.py b/py/better_combos.py index c3f1c36..46f3efa 100644 --- a/py/better_combos.py +++ b/py/better_combos.py @@ -106,6 +106,17 @@ class LoraLoaderWithImages(LoraLoader): populate_items(names, "loras") return types + @classmethod + def VALIDATE_INPUTS(s, lora_name): + types = super().INPUT_TYPES() + names = types["required"]["lora_name"][0] + + name = lora_name["content"] + if name in names: + return True + else: + return f"Lora not found: {name}" + def load_lora(self, **kwargs): kwargs["lora_name"] = kwargs["lora_name"]["content"] return super().load_lora(**kwargs) @@ -119,6 +130,17 @@ class CheckpointLoaderSimpleWithImages(CheckpointLoaderSimple): populate_items(names, "checkpoints") return types + @classmethod + def VALIDATE_INPUTS(s, ckpt_name): + types = super().INPUT_TYPES() + names = types["required"]["ckpt_name"][0] + + name = ckpt_name["content"] + if name in names: + return True + else: + return f"Checkpoint not found: {name}" + def load_checkpoint(self, **kwargs): kwargs["ckpt_name"] = kwargs["ckpt_name"]["content"] return super().load_checkpoint(**kwargs)