mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] modify g content
This commit is contained in:
parent
e5c89b4ca3
commit
ee620a01bb
10 changed files with 387 additions and 22 deletions
|
|
@ -1,14 +1,16 @@
|
|||
global_config:
|
||||
thread_pool_max_count: 5
|
||||
language: en
|
||||
max_workers: 5
|
||||
dash_scope_apikey:
|
||||
open_ai_apikey:
|
||||
language: en
|
||||
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
|
||||
|
|
|
|||
0
memory_scope/chat_v2/__init__.py
Normal file
0
memory_scope/chat_v2/__init__.py
Normal file
17
memory_scope/chat_v2/base_memory_chat.py
Normal file
17
memory_scope/chat_v2/base_memory_chat.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
"""
|
||||
:param query:
|
||||
:return:
|
||||
"""
|
||||
|
||||
def run(self):
|
||||
pass
|
||||
4
memory_scope/chat_v2/base_memory_service.py
Normal file
4
memory_scope/chat_v2/base_memory_service.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
class BaseMemoryService(object):
|
||||
def __init__(self, **kwargs):
|
||||
|
||||
self.kwargs = kwargs
|
||||
83
memory_scope/chat_v2/cli_memory_chat.py
Normal file
83
memory_scope/chat_v2/cli_memory_chat.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
import datetime
|
||||
|
||||
import questionary
|
||||
from rich.console import Console
|
||||
|
||||
from .memory_chat import MemoryChat
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from scheme.message import Message
|
||||
|
||||
|
||||
class CliMemoryChat(MemoryChat):
|
||||
|
||||
USER_COMMANDS = {
|
||||
"/exit": "exit the CLI",
|
||||
"/memory": "print the current contents of agent memory",
|
||||
"/retrieve": "retrieve related memory",
|
||||
"/log": "log chat progress",
|
||||
# TODO add more commands
|
||||
}
|
||||
|
||||
def chat_with_memory(self, query): # for testing
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
message = Message(
|
||||
role=MessageRoleEnum.USER, content=query, time_created=time_created
|
||||
)
|
||||
messages = [message]
|
||||
return self.generation_model.call(messages=messages, stream=True)
|
||||
|
||||
def retrieve_all(self): # for testing
|
||||
return "memory 1. 2. 3."
|
||||
|
||||
def run(self):
|
||||
console = Console()
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Try again!")
|
||||
continue
|
||||
|
||||
# Handle CLI commands
|
||||
if query.startswith("/"):
|
||||
if query.lower() == "/exit":
|
||||
break
|
||||
elif query.lower() == "/memory":
|
||||
console.print(self.memory_service.retrieve_all())
|
||||
elif query.lower() == "/help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
|
||||
continue
|
||||
|
||||
while True:
|
||||
try:
|
||||
# with console.status("[bold cyan]Thinking..."):
|
||||
for msg in self.chat_with_memory(query=query):
|
||||
console.print(msg.delta, end="")
|
||||
console.print()
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
console.print("User interrupt occurred.")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
except Exception as e:
|
||||
console.print(
|
||||
f"An exception occurred when running chat_with_memory(): {e}"
|
||||
)
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
23
memory_scope/chat_v2/global_context.py
Normal file
23
memory_scope/chat_v2/global_context.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
import pydantic
|
||||
|
||||
from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
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")
|
||||
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")
|
||||
language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en")
|
||||
|
||||
|
||||
G_CONTEXT = GlobalContext()
|
||||
67
memory_scope/chat_v2/memory_chat.py
Normal file
67
memory_scope/chat_v2/memory_chat.py
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
import datetime
|
||||
from typing import List
|
||||
|
||||
from .base_memory_chat import BaseMemoryChat
|
||||
from .global_context import GLOBAL_CONTEXT
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from models.base_model import BaseModel
|
||||
from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
|
||||
from scheme.message import Message
|
||||
from .memory_service import MemoryService
|
||||
|
||||
|
||||
class MemoryChat(BaseMemoryChat):
|
||||
|
||||
def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
|
||||
self.generation_model_name: str = generation_model
|
||||
self.history_msg_count: int = history_msg_count
|
||||
|
||||
self._generation_model: BaseModel | None = None
|
||||
self.history_message_list: List[Message] = []
|
||||
|
||||
@property
|
||||
def generation_model(self):
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict[
|
||||
self.generation_model_name
|
||||
]
|
||||
return self._generation_model
|
||||
|
||||
@staticmethod
|
||||
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language]
|
||||
if related_memories:
|
||||
memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language]
|
||||
system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories)
|
||||
return Message(
|
||||
role=MessageRoleEnum.SYSTEM,
|
||||
content=system_prompt.strip(),
|
||||
time_created=time_created,
|
||||
)
|
||||
|
||||
def chat_with_memory(self, query: str):
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
new_message: Message = Message(
|
||||
role=MessageRoleEnum.USER, content=query, time_created=time_created
|
||||
)
|
||||
related_memories: List[str] = self.memory_service.retrieve(message=new_message)
|
||||
system_message = self.get_system_prompt(related_memories, time_created)
|
||||
self.history_message_list.append(new_message)
|
||||
self.history_message_list = self.history_message_list[-self.history_msg_count :]
|
||||
all_messages = [system_message] + self.history_message_list
|
||||
# TODO at xian zhe
|
||||
return self.generation_model.call(messages=all_messages, stream=True)
|
||||
|
||||
def run(self):
|
||||
self.memory_service.start_memory_backend()
|
||||
while True:
|
||||
query = input("wait for input:")
|
||||
if query in ["stop", "停止"]:
|
||||
break
|
||||
self.chat_with_memory(query=query)
|
||||
70
memory_scope/chat_v2/memory_service.py
Normal file
70
memory_scope/chat_v2/memory_service.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
from constants.common_constants import RELATED_MEMORIES
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from scheme.message import Message
|
||||
from utils.pipeline import Pipeline
|
||||
from .base_memory_service import BaseMemoryService
|
||||
|
||||
|
||||
class MemoryService(BaseMemoryService):
|
||||
def __init__(
|
||||
self,
|
||||
chat_name: str,
|
||||
retrieve_pipeline: str,
|
||||
retrieve_all_pipeline: str,
|
||||
summary_short_pipeline: str,
|
||||
summary_long_pipeline: str,
|
||||
summary_short_interval_time: int = 60,
|
||||
summary_short_minimum_count: int = 5,
|
||||
summary_long_interval_time: int = 60 * 5,
|
||||
summary_long_minimum_count: int = 5 * 5,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.retrieve_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE,
|
||||
pipeline_str=retrieve_pipeline,
|
||||
)
|
||||
|
||||
self.retrieve_all_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
|
||||
pipeline_str=retrieve_all_pipeline,
|
||||
)
|
||||
|
||||
self.summary_short_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
|
||||
pipeline_str=summary_short_pipeline,
|
||||
loop_interval_time=summary_short_interval_time,
|
||||
loop_minimum_count=summary_short_minimum_count,
|
||||
)
|
||||
|
||||
self.summary_long_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
|
||||
pipeline_str=summary_long_pipeline,
|
||||
loop_interval_time=summary_long_interval_time,
|
||||
loop_minimum_count=summary_long_minimum_count,
|
||||
)
|
||||
|
||||
def retrieve(self, message: Message):
|
||||
self.retrieve_pipeline.submit_message(message, with_lock=False)
|
||||
self.summary_short_pipeline.submit_message(message)
|
||||
self.summary_long_pipeline.submit_message(message)
|
||||
return self.retrieve_pipeline.run(RELATED_MEMORIES)
|
||||
|
||||
def retrieve_all(self):
|
||||
return self.retrieve_all_pipeline.run(RELATED_MEMORIES)
|
||||
|
||||
def start_memory_backend(self):
|
||||
self.summary_short_pipeline.start_loop_run()
|
||||
self.summary_long_pipeline.start_loop_run()
|
||||
|
||||
def get_worker_list(self) -> list:
|
||||
worker_set = set()
|
||||
worker_set.update(self.retrieve_pipeline.worker_set)
|
||||
worker_set.update(self.retrieve_all_pipeline.worker_set)
|
||||
worker_set.update(self.summary_short_pipeline.worker_set)
|
||||
worker_set.update(self.summary_long_pipeline.worker_set)
|
||||
return sorted(worker_set)
|
||||
101
memory_scope/cli_job.py
Normal file
101
memory_scope/cli_job.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
import json
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
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))
|
||||
|
||||
@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 init_global_content_by_config(self):
|
||||
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()
|
||||
|
||||
# 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 model_config in
|
||||
|
||||
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config)
|
||||
|
||||
# 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"])
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
with GLOBAL_CONTEXT.thread_pool:
|
||||
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
|
||||
memory_chat.run()
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
import re
|
||||
from importlib import import_module
|
||||
from datetime import datetime
|
||||
from importlib import import_module
|
||||
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
|
||||
def under_line_to_hump(underline_str):
|
||||
|
|
@ -10,26 +10,24 @@ def under_line_to_hump(underline_str):
|
|||
return sub[0:1].upper() + sub[1:]
|
||||
|
||||
|
||||
def init_instance_by_config(
|
||||
config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs
|
||||
):
|
||||
clazz_path = config.pop("clazz")
|
||||
if not clazz_path:
|
||||
raise RuntimeError("empty clazz_path!")
|
||||
clazz_name_split = clazz_path.split(".")
|
||||
clazz_name: str = clazz_name_split[-1]
|
||||
if suffix_name and not clazz_name.endswith(suffix_name):
|
||||
clazz_name = f"{clazz_name}_{suffix_name}"
|
||||
def init_instance_by_config(config: dict, default_class_path: str = "", suffix_name: str = "", **kwargs):
|
||||
class_name = config.pop("class")
|
||||
if not class_name:
|
||||
raise RuntimeError("empty class_name!")
|
||||
|
||||
# 构造path
|
||||
clazz_paths = []
|
||||
if default_clazz_path:
|
||||
clazz_paths.append(default_clazz_path)
|
||||
clazz_paths.extend(clazz_name_split[:-1])
|
||||
clazz_paths.append(clazz_name)
|
||||
module = import_module(".".join(clazz_paths))
|
||||
class_name_split = class_name.split(".")
|
||||
class_name: str = class_name_split[-1]
|
||||
if suffix_name and not class_name.lower().endswith(suffix_name.lower()):
|
||||
class_name = f"{class_name}_{suffix_name}"
|
||||
class_name_split[-1] = class_name
|
||||
|
||||
cls_name = under_line_to_hump(clazz_name)
|
||||
class_paths = []
|
||||
if default_class_path:
|
||||
class_paths.append(default_class_path)
|
||||
class_paths.extend(class_name_split)
|
||||
module = import_module(".".join(class_paths))
|
||||
|
||||
cls_name = under_line_to_hump(class_name)
|
||||
return getattr(module, cls_name)(**config, **kwargs)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue