add custom worker

* udpate CLI readme

* add rewrite worker

* fix new worker prompt & modify unittest

* feat: Resolve conflict, auto committed by CodeFlow
This commit is contained in:
fuqingxu.fqx 2024-08-29 12:23:54 +08:00 committed by jinli.yl
parent a41e616650
commit 4e9cfe08c7
13 changed files with 549 additions and 71 deletions

Binary file not shown.

Before

Width:  |  Height:  |  Size: 197 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 251 KiB

View file

@ -0,0 +1,49 @@
# 自定义 Operator 和 Worker
1. 在 `contrib` 路径下创建新worker命名为 `example_query_worker.py`:
```bash
vim memoryscope/contrib/example_query_worker.py
```
2. 写入新的自定义worker的程序注意`class`的命名需要与文件名保持一致,为`ExampleQueryWorker`
```python
import datetime
from memoryscope.constants.common_constants import QUERY_WITH_TS
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
class ExampleQueryWorker(MemoryBaseWorker):
def _run(self):
timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
assert "query" in self.chat_kwargs
query = self.chat_kwargs["query"]
if not query:
query = ""
else:
query = query.strip() + "\n You must add a `meow~` at the end of each of your answer."
# Store the determined query and its timestamp in the context
self.set_workflow_context(QUERY_WITH_TS, (query, timestamp))
```
3. 创建yaml启动文件复制demo_config.yaml
```
cp memoryscope/core/config/demo_config.yaml examples/advance/replacement.yaml
vim examples/advance/replacement.yaml
```
4. 在最下面插入新worker的定义并且取代之前的默认`set_query`worker
```
set_query_meow:
class: contrib.example_query_worker
```
5. 验证:
```
python quick-start-demo.py --config examples/advance/replacement.yaml
```

View file

