diff --git a/config/demo_config.yaml b/config/demo_config.yaml index b43b11f9..919c994f 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -51,6 +51,7 @@ worker: retrieve_expired_top_k: 0 extract_time: class: memory.worker.read.extract_time_worker + generation_model: dashscope_generation generation_model_top_k: 1 semantic_rank: class: memory.worker.read.semantic_rank_worker diff --git a/config/demo_config_no_stream.yaml b/config/demo_config_no_stream.yaml new file mode 100644 index 00000000..b9154131 --- /dev/null +++ b/config/demo_config_no_stream.yaml @@ -0,0 +1,148 @@ +global_config: + language: cn + max_workers: 5 + +memory_chat: + cli_memory_chat: + class: chat.cli_memory_chat + stream: false + memory_service: memory_chat_service + generation_model: dashscope_generation + +memory_service: + memory_chat_service: + class: memory.service.chat_memory_service + history_msg_count: 32 + contextual_msg_count: 6 + memory_operations: + read_message: + class: memory.operation.read_message + description: "read session messages of the user" + read_memory: + class: memory.operation.read_memory + workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank + description: "read related memories of the user" + list_memory: + class: memory.operation.read_memory + workflow: set_query,retrieve_memory2,print_memory + description: "read all memories of the user" + write_memory: + class: memory.operation.write_memory + workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory + description: "write observation memories of the user" + interval_time: 5 +# summary_memory: +# class: memory.operation.summary_memory +# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory +# description: "summary observation memories of the user" +# interval_time: 60 + +worker: + dummy: + class: memory.worker.dummy_worker + generation_model: dashscope_generation + embedding_model: dashscope_embedding + rank_model: dashscope_rank + set_query: + class: memory.worker.read.set_query_worker + retrieve_memory1: + class: memory.worker.read.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_pf_top_k: 100 + retrieve_expired_top_k: 0 + extract_time: + class: memory.worker.read.extract_time_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + semantic_rank: + class: memory.worker.read.semantic_rank_worker + fuse_rerank: + class: memory.worker.read.fuse_rerank_worker + fuse_score_threshold: 0.1 + fuse_ratio_dict: + conversation: 0.5 + observation: 1 + obs_customized: 1.2 + insight: 2.0 + fuse_time_ratio: 2.0 + fuse_rerank_top_k: 10 + retrieve_memory2: + class: memory.worker.read.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_pf_top_k: 100 + retrieve_expired_top_k: 100 + print_memory: + class: memory.worker.read.print_memory_worker + info_filter: + class: memory.worker.write.info_filter_worker + generation_model: dashscope_generation + info_filter_msg_max_size: 200 + generation_model_top_k: 1 + load_memory1: + class: memory.worker.write.load_memory_worker + retrieve_not_reflected_top_k: 0 + retrieve_not_updated_top_k: 0 + retrieve_insight_top_k: 0 + today_obs_top_k: 100 + get_observation: + class: memory.worker.write.get_observation_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + get_observation_with_time: + class: memory.worker.write.get_observation_with_time_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + contra_repeat: + class: memory.worker.write.contra_repeat_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + retrieve_top_k: 30 + contra_repeat_max_count: 50 + store_memory: + class: memory.worker.write.store_memory_worker + store_key: all + load_memory2: + class: memory.worker.write.load_memory_worker + retrieve_not_reflected_top_k: 100 + retrieve_not_updated_top_k: 100 + retrieve_insight_top_k: 100 + today_obs_top_k: 0 + get_reflection_subject: + class: memory.worker.summary.get_reflection_subject_worker + retrieve_top_k: 100 + reflect_obs_cnt_threshold: 32 + generation_model_top_k: 1 + update_insight: + class: memory.worker.summary.update_insight_worker + update_insight_threshold: 0.1 + generation_model_top_k: 1 + update_insight_max_thread: 10 + long_contra_repeat: + class: memory.worker.summary.long_contra_repeat_worker + long_contra_repeat_top_k: 2 + long_contra_repeat_threshold: 0.1 + generation_model_top_k: 1 + +models: + dashscope_generation: + class: models.llama_index_generation_model + module_name: dashscope_generation + model_name: qwen-max + dashscope_embedding: + class: models.llama_index_embedding_model + module_name: dashscope_embedding + model_name: text-embedding-v2 + dashscope_rank: + class: models.llama_index_rank_model + module_name: dashscope_rank + model_name: gte-rerank + +memory_store: + class: storage.llama_index_es_memory_store + embedding_model: dashscope_embedding + index_name: memory_index + es_url: http://localhost:9200 + use_hybrid: false + +monitor: + class: storage.dummy_monitor \ No newline at end of file diff --git a/config/test_config.yaml b/config/test_config.yaml index 2668a5e7..a6f71411 100644 --- a/config/test_config.yaml +++ b/config/test_config.yaml @@ -81,3 +81,9 @@ worker: fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 +memory_store: + class: storage.llama_index_es_memory_store + embedding_model: dashscope_embedding + index_name: memory_index + es_url: http://11.160.132.46:9200 + use_hybrid: false \ No newline at end of file diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 585ff831..f3bdbb01 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -76,7 +76,7 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = G_CONTEXT.model_dict[self._generation_model] return self._generation_model - def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: + def chat_with_memory(self, query: str, remember_response:bool=False) -> ModelResponse | ModelResponseGen: new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) self.memory_service.add_messages(new_message) @@ -98,7 +98,18 @@ class CliMemoryChat(BaseMemoryChat): # add new_message messages.append(new_message) self.logger.info(f"messages={messages}") - return self.generation_model.call(messages=messages, stream=self.stream) + + # call LLM. in stream mode, return generator. in non-stream mode, return response. + generated = self.generation_model.call(messages=messages, stream=self.stream) + + # in non-stream mode, remember the response if user demand to do so. + if remember_response: + assert not self.stream + generated.message.role_name = self.assistant_name + self.memory_service.add_messages(generated.message) + + # return response or generator + return generated @staticmethod def parse_query_command(query: str): @@ -189,6 +200,7 @@ class CliMemoryChat(BaseMemoryChat): model_response = self.chat_with_memory(query=query) questionary.print(model_response.message.content) + # add response to memory model_response.message.role_name = self.assistant_name self.memory_service.add_messages(model_response.message) diff --git a/memory_scope/cli.py b/memory_scope/cli.py index c11a58bb..d6f8b050 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -8,6 +8,7 @@ from typing import Dict, Any import fire import yaml +import atexit from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.enumeration.model_enum import ModelEnum @@ -16,8 +17,7 @@ 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 CliJob(object): +class MemoryScope(object): def __init__(self): self.config: Dict[str, Any] = {} @@ -32,6 +32,14 @@ class CliJob(object): else: raise RuntimeError("not supported config file type!") self.init_global_content_by_config() + atexit.register(self.shutdown) # register clean up function + return self + + def shutdown(self): + print('Gracefully executing the shutdown function...') + G_CONTEXT.memory_store.close() + G_CONTEXT.monitor.close() + G_CONTEXT.thread_pool.shutdown() def set_global_config(self): G_CONTEXT.global_config = global_config = self.config["global_config"] @@ -56,6 +64,8 @@ class CliJob(object): G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) # init vector_store + if "memory_store" not in self.config: + raise RuntimeError("memory_store config is required!") memory_store_config = self.config["memory_store"] embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model) @@ -66,6 +76,14 @@ class CliJob(object): # set worker config G_CONTEXT.worker_config = self.config["worker"] + def get_default_service(self): + return list(G_CONTEXT.memory_service_dict.values())[0] + + def get_default_chat_handle(self): + return list(G_CONTEXT.memory_chat_dict.values())[0] + +class CliJob(MemoryScope): + def run(self, config: str): self.load_config(config) @@ -73,10 +91,6 @@ class CliJob(object): memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() - G_CONTEXT.memory_store.close() - G_CONTEXT.monitor.close() - G_CONTEXT.thread_pool.shutdown() - if __name__ == "__main__": cli_job = CliJob() diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 657b8ada..12055aeb 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -9,6 +9,9 @@ from memory_scope.utils.tool_functions import prompt_to_msg class InfoFilterWorker(MemoryBaseWorker): + """ + This worker will filter and modify `self.chat_messages`, preserving only the messages that contain important information. + """ FILE_PATH: str = __file__ def _run(self): diff --git a/tests/operations/init_test.py b/tests/operations/init_test.py new file mode 100644 index 00000000..a6c3649d --- /dev/null +++ b/tests/operations/init_test.py @@ -0,0 +1,10 @@ +def validate_path(): + import os, sys + + os.path.dirname(__file__) + root_dir_assume = os.path.abspath(os.path.dirname(__file__) + "/../..") + os.chdir(root_dir_assume) + sys.path.append(root_dir_assume) + + +validate_path() # validate path so you can run from base directory diff --git a/tests/operations/test_operation.py b/tests/operations/test_operation.py new file mode 100644 index 00000000..0ad9571f --- /dev/null +++ b/tests/operations/test_operation.py @@ -0,0 +1,21 @@ +import init_test +from memory_scope.cli import MemoryScope +from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.scheme.message import Message + +ms = MemoryScope().load_config("config/demo_config_no_stream.yaml") +memory_service = ms.get_default_service() +memory_chat = ms.get_default_chat_handle() + +# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。") +# memory_service.add_messages(new_message) + +res:Message = memory_chat.chat_with_memory(query="我的爱好是弹琴。", remember_response=True) +print(res.message.content) + +res:Message = memory_chat.chat_with_memory(query="昨天弹出一个光粒,消灭了星系0x4be。", remember_response=True) +print(res.message.content) + +res:Message = memory_chat.chat_with_memory(query="今天弹出一个二向箔,消灭了星系0xa2e。", remember_response=True) +print(res.message.content) +