【feature】modify memoryscope enter logic & optimize es batch insert

This commit is contained in:
jinli.yl 2024-08-19 17:13:55 +08:00
parent 88946c19e2
commit bfdc2c7564
7 changed files with 101 additions and 98 deletions

View file

@ -1,11 +1,12 @@
from typing import Optional, Union, Sequence
import agentscope
import sys
import os
from agentscope.agents import AgentBase, UserAgent
from agentscope.message import Msg
from memoryscope import MemoryScope, Arguments
class MemoryScopeAgent(AgentBase):
def __init__(self, name: str, arguments: Arguments, **kwargs) -> None:
# Disable AgentScope memory and use MemoryScope memory instead
@ -13,7 +14,6 @@ class MemoryScopeAgent(AgentBase):
# Create a memory client in MemoryScope
self.memory_scope = MemoryScope(arguments=arguments)
self.memory_scope.init_context_by_config()
self.memory_chat = self.memory_scope.default_memory_chat
def reply(self, x: Optional[Union[Msg, Sequence[Msg]]] = None) -> Msg:
@ -46,7 +46,6 @@ def main():
generation_model="qwen-max",
embedding_backend="dashscope_embedding",
embedding_model="text-embedding-v2",
use_dummy_ranker=False,
rank_backend="dashscope_rank",
rank_model="gte-rerank"
)
@ -79,4 +78,4 @@ def main():
if __name__ == "__main__":
main()
main()

View file

@ -1,7 +1,7 @@
import os
import sys
from typing import Optional, Union, Sequence, Literal, Dict, List, Any, Tuple
from autogen import Agent, ConversableAgent, UserProxyAgent, config_list_from_json
from typing import Optional, Union, Literal, Dict, List, Any, Tuple
from autogen import Agent, ConversableAgent, UserProxyAgent
from memoryscope import MemoryScope, Arguments
@ -25,7 +25,6 @@ class MemoryScopeAgent(ConversableAgent):
# Create a memory client in MemoryScope
self.memory_scope = MemoryScope(arguments=arguments)
self.memory_scope.init_context_by_config()
self.memory_chat = self.memory_scope.default_memory_chat
self.register_reply([Agent, None], MemoryScopeAgent.generate_reply_with_memory,remove_other_reply_funcs=True)
@ -63,7 +62,6 @@ def main():
generation_model="qwen-max",
embedding_backend="dashscope_embedding",
embedding_model="text-embedding-v2",
use_dummy_ranker=False,
rank_backend="dashscope_rank",
rank_model="gte-rerank"
)
@ -79,4 +77,4 @@ def main():
if __name__ == "__main__":
main()
main()

View file