@ -0,0 +1,185 @@
global:
language: en
thread_pool_max_workers: 5
logger_name: memoryscope
logger_name_time_suffix: "%Y%m%d_%H%M%S"
logger_to_screen: false
enable_ranker: false
enable_today_contra_repeat: true
enable_long_contra_repeat: false
output_memory_max_count: 20
memory_chat:
cli_memory_chat:
class: core.chat.cli_memory_chat
memory_service: memoryscope_service
generation_model: generation_model
stream: true
memory_service:
memoryscope_service:
class: core.service.memory_scope_service
human_name: user
assistant_name: AI
memory_operations:
read_message:
class: core.operation.frontend_operation
workflow: read_message
description: "read short memory"
retrieve_memory:
class: core.operation.frontend_operation
workflow: set_query_meow,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank
description: "retrieve long-term memory"
list_memory:
class: core.operation.frontend_operation
workflow: set_query,retrieve_top_memory,print_memory
description: "read all long-term memory of the user, use `refresh_time=5` to refresh screen every 5 seconds."
delete_memory:
class: core.operation.frontend_operation
workflow: set_query,retrieve_all_memory,delete_memory
description: "delete a single long-term memory"
delete_all:
class: core.operation.frontend_operation
workflow: set_query,retrieve_all_memory,delete_all
description: "delete all long-term memory"
add_memory:
class: core.operation.frontend_operation
workflow: add_memory
description: "add a single observation"
consolidate_memory:
class: core.operation.consolidate_memory_op
workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory
description: "summary user's observation memory, run backend."
interval_time: 1
reflect_and_reconsolidate:
class: core.operation.backend_operation
workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory
description: "summary user's insight memory, run backend."
interval_time: 15
worker:
dummy:
class: core.worker.dummy_worker
generation_model: generation_model
embedding_model: embedding_model
rank_model: rank_model
read_message:
class: core.worker.frontend.read_message_worker
set_query:
class: core.worker.frontend.set_query_worker
set_query_meow:
class: contrib.example_query_worker
generation_model: generation_model
retrieve_obs_ins:
class: core.worker.frontend.retrieve_memory_worker
retrieve_obs_top_k: 100
retrieve_ins_top_k: 100
extract_time:
class: core.worker.frontend.extract_time_worker
generation_model: generation_model
semantic_rank:
class: core.worker.frontend.semantic_rank_worker
rank_model: rank_model
fuse_rerank:
class: core.worker.frontend.fuse_rerank_worker
fuse_score_threshold: 0.01
fuse_ratio_dict:
conversation: 0.5
observation: 1
obs_customized: 1.2
insight: 2.0
fuse_time_ratio: 2.0
retrieve_top_memory:
class: core.worker.frontend.retrieve_memory_worker
retrieve_obs_top_k: 100
retrieve_ins_top_k: 100
retrieve_expired_top_k: 100
print_memory:
class: core.worker.frontend.print_memory_worker
retrieve_all_memory:
class: core.worker.frontend.retrieve_memory_worker
retrieve_obs_top_k: 1000
retrieve_ins_top_k: 1000
retrieve_expired_top_k: 1000
delete_memory:
class: core.worker.backend.update_memory_worker
method: delete_memory
delete_all:
class: core.worker.backend.update_memory_worker
method: delete_all
add_memory:
class: core.worker.backend.update_memory_worker
method: from_query
info_filter:
class: core.worker.backend.info_filter_worker
generation_model: generation_model
load_today_memory:
class: core.worker.backend.load_memory_worker
retrieve_today_top_k: 100
get_observation:
class: core.worker.backend.get_observation_worker
generation_model: generation_model
get_observation_with_time:
class: core.worker.backend.get_observation_with_time_worker
generation_model: generation_model
contra_repeat:
class: core.worker.backend.contra_repeat_worker
generation_model: generation_model
store_memory:
class: core.worker.backend.update_memory_worker
method: from_memory_key
memory_key: all
load_obs_and_insight:
class: core.worker.backend.load_memory_worker
retrieve_not_reflected_top_k: 100
retrieve_not_updated_top_k: 100
retrieve_insight_top_k: 100
get_reflection_subject:
class: core.worker.backend.get_reflection_subject_worker
generation_model: generation_model
reflect_obs_cnt_threshold: 5
update_insight:
class: core.worker.backend.update_insight_worker
generation_model: generation_model
rank_model: rank_model
embedding_model: embedding_model
update_insight_threshold: 0.01
enable_parallel: false
long_contra_repeat:
class: core.worker.backend.long_contra_repeat_worker
generation_model: generation_model
long_contra_repeat_threshold: 0.5
model:
generation_model:
class: core.models.llama_index_generation_model
module_name: dashscope_generation
model_name: qwen-max
max_tokens: 2000
temperature: 0.01
embedding_model:
class: core.models.llama_index_embedding_model
module_name: dashscope_embedding
model_name: text-embedding-v2
rank_model:
class: core.models.llama_index_rank_model
module_name: dashscope_rank
model_name: gte-rerank
top_n: 500
memory_store:
class: core.storage.llama_index_es_memory_store
embedding_model: embedding_model
index_name: memory_index
es_url: http://localhost:9200
retrieve_mode: dense
monitor:
class: core.storage.dummy_monitor

View file

@ -21,7 +21,7 @@ class MemoryScopeAgent(AgentBase):
response = self.memory_chat.chat_with_memory(query=x.content)
# Wrap the response in a message object in AgentScope
msg = Msg(name=self.name, content=response.message.content, role="Assistant")
msg = Msg(name=self.name, content=response.message.content, role="assistant")
# Print/speak the message in this agent's voice
self.speak(msg)

57
examples/cli/README_ZH.md Normal file
View file

