From 843ef15e1353849af0dd5bef38050c72a776e3b0 Mon Sep 17 00:00:00 2001 From: Gabriel Almeida Date: Fri, 7 Apr 2023 18:36:42 -0300 Subject: [PATCH] feat: added memoize to cache functions --- src/backend/langflow/api/endpoints.py | 4 +-- src/backend/langflow/cache/utils.py | 28 +++++++++++++++ src/backend/langflow/interface/run.py | 52 ++++++++++++++++++++++++++- 3 files changed, 81 insertions(+), 3 deletions(-) diff --git a/src/backend/langflow/api/endpoints.py b/src/backend/langflow/api/endpoints.py index 22f548156..b8290e691 100644 --- a/src/backend/langflow/api/endpoints.py +++ b/src/backend/langflow/api/endpoints.py @@ -3,7 +3,7 @@ from typing import Any, Dict from fastapi import APIRouter, HTTPException -from langflow.interface.run import process_graph +from langflow.interface.run import process_graph_cached from langflow.interface.types import build_langchain_types_dict # build router @@ -19,7 +19,7 @@ def get_all(): @router.post("/predict") def get_load(data: Dict[str, Any]): try: - return process_graph(data) + return process_graph_cached(data) except Exception as e: # Log stack trace logger.exception(e) diff --git a/src/backend/langflow/cache/utils.py b/src/backend/langflow/cache/utils.py index 3c416f4d7..f04f73761 100644 --- a/src/backend/langflow/cache/utils.py +++ b/src/backend/langflow/cache/utils.py @@ -6,6 +6,34 @@ import tempfile from pathlib import Path import dill # type: ignore +import functools +from collections import OrderedDict + + +def memoize(maxsize=128): + cache = OrderedDict() + + def decorator(func): + @functools.wraps(func) + def wrapper(*args, **kwargs): + key = (func.__name__, args, frozenset(kwargs.items())) + if key not in cache: + result = func(*args, **kwargs) + cache[key] = result + if len(cache) > maxsize: + cache.popitem(last=False) + else: + result = cache[key] + return result + + def clear_cache(): + cache.clear() + + wrapper.clear_cache = clear_cache + return wrapper + + return decorator + PREFIX = "langflow_cache" diff --git a/src/backend/langflow/interface/run.py b/src/backend/langflow/interface/run.py index 8f7765ef2..7b4a9e58c 100644 --- a/src/backend/langflow/interface/run.py +++ b/src/backend/langflow/interface/run.py @@ -2,7 +2,7 @@ import contextlib import io from typing import Any, Dict -from langflow.cache.utils import compute_hash, load_cache +from langflow.cache.utils import compute_hash, load_cache, memoize from langflow.graph.graph import Graph from langflow.interface import loading from langflow.utils.logger import logger @@ -22,6 +22,32 @@ def load_langchain_object(data_graph, is_first_message=False): return computed_hash, langchain_object +def load_or_build_langchain_object(data_graph, is_first_message=False): + """ + Load langchain object from cache if it exists, otherwise build it. + """ + if is_first_message: + build_langchain_object_with_caching.clear_cache() + return build_langchain_object_with_caching(data_graph) + + +@memoize(maxsize=1) +def build_langchain_object_with_caching(data_graph): + """ + Build langchain object from data_graph. + """ + + logger.debug("Building langchain object") + nodes = data_graph["nodes"] + # Add input variables + # nodes = payload.extract_input_variables(nodes) + # Nodes, edges and root node + edges = data_graph["edges"] + graph = Graph(nodes, edges) + + return graph.build() + + def build_langchain_object(data_graph): """ Build langchain object from data_graph. @@ -72,6 +98,30 @@ def process_graph(data_graph: Dict[str, Any]): return {"result": str(result), "thought": thought.strip()} +def process_graph_cached(data_graph: Dict[str, Any]): + """ + Process graph by extracting input variables and replacing ZeroShotPrompt + with PromptTemplate,then run the graph and return the result and thought. + """ + # Load langchain object + message = data_graph.pop("message", "") + is_first_message = len(data_graph.get("chatHistory", [])) == 0 + langchain_object = load_or_build_langchain_object(data_graph, is_first_message) + logger.debug("Loaded langchain object") + + if langchain_object is None: + # Raise user facing error + raise ValueError( + "There was an error loading the langchain_object. Please, check all the nodes and try again." + ) + + # Generate result and thought + logger.debug("Generating result and thought") + result, thought = get_result_and_thought_using_graph(langchain_object, message) + logger.debug("Generated result and thought") + return {"result": str(result), "thought": thought.strip()} + + def get_memory_key(langchain_object): """ Given a LangChain object, this function retrieves the current memory key from the object's memory attribute.