Only unload other models if unload is checked

This commit is contained in:
Zuellni
2024-02-17 16:53:00 +01:00
parent 8c4b9d058f
commit 217304e0ef
+6 -5
View File
@@ -61,8 +61,6 @@ class Loader:
return (self,) return (self,)
def load(self): def load(self):
unload_all_models()
if self.ckpt and self.cache and self.tokenizer and self.generator: if self.ckpt and self.cache and self.tokenizer and self.generator:
return return
@@ -151,6 +149,9 @@ class Generator:
if not text: if not text:
return ("",) return ("",)
if unload:
unload_all_models()
model.load() model.load()
input = model.tokenizer.encode(text, encode_special_tokens=True) input = model.tokenizer.encode(text, encode_special_tokens=True)
input_len = input.shape[-1] input_len = input.shape[-1]
@@ -198,9 +199,6 @@ class Generator:
f"({input_len} context, {tokens} tokens, {speed}t/s)", f"({input_len} context, {tokens} tokens, {speed}t/s)",
) )
if unload:
model.unload()
if id and info and "workflow" in info: if id and info and "workflow" in info:
nodes = info["workflow"]["nodes"] nodes = info["workflow"]["nodes"]
node = next((n for n in nodes if str(n["id"]) == id), None) node = next((n for n in nodes if str(n["id"]) == id), None)
@@ -208,6 +206,9 @@ class Generator:
if node: if node:
node["widgets_values"] = [output] node["widgets_values"] = [output]
if unload:
model.unload()
return (output,) return (output,)