@ -0,0 +1,57 @@
# MemoryScope 的命令行接口
## 使用方法
MemoryScope 可以通过两种不同的方式启动:
### 1. 使用 YAML 配置文件
如果您更喜欢通过 YAML 文件配置设置,可以通过提供配置文件的路径来实现:
```bash
memoryscope --config_path=memoryscope/core/config/demo_config.yaml
```
### 2. 使用命令行参数
或者,您可以直接在命令行上指定所有参数:
```
# 中文
memoryscope --language="cn" \
--memory_chat_class="cli_memory_chat" \
--human_name="用户" \
--assistant_name="AI" \
--generation_backend="dashscope_generation" \
--generation_model="qwen-max" \
--embedding_backend="dashscope_embedding" \
--embedding_model="text-embedding-v2" \
--enable_ranker=True \
--rank_backend="dashscope_rank" \
--rank_model="gte-rerank"
# 英文
memoryscope --language="en" \
--memory_chat_class="cli_memory_chat" \
--human_name="User" \
--assistant_name="AI" \
--generation_backend="openai_generation" \
--generation_model="gpt-4o" \
--embedding_backend="openai_embedding" \
--embedding_model="text-embedding-3-small" \
--enable_ranker=False
```
以下是可以通过任一方法设置的可用选项:
- `--language`: 对话中使用的语言。
- `--memory_chat_class`: 管理聊天记录的类名。
- `--human_name`: 人类用户的名字。
- `--assistant_name`: AI 助手的名字。
- `--generation_backend`: 用于生成回复的后端。
- `--generation_model`: 用于生成回复的模型。
- `--embedding_backend`: 用于文本嵌入的后端。
- `--embedding_model`: 用于创建文本嵌入的模型。
- `--enable_ranker`: 一个布尔值,指示是否使用排名器(默认为 False
- `--rank_backend`: 用于排名回复的后端。
- `--rank_model`: 用于排名回复的模型。

View file

@ -0,0 +1,86 @@
import datetime
from memoryscope.constants.common_constants import QUERY_WITH_TS
from memoryscope.constants.language_constants import NONE_WORD
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
class ExampleQueryWorker(MemoryBaseWorker):
# NOTE: If you want to utilize the capabilities of the prompt handler, please be sure to include this sentence.
FILE_PATH: str = __file__
def _parse_params(self, **kwargs):
self.rewrite_history_count: int = kwargs.get("rewrite_history_count", 2)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
def rewrite_query(self, query: str) -> str:
chat_messages = self.chat_messages_scatter
if len(chat_messages) <= 1:
return query
if chat_messages[-1].role == MessageRoleEnum.USER:
chat_messages = chat_messages[:-1]
chat_messages = chat_messages[-self.rewrite_history_count:]
# get context
context_list = []
for message in chat_messages:
context = message.content
if len(context) > 200:
context = context[:100] + context[-100:]
if message.role == MessageRoleEnum.USER:
context_list.append(f"{self.target_name}: {context}")
elif message.role == MessageRoleEnum.ASSISTANT:
context_list.append(f"Assistant: {context}")
if not context_list:
return query
system_prompt = self.prompt_handler.rewrite_query_system
user_query = self.prompt_handler.rewrite_query_query.format(query=query,
context="\n".join(context_list))
rewrite_query_message = self.prompt_to_msg(system_prompt=system_prompt,
few_shot="",
user_query=user_query)
self.logger.info(f"rewrite_query_message={rewrite_query_message}")
# Invoke the LLM to generate a response
response = self.generation_model.call(messages=rewrite_query_message,
**self.generation_model_kwargs)
# Handle empty or unsuccessful responses
if not response.status or not response.message.content:
return query
response_text = response.message.content
self.logger.info(f"rewrite_query.response_text={response_text}")
if not response_text or response_text.lower() == self.get_language_value(NONE_WORD):
return query
return response_text
def _run(self):
query = "" # Default query value
timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
if "query" in self.chat_kwargs:
# set query if exists
query = self.chat_kwargs["query"]
if not query:
query = ""
query = query.strip()
# set ts if exists
_timestamp = self.chat_kwargs.get("timestamp")
if _timestamp and isinstance(_timestamp, int):
timestamp = _timestamp
if self.rewrite_history_count > 0:
t_query = self.rewrite_query(query=query)
if t_query:
query = t_query
# Store the determined query and its timestamp in the context
self.set_workflow_context(QUERY_WITH_TS, (query, timestamp))

View file

@ -0,0 +1,21 @@
rewrite_query_system:
cn: |
任务: 消除指代问题并重写
要求: 检查提供的问题是否存在指代。如果存在指代,通过上下文信息重写问题,使其信息充足,能够单独回答。如果没有指代问题,则回答“无”。
en: |
Task: Eliminate referencing issues and rewrite
Requirements: Check the provided questions for any references. If references exist, rewrite the questions using contextual information to make them sufficiently informative so they can be answered independently. If there are no referencing issues, respond with "None".
rewrite_query_query:
cn: |
上下文:
{context}
问题:{query}
重写:
en: |
Context:
{context}
Question: {query}
Rewrite:

View file

@ -1,6 +1,7 @@
from dataclasses import dataclass, field
from typing import Literal, Dict
@dataclass
class Arguments(object):
language: Literal["cn", "en"] = field(default="cn", metadata={"help": "support en & cn now"})
@ -32,7 +33,7 @@ class Arguments(object):
generation_backend: str = field(default="dashscope_generation", metadata={
"help": "global generation backend: openai_generation, dashscope_generation, etc."})
generation_model: str = field(default="gpt-4o", metadata={
generation_model: str = field(default="qwen-max", metadata={
"help": "global generation model: gpt-4o, gpt-4o-mini, gpt-4-turbo, qwen-max, etc."})
generation_params: dict = field(default_factory=lambda: {}, metadata={
@ -41,7 +42,7 @@ class Arguments(object):
embedding_backend: str = field(default="dashscope_generation", metadata={
"help": "global embedding backend: openai_embedding, dashscope_embedding, etc."})
embedding_model: str = field(default="text-embedding-3-small", metadata={
embedding_model: str = field(default="text-embedding-v2", metadata={
"help": "global embedding model: text-embedding-3-large, text-embedding-3-small, text-embedding-ada-002, "
"text-embedding-v2, etc."})

View file

@ -91,7 +91,7 @@ class MemoryScope(ConfigManager):
self.close()
@property
def content(self):
def context(self):
return self._context
@property

View file

@ -11,6 +11,7 @@ from memoryscope.scheme.message import Message
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
from memoryscope.core.utils.logger import Logger
class LlamaIndexGenerationModel(BaseModel):
"""
This class represents a generation model within the LlamaIndex framework,
@ -73,7 +74,7 @@ class LlamaIndexGenerationModel(BaseModel):
model_response.message.content += delta
model_response.delta = response.delta
yield model_response
self.logger.info(self.logger.format_chat_message(model_response))
return gen()
else:
if isinstance(call_result, CompletionResponse):

View file

@ -31,8 +31,6 @@ class TestWorkersCn(unittest.TestCase):
enable_ranker=True,
)
self.ms = MemoryScope(arguments=self.arguments)
config = self.ms.dump_config()
self.ms.logger.info(f"config=\n{config}")
def tearDown(self):
self.ms.close()
@ -42,12 +40,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
query = "明天我去上海出差"
query_timestamp = int(datetime.datetime.now().timestamp())
@ -57,17 +56,18 @@ class TestWorkersCn(unittest.TestCase):
result = worker.get_workflow_context(EXTRACT_TIME_DICT)
worker.logger.info(f"result={result}")
# @unittest.skip
@unittest.skip
def test_info_filter(self):
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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name),
@ -96,12 +96,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗",
@ -144,12 +145,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name),
@ -178,12 +180,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。", role_name=self.arguments.human_name),
@ -209,12 +212,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术", role_name=self.arguments.human_name),
@ -238,12 +242,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
@ -298,12 +303,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name),
@ -334,12 +340,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="用户喜欢打王者荣耀", role_name=self.arguments.human_name),
@ -356,12 +363,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name),
@ -375,3 +383,34 @@ class TestWorkersCn(unittest.TestCase):
result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")
# @unittest.skip
def test_example_query_worker(self):
name = "example_query_worker"
worker: MemoryBaseWorker = init_instance_by_config(
config={
"class": "contrib.example_query_worker",
"generation_model": "generation_model",
},
name=name,
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name,
"chat_kwargs": {"query": "我一直很爱他们"}},
context_lock=None,
memoryscope_context=self.ms.context,
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我的两个孩子分别叫小明和小红",
role_name=self.arguments.human_name),
Message(role=MessageRoleEnum.ASSISTANT.value,
content="很高兴认识您和您的家庭成员!小明和小红是非常通俗且好听的名字。",
role_name=self.arguments.assistant_name),
Message(role=MessageRoleEnum.USER.value, content="我一直很爱他们", role_name=self.arguments.human_name),
]
worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages)
worker.run()
result = worker.get_workflow_context(QUERY_WITH_TS)
worker.logger.info(f"result={result}")

View file

@ -3,7 +3,7 @@ import unittest
from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \
MEMORYSCOPE_CONTEXT
MEMORYSCOPE_CONTEXT, TARGET_NAME, CHAT_MESSAGES_SCATTER
from memoryscope.core.config.arguments import Arguments
from memoryscope.core.memoryscope import MemoryScope
from memoryscope.core.utils.tool_functions import init_instance_by_config
@ -17,7 +17,7 @@ class TestWorkersEn(unittest.TestCase):
"""Tests for LLIEmbedding"""
def setUp(self):
arguments = Arguments(
self.arguments = Arguments(
language="en",
human_name="user",
assistant_name="AI",
@ -29,9 +29,7 @@ class TestWorkersEn(unittest.TestCase):
rank_backend="dashscope_rank",
rank_model="gte-rerank",
)
self.ms = MemoryScope(arguments=arguments)
config = self.ms.dump_config()
self.ms.logger.info(f"config=\n{config}")
self.ms = MemoryScope(arguments=self.arguments)
def tearDown(self):
self.ms.close()
@ -41,12 +39,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
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 +60,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."),
@ -87,12 +87,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
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 +136,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
# FIXME Does the appearance of 'am' indicate the presence of a time keyword?
chat_messages = [
@ -164,12 +166,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value,
@ -206,12 +209,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value,
@ -238,12 +242,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"),
@ -296,12 +301,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),
@ -333,12 +339,13 @@ 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)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="Users like to play King of Glory"),
@ -355,12 +362,13 @@ 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, TARGET_NAME: self.arguments.human_name},
context_lock=None,
thread_pool=self.ms._context.thread_pool)
memoryscope_context=self.ms.context,
thread_pool=self.ms.context.thread_pool)
nodes = [
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),
@ -374,3 +382,34 @@ class TestWorkersEn(unittest.TestCase):
result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")
# @unittest.skip
def test_example_query_worker(self):
name = "example_query_worker"
worker: MemoryBaseWorker = init_instance_by_config(
config={
"class": "contrib.example_query_worker",
"generation_model": "generation_model",
},
name=name,
context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name,
"chat_kwargs": {"query": "I have always loved them."}},
context_lock=None,
memoryscope_context=self.ms.context,
thread_pool=self.ms._context.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="My two children are named Xiaoming and Xiaohong.",
role_name=self.arguments.human_name),
Message(role=MessageRoleEnum.ASSISTANT.value,
content="I am very pleased to meet you and your family members! Xiaoming and Xiaohong are very pleasant names.",
role_name=self.arguments.assistant_name),
Message(role=MessageRoleEnum.USER.value, content="I have always loved them.", role_name=self.arguments.human_name),
]
worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages)
worker.run()
result = worker.get_workflow_context(QUERY_WITH_TS)
worker.logger.info(f"result={result}")