FEGenStringNBus
This commit is contained in:
+12
-4
@@ -3,9 +3,12 @@ from .utils.ask_stream import openai_ask_background
|
||||
from .config.configs import config
|
||||
from .utils.any_hack import any
|
||||
from .utils.node_defs import FEAlwaysChangeNode
|
||||
from .utils.cached_fn_result import CachedFunctionResult
|
||||
from .utils.job import Job
|
||||
|
||||
|
||||
class FEGenStringNBus(FEAlwaysChangeNode):
|
||||
global_cached = CachedFunctionResult()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -14,6 +17,7 @@ class FEGenStringNBus(FEAlwaysChangeNode):
|
||||
"pkws": (any,),
|
||||
"api": ("STRING", {"default": config["nbus_api"], "multiline": False}),
|
||||
"model": ("STRING", {"default": "sd_axl3plus_v2", "multiline": False}),
|
||||
"cache_charge": ("BOOLEAN", {"default": True})
|
||||
},
|
||||
"optional": {
|
||||
"background": (any,),
|
||||
@@ -27,7 +31,7 @@ class FEGenStringNBus(FEAlwaysChangeNode):
|
||||
CATEGORY = CATE_GEN
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def query(self, pkws, api, model, background=None):
|
||||
def query(self, pkws, api, model, cache_charge: bool, background=None):
|
||||
if background is None or not background:
|
||||
background = {}
|
||||
if isinstance(background, list):
|
||||
@@ -38,9 +42,13 @@ class FEGenStringNBus(FEAlwaysChangeNode):
|
||||
api = api[0]
|
||||
if isinstance(model, list):
|
||||
model = model[0]
|
||||
|
||||
if isinstance(cache_charge, list):
|
||||
cache_charge = cache_charge[0]
|
||||
print("generating prompt...", end="")
|
||||
result, pkw, _, _ = openai_ask_background(
|
||||
model_name=model, background=background, pkw_list=pkws, api=api)
|
||||
if cache_charge:
|
||||
fn = Job(target=openai_ask_background, model_name=model, background=background, pkw_list=pkws, api=api)
|
||||
result, pkw, _, _ = FEGenStringNBus.global_cached.queue(model, fn)
|
||||
else:
|
||||
result, pkw, _, _ = openai_ask_background(model_name=model, background=background, pkw_list=pkws, api=api)
|
||||
print("done")
|
||||
return (result, pkw,)
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
from .job import Job
|
||||
from typing import Callable, Union
|
||||
|
||||
|
||||
class CachedFunctionResult:
|
||||
def __init__(self):
|
||||
self.cache_box = {}
|
||||
|
||||
def run_fn(self, fn: Union[Callable, Job]):
|
||||
if isinstance(fn, Job):
|
||||
fn.start()
|
||||
return fn
|
||||
j = Job(target=fn)
|
||||
j.start()
|
||||
return j
|
||||
|
||||
def instant_run_fn(self, fn: Union[Callable, Job]):
|
||||
if isinstance(fn, Job):
|
||||
return fn._target(**fn._kwargs)
|
||||
return fn()
|
||||
|
||||
def queue(self, name: str, fn: Union[Callable, Job]):
|
||||
if name not in self.cache_box or len(self.cache_box[name]) == 0:
|
||||
self.cache_box.setdefault(name, [])
|
||||
self.cache_box[name].append(self.run_fn(fn))
|
||||
return self.instant_run_fn(fn)
|
||||
else:
|
||||
self.cache_box[name].append(self.run_fn(fn))
|
||||
return self.cache_box[name].pop(0).get_result(True)
|
||||
Reference in New Issue
Block a user