From 4fdebee1573c929fd2332ade1c6e5d16096dd02f Mon Sep 17 00:00:00 2001 From: fexli Date: Wed, 13 Mar 2024 14:40:42 +0800 Subject: [PATCH] FEGenStringNBus --- FEGenStringNBus.py | 16 ++++++++++++---- utils/cached_fn_result.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 4 deletions(-) create mode 100644 utils/cached_fn_result.py diff --git a/FEGenStringNBus.py b/FEGenStringNBus.py index 5972253..7d52d31 100644 --- a/FEGenStringNBus.py +++ b/FEGenStringNBus.py @@ -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,) diff --git a/utils/cached_fn_result.py b/utils/cached_fn_result.py new file mode 100644 index 0000000..2428770 --- /dev/null +++ b/utils/cached_fn_result.py @@ -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)