mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] add dummy worker
This commit is contained in:
parent
ee620a01bb
commit
0505b04781
17 changed files with 428 additions and 115 deletions
|
|
@ -3,34 +3,48 @@ global_config:
|
|||
max_workers: 5
|
||||
dash_scope_apikey:
|
||||
open_ai_apikey:
|
||||
chat_list:
|
||||
- memory_chat
|
||||
memory_chat:
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
memory_chat_service:
|
||||
class: memory.base_memory_service
|
||||
history_msg_count: 5
|
||||
memory_operations:
|
||||
- name: read_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank
|
||||
work_type: frontend
|
||||
- name: list_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: frontend
|
||||
- name: extract_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 60
|
||||
min_count: 5
|
||||
- name: reflect_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 300
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.base_memory_service
|
||||
history_msg_count: 5
|
||||
memory_operations:
|
||||
read_memory:
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank
|
||||
work_type: frontend
|
||||
list_memory:
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: frontend
|
||||
extract_memory:
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 60
|
||||
min_count: 5
|
||||
reflect_memory:
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 300
|
||||
models:
|
||||
dashscope_generation:
|
||||
clazz: models.llama_index_generation_model
|
||||
module_name: DashScope
|
||||
model_name: qwen-max
|
||||
dashscope_embedding:
|
||||
clazz: models.base_embedding_model
|
||||
module_name: DashScopeEmbedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
clazz: models.base_rank_model
|
||||
module_name: DashScopeRerank
|
||||
model_name: gte-rerank
|
||||
vector_store:
|
||||
clazz: storage.base_vector_store
|
||||
index_name: memory_test
|
||||
|
|
@ -39,21 +53,8 @@ monitor:
|
|||
clazz: storage.base_monitor
|
||||
index_name: memory_test
|
||||
workers:
|
||||
- name: update_insight
|
||||
update_insight:
|
||||
clazz: worker.summary_long.update_insight
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
models:
|
||||
- name: dashscope_generation
|
||||
clazz: models.llama_index_generation_model
|
||||
module_name: DashScope
|
||||
model_name: qwen-max
|
||||
- name: dashscope_embedding
|
||||
clazz: models.base_embedding_model
|
||||
module_name: DashScopeEmbedding
|
||||
model_name: text-embedding-v2
|
||||
- name: dashscope_rank
|
||||
clazz: models.base_rank_model
|
||||
module_name: DashScopeRerank
|
||||
model_name: gte-rerank
|
||||
rank_model: dashscope_rank
|
||||
|
|
@ -2,7 +2,7 @@ from abc import ABCMeta, abstractmethod
|
|||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
def __init__(self, memory_service: str, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5,15 +5,20 @@ import pydantic
|
|||
|
||||
from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from memory_scope.memory.base_memory_service import BaseMemoryService
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class GlobalContext(pydantic.BaseModel):
|
||||
global_config: Dict[str, Any] = pydantic.Field({}, description="global configs")
|
||||
model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="global model_dict")
|
||||
memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="global memory_chat_dict")
|
||||
global_config: Dict[str, Any] = pydantic.Field({}, description="global config")
|
||||
worker_config: Dict[str, Any] = pydantic.Field({}, description="worker config")
|
||||
|
||||
memory_service_dict: Dict[str, BaseMemoryService] = pydantic.Field({}, description="memory_service dict")
|
||||
model_dict: Dict[str, BaseModel] = pydantic.Field({}, description="model dict")
|
||||
memory_chat_dict: Dict[str, BaseMemoryChat] = pydantic.Field({}, description="memory_chat dict")
|
||||
|
||||
vector_store: BaseVectorStore | None = pydantic.Field(None, description="global vector_store")
|
||||
monitor: BaseMonitor | None = pydantic.Field(None, description="global monitor")
|
||||
thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool")
|
||||
|
|
|
|||
|
|
@ -1,5 +1,3 @@
|
|||
import json
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
|
|
@ -7,12 +5,8 @@ import yaml
|
|||
|
||||
from chat_v2.global_context import G_CONTEXT
|
||||
from enumeration.language_enum import LanguageEnum
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from utils.logger import Logger
|
||||
from utils.tool_functions import (
|
||||
complete_config_name,
|
||||
init_instance_by_config,
|
||||
)
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class CliJob(object):
|
||||
|
|
@ -20,82 +14,54 @@ class CliJob(object):
|
|||
def __init__(self, config_path: str, config_suffix: str = ".yaml"):
|
||||
self.config_path: str = config_path
|
||||
self.config_suffix: str = config_suffix
|
||||
|
||||
self.config: Dict[str, Any] = {}
|
||||
self.global_config: Dict[str, Any] = {}
|
||||
|
||||
self.logger: Logger = Logger.get_logger("memory_chat")
|
||||
|
||||
def init_model(self, model_name: str):
|
||||
if not model_name or model_name in G_CONTEXT.model_dict:
|
||||
return
|
||||
|
||||
with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f:
|
||||
model_config = json.load(f)
|
||||
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config)
|
||||
|
||||
def init_workers(self):
|
||||
"""load worker config & init workers"""
|
||||
worker_config_name: str = self.config["workers"]
|
||||
with open(
|
||||
os.path.join(self.config_base_dir, complete_config_name(worker_config_name))
|
||||
) as f:
|
||||
worker_config_dict = json.load(f)
|
||||
|
||||
for worker_name, worker_config in worker_config_dict.items():
|
||||
if worker_name not in self.worker_chat_dict:
|
||||
continue
|
||||
|
||||
chat_name_list = self.worker_chat_dict[worker_name]
|
||||
for chat_name in chat_name_list:
|
||||
if chat_name not in GLOBAL_CONTEXT.worker_dict:
|
||||
GLOBAL_CONTEXT.worker_dict[chat_name] = {}
|
||||
GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = (
|
||||
init_instance_by_config(
|
||||
worker_config,
|
||||
suffix_name="worker",
|
||||
**GLOBAL_CONTEXT.global_configs,
|
||||
)
|
||||
)
|
||||
|
||||
self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value))
|
||||
self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value))
|
||||
self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value))
|
||||
self.logger: Logger = Logger.get_logger("cli_job")
|
||||
|
||||
@staticmethod
|
||||
def set_global_config():
|
||||
# TODO at sen, set global_configs & set apikey into env
|
||||
G_CONTEXT.language = LanguageEnum(G_CONTEXT.global_configs["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(G_CONTEXT.global_configs["max_workers"]))
|
||||
def set_global_config(global_config: Dict[str, Any]):
|
||||
""" set global_configs & set apikey into env
|
||||
:return:
|
||||
TODO at sen
|
||||
"""
|
||||
G_CONTEXT.global_config = global_config
|
||||
G_CONTEXT.language = LanguageEnum(global_config["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
|
||||
|
||||
def init_global_content_by_config(self):
|
||||
# load config
|
||||
config_path = self.config_path
|
||||
if not self.config_path.endswith(self.config_suffix):
|
||||
config_path += self.config_suffix
|
||||
|
||||
with open(config_path) as f:
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
|
||||
G_CONTEXT.global_configs = self.global_config = self.config["global_configs"]
|
||||
self.set_global_config()
|
||||
# set global_config
|
||||
self.set_global_config(self.config["global_config"])
|
||||
|
||||
# init memory_chat
|
||||
for chat_name in self.global_config["chat_list"]:
|
||||
memory_chat_config = self.config[chat_name]
|
||||
G_CONTEXT.memory_chat_dict[chat_name] = init_instance_by_config(memory_chat_config, chat_name=chat_name)
|
||||
for name, conf in self.config["memory_chat"].items():
|
||||
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
for model_config in
|
||||
# set memory_service
|
||||
for name, conf in self.config["memory_service"].items():
|
||||
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config)
|
||||
# init models
|
||||
for name, conf in self.config["models"].items():
|
||||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# TODO no db and monitor now
|
||||
GLOBAL_CONTEXT.vector_store = init_instance_by_config(
|
||||
self.config["vector_store"]
|
||||
)
|
||||
GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
# init vector_store
|
||||
G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"])
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["workers"]
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
with GLOBAL_CONTEXT.thread_pool:
|
||||
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
|
||||
with G_CONTEXT.thread_pool:
|
||||
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
memory_chat.run()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
RESULT = "result"
|
||||
|
||||
CHAT_MESSAGES = "chat_messages"
|
||||
|
||||
RELATED_MEMORIES = "related_memories"
|
||||
|
||||
MESSAGES = "messages"
|
||||
|
|
|
|||
6
memory_scope/memory/base_memory_service.py
Normal file
6
memory_scope/memory/base_memory_service.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from abc import ABCMeta
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
0
memory_scope/memory/worker/__init__.py
Normal file
0
memory_scope/memory/worker/__init__.py
Normal file
73
memory_scope/memory/worker/base_worker.py
Normal file
73
memory_scope/memory/worker/base_worker.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
from typing import Any, Dict
|
||||
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
||||
def __init__(self, raise_exception: bool = True, **kwargs):
|
||||
super(BaseWorker, self).__init__(**kwargs)
|
||||
# 异常是否继续执行
|
||||
self.raise_exception: bool = raise_exception
|
||||
|
||||
# True 为正常运行,False会结束整个pipeline
|
||||
self.continue_run: bool = True
|
||||
|
||||
# 短name
|
||||
self._name_simple: str = ""
|
||||
|
||||
# 是否多线程环境
|
||||
self.is_multi_thread: bool = False
|
||||
|
||||
# pipeline 上下文
|
||||
self.context_dict: Dict[str, Any] | None = None
|
||||
self.context_lock = None
|
||||
|
||||
# 日志
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
# worker 参数
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
def _run(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
self.logger.info(f"----- Begin {self.name_simple} -----")
|
||||
with Timer(self.name_simple, log_time=False) as t:
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
else:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as e:
|
||||
self.logger.exception(f"run {self.name_simple} failed! args={e.args}")
|
||||
|
||||
self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----")
|
||||
|
||||
def set_context_dict(self, context_dict: dict, context_lock=None):
|
||||
self.context_dict = context_dict
|
||||
if context_lock is not None:
|
||||
self.context_lock = context_lock
|
||||
self.is_multi_thread = True
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context_dict.get(key, default)
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
# add lock to multi thread
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
else:
|
||||
self.context_dict[key] = value
|
||||
|
||||
def __getattr__(self, key):
|
||||
return self.kwargs[key]
|
||||
|
||||
@property
|
||||
def name_simple(self) -> str:
|
||||
if not self._name_simple:
|
||||
self._name_simple = self.__class__.__name__.replace("Worker", "")
|
||||
return self._name_simple
|
||||
6
memory_scope/memory/worker/dummy_worker.py
Normal file
6
memory_scope/memory/worker/dummy_worker.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
pass
|
||||
70
memory_scope/memory/worker/memory_base_worker.py
Normal file
70
memory_scope/memory/worker/memory_base_worker.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
from typing import List
|
||||
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from models.base_model import BaseModel
|
||||
from scheme.message import Message
|
||||
from storage.base_monitor import BaseMonitor
|
||||
from storage.base_vector_store import BaseVectorStore
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
def __init__(self,
|
||||
embedding_model: str,
|
||||
generation_model: str,
|
||||
rank_model: str,
|
||||
**kwargs):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
self.embedding_model_name: str = embedding_model
|
||||
self.generation_model_name: str = generation_model
|
||||
self.rank_model_name: str = rank_model
|
||||
|
||||
self._embedding_model: BaseModel | None = None
|
||||
self._generation_model: BaseModel | None = None
|
||||
self._rank_model: BaseModel | None = None
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
@property
|
||||
def messages(self) -> List[Message]:
|
||||
return self.get_context(MESSAGES)
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value):
|
||||
self.set_context(MESSAGES, value)
|
||||
|
||||
@property
|
||||
def chat_name(self):
|
||||
return self.get_context(CHAT_NAME)
|
||||
|
||||
@property
|
||||
def embedding_model(self):
|
||||
if self._embedding_model is None:
|
||||
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self):
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self):
|
||||
if self._rank_model is None:
|
||||
self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name)
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self):
|
||||
if self._vector_store is None:
|
||||
self._vector_store = GLOBAL_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
||||
@property
|
||||
def monitor(self):
|
||||
if self._monitor is None:
|
||||
self._monitor = GLOBAL_CONTEXT.monitor
|
||||
return self._monitor
|
||||
0
memory_scope/memory/workflow/__init__.py
Normal file
0
memory_scope/memory/workflow/__init__.py
Normal file
43
memory_scope/memory/workflow/backend_v1_workflow.py
Normal file
43
memory_scope/memory/workflow/backend_v1_workflow.py
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
import time
|
||||
|
||||
from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES
|
||||
from memory_scope.memory.workflow.base_workflow import BaseWorkflow
|
||||
|
||||
|
||||
class BackendV1Workflow(BaseWorkflow):
|
||||
|
||||
def __init__(self, interval_time: int, min_count: int, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.interval_time: int = interval_time
|
||||
self.min_count: int = min_count
|
||||
|
||||
@property
|
||||
def not_memorized_size(self):
|
||||
return sum([not x.memorized for x in self.chat_messages])
|
||||
|
||||
def set_memorized(self):
|
||||
for msg in self.chat_messages:
|
||||
msg.memorized = True
|
||||
|
||||
def _loop(self):
|
||||
while self.loop_switch:
|
||||
time.sleep(self.interval_time)
|
||||
if self.not_memorized_size < self.min_count:
|
||||
continue
|
||||
|
||||
self.context[CHAT_MESSAGES] = self.chat_messages
|
||||
self.__call__()
|
||||
self.context.clear()
|
||||
self.set_memorized()
|
||||
|
||||
def start_loop_run(self):
|
||||
if not self.loop_switch:
|
||||
self.loop_switch = True
|
||||
return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop)
|
||||
|
||||
def run_workflow(self):
|
||||
self.context[CHAT_MESSAGES] = self.chat_messages
|
||||
self.__call__()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
return result
|
||||
126
memory_scope/memory/workflow/base_workflow.py
Normal file
126
memory_scope/memory/workflow/base_workflow.py
Normal file
|
|
@ -0,0 +1,126 @@
|
|||
import re
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from memory_scope.chat_v2.global_context import G_CONTEXT
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class BaseWorkflow(object):
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
workflow: str,
|
||||
thread_pool: ThreadPoolExecutor,
|
||||
chat_messages: List[Message],
|
||||
max_history_message_count: int,
|
||||
**kwargs):
|
||||
|
||||
self.name: str = name
|
||||
self.workflow: str = workflow
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
self.max_history_message_count: int = max_history_message_count
|
||||
self.kwargs = kwargs
|
||||
|
||||
self.workflow_worker_list: List[List[List[str]]] = []
|
||||
self.worker_dict: Dict[str, BaseWorker | bool] = {}
|
||||
self.context: Dict[str, Any] = {}
|
||||
self.context_lock = threading.Lock()
|
||||
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
if self.workflow:
|
||||
self._parse_workflow()
|
||||
self._print_workflow()
|
||||
|
||||
def _parse_workflow(self):
|
||||
# re-match e.g., [a|b],c,[d,e,f|g,h],j
|
||||
pattern = r'(\[[^\]]*\]|[^,]+)'
|
||||
workflow_split = re.findall(pattern, self.workflow)
|
||||
for workflow_part in workflow_split:
|
||||
# e.g., [d,e,f|g,h]
|
||||
workflow_part = workflow_part.strip()
|
||||
if '[' in workflow_part or ']' in workflow_part:
|
||||
workflow_part = workflow_part.replace('[', '').replace(']', '')
|
||||
|
||||
# e.g., ["d,e,f", "g,h"]
|
||||
line_split = [x.strip() for x in workflow_part.split("|") if x]
|
||||
if len(line_split) <= 0:
|
||||
continue
|
||||
|
||||
# is under multi thread cond
|
||||
is_multi_thread: bool = len(line_split) > 1
|
||||
|
||||
# e.g., ["d","e","f"]
|
||||
line_split_split: List[List[str]] = []
|
||||
for sub_line_split in line_split:
|
||||
sub_split = [x.strip() for x in sub_line_split.split(",")]
|
||||
line_split_split.append(sub_split)
|
||||
# add workers
|
||||
for sub_item in sub_split:
|
||||
self.worker_dict[sub_item] = is_multi_thread
|
||||
self.workflow_worker_list.append(line_split_split)
|
||||
|
||||
def _print_workflow(self):
|
||||
self.logger.info(f"----- print_workflow_{self.name}_begin -----")
|
||||
i: int = 0
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
for w in workflow_part[0]:
|
||||
self.logger.info(f"stage{i}: {w}")
|
||||
i += 1
|
||||
else:
|
||||
for w_zip in zip_longest(*workflow_part, fillvalue="-"):
|
||||
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
|
||||
i += 1
|
||||
for w in w_zip:
|
||||
if w == "-":
|
||||
continue
|
||||
self.logger.info(f"----- print_workflow_{self.name}_end -----")
|
||||
|
||||
def init_workers(self):
|
||||
for name in list(self.worker_dict.keys()):
|
||||
if name not in G_CONTEXT.worker_config:
|
||||
raise RuntimeError(f"worker={name} is not exists in worker_config!")
|
||||
|
||||
self.worker_dict[name] = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
name=name,
|
||||
is_multi_thread=self.worker_dict[name],
|
||||
context=self.context,
|
||||
context_lock=self.context_lock)
|
||||
|
||||
def _run_sub_workflow(self, worker_list: List[str]) -> bool:
|
||||
for name in worker_list:
|
||||
worker = self.worker_dict[name]
|
||||
worker.run()
|
||||
if not worker.continue_run:
|
||||
return False
|
||||
return True
|
||||
|
||||
def run_workflow(self):
|
||||
with Timer(f"run_workflow_{self.name}"):
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
if not self._run_sub_workflow(workflow_part[0]):
|
||||
break
|
||||
else:
|
||||
t_list = []
|
||||
for sub_workflow in workflow_part:
|
||||
t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow))
|
||||
|
||||
flag = True
|
||||
for future in as_completed(t_list):
|
||||
if not future.result():
|
||||
flag = False
|
||||
break
|
||||
if not flag:
|
||||
break
|
||||
12
memory_scope/memory/workflow/frontend_workflow.py
Normal file
12
memory_scope/memory/workflow/frontend_workflow.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES
|
||||
from memory_scope.memory.workflow.base_workflow import BaseWorkflow
|
||||
|
||||
|
||||
class FrontendWorkflow(BaseWorkflow):
|
||||
|
||||
def run_workflow(self):
|
||||
self.context[CHAT_MESSAGES] = self.chat_messages[:1 + self.max_history_message_count]
|
||||
self.__call__()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
return result
|
||||
|
|
@ -1,4 +1,3 @@
|
|||
from utils.registry import Registry
|
||||
from memory_scope.utils.registry import Registry
|
||||
|
||||
# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"]
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import time
|
|||
from abc import abstractmethod, ABCMeta
|
||||
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from . import MODEL_REGISTRY
|
||||
from memory_scope.models import MODEL_REGISTRY
|
||||
from .response import ModelResponse, ModelResponseGen
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
|
|
|
|||
|
|
@ -6,4 +6,6 @@ class Message(BaseModel):
|
|||
|
||||
content: str = Field(..., description="The body of the message")
|
||||
|
||||
time_created: int = Field("", description="Timestamp when the message was created")
|
||||
time_created: int = Field(..., description="Timestamp when the message was created")
|
||||
|
||||
memorized: bool = Field(False, description="indicate whether message is memorized")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue