[dev] add dummy worker

This commit is contained in:
jinli.yl 2024-06-26 10:07:45 +08:00
parent ee620a01bb
commit 0505b04781
17 changed files with 428 additions and 115 deletions

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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()

View file

@ -1,3 +1,7 @@
RESULT = "result"
CHAT_MESSAGES = "chat_messages"
RELATED_MEMORIES = "related_memories"
MESSAGES = "messages"

View file

@ -0,0 +1,6 @@
from abc import ABCMeta
class BaseMemoryService(metaclass=ABCMeta):
def __init__(self, **kwargs):
self.kwargs = kwargs

View file

View 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

View file

@ -0,0 +1,6 @@
from memory_base_worker import MemoryBaseWorker
class DummyWorker(MemoryBaseWorker):
def _run(self):
pass

View 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

View file

View 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

View 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

View 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

View file

@ -1,4 +1,3 @@
from utils.registry import Registry
from memory_scope.utils.registry import Registry
# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"]
MODEL_REGISTRY = Registry("models")

View file

@ -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

View file

@ -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")