Only unload other models if unload is checked
This commit is contained in:
+6
-5
@@ -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,)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user