mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-07 08:26:06 +00:00
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:
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 |
49
examples/advance/custom_operator.md
Normal file
49
examples/advance/custom_operator.md
Normal 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
|
||||
```
|
||||
185
examples/advance/replacement.yaml
Normal file
185
examples/advance/replacement.yaml
Normal 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
|
||||
|
|
@ -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
57
examples/cli/README_ZH.md
Normal 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`: 用于排名回复的模型。
|
||||
86
memoryscope/contrib/example_query_worker.py
Normal file
86
memoryscope/contrib/example_query_worker.py
Normal 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))
|
||||
21
memoryscope/contrib/example_query_worker.yaml
Normal file
21
memoryscope/contrib/example_query_worker.yaml
Normal 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:
|
||||
|
||||
|
|
@ -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."})
|
||||
|
||||
|
|
|
|||
|
|
@ -91,7 +91,7 @@ class MemoryScope(ConfigManager):
|
|||
self.close()
|
||||
|
||||
@property
|
||||
def content(self):
|
||||
def context(self):
|
||||
return self._context
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
Loading…
Add table
Reference in a new issue