@ -38,7 +38,7 @@ class Arguments(object):
generation_backend: str = field(default="openai_generation", metadata={
"help": "global generation backend: openai_generation, dashscope_generation, etc."})
generation_model: str = field(default="gpt-4o-mini", metadata={
generation_model: str = field(default="gpt-4o", metadata={
"help": "global generation model: gpt-4o, gpt-4o-mini, gpt-4-turbo, qwen-max, etc."})
generation_params: dict = field(default_factory=lambda: {}, metadata={

View file

@ -14,15 +14,15 @@ class MemoryScope(ConfigManager):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.context: MemoryscopeContext = MemoryscopeContext()
self.init_context_by_config()
self._context: MemoryscopeContext = MemoryscopeContext()
self._init_context_by_config()
def init_context_by_config(self):
def _init_context_by_config(self):
# set global config
global_conf = self.config["global"]
self.context.language = LanguageEnum(global_conf["language"])
self.context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"])
self.context.meta_data.update({
self._context.language = LanguageEnum(global_conf["language"])
self._context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"])
self._context.meta_data.update({
"enable_ranker": global_conf["enable_ranker"],
"enable_today_contra_repeat": global_conf["enable_today_contra_repeat"],
"enable_long_contra_repeat": global_conf["enable_long_contra_repeat"],
@ -42,63 +42,68 @@ class MemoryScope(ConfigManager):
memory_chat_conf_dict = self.config["memory_chat"]
if memory_chat_conf_dict:
for name, conf in memory_chat_conf_dict.items():
self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
self._context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self._context)
# set memory_service
memory_service_conf_dict = self.config["memory_service"]
assert memory_service_conf_dict
for name, conf in memory_service_conf_dict.items():
self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
self._context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self._context)
# init model
model_conf_dict = self.config["model"]
assert model_conf_dict
for name, conf in model_conf_dict.items():
self.context.model_dict[name] = init_instance_by_config(conf, name=name)
self._context.model_dict[name] = init_instance_by_config(conf, name=name)
# init memory_store
memory_store_conf = self.config["memory_store"]
assert memory_store_conf
emb_model_name: str = memory_store_conf[ModelEnum.EMBEDDING_MODEL.value]
embedding_model = self.context.model_dict[emb_model_name]
self.context.memory_store = init_instance_by_config(memory_store_conf, embedding_model=embedding_model)
embedding_model = self._context.model_dict[emb_model_name]
self._context.memory_store = init_instance_by_config(memory_store_conf, embedding_model=embedding_model)
# init monitor
monitor_conf = self.config["monitor"]
if monitor_conf:
self.context.monitor = init_instance_by_config(monitor_conf)
self._context.monitor = init_instance_by_config(monitor_conf)
# set worker config
self.context.worker_conf_dict = self.config["worker"]
self._context.worker_conf_dict = self.config["worker"]
def close(self):
# wait service to stop
for _, service in self.context.memory_service_dict.items():
for _, service in self._context.memory_service_dict.items():
service.stop_backend_service(wait_service=True)
self.context.thread_pool.shutdown()
self._context.thread_pool.shutdown()
self.context.memory_store.close()
self._context.memory_store.close()
if self.context.monitor:
self.context.monitor.close()
if self._context.monitor:
self._context.monitor.close()
self.logger.close()
def __enter__(self):
self.init_context_by_config()
return self
def __exit__(self, exc_type, exc_val, exc_tb):
if exc_type is not None:
self.logger.warning(f"An exception occurred: {exc_type.__name__}: {exc_val}\n{exc_tb}")
self.close()
@property
def content(self):
return self._context
@property
def memory_chat_dict(self):
return self.context.memory_chat_dict
return self._context.memory_chat_dict
@property
def memory_service_dict(self):
return self.context.memory_service_dict
return self._context.memory_service_dict
@property
def default_memory_chat(self) -> BaseMemoryChat:

View file

@ -104,14 +104,15 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
return [self._text_node_2_memory_node(n) for n in text_nodes]
def batch_insert(self, nodes: List[MemoryNode]):
# TODO batch insert
for node in nodes:
self.insert(node)
self.index.insert_nodes([self._memory_node_2_text_node(node) for node in nodes])
def batch_update(self, nodes: List[MemoryNode], update_embedding: bool = True):
# TODO batch_update
for node in nodes:
self.update(node, update_embedding=update_embedding)
if update_embedding:
for node in nodes:
node.vector = []
self.batch_delete(nodes)
self.batch_insert(nodes)
def batch_delete(self, nodes: List[MemoryNode]):
# TODO batch_delete

View file

@ -43,12 +43,12 @@ class TestWorkersCn(unittest.TestCase):
name = "extract_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
query = "明天我去上海出差"
query_timestamp = int(datetime.datetime.now().timestamp())
@ -63,12 +63,12 @@ class TestWorkersCn(unittest.TestCase):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name),
@ -97,12 +97,12 @@ class TestWorkersCn(unittest.TestCase):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗",
@ -145,12 +145,12 @@ class TestWorkersCn(unittest.TestCase):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name),
@ -179,12 +179,12 @@ class TestWorkersCn(unittest.TestCase):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。", role_name=self.arguments.human_name),
@ -210,12 +210,12 @@ class TestWorkersCn(unittest.TestCase):
name = "get_observation_with_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术", role_name=self.arguments.human_name),
@ -239,12 +239,12 @@ class TestWorkersCn(unittest.TestCase):
name = "contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
@ -299,12 +299,12 @@ class TestWorkersCn(unittest.TestCase):
name = "get_reflection_subject"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name),
@ -335,12 +335,12 @@ class TestWorkersCn(unittest.TestCase):
name = "update_insight"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context=reflection_worker.context,
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="用户喜欢打王者荣耀", role_name=self.arguments.human_name),
@ -357,12 +357,12 @@ class TestWorkersCn(unittest.TestCase):
name = "long_contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name},
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name),

View file

@ -41,12 +41,12 @@ class TestWorkersEn(unittest.TestCase):
name = "extract_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
query = "I will be on a business trip to Shanghai tomorrow."
query_timestamp = int(datetime.datetime.now().timestamp())
@ -61,12 +61,12 @@ class TestWorkersEn(unittest.TestCase):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."),
@ -87,12 +87,12 @@ class TestWorkersEn(unittest.TestCase):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"),
@ -135,12 +135,12 @@ class TestWorkersEn(unittest.TestCase):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
# FIXME Does the appearance of 'am' indicate the presence of a time keyword?
chat_messages = [
@ -164,12 +164,12 @@ class TestWorkersEn(unittest.TestCase):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value,
@ -206,12 +206,12 @@ class TestWorkersEn(unittest.TestCase):
name = "get_observation_with_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value,
@ -238,12 +238,12 @@ class TestWorkersEn(unittest.TestCase):
name = "contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"),
@ -296,12 +296,12 @@ class TestWorkersEn(unittest.TestCase):
name = "get_reflection_subject"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),
@ -333,12 +333,12 @@ class TestWorkersEn(unittest.TestCase):
name = "update_insight"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context=reflection_worker.context,
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="Users like to play King of Glory"),
@ -355,12 +355,12 @@ class TestWorkersEn(unittest.TestCase):
name = "long_contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=self.ms.context.worker_conf_dict[name],
config=self.ms._context.worker_conf_dict[name],
name=name,
is_multi_thread=False,
context={MEMORYSCOPE_CONTEXT: self.ms.context},
context={MEMORYSCOPE_CONTEXT: self.ms._context},
context_lock=None,
thread_pool=self.ms.context.thread_pool)
thread_pool=self.ms._context.thread_pool)
nodes = [
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),