mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-10 22:41:06 +00:00
add test interface
This commit is contained in:
parent
907a8f1815
commit
6cb54996bf
8 changed files with 223 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
148
config/demo_config_no_stream.yaml
Normal file
148
config/demo_config_no_stream.yaml
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
10
tests/operations/init_test.py
Normal file
10
tests/operations/init_test.py
Normal file
|
|
@ -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
|
||||
21
tests/operations/test_operation.py
Normal file
21
tests/operations/test_operation.py
Normal file
|
|
@ -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)
|
||||
|
||||
Loading…
Add table
Reference in a new issue