mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
commit
fba0a00802
133 changed files with 5352 additions and 3158 deletions
85
.pre-commit-config.yaml
Normal file
85
.pre-commit-config.yaml
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
repos:
|
||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||
rev: v6.0.0
|
||||
hooks:
|
||||
- id: check-ast
|
||||
exclude: ^(test/|cookbook/)
|
||||
- id: check-yaml
|
||||
- id: check-xml
|
||||
- id: check-toml
|
||||
- id: check-json
|
||||
- id: detect-private-key
|
||||
- id: trailing-whitespace
|
||||
- repo: https://github.com/asottile/add-trailing-comma
|
||||
rev: v4.0.0
|
||||
hooks:
|
||||
- id: add-trailing-comma
|
||||
exclude: ^(test/|cookbook/)
|
||||
- repo: https://github.com/psf/black
|
||||
rev: 25.9.0
|
||||
hooks:
|
||||
- id: black
|
||||
exclude: ^(test/|cookbook/)
|
||||
args: [--line-length=120]
|
||||
- repo: https://github.com/PyCQA/flake8
|
||||
rev: 7.3.0
|
||||
hooks:
|
||||
- id: flake8
|
||||
exclude: ^(test/|cookbook/)
|
||||
args: [
|
||||
"--extend-ignore=E203",
|
||||
"--max-line-length=120"
|
||||
]
|
||||
- repo: https://github.com/pylint-dev/pylint
|
||||
rev: v4.0.2
|
||||
hooks:
|
||||
- id: pylint
|
||||
exclude:
|
||||
(?x)(
|
||||
^docs
|
||||
| ^test/
|
||||
| ^cookbook/
|
||||
| pb2\.py$
|
||||
| grpc\.py$
|
||||
| \.demo$
|
||||
| \.md$
|
||||
| \.html$
|
||||
)
|
||||
args: [
|
||||
--disable=W0511,
|
||||
--disable=W0718,
|
||||
--disable=W0122,
|
||||
--disable=C0103,
|
||||
--disable=R0913,
|
||||
--disable=R0917,
|
||||
--disable=E0401,
|
||||
--disable=E1101,
|
||||
--disable=E1111,
|
||||
--disable=C0415,
|
||||
--disable=W0603,
|
||||
--disable=R1705,
|
||||
--disable=R0914,
|
||||
--disable=E0601,
|
||||
--disable=W0602,
|
||||
--disable=W0604,
|
||||
--disable=R0801,
|
||||
--disable=R0902,
|
||||
--disable=R0903,
|
||||
--disable=R0904,
|
||||
--disable=C0123,
|
||||
--disable=W0231,
|
||||
--disable=W1113,
|
||||
--disable=W0221,
|
||||
--disable=R0401,
|
||||
--disable=W0632,
|
||||
--disable=W0123,
|
||||
--disable=C3001,
|
||||
--disable=R1702,
|
||||
--disable=R0912,
|
||||
--max-line-length=120,
|
||||
]
|
||||
- repo: https://github.com/regebro/pyroma
|
||||
rev: "5.0"
|
||||
hooks:
|
||||
- id: pyroma
|
||||
args: [--min=10, .]
|
||||
40
README.md
40
README.md
|
|
@ -3,8 +3,8 @@
|
|||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.12+-blue" alt="Python Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-v0.1.10.x-blue?logo=pypi" alt="PyPI Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.10+-blue" alt="Python Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-0.2.0.0-blue?logo=pypi" alt="PyPI Version"></a>
|
||||
<a href="./LICENSE"><img src="https://img.shields.io/badge/license-Apache--2.0-black" alt="License"></a>
|
||||
<a href="https://github.com/agentscope-ai/ReMe"><img src="https://img.shields.io/github/stars/modelscope/ReMe?style=social" alt="GitHub Stars"></a>
|
||||
</p>
|
||||
|
|
@ -72,22 +72,14 @@ Learn more about how to use tool memory from [tool memory](docs/tool_memory/tool
|
|||
|
||||
## 📰 Latest Updates
|
||||
|
||||
- **[2025-10]** 🚀 ReMe v0.1.10 released! Core enhancement: direct Python import support. You can now use ReMe without starting an HTTP or MCP service - simply `from reme_ai import ReMeApp` and call methods directly in your Python code.
|
||||
- **[2025-10]** 🔧 Tool Memory support is now available! Enables data-driven tool selection and parameter optimization through historical performance tracking. Check out the [Tool Memory Guide](docs/tool_memory/tool_memory.md) and [benchmark results](docs/tool_memory/tool_bench.md).
|
||||
- **[2025-09]** 🎉 ReMe v0.1.9 has been officially released, adding support for asynchronous operations. It has also been
|
||||
integrated into the memory service of agentscope-runtime.
|
||||
- **[2025-09]** 🎉 ReMe v0.1 officially released, integrating task memory and personal memory. If you want to use the
|
||||
original memoryscope project, you can find it
|
||||
in [MemoryScope](https://github.com/agentscope-ai/ReMe/tree/memoryscope_branch).
|
||||
- **[2025-09]** 🧪 We validated the effectiveness of task memory extraction and reuse in agents in appworld, bfcl(v3),
|
||||
and frozenlake environments. For more information,
|
||||
check [appworld exp](docs/cookbook/appworld/quickstart.md), [bfcl exp](docs/cookbook/bfcl/quickstart.md),
|
||||
and [frozenlake exp](docs/cookbook/frozenlake/quickstart.md).
|
||||
- **[2025-08]** 🚀 MCP protocol support is now available -> [MCP Quick Start](docs/mcp_quick_start.md).
|
||||
- **[2025-06]** 🚀 Multiple backend vector storage support (Elasticsearch &
|
||||
ChromaDB) -> [Vector DB quick start](docs/vector_store_api_guide.md).
|
||||
- **[2024-09]** 🧠 [MemoryScope](https://github.com/agentscope-ai/ReMe/tree/memoryscope_branch) v0.1 released,
|
||||
personalized and time-aware memory storage and usage.
|
||||
- **[2025-10]** 🚀 Direct Python import support: use `from reme_ai import ReMeApp` without HTTP/MCP service
|
||||
- **[2025-10]** 🔧 Tool Memory: data-driven tool selection and parameter optimization ([Guide](docs/tool_memory/tool_memory.md))
|
||||
- **[2025-09]** 🎉 Async operations support, integrated into agentscope-runtime
|
||||
- **[2025-09]** 🎉 Task memory and personal memory integration
|
||||
- **[2025-09]** 🧪 Validated effectiveness in appworld, bfcl(v3), and frozenlake ([Experiments](docs/cookbook))
|
||||
- **[2025-08]** 🚀 MCP protocol support ([Quick Start](docs/mcp_quick_start.md))
|
||||
- **[2025-06]** 🚀 Multiple backend vector storage (Elasticsearch & ChromaDB) ([Guide](docs/vector_store_api_guide.md))
|
||||
- **[2024-09]** 🧠 Personalized and time-aware memory storage
|
||||
|
||||
---
|
||||
|
||||
|
|
@ -194,7 +186,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Retriever: Get relevant memories
|
||||
result = await app.async_execute(
|
||||
name="retrieve_task_memory",
|
||||
|
|
@ -327,7 +319,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Memory Retrieval: Get personal memory fragments
|
||||
result = await app.async_execute(
|
||||
name="retrieve_personal_memory",
|
||||
|
|
@ -477,7 +469,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Generate usage guidelines from history
|
||||
result = await app.async_execute(
|
||||
name="summary_tool_memory",
|
||||
|
|
@ -485,7 +477,7 @@ async def main():
|
|||
tool_names="web_search"
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Retrieve tool guidelines before use
|
||||
result = await app.async_execute(
|
||||
name="retrieve_tool_memory",
|
||||
|
|
@ -651,7 +643,7 @@ async def main():
|
|||
path="./docs/library/"
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Query relevant memories
|
||||
result = await app.async_execute(
|
||||
name="retrieve_task_memory",
|
||||
|
|
@ -679,7 +671,7 @@ We tested ReMe on Appworld using qwen3-8b:
|
|||
| with ReMe | 0.109 **(+2.6%)** | 0.175 **(+3.5%)** | 0.281 **(+5.3%)** |
|
||||
|
||||
Pass@K measures the probability that at least one of the K generated samples successfully completes the task (
|
||||
score=1).
|
||||
score=1).
|
||||
The current experiment uses an internal AppWorld environment, which may have slight differences.
|
||||
|
||||
You can find more details on reproducing the experiment in [quickstart.md](docs/cookbook/appworld/quickstart.md).
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
# flake8: noqa: E402, E501
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
os.environ["APPWORLD_ROOT"] = "."
|
||||
|
|
@ -7,7 +9,6 @@ from dotenv import load_dotenv
|
|||
|
||||
load_dotenv("../../../.env")
|
||||
|
||||
import re
|
||||
import time
|
||||
import json
|
||||
import ray
|
||||
|
|
@ -18,26 +19,28 @@ from jinja2 import Template
|
|||
from loguru import logger
|
||||
from openai import OpenAI
|
||||
|
||||
from prompt import PROMPT_TEMPLATE, PROMPT_TEMPLATE_WITH_EXPERIENCE
|
||||
from prompt import PROMPT_TEMPLATE_WITH_EXPERIENCE
|
||||
|
||||
|
||||
@ray.remote
|
||||
class AppworldReactAgent:
|
||||
"""A minimal ReAct Agent for AppWorld tasks."""
|
||||
|
||||
def __init__(self,
|
||||
index: int,
|
||||
task_ids: List[str],
|
||||
experiment_name: str,
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.9,
|
||||
max_interactions: int = 30,
|
||||
max_response_size: int = 2048,
|
||||
num_runs: int = 1,
|
||||
use_task_memory: bool = False,
|
||||
make_task_memory: bool = False,
|
||||
api_url: str = "http://0.0.0.0:8002/",
|
||||
workspace_id: str="appworld_v1"):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
task_ids: List[str],
|
||||
experiment_name: str,
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.9,
|
||||
max_interactions: int = 30,
|
||||
max_response_size: int = 2048,
|
||||
num_runs: int = 1,
|
||||
use_task_memory: bool = False,
|
||||
make_task_memory: bool = False,
|
||||
api_url: str = "http://0.0.0.0:8002/",
|
||||
workspace_id: str = "appworld_v1",
|
||||
):
|
||||
|
||||
self.index: int = index
|
||||
self.task_ids: List[str] = task_ids
|
||||
|
|
@ -62,7 +65,8 @@ class AppworldReactAgent:
|
|||
messages=messages,
|
||||
temperature=self.temperature,
|
||||
extra_body={"enable_thinking": False},
|
||||
seed=0)
|
||||
seed=0,
|
||||
)
|
||||
|
||||
return response.choices[0].message.content
|
||||
|
||||
|
|
@ -72,13 +76,17 @@ class AppworldReactAgent:
|
|||
|
||||
return "call llm error"
|
||||
|
||||
def prompt_messages(self,world: AppWorld) -> list[dict]:
|
||||
def prompt_messages(self, world: AppWorld) -> list[dict]:
|
||||
if self.use_task_memory:
|
||||
task_memory = self.get_task_memory(world.task.instruction)
|
||||
logger.info(f"loaded task_memory: {task_memory}")
|
||||
dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction, "experience": task_memory}
|
||||
dictionary = {
|
||||
"supervisor": world.task.supervisor,
|
||||
"instruction": world.task.instruction,
|
||||
"experience": task_memory,
|
||||
}
|
||||
else:
|
||||
dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction ,"experience": ""}
|
||||
dictionary = {"supervisor": world.task.supervisor, "instruction": world.task.instruction, "experience": ""}
|
||||
print(dictionary)
|
||||
prompt = Template(PROMPT_TEMPLATE_WITH_EXPERIENCE.lstrip()).render(dictionary)
|
||||
messages: list[dict] = []
|
||||
|
|
@ -96,7 +104,7 @@ class AppworldReactAgent:
|
|||
# messages.append({"role": role_type, "content": None})
|
||||
# last_start = match.span()[1]
|
||||
# messages[-1]["content"] = prompt[last_start:]
|
||||
messages.append({"role":"user", "content":prompt})
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
return messages
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -122,7 +130,7 @@ class AppworldReactAgent:
|
|||
output = world.execute(code)
|
||||
if len(output) > self.max_response_size:
|
||||
# logger.warning(f"output exceed max size={len(output)}")
|
||||
output = output[:self.max_response_size]
|
||||
output = output[: self.max_response_size]
|
||||
history.append({"role": "user", "content": output})
|
||||
|
||||
if world.task_completed():
|
||||
|
|
@ -164,7 +172,7 @@ class AppworldReactAgent:
|
|||
json={
|
||||
"workspace_id": self.workspace_id,
|
||||
"query": query,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
|
|
@ -181,26 +189,28 @@ class AppworldReactAgent:
|
|||
if not result:
|
||||
print("No results to summarize")
|
||||
return
|
||||
|
||||
|
||||
# Prepare trajectories from results
|
||||
trajectories = []
|
||||
for r in result:
|
||||
if "task_history" in r:
|
||||
trajectories.append({
|
||||
"messages": r["task_history"],
|
||||
"score": float(r.get("uplift_score", 0.0))
|
||||
})
|
||||
|
||||
trajectories.append(
|
||||
{
|
||||
"messages": r["task_history"],
|
||||
"score": float(r.get("uplift_score", 0.0)),
|
||||
},
|
||||
)
|
||||
|
||||
if not trajectories:
|
||||
print("No trajectories to summarize")
|
||||
return
|
||||
|
||||
|
||||
response = requests.post(
|
||||
url=f"{self.api_url}summary_task_memory",
|
||||
json={
|
||||
"workspace_id": self.workspace_id,
|
||||
"trajectories": trajectories
|
||||
}
|
||||
"trajectories": trajectories,
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
|
|
@ -222,4 +232,4 @@ def main():
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
# flake8: noqa: E402, E501
|
||||
# This is a basic prompt template containing all the necessary onboarding information to solve AppWorld tasks. It explains the role of the agent and the supervisor, how to explore the API documentation, how to operate the interactive coding environment and call APIs via a simple task, and provides key instructions and disclaimers.
|
||||
|
||||
# You can adapt it as needed by your agent. You can also choose to bypass API docs app and build your own API retrieval, e.g., for FullCodeRefl, IPFunCall, etc, we asked an LLM to predict relevant APIs separately and put its documentation directly in the prompt.
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
# flake8: noqa: E402
|
||||
import os
|
||||
import time
|
||||
import requests
|
||||
|
||||
import ray
|
||||
import requests
|
||||
from ray import logger
|
||||
|
||||
os.environ["APPWORLD_ROOT"] = "."
|
||||
|
|
@ -35,7 +36,7 @@ def delete_workspace(workspace_id: str, api_url: str = "http://0.0.0.0:8002/"):
|
|||
json={
|
||||
"workspace_id": workspace_id,
|
||||
"action": "delete",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -51,7 +52,7 @@ def dump_memory(workspace_id: str, path: str = "./", api_url: str = "http://0.0.
|
|||
"workspace_id": workspace_id,
|
||||
"action": "dump",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -67,7 +68,7 @@ def load_memory(workspace_id: str, path: str = "docs/library", api_url: str = "h
|
|||
"workspace_id": workspace_id,
|
||||
"action": "load",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -75,7 +76,16 @@ def load_memory(workspace_id: str, path: str = "docs/library", api_url: str = "h
|
|||
print(f"Memory loaded from {path}")
|
||||
|
||||
|
||||
def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_runs: int = 1, use_task_memory: bool = False, make_task_memory: bool = False, workspace_id: str="appworld_v1", api_url: str = "http://0.0.0.0:8002/") :
|
||||
def run_agent(
|
||||
dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
max_workers: int,
|
||||
num_runs: int = 1,
|
||||
use_task_memory: bool = False,
|
||||
make_task_memory: bool = False,
|
||||
workspace_id: str = "appworld_v1",
|
||||
api_url: str = "http://0.0.0.0:8002/",
|
||||
):
|
||||
experiment_name = dataset_name + "_" + experiment_suffix
|
||||
path: Path = Path(f"./exp_result")
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
|
@ -93,14 +103,16 @@ def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_r
|
|||
for i in range(max_workers):
|
||||
# Assign tasks to each worker, ensuring each task runs num_runs times
|
||||
worker_task_ids = task_ids[i::max_workers]
|
||||
actor = AppworldReactAgent.remote(index=i,
|
||||
task_ids=worker_task_ids,
|
||||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=make_task_memory,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url)
|
||||
actor = AppworldReactAgent.remote(
|
||||
index=i,
|
||||
task_ids=worker_task_ids,
|
||||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=make_task_memory,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url,
|
||||
)
|
||||
future = actor.execute.remote()
|
||||
future_list.append(future)
|
||||
time.sleep(1)
|
||||
|
|
@ -119,14 +131,16 @@ def run_agent(dataset_name: str, experiment_suffix: str, max_workers: int, num_r
|
|||
|
||||
else:
|
||||
for index, task_id in enumerate(task_ids):
|
||||
agent = AppworldReactAgent(index=index,
|
||||
task_ids=[task_id],
|
||||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=make_task_memory,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url)
|
||||
agent = AppworldReactAgent(
|
||||
index=index,
|
||||
task_ids=[task_id],
|
||||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=make_task_memory,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url,
|
||||
)
|
||||
task_results = agent.execute()
|
||||
if isinstance(task_results, list):
|
||||
result.extend(task_results)
|
||||
|
|
@ -140,14 +154,14 @@ def main():
|
|||
num_runs = 1 # Run each task once
|
||||
workspace_id = "appworld"
|
||||
api_url = "http://0.0.0.0:8002/"
|
||||
|
||||
|
||||
if max_workers > 1:
|
||||
ray.init(num_cpus=8)
|
||||
|
||||
|
||||
# Clean up workspace before starting
|
||||
logger.info("Deleting workspace...")
|
||||
delete_workspace(workspace_id=workspace_id, api_url=api_url)
|
||||
|
||||
|
||||
# First run to build task memories
|
||||
logger.info("Start load experiments to build task memories")
|
||||
load_memory(workspace_id=workspace_id, api_url=api_url)
|
||||
|
|
@ -160,19 +174,30 @@ def main():
|
|||
|
||||
# Run experiments with task memory
|
||||
logger.info("Start running experiments with task memory")
|
||||
run_agent(dataset_name="dev", experiment_suffix=f"with-memory",
|
||||
max_workers=max_workers, num_runs=1,
|
||||
use_task_memory=True, make_task_memory=False,
|
||||
workspace_id=workspace_id, api_url=api_url)
|
||||
run_agent(
|
||||
dataset_name="dev",
|
||||
experiment_suffix=f"with-memory",
|
||||
max_workers=max_workers,
|
||||
num_runs=1,
|
||||
use_task_memory=True,
|
||||
make_task_memory=False,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url,
|
||||
)
|
||||
|
||||
# Run experiments without task memory
|
||||
logger.info("Start running experiments without task memory")
|
||||
run_agent(dataset_name="dev", experiment_suffix=f"no-memory",
|
||||
max_workers=max_workers, num_runs=1,
|
||||
use_task_memory=False, make_task_memory=False,
|
||||
workspace_id=workspace_id, api_url=api_url)
|
||||
|
||||
run_agent(
|
||||
dataset_name="dev",
|
||||
experiment_suffix=f"no-memory",
|
||||
max_workers=max_workers,
|
||||
num_runs=1,
|
||||
use_task_memory=False,
|
||||
make_task_memory=False,
|
||||
workspace_id=workspace_id,
|
||||
api_url=api_url,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import json
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
|
||||
|
||||
|
|
@ -24,7 +24,7 @@ def calculate_best_at_k(scores: list, k: int) -> float:
|
|||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
group = scores[i : i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
|
@ -36,8 +36,8 @@ def calculate_pass_at_k(scores: list, k: int) -> float:
|
|||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
is_pass = 1.0 if max(group) >=1.0 else 0.0
|
||||
group = scores[i : i + k]
|
||||
is_pass = 1.0 if max(group) >= 1.0 else 0.0
|
||||
group_maxs.append(is_pass)
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
|
@ -61,7 +61,7 @@ def get_possible_k_values(total_runs: int) -> list:
|
|||
|
||||
|
||||
def run_exp_statistic():
|
||||
path: Path = Path(f"./exp_result")
|
||||
path: Path = Path("./exp_result")
|
||||
|
||||
# Store results for all experiments
|
||||
all_results = {}
|
||||
|
|
@ -134,10 +134,10 @@ def run_exp_statistic():
|
|||
# Create and display table
|
||||
if all_results:
|
||||
df = pd.DataFrame(list(all_results.values()))
|
||||
df = df.set_index('file')
|
||||
df = df.set_index("file")
|
||||
|
||||
# Sort columns by the number in column name (best@8, best@4, best@2, best@1)
|
||||
pass_columns = [col for col in df.columns if col.startswith('pass@')]
|
||||
pass_columns = [col for col in df.columns if col.startswith("pass@")]
|
||||
# best_columns = [col for col in df.columns]
|
||||
pass_columns.sort(key=lambda x: x, reverse=False)
|
||||
df = df[pass_columns]
|
||||
|
|
@ -157,4 +157,4 @@ def run_exp_statistic():
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_exp_statistic()
|
||||
run_exp_statistic()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
# flake8: noqa: E402
|
||||
import os
|
||||
|
||||
os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl"
|
||||
|
|
@ -6,7 +7,6 @@ from dotenv import load_dotenv
|
|||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
import re
|
||||
import time
|
||||
import json
|
||||
import ray
|
||||
|
|
@ -29,7 +29,7 @@ from bfcl_utils import (
|
|||
extract_single_turn_response,
|
||||
extract_multi_turn_responses,
|
||||
capture_and_print_score_files,
|
||||
create_error_response
|
||||
create_error_response,
|
||||
)
|
||||
from bfcl_eval.model_handler.api_inference.qwen import QwenAPIHandler
|
||||
from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import (
|
||||
|
|
@ -47,30 +47,33 @@ from bfcl_eval.utils import (
|
|||
load_file,
|
||||
)
|
||||
|
||||
|
||||
@ray.remote
|
||||
class BFCLAgent:
|
||||
"""A minimal ReAct Agent for BFCL-v3(multi-turn) tasks."""
|
||||
|
||||
def __init__(self,
|
||||
index: int,
|
||||
task_ids: List[str],
|
||||
experiment_name: str,
|
||||
data_path: str = os.getenv("BFCL_DATA_PATH"),
|
||||
answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")),
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.9,
|
||||
max_interactions: int = 30,
|
||||
max_response_size: int = 2000,
|
||||
num_runs: int = 1,
|
||||
enable_thinking: bool = False,
|
||||
use_memory: bool = False,
|
||||
use_memory_addition: bool = False,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
memory_base_url: str = "http://0.0.0.0:8001/",
|
||||
memory_workspace_id: str = "bfcl_8b_0725"):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
task_ids: List[str],
|
||||
experiment_name: str,
|
||||
data_path: str = os.getenv("BFCL_DATA_PATH"),
|
||||
answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")),
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.9,
|
||||
max_interactions: int = 30,
|
||||
max_response_size: int = 2000,
|
||||
num_runs: int = 1,
|
||||
enable_thinking: bool = False,
|
||||
use_memory: bool = False,
|
||||
use_memory_addition: bool = False,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
memory_base_url: str = "http://0.0.0.0:8001/",
|
||||
memory_workspace_id: str = "bfcl_8b_0725",
|
||||
):
|
||||
|
||||
self.index: int = index
|
||||
self.task_ids: List[str] = task_ids
|
||||
|
|
@ -92,14 +95,14 @@ class BFCLAgent:
|
|||
self.utility_threshold: float = utility_threshold
|
||||
self.memory_base_url: str = memory_base_url
|
||||
self.memory_workspace_id: str = memory_workspace_id
|
||||
|
||||
|
||||
self.history: List[List[List[dict]]] = [[] for _ in range(num_runs)]
|
||||
self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_runs)]
|
||||
self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)]
|
||||
self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)]
|
||||
self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_runs)]
|
||||
self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_runs)]
|
||||
|
||||
|
||||
for run_id in range(num_runs):
|
||||
for task_index in range(len(task_ids)):
|
||||
self.init_state(run_id, task_index)
|
||||
|
|
@ -113,7 +116,7 @@ class BFCLAgent:
|
|||
if self.use_memory:
|
||||
query = msg["content"]
|
||||
response = self.get_memory(query)
|
||||
|
||||
|
||||
if len(response["metadata"]["memory_list"]):
|
||||
self.retrieved_memory_list[run_id].append(response["metadata"]["memory_list"])
|
||||
exp: str = response["answer"]
|
||||
|
|
@ -129,22 +132,25 @@ class BFCLAgent:
|
|||
def get_query_with_memory(self, query: str, memory: str):
|
||||
return {
|
||||
"role": "user",
|
||||
"content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory
|
||||
"content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory,
|
||||
}
|
||||
|
||||
def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float):
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"messages": task_history,
|
||||
"score": reward
|
||||
"score": reward,
|
||||
}
|
||||
|
||||
def get_memory(self, query: str):
|
||||
response = requests.post(url=self.memory_base_url + "retrieve_task_memory", json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"query": query,
|
||||
"top_k": 5
|
||||
})
|
||||
response = requests.post(
|
||||
url=self.memory_base_url + "retrieve_task_memory",
|
||||
json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"query": query,
|
||||
"top_k": 5,
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.info(response.text)
|
||||
|
|
@ -153,33 +159,42 @@ class BFCLAgent:
|
|||
response = response.json()
|
||||
logger.info(f"query: {query}, response: {response}")
|
||||
return response
|
||||
|
||||
|
||||
def add_memory(self, trajectories):
|
||||
response = requests.post(url=self.memory_base_url + "summary_task_memory", json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"trajectories": trajectories,
|
||||
})
|
||||
response = requests.post(
|
||||
url=self.memory_base_url + "summary_task_memory",
|
||||
json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"trajectories": trajectories,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
response = response.json()
|
||||
logger.info(f"add new memorys: {response["metadata"]["memory_list"]}")
|
||||
|
||||
def update_memory_information(self, memory_list, update_utility: bool=False):
|
||||
response = requests.post(url=self.memory_base_url + "record_task_memory", json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"memory_dicts": memory_list,
|
||||
"update_utility": update_utility
|
||||
})
|
||||
response.raise_for_status()
|
||||
logger.info(response.json())
|
||||
|
||||
def delete_memory(self):
|
||||
response = requests.post(url=self.memory_base_url + "delete_task_memory", json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"freq_threshold": self.freq_threshold,
|
||||
"utility_threshold": self.utility_threshold
|
||||
})
|
||||
logger.info(f'add new memorys: {response["metadata"]["memory_list"]}')
|
||||
|
||||
def update_memory_information(self, memory_list, update_utility: bool = False):
|
||||
response = requests.post(
|
||||
url=self.memory_base_url + "record_task_memory",
|
||||
json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"memory_dicts": memory_list,
|
||||
"update_utility": update_utility,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
logger.info(response.json())
|
||||
|
||||
def delete_memory(self):
|
||||
response = requests.post(
|
||||
url=self.memory_base_url + "delete_task_memory",
|
||||
json={
|
||||
"workspace_id": self.memory_workspace_id,
|
||||
"freq_threshold": self.freq_threshold,
|
||||
"utility_threshold": self.utility_threshold,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
def call_llm(self, messages: list, tool_schemas: list[dict]) -> str:
|
||||
for i in range(100):
|
||||
try:
|
||||
|
|
@ -200,10 +215,12 @@ class BFCLAgent:
|
|||
return out_msg.model_dump(exclude_unset=True, exclude_none=True)
|
||||
else:
|
||||
reasoning_content = "" # Complete reasoning process
|
||||
answer_content = "" # Define complete response
|
||||
tool_info = [] # Store tool invocation information
|
||||
is_answering = False # Determine whether the reasoning process has finished and response has started
|
||||
|
||||
answer_content = "" # Define complete response
|
||||
tool_info = [] # Store tool invocation information
|
||||
is_answering = (
|
||||
False # Determine whether the reasoning process has finished and response has started
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
if not chunk.choices:
|
||||
# Handle usage information
|
||||
|
|
@ -211,36 +228,43 @@ class BFCLAgent:
|
|||
else:
|
||||
delta = chunk.choices[0].delta
|
||||
# Handle AI's thought process (chain reasoning)
|
||||
if hasattr(delta, 'reasoning_content') and delta.reasoning_content is not None:
|
||||
if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None:
|
||||
reasoning_content += delta.reasoning_content
|
||||
|
||||
|
||||
# Handle final response content
|
||||
else:
|
||||
if not is_answering: # Print title when entering the response phase for the first time
|
||||
is_answering = True
|
||||
if delta.content is not None:
|
||||
answer_content += delta.content
|
||||
|
||||
|
||||
# Handle tool invocation information (support parallel tool calls)
|
||||
if delta.tool_calls is not None:
|
||||
for tool_call in delta.tool_calls:
|
||||
index = tool_call.index # Tool call index, used for parallel calls
|
||||
|
||||
|
||||
# Dynamically expand tool information storage list
|
||||
while len(tool_info) <= index:
|
||||
tool_info.append({"id": "", "type": "function", "index": index, "function": { "name": "", "arguments": "" }})
|
||||
|
||||
tool_info.append(
|
||||
{
|
||||
"id": "",
|
||||
"type": "function",
|
||||
"index": index,
|
||||
"function": {"name": "", "arguments": ""},
|
||||
},
|
||||
)
|
||||
|
||||
# Collect tool call ID (used for subsequent function calls)
|
||||
if tool_call.id:
|
||||
tool_info[index]['id'] += tool_call.id
|
||||
|
||||
tool_info[index]["id"] += tool_call.id
|
||||
|
||||
# Collect function name (used for subsequent routing to specific functions)
|
||||
if tool_call.function and tool_call.function.name:
|
||||
tool_info[index]['function']['name'] += tool_call.function.name
|
||||
|
||||
tool_info[index]["function"]["name"] += tool_call.function.name
|
||||
|
||||
# Collect function parameters (in JSON string format, need subsequent parsing)
|
||||
if tool_call.function and tool_call.function.arguments:
|
||||
tool_info[index]['function']['arguments'] += tool_call.function.arguments
|
||||
tool_info[index]["function"]["arguments"] += tool_call.function.arguments
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": answer_content,
|
||||
|
|
@ -269,29 +293,35 @@ class BFCLAgent:
|
|||
Dict containing next message and tools if applicable
|
||||
"""
|
||||
try:
|
||||
if not messages:
|
||||
if not messages:
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
|
||||
if messages[-1]["role"] != "assistant":
|
||||
return create_error_response(
|
||||
"Last message must be from assistant"
|
||||
"Last message must be from assistant",
|
||||
)
|
||||
|
||||
if "tool_calls" in messages[-1] and len(messages[-1]["tool_calls"]) > 0:
|
||||
try:
|
||||
tool_calls = messages[-1]["tool_calls"]
|
||||
decoded_calls = self._convert_tool_calls_to_execution_format(
|
||||
tool_calls
|
||||
tool_calls,
|
||||
)
|
||||
# decoded_calls:[function(param=xxx)]
|
||||
print(f"decoded_calls: {decoded_calls}")
|
||||
if is_empty_execute_response(decoded_calls):
|
||||
warnings.warn(
|
||||
f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}"
|
||||
f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}",
|
||||
)
|
||||
return handle_user_turn(
|
||||
self.original_test_entry[run_id][index],
|
||||
self.current_turn[run_id][index],
|
||||
)
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
return handle_tool_calls(
|
||||
tool_calls, decoded_calls, self.original_test_entry[run_id][index], self.current_turn[run_id][index]
|
||||
tool_calls,
|
||||
decoded_calls,
|
||||
self.original_test_entry[run_id][index],
|
||||
self.current_turn[run_id][index],
|
||||
)
|
||||
except Exception as e:
|
||||
warnings.warn(f"Errors during tool invocation: {str(e)}")
|
||||
|
|
@ -303,7 +333,8 @@ class BFCLAgent:
|
|||
return create_error_response(f"Failed to process request: {str(e)}")
|
||||
|
||||
def _convert_tool_calls_to_execution_format(
|
||||
self, tool_calls: List[Dict[str, Any]]
|
||||
self,
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
) -> List[str]:
|
||||
"""
|
||||
Convert OpenAI format tool calls to execution format.
|
||||
|
|
@ -334,7 +365,7 @@ class BFCLAgent:
|
|||
execution_list.append(f"{function_name}()")
|
||||
|
||||
return execution_list
|
||||
|
||||
|
||||
def get_reward(self, run_id, index) -> float:
|
||||
try:
|
||||
if not self.history[run_id][index] or not self.original_test_entry[run_id][index]:
|
||||
|
|
@ -342,7 +373,8 @@ class BFCLAgent:
|
|||
|
||||
model_name = "env_handler"
|
||||
handler = QwenAPIHandler(
|
||||
model_name, temperature=1.0
|
||||
model_name,
|
||||
temperature=1.0,
|
||||
) # FIXME: magic number
|
||||
|
||||
model_result_data = self._convert_conversation_to_eval_format(run_id, index)
|
||||
|
|
@ -351,23 +383,28 @@ class BFCLAgent:
|
|||
|
||||
state = {"leaderboard_table": {}}
|
||||
record_cost_latency(
|
||||
state["leaderboard_table"], model_name, [model_result_data]
|
||||
state["leaderboard_table"],
|
||||
model_name,
|
||||
[model_result_data],
|
||||
)
|
||||
|
||||
if is_relevance_or_irrelevance(self.categories[index]):
|
||||
accuracy, _ = self._eval_relevance_test(
|
||||
handler, model_result_data, prompt_data, model_name, self.category
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
model_name,
|
||||
self.category,
|
||||
)
|
||||
else:
|
||||
# Find the corresponding possible answer file
|
||||
|
||||
possible_answer_file = find_file_with_suffix(
|
||||
self.answer_path, self.categories[index]
|
||||
self.answer_path,
|
||||
self.categories[index],
|
||||
)
|
||||
possible_answer = load_file(possible_answer_file, sort_by_id=True)
|
||||
possible_answer = [
|
||||
item for item in possible_answer if item["id"] == self.task_ids[index]
|
||||
]
|
||||
possible_answer = [item for item in possible_answer if item["id"] == self.task_ids[index]]
|
||||
if is_multi_turn(self.categories[index]):
|
||||
accuracy, _ = self._eval_multi_turn_test(
|
||||
handler,
|
||||
|
|
@ -396,7 +433,7 @@ class BFCLAgent:
|
|||
|
||||
traceback.print_exc()
|
||||
return 0
|
||||
|
||||
|
||||
def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert conversation history to evaluation format.
|
||||
|
|
@ -422,7 +459,7 @@ class BFCLAgent:
|
|||
}
|
||||
|
||||
return model_result_data
|
||||
|
||||
|
||||
def _eval_multi_turn_test(
|
||||
self,
|
||||
handler,
|
||||
|
|
@ -458,10 +495,13 @@ class BFCLAgent:
|
|||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir, model_name, test_category, "multi_turn"
|
||||
score_dir,
|
||||
model_name,
|
||||
test_category,
|
||||
"multi_turn",
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
|
||||
def _eval_single_turn_test(
|
||||
self,
|
||||
handler,
|
||||
|
|
@ -504,7 +544,10 @@ class BFCLAgent:
|
|||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir, model_name, test_category, "single_turn"
|
||||
score_dir,
|
||||
model_name,
|
||||
test_category,
|
||||
"single_turn",
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
|
|
@ -516,37 +559,42 @@ class BFCLAgent:
|
|||
try:
|
||||
start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
for i in range(self.max_interactions):
|
||||
llm_output = self.call_llm(self.history[run_id][task_index], self.tool_schema[run_id][task_index])
|
||||
llm_output = self.call_llm(
|
||||
self.history[run_id][task_index],
|
||||
self.tool_schema[run_id][task_index],
|
||||
)
|
||||
self.history[run_id][task_index].append(llm_output)
|
||||
|
||||
env_output = self.env_step(run_id, task_index, self.history[run_id][task_index])
|
||||
# Possible env_output returns after environment interaction:
|
||||
# 1. Triggers a query with available tools list: {"messages": [{"role": "user", "content": user_query}], "tools": tools}
|
||||
# 1. Triggers a query with available tools list: {"messages": [{"role": "user", "content": user_query}], "tools": tools}
|
||||
# 2. Returns tool invocation result: {"messages": [{"role": "tool", "content": {<execution_results>}, 'tool_call_id': 'chatcmpl-tool-xxx'}]}
|
||||
# <execution_results>: when success, returns result dicts, e.g., {"travel_cost_list": [1140.0]}, when error, returns error message, e.g., {"error": "cd: temporary: No such directory. You cannot use path to change directory."}
|
||||
# 3. Conversation completion: {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
|
||||
# 4. Program error: {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
|
||||
|
||||
|
||||
# tool_list update
|
||||
if "tools" in env_output:
|
||||
self.tool_schema[run_id][task_index] = extract_tool_schema(env_output["tools"])
|
||||
|
||||
new_tool_calls=[]
|
||||
new_tool_call_ids=[]
|
||||
|
||||
new_tool_calls = []
|
||||
new_tool_call_ids = []
|
||||
next_user_msg = ""
|
||||
for idx, msg in enumerate(env_output.get("messages", [])):
|
||||
if msg["role"] == "tool" and len(msg["content"])>0:
|
||||
if msg["role"] == "tool" and len(msg["content"]) > 0:
|
||||
new_tool_calls.append(msg.get("content", ""))
|
||||
new_tool_call_ids.append(msg.get("tool_call_id", ""))
|
||||
elif msg["role"] == "user":
|
||||
next_user_msg = msg.get("content", "")
|
||||
self.current_turn[run_id][task_index] += 1
|
||||
else: # for env role messages
|
||||
else: # for env role messages
|
||||
next_user_msg = msg.get("content", "")
|
||||
|
||||
|
||||
if new_tool_calls:
|
||||
for idx, call in enumerate(new_tool_calls):
|
||||
self.history[run_id][task_index].append({"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]})
|
||||
self.history[run_id][task_index].append(
|
||||
{"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]},
|
||||
)
|
||||
else:
|
||||
self.history[run_id][task_index].append({"role": "user", "content": next_user_msg})
|
||||
|
||||
|
|
@ -557,14 +605,16 @@ class BFCLAgent:
|
|||
|
||||
reward = self.get_reward(run_id, task_index)
|
||||
if self.use_memory:
|
||||
if reward == 1 and self.use_memory_addition: # selectively add memories when succeed
|
||||
new_traj_list = [self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward)]
|
||||
self.add_memory(new_traj_list)
|
||||
|
||||
if reward == 1 and self.use_memory_addition: # selectively add memories when succeed
|
||||
new_traj_list = [
|
||||
self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward),
|
||||
]
|
||||
self.add_memory(new_traj_list)
|
||||
|
||||
# update the freq & utility attributes of retrieved memories
|
||||
update_utility: bool = (reward == 1)
|
||||
update_utility: bool = reward == 1
|
||||
self.update_memory_information(self.retrieved_memory_list[run_id][task_index], update_utility)
|
||||
|
||||
|
||||
counter += 1
|
||||
if self.use_memory_deletion and counter % self.delete_freq == 0:
|
||||
self.delete_memory()
|
||||
|
|
@ -594,14 +644,15 @@ class BFCLAgent:
|
|||
"""
|
||||
return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]"
|
||||
|
||||
|
||||
def main():
|
||||
with open(os.getenv("BFCL_DATA_PATH"), "r", encoding="utf-8") as f:
|
||||
task_ids = [json.loads(l)["id"] for l in f]
|
||||
dataset_name = "dev"
|
||||
agent = BFCLAgent(
|
||||
index=0,
|
||||
task_id=task_ids[0],
|
||||
experiment_name=f"zouying_{dataset_name}",
|
||||
index=0,
|
||||
task_id=task_ids[0],
|
||||
experiment_name=f"zouying_{dataset_name}",
|
||||
)
|
||||
result = agent.execute()
|
||||
logger.info(f"result={json.dumps(result)}")
|
||||
|
|
|
|||
|
|
@ -1,26 +1,20 @@
|
|||
import json
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Any
|
||||
|
||||
from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI
|
||||
from bfcl_eval.constants.default_prompts import (
|
||||
DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
|
||||
)
|
||||
from bfcl_eval.model_handler.model_style import ModelStyle
|
||||
from bfcl_eval.model_handler.utils import (
|
||||
convert_to_function_call,
|
||||
convert_to_tool,
|
||||
default_decode_ast_prompting,
|
||||
default_decode_execute_prompting,
|
||||
format_execution_results_prompting,
|
||||
func_doc_language_specific_pre_processing,
|
||||
retry_with_backoff,
|
||||
system_prompt_pre_processing_chat_model,
|
||||
)
|
||||
from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI
|
||||
from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import (
|
||||
execute_multi_turn_func_call,
|
||||
)
|
||||
from bfcl_eval.model_handler.model_style import ModelStyle
|
||||
from bfcl_eval.model_handler.utils import (
|
||||
convert_to_tool,
|
||||
default_decode_execute_prompting,
|
||||
func_doc_language_specific_pre_processing,
|
||||
)
|
||||
|
||||
|
||||
def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]:
|
||||
|
|
@ -44,8 +38,10 @@ def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]:
|
|||
return data
|
||||
raise ValueError(f"Test case id '{test_id}' not found in {data_path}")
|
||||
|
||||
|
||||
def handle_user_turn(
|
||||
test_entry: Dict[str, Any], current_turn: int
|
||||
test_entry: Dict[str, Any],
|
||||
current_turn: int,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle user turn by returning appropriate content from test_entry["question"].
|
||||
|
|
@ -67,19 +63,17 @@ def handle_user_turn(
|
|||
if str(current_turn) in holdout_function:
|
||||
test_entry["function"].extend(holdout_function[str(current_turn)])
|
||||
tools = compile_tools(test_entry)
|
||||
assert (
|
||||
len(questions[current_turn]) == 0
|
||||
), "Holdout turn should not have user message."
|
||||
assert len(questions[current_turn]) == 0, "Holdout turn should not have user message."
|
||||
current_turn_message = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
|
||||
}
|
||||
},
|
||||
]
|
||||
return create_user_response(current_turn_message, tools)
|
||||
if current_turn >= len(questions):
|
||||
return create_completion_response()
|
||||
|
||||
|
||||
current_turn_message = questions[current_turn]
|
||||
|
||||
return create_user_response(current_turn_message, tools)
|
||||
|
|
@ -87,6 +81,7 @@ def handle_user_turn(
|
|||
except Exception as e:
|
||||
return create_error_response(f"Failed to process user message: {str(e)}")
|
||||
|
||||
|
||||
def handle_tool_calls(
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
decoded_calls: list[str],
|
||||
|
|
@ -111,9 +106,7 @@ def handle_tool_calls(
|
|||
involved_classes=test_entry["involved_classes"],
|
||||
model_name="env_handler",
|
||||
test_entry_id=test_entry["id"],
|
||||
long_context=(
|
||||
"long_context" in test_entry["id"] or "composite" in test_entry["id"]
|
||||
),
|
||||
long_context=("long_context" in test_entry["id"] or "composite" in test_entry["id"]),
|
||||
is_evaL_run=False,
|
||||
)
|
||||
# print('execution_results in handler_tool_calls:', execution_results)
|
||||
|
|
@ -138,9 +131,11 @@ def compile_tools(test_entry: dict) -> list:
|
|||
tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions)
|
||||
|
||||
return tools
|
||||
|
||||
|
||||
|
||||
def create_tool_response(
|
||||
tool_calls: List[Dict[str, Any]], execution_results: List[str]
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
execution_results: List[str],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for tool calls.
|
||||
|
|
@ -159,13 +154,15 @@ def create_tool_response(
|
|||
"role": "tool",
|
||||
"content": result,
|
||||
"tool_call_id": tool_call.get("id", f"call_{i}"),
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
return {"messages": tool_messages}
|
||||
|
||||
|
||||
def create_user_response(
|
||||
question_turn: List[Dict[str, Any]], tools: List[Dict[str, Any]]
|
||||
question_turn: List[Dict[str, Any]],
|
||||
tools: List[Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response containing user message.
|
||||
|
|
@ -173,7 +170,7 @@ def create_user_response(
|
|||
Args:
|
||||
question_turn: List of messages for current turn
|
||||
tools: List of available tools
|
||||
|
||||
|
||||
Returns:
|
||||
Response containing user message and tools
|
||||
"""
|
||||
|
|
@ -185,6 +182,7 @@ def create_user_response(
|
|||
|
||||
return {"messages": [{"role": "user", "content": user_content}], "tools": tools}
|
||||
|
||||
|
||||
def create_completion_response() -> Dict[str, Any]:
|
||||
"""
|
||||
Create response indicating conversation completion.
|
||||
|
|
@ -194,6 +192,7 @@ def create_completion_response() -> Dict[str, Any]:
|
|||
"""
|
||||
return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
|
||||
|
||||
|
||||
def create_error_response(error_message: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for error conditions.
|
||||
|
|
@ -206,6 +205,7 @@ def create_error_response(error_message: str) -> Dict[str, Any]:
|
|||
"""
|
||||
return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
|
||||
|
||||
|
||||
def decode_execute(result):
|
||||
"""
|
||||
Decode execute results for compatibility with evaluation framework.
|
||||
|
|
@ -218,6 +218,7 @@ def decode_execute(result):
|
|||
"""
|
||||
return default_decode_execute_prompting(result)
|
||||
|
||||
|
||||
def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
|
||||
"""
|
||||
Extract single-turn response from conversation messages.
|
||||
|
|
@ -234,7 +235,7 @@ def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
|
|||
formatted_calls = []
|
||||
for tool_call in message["tool_calls"]:
|
||||
formatted_call = format_single_tool_call_for_eval(
|
||||
tool_call
|
||||
tool_call,
|
||||
)
|
||||
if formatted_call:
|
||||
formatted_calls.append(formatted_call)
|
||||
|
|
@ -243,9 +244,10 @@ def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
|
|||
return message["content"]
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
|
||||
def extract_multi_turn_responses(
|
||||
messages: List[Dict[str, Any]]
|
||||
messages: List[Dict[str, Any]],
|
||||
) -> List[List[str]]:
|
||||
"""
|
||||
Extract multi-turn responses from conversation messages.
|
||||
|
|
@ -275,7 +277,7 @@ def extract_multi_turn_responses(
|
|||
if "tool_calls" in assistant_msg and assistant_msg["tool_calls"]:
|
||||
for tool_call in assistant_msg["tool_calls"]:
|
||||
formatted_call = format_single_tool_call_for_eval(
|
||||
tool_call
|
||||
tool_call,
|
||||
)
|
||||
if formatted_call:
|
||||
current_turn_responses.append(formatted_call)
|
||||
|
|
@ -292,10 +294,11 @@ def extract_multi_turn_responses(
|
|||
|
||||
return turns_data
|
||||
|
||||
|
||||
def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Format a single tool call into string representation for evaluation.
|
||||
|
||||
|
||||
Args:
|
||||
tool_call: Single tool call in OpenAI format
|
||||
|
||||
|
|
@ -315,11 +318,15 @@ def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str:
|
|||
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
|
||||
return f"{function_name}({args_str})"
|
||||
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
return f"{function_name}()"
|
||||
|
||||
|
||||
def capture_and_print_score_files(
|
||||
score_dir: Path, model_name: str, test_category: str, eval_type: str
|
||||
score_dir: Path,
|
||||
model_name: str,
|
||||
test_category: str,
|
||||
eval_type: str,
|
||||
):
|
||||
"""
|
||||
Capture and print contents of score files written to score_dir.
|
||||
|
|
@ -360,8 +367,10 @@ def capture_and_print_score_files(
|
|||
parsed = json.loads(line)
|
||||
formatted_lines.append(
|
||||
json.dumps(
|
||||
parsed, ensure_ascii=False, indent=2
|
||||
)
|
||||
parsed,
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
),
|
||||
)
|
||||
content = "\n".join(formatted_lines)
|
||||
except json.JSONDecodeError:
|
||||
|
|
@ -379,7 +388,8 @@ def capture_and_print_score_files(
|
|||
except Exception as e:
|
||||
print(f"Error capturing evaluation result files: {str(e)}")
|
||||
|
||||
|
||||
def extract_tool_schema(tools):
|
||||
for i in range(len(tools)):
|
||||
tools[i]['function'].pop("response")
|
||||
tools[i]["function"].pop("response")
|
||||
return tools
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import json
|
||||
import requests
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Any
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Any
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def load_task_case(data_path: str, task_id: str | None) -> Dict[str, Any]:
|
||||
|
|
@ -34,32 +35,33 @@ def get_tool_prompt(tools):
|
|||
tool_prompt = "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>"
|
||||
for tool in tools:
|
||||
tool_prompt += "\n" + json.dumps(tool)
|
||||
tool_prompt += "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call>"
|
||||
tool_prompt += '\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{"name": <function-name>, "arguments": <args-json-object>}\n</tool_call>'
|
||||
return tool_prompt
|
||||
|
||||
|
||||
def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[List[Any]]:
|
||||
"""
|
||||
根据task_id字段对trajectories进行分组
|
||||
|
||||
|
||||
Args:
|
||||
jsonl_entries: JSONL条目列表
|
||||
|
||||
|
||||
Returns:
|
||||
List[List[Any]]: 按task_id分组的trajectory列表
|
||||
"""
|
||||
# 按task_id分组
|
||||
grouped = defaultdict(list)
|
||||
|
||||
|
||||
for entry in jsonl_entries:
|
||||
task_id = entry.get("task_id", "")
|
||||
taks_case = load_task_case("data/multiturn_data_base.jsonl", task_id)
|
||||
tools = taks_case.get("tools", [{}])
|
||||
from bfcl_utils import extract_tool_schema
|
||||
|
||||
tool_schema = extract_tool_schema(tools)
|
||||
entry["task_history"][0]["content"] += get_tool_prompt(tool_schema)
|
||||
grouped[task_id].append(entry)
|
||||
|
||||
|
||||
# 对每组只保留最大和最小reward的两个
|
||||
filtered_groups = []
|
||||
for key, trajectories in grouped.items():
|
||||
|
|
@ -75,33 +77,36 @@ def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[L
|
|||
min_reward_traj = trajectories[0] # 最小reward
|
||||
max_reward_traj = trajectories[-1] # 最大reward
|
||||
filtered_groups.append([min_reward_traj, max_reward_traj])
|
||||
|
||||
|
||||
return filtered_groups
|
||||
|
||||
|
||||
def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: str) -> Dict[str, Any]:
|
||||
"""
|
||||
将trajectories发送到summarizer服务
|
||||
|
||||
|
||||
Args:
|
||||
trajectories: trajectory列表
|
||||
service_url: 服务URL
|
||||
workspace_id: 工作空间ID
|
||||
|
||||
|
||||
Returns:
|
||||
响应结果
|
||||
"""
|
||||
trajectory_dicts = [{
|
||||
"task_id": traj["task_id"],
|
||||
"messages": traj["task_history"],
|
||||
"score": traj["reward"]
|
||||
} for traj in trajectories]
|
||||
trajectory_dicts = [
|
||||
{
|
||||
"task_id": traj["task_id"],
|
||||
"messages": traj["task_history"],
|
||||
"score": traj["reward"],
|
||||
}
|
||||
for traj in trajectories
|
||||
]
|
||||
|
||||
request_data = {
|
||||
"traj_list": trajectory_dicts,
|
||||
"workspace_id": workspace_id
|
||||
"workspace_id": workspace_id,
|
||||
}
|
||||
|
||||
|
||||
try:
|
||||
response = requests.post(f"{service_url}/summarizer", json=request_data)
|
||||
response.raise_for_status()
|
||||
|
|
@ -110,31 +115,33 @@ def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id:
|
|||
return {"error": str(e), "trajectories_count": len(trajectories)}
|
||||
|
||||
|
||||
def process_trajectories_with_threads(grouped_trajectories: List[List[Any]],
|
||||
service_url: str,
|
||||
workspace_id: str,
|
||||
n_threads: int = 4) -> List[Dict[str, Any]]:
|
||||
def process_trajectories_with_threads(
|
||||
grouped_trajectories: List[List[Any]],
|
||||
service_url: str,
|
||||
workspace_id: str,
|
||||
n_threads: int = 4,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
使用多线程处理trajectories组
|
||||
|
||||
|
||||
Args:
|
||||
grouped_trajectories: 按task_id分组的trajectory列表
|
||||
service_url: summarizer服务URL
|
||||
workspace_id: 工作空间ID
|
||||
n_threads: 线程数
|
||||
|
||||
|
||||
Returns:
|
||||
所有结果列表
|
||||
"""
|
||||
results = []
|
||||
|
||||
|
||||
with ThreadPoolExecutor(max_workers=n_threads) as executor:
|
||||
# 提交所有任务
|
||||
future_to_group = {
|
||||
executor.submit(post_to_summarizer, group, service_url, workspace_id): i
|
||||
executor.submit(post_to_summarizer, group, service_url, workspace_id): i
|
||||
for i, group in enumerate(grouped_trajectories)
|
||||
}
|
||||
|
||||
|
||||
# 收集结果
|
||||
for future in as_completed(future_to_group):
|
||||
group_index = future_to_group[future]
|
||||
|
|
@ -143,16 +150,18 @@ def process_trajectories_with_threads(grouped_trajectories: List[List[Any]],
|
|||
result["group_index"] = group_index
|
||||
result["group_size"] = len(grouped_trajectories[group_index])
|
||||
results.append(result)
|
||||
print(f"✅ Group {group_index} processed: {result.get('experience_list', 0) if 'experience_list' in result else 'error'}")
|
||||
print(
|
||||
f"✅ Group {group_index} processed: {result.get('experience_list', 0) if 'experience_list' in result else 'error'}",
|
||||
)
|
||||
except Exception as e:
|
||||
error_result = {
|
||||
"group_index": group_index,
|
||||
"group_size": len(grouped_trajectories[group_index]),
|
||||
"error": str(e)
|
||||
"error": str(e),
|
||||
}
|
||||
results.append(error_result)
|
||||
print(f"❌ Group {group_index} failed: {e}")
|
||||
|
||||
|
||||
return results
|
||||
|
||||
|
||||
|
|
@ -160,20 +169,20 @@ def main():
|
|||
"""
|
||||
主函数,支持命令行参数
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description='Convert JSONL to experiences using experience maker service')
|
||||
parser.add_argument('--jsonl_file', type=str, required=True, help='Path to the JSONL file')
|
||||
parser.add_argument('--service_url', type=str, default='http://localhost:8001', help='Experience maker service URL')
|
||||
parser.add_argument('--workspace_id', type=str, required=True, help='Workspace ID for the experience')
|
||||
parser.add_argument('--output_file', type=str, help='Output file to save results (optional)')
|
||||
parser.add_argument('--n_threads', type=int, default=4, help='Number of threads for processing')
|
||||
|
||||
parser = argparse.ArgumentParser(description="Convert JSONL to experiences using experience maker service")
|
||||
parser.add_argument("--jsonl_file", type=str, required=True, help="Path to the JSONL file")
|
||||
parser.add_argument("--service_url", type=str, default="http://localhost:8001", help="Experience maker service URL")
|
||||
parser.add_argument("--workspace_id", type=str, required=True, help="Workspace ID for the experience")
|
||||
parser.add_argument("--output_file", type=str, help="Output file to save results (optional)")
|
||||
parser.add_argument("--n_threads", type=int, default=4, help="Number of threads for processing")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
print(f"Processing JSONL file: {args.jsonl_file}")
|
||||
print(f"Service URL: {args.service_url}")
|
||||
print(f"Workspace ID: {args.workspace_id}")
|
||||
print(f"Threads: {args.n_threads}")
|
||||
|
||||
|
||||
# 读取JSONL文件
|
||||
try:
|
||||
with open(args.jsonl_file, "r") as f:
|
||||
|
|
@ -182,31 +191,30 @@ def main():
|
|||
except Exception as e:
|
||||
print(f"Error reading JSONL file: {e}")
|
||||
return
|
||||
|
||||
|
||||
# 分组处理
|
||||
grouped_trajectories = group_trajectories_by_task_id(data)
|
||||
print(f"Total groups: {len(grouped_trajectories)}")
|
||||
|
||||
|
||||
# 多线程处理
|
||||
results = process_trajectories_with_threads(
|
||||
grouped_trajectories,
|
||||
args.service_url,
|
||||
grouped_trajectories,
|
||||
args.service_url,
|
||||
args.workspace_id,
|
||||
n_threads=args.n_threads
|
||||
n_threads=args.n_threads,
|
||||
)
|
||||
|
||||
print(f"Processed {len(results)} groups")
|
||||
|
||||
# 统计结果
|
||||
success_count = sum(1 for r in results if 'error' not in r)
|
||||
error_count = len(results) - success_count
|
||||
total_experiences = sum(len(r.get('experiences', [])) for r in results if 'experiences' in r)
|
||||
|
||||
|
||||
print(f"Processed {len(results)} groups")
|
||||
|
||||
# 统计结果
|
||||
success_count = sum(1 for r in results if "error" not in r)
|
||||
error_count = len(results) - success_count
|
||||
total_experiences = sum(len(r.get("experiences", [])) for r in results if "experiences" in r)
|
||||
|
||||
print(f"✅ Success: {success_count}")
|
||||
print(f"❌ Errors: {error_count}")
|
||||
print(f"📊 Total experiences created: {total_experiences}")
|
||||
|
||||
|
||||
# 保存结果到文件
|
||||
if args.output_file:
|
||||
try:
|
||||
|
|
@ -217,10 +225,10 @@ def main():
|
|||
"success_count": success_count,
|
||||
"error_count": error_count,
|
||||
"total_experiences": total_experiences,
|
||||
"results": results
|
||||
"results": results,
|
||||
}
|
||||
|
||||
with open(args.output_file, 'w') as f:
|
||||
|
||||
with open(args.output_file, "w") as f:
|
||||
json.dump(summary, f, indent=2)
|
||||
print(f"Results saved to: {args.output_file}")
|
||||
except Exception as e:
|
||||
|
|
@ -231,6 +239,7 @@ def main():
|
|||
if __name__ == "__main__":
|
||||
# 检查是否有命令行参数
|
||||
import sys
|
||||
|
||||
if len(sys.argv) > 1:
|
||||
# 使用新的命令行接口
|
||||
main()
|
||||
|
|
@ -239,15 +248,15 @@ if __name__ == "__main__":
|
|||
print("Running in compatibility mode...")
|
||||
with open("exp_result/qwen-max-2025-01-25/no_think/bfcl-multi-turn-base-train50_wo-exp.jsonl", "r") as f:
|
||||
data = [json.loads(line) for line in f]
|
||||
|
||||
|
||||
# 分组
|
||||
grouped_trajectories = group_trajectories_by_task_id(data)
|
||||
print(f"Total groups: {len(grouped_trajectories)}")
|
||||
|
||||
|
||||
results = process_trajectories_with_threads(
|
||||
grouped_trajectories,
|
||||
"http://localhost:8001",
|
||||
grouped_trajectories,
|
||||
"http://localhost:8001",
|
||||
"bfcl_train50_qwen_max_2025_01_25_extract_compare_validate",
|
||||
n_threads=4
|
||||
n_threads=4,
|
||||
)
|
||||
print(f"Processed {len(results)} groups")
|
||||
print(f"Processed {len(results)} groups")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import json
|
||||
import requests
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Any
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Any
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def load_task_case(data_path: str, task_id: str | None) -> Dict[str, Any]:
|
||||
|
|
@ -36,31 +37,32 @@ def get_tool_prompt(tools):
|
|||
tool_prompt = "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>"
|
||||
for tool in tools:
|
||||
tool_prompt += "\n" + json.dumps(tool)
|
||||
tool_prompt += "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call>"
|
||||
tool_prompt += '\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{"name": <function-name>, "arguments": <args-json-object>}\n</tool_call>'
|
||||
return tool_prompt
|
||||
|
||||
|
||||
def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[List[Any]]:
|
||||
"""
|
||||
group trajectories by task_id
|
||||
|
||||
|
||||
Args:
|
||||
jsonl_entries: JSONL entry list
|
||||
|
||||
|
||||
Returns:
|
||||
List[List[Any]]: trajectory list grouped by task_id
|
||||
"""
|
||||
grouped = defaultdict(list)
|
||||
|
||||
|
||||
for entry in jsonl_entries:
|
||||
task_id = entry.get("task_id", "")
|
||||
taks_case = load_task_case("data/multiturn_data_base.jsonl", task_id)
|
||||
tools = taks_case.get("tools", [{}])
|
||||
from bfcl_utils import extract_tool_schema
|
||||
|
||||
tool_schema = extract_tool_schema(tools)
|
||||
entry["task_history"][0]["content"] += get_tool_prompt(tool_schema)
|
||||
grouped[task_id].append(entry)
|
||||
|
||||
|
||||
# retain only the two with the highest and lowest rewards
|
||||
filtered_groups = []
|
||||
for key, trajectories in grouped.items():
|
||||
|
|
@ -76,22 +78,25 @@ def group_trajectories_by_task_id(jsonl_entries: List[Dict[str, Any]]) -> List[L
|
|||
min_reward_traj = trajectories[0] # highest reward
|
||||
max_reward_traj = trajectories[-1] # lowest reward
|
||||
filtered_groups.append([min_reward_traj, max_reward_traj])
|
||||
|
||||
|
||||
return filtered_groups
|
||||
|
||||
|
||||
def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id: str) -> Dict[str, Any]:
|
||||
trajectory_dicts = [{
|
||||
"task_id": traj["task_id"],
|
||||
"messages": traj["task_history"],
|
||||
"score": traj["reward"]
|
||||
} for traj in trajectories]
|
||||
trajectory_dicts = [
|
||||
{
|
||||
"task_id": traj["task_id"],
|
||||
"messages": traj["task_history"],
|
||||
"score": traj["reward"],
|
||||
}
|
||||
for traj in trajectories
|
||||
]
|
||||
|
||||
request_data = {
|
||||
"trajectories": trajectory_dicts,
|
||||
"workspace_id": workspace_id
|
||||
"workspace_id": workspace_id,
|
||||
}
|
||||
|
||||
|
||||
try:
|
||||
response = requests.post(f"{service_url}/summary_task_memory", json=request_data)
|
||||
response.raise_for_status()
|
||||
|
|
@ -100,30 +105,32 @@ def post_to_summarizer(trajectories: List[Any], service_url: str, workspace_id:
|
|||
return {"error": str(e), "trajectories_count": len(trajectories)}
|
||||
|
||||
|
||||
def process_trajectories_with_threads(grouped_trajectories: List[List[Any]],
|
||||
service_url: str,
|
||||
workspace_id: str,
|
||||
n_threads: int = 4) -> List[Dict[str, Any]]:
|
||||
def process_trajectories_with_threads(
|
||||
grouped_trajectories: List[List[Any]],
|
||||
service_url: str,
|
||||
workspace_id: str,
|
||||
n_threads: int = 4,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
use threads to process trajectories
|
||||
|
||||
|
||||
Args:
|
||||
grouped_trajectories: group trajectory list by task_id
|
||||
service_url: memory summarizer service URL
|
||||
workspace_id: workspace ID
|
||||
n_threads: number of threads
|
||||
|
||||
|
||||
Returns:
|
||||
all results
|
||||
"""
|
||||
results = []
|
||||
|
||||
|
||||
with ThreadPoolExecutor(max_workers=n_threads) as executor:
|
||||
future_to_group = {
|
||||
executor.submit(post_to_summarizer, group, service_url, workspace_id): i
|
||||
executor.submit(post_to_summarizer, group, service_url, workspace_id): i
|
||||
for i, group in enumerate(grouped_trajectories)
|
||||
}
|
||||
|
||||
|
||||
for future in as_completed(future_to_group):
|
||||
group_index = future_to_group[future]
|
||||
try:
|
||||
|
|
@ -131,58 +138,60 @@ def process_trajectories_with_threads(grouped_trajectories: List[List[Any]],
|
|||
result["group_index"] = group_index
|
||||
result["group_size"] = len(grouped_trajectories[group_index])
|
||||
results.append(result)
|
||||
print(f"✅ Group {group_index} processed: {result["metadata"].get('memory_list', 0) if 'memory_list' in result["metadata"] else 'error'}")
|
||||
print(
|
||||
f'✅ Group {group_index} processed: {result["metadata"].get("memory_list", 0) if "memory_list" in result["metadata"] else "error"}',
|
||||
)
|
||||
except Exception as e:
|
||||
error_result = {
|
||||
"group_index": group_index,
|
||||
"group_size": len(grouped_trajectories[group_index]),
|
||||
"error": str(e)
|
||||
"error": str(e),
|
||||
}
|
||||
results.append(error_result)
|
||||
print(f"❌ Group {group_index} failed: {e}")
|
||||
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='Convert JSONL to memories using ReMe service')
|
||||
parser.add_argument('--jsonl_file', type=str, required=True, help='Path to the JSONL file')
|
||||
parser.add_argument('--service_url', type=str, default='http://localhost:8001', help='ReMe service URL')
|
||||
parser.add_argument('--workspace_id', type=str, required=True, help='Workspace ID for the task memory pool')
|
||||
parser.add_argument('--output_file', type=str, help='Output file to save results (optional)')
|
||||
parser.add_argument('--n_threads', type=int, default=4, help='Number of threads for processing')
|
||||
|
||||
parser = argparse.ArgumentParser(description="Convert JSONL to memories using ReMe service")
|
||||
parser.add_argument("--jsonl_file", type=str, required=True, help="Path to the JSONL file")
|
||||
parser.add_argument("--service_url", type=str, default="http://localhost:8001", help="ReMe service URL")
|
||||
parser.add_argument("--workspace_id", type=str, required=True, help="Workspace ID for the task memory pool")
|
||||
parser.add_argument("--output_file", type=str, help="Output file to save results (optional)")
|
||||
parser.add_argument("--n_threads", type=int, default=4, help="Number of threads for processing")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
print(f"Processing JSONL file: {args.jsonl_file}")
|
||||
print(f"Service URL: {args.service_url}")
|
||||
print(f"Workspace ID: {args.workspace_id}")
|
||||
print(f"Threads: {args.n_threads}")
|
||||
|
||||
|
||||
with open(args.jsonl_file, "r") as f:
|
||||
data = [json.loads(line) for line in f]
|
||||
print(f"Loaded {len(data)} entries from JSONL file")
|
||||
|
||||
|
||||
grouped_trajectories = group_trajectories_by_task_id(data)
|
||||
print(f"Total groups: {len(grouped_trajectories)}")
|
||||
|
||||
|
||||
results = process_trajectories_with_threads(
|
||||
grouped_trajectories,
|
||||
args.service_url,
|
||||
grouped_trajectories,
|
||||
args.service_url,
|
||||
args.workspace_id,
|
||||
n_threads=args.n_threads
|
||||
n_threads=args.n_threads,
|
||||
)
|
||||
|
||||
|
||||
print(f"Processed {len(results)} groups")
|
||||
|
||||
success_count = sum(1 for r in results if 'error' not in r)
|
||||
|
||||
success_count = sum(1 for r in results if "error" not in r)
|
||||
error_count = len(results) - success_count
|
||||
total_memories = sum(len(r["metadata"].get('memory_list', [])) for r in results if 'memory_list' in r["metadata"])
|
||||
total_memories = sum(len(r["metadata"].get("memory_list", [])) for r in results if "memory_list" in r["metadata"])
|
||||
|
||||
print(f"✅ Success: {success_count}")
|
||||
print(f"❌ Errors: {error_count}")
|
||||
print(f"📊 Total task memories created: {total_memories}")
|
||||
|
||||
|
||||
if args.output_file:
|
||||
try:
|
||||
summary = {
|
||||
|
|
@ -192,10 +201,10 @@ def main():
|
|||
"success_count": success_count,
|
||||
"error_count": error_count,
|
||||
"total_task_memories": total_memories,
|
||||
"results": results
|
||||
"results": results,
|
||||
}
|
||||
|
||||
with open(args.output_file, 'w') as f:
|
||||
|
||||
with open(args.output_file, "w") as f:
|
||||
json.dump(summary, f, indent=2)
|
||||
print(f"Results saved to: {args.output_file}")
|
||||
except Exception as e:
|
||||
|
|
@ -204,20 +213,21 @@ def main():
|
|||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
if len(sys.argv) > 1:
|
||||
main()
|
||||
else:
|
||||
print("Running in compatibility mode...")
|
||||
with open("exp_result/qwen3-14b/no_think/bfcl-multi-turn-base_wo-exp.jsonl", "r") as f:
|
||||
data = [json.loads(line) for line in f]
|
||||
|
||||
|
||||
grouped_trajectories = group_trajectories_by_task_id(data)
|
||||
print(f"Total groups: {len(grouped_trajectories)}")
|
||||
|
||||
|
||||
results = process_trajectories_with_threads(
|
||||
grouped_trajectories,
|
||||
"http://localhost:8001",
|
||||
grouped_trajectories,
|
||||
"http://localhost:8001",
|
||||
"bfcl_test",
|
||||
n_threads=4
|
||||
n_threads=4,
|
||||
)
|
||||
print(f"Processed {len(results)} groups")
|
||||
print(f"Processed {len(results)} groups")
|
||||
|
|
|
|||
|
|
@ -1,27 +1,26 @@
|
|||
import json
|
||||
with open("../../file_vector_store/bfcl_test.jsonl", 'r') as f:
|
||||
|
||||
with open("../../file_vector_store/bfcl_test.jsonl", "r") as f:
|
||||
bfcl = [json.loads(line) for line in f]
|
||||
|
||||
|
||||
new_bfcl = []
|
||||
for exp in bfcl:
|
||||
new_exp = {}
|
||||
new_exp["workspace_id"] = exp["workspace_id"]
|
||||
new_exp["memory_id"] = exp["unique_id"]
|
||||
new_exp["memory_type"] = exp["metadata"]["memory_type"]
|
||||
|
||||
|
||||
new_exp["when_to_use"] = exp["content"]
|
||||
new_exp["content"] = exp["metadata"]["content"]
|
||||
new_exp["score"] = exp["metadata"]["score"]
|
||||
|
||||
new_exp["score"] = exp["metadata"]["score"]
|
||||
|
||||
new_exp["time_created"] = exp["metadata"]["time_created"]
|
||||
new_exp["time_modified"] = exp["metadata"]["time_modified"]
|
||||
new_exp["author"] = exp["metadata"]["author"]
|
||||
|
||||
new_exp["metadata"]= exp["metadata"]["metadata"]
|
||||
|
||||
new_exp["metadata"] = exp["metadata"]["metadata"]
|
||||
|
||||
new_bfcl.append(new_exp)
|
||||
|
||||
|
||||
|
||||
with open('../../library/bfcl_test.jsonl', 'w', encoding='utf-8') as f:
|
||||
f.writelines(json.dumps(item, ensure_ascii=False) + '\n' for item in new_bfcl)
|
||||
with open("../../library/bfcl_test.jsonl", "w", encoding="utf-8") as f:
|
||||
f.writelines(json.dumps(item, ensure_ascii=False) + "\n" for item in new_bfcl)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
import os
|
||||
import time
|
||||
|
||||
import ray
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# from ray import logger
|
||||
from loguru import logger
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
import json
|
||||
|
|
@ -15,26 +15,30 @@ from pathlib import Path
|
|||
from bfcl_agent import BFCLAgent
|
||||
|
||||
|
||||
def run_agent(dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
max_workers: int,
|
||||
num_runs: int = 4,
|
||||
model_name: str = "qwen3-8b",
|
||||
data_path: str = "data/multiturn_data_base_val.jsonl",
|
||||
answer_path: Path = Path("data/possible_answer"),
|
||||
use_memory: bool = False,
|
||||
use_memory_addition: bool = True,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
enable_thinking: bool = False,
|
||||
memory_base_url: str = "http://0.0.0.0:8001/",
|
||||
memory_workspace_id: str = "bfcl_test"):
|
||||
def run_agent(
|
||||
dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
max_workers: int,
|
||||
num_runs: int = 4,
|
||||
model_name: str = "qwen3-8b",
|
||||
data_path: str = "data/multiturn_data_base_val.jsonl",
|
||||
answer_path: Path = Path("data/possible_answer"),
|
||||
use_memory: bool = False,
|
||||
use_memory_addition: bool = True,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
enable_thinking: bool = False,
|
||||
memory_base_url: str = "http://0.0.0.0:8001/",
|
||||
memory_workspace_id: str = "bfcl_test",
|
||||
):
|
||||
experiment_name = dataset_name + "_" + experiment_suffix
|
||||
path: Path = Path(f"./exp_result/{model_name}/with_think" if enable_thinking else f"./exp_result/{model_name}/no_think")
|
||||
path: Path = Path(
|
||||
f"./exp_result/{model_name}/with_think" if enable_thinking else f"./exp_result/{model_name}/no_think",
|
||||
)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
task_ids = [json.loads(l)["id"] for l in f]
|
||||
|
||||
|
|
@ -48,7 +52,7 @@ def run_agent(dataset_name: str,
|
|||
future_list: list = []
|
||||
for i in range(max_workers):
|
||||
actor = BFCLAgent.remote(
|
||||
index=i,
|
||||
index=i,
|
||||
task_ids=task_ids[i::max_workers],
|
||||
experiment_name=experiment_name,
|
||||
data_path=data_path,
|
||||
|
|
@ -63,7 +67,7 @@ def run_agent(dataset_name: str,
|
|||
utility_threshold=utility_threshold,
|
||||
enable_thinking=enable_thinking,
|
||||
memory_base_url=memory_base_url,
|
||||
memory_workspace_id=memory_workspace_id
|
||||
memory_workspace_id=memory_workspace_id,
|
||||
)
|
||||
future = actor.execute.remote()
|
||||
future_list.append(future)
|
||||
|
|
@ -80,7 +84,7 @@ def run_agent(dataset_name: str,
|
|||
|
||||
logger.info(f"{i + 1}/{len(task_ids)} complete")
|
||||
dump_file()
|
||||
|
||||
|
||||
|
||||
def main():
|
||||
max_workers = 4
|
||||
|
|
@ -94,11 +98,11 @@ def main():
|
|||
ray.init(num_cpus=max_workers)
|
||||
for run_id in range(num_runs):
|
||||
run_agent(
|
||||
dataset_name="bfcl-multi-turn-base",
|
||||
dataset_name="bfcl-multi-turn-base",
|
||||
experiment_suffix=f"wo-exp",
|
||||
model_name="qwen3-8b",
|
||||
max_workers=max_workers,
|
||||
num_runs=1,
|
||||
max_workers=max_workers,
|
||||
num_runs=1,
|
||||
data_path="data/multiturn_data_base_val.jsonl",
|
||||
answer_path=Path("data/possible_answer"),
|
||||
enable_thinking=False,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import json
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
|
||||
|
||||
|
|
@ -24,7 +24,7 @@ def calculate_best_at_k(scores: list, k: int) -> float:
|
|||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
group = scores[i : i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
|
@ -36,8 +36,8 @@ def calculate_pass_at_k(scores: list, k: int) -> float:
|
|||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
is_pass = 1.0 if max(group) >=1.0 else 0.0
|
||||
group = scores[i : i + k]
|
||||
is_pass = 1.0 if max(group) >= 1.0 else 0.0
|
||||
group_maxs.append(is_pass)
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
|
@ -133,7 +133,7 @@ def run_exp_statistic():
|
|||
# Create and display table
|
||||
if all_results:
|
||||
df = pd.DataFrame(list(all_results.values()))
|
||||
df = df.set_index('file')
|
||||
df = df.set_index("file")
|
||||
|
||||
# Sort columns by the number in column name (best@8, best@4, best@2, best@1)
|
||||
# best_columns = [col for col in df.columns if col.startswith('best@')]
|
||||
|
|
@ -156,4 +156,4 @@ def run_exp_statistic():
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_exp_statistic()
|
||||
run_exp_statistic()
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
import argparse
|
||||
import json
|
||||
import random
|
||||
import argparse
|
||||
|
||||
|
||||
def split_jsonl(input_file, train_file, val_file, ratio=0.8):
|
||||
with open(input_file, 'r', encoding='utf-8') as f:
|
||||
with open(input_file, "r", encoding="utf-8") as f:
|
||||
data = [json.loads(line) for line in f]
|
||||
random.shuffle(data)
|
||||
|
||||
|
|
@ -11,19 +12,20 @@ def split_jsonl(input_file, train_file, val_file, ratio=0.8):
|
|||
train_data = data[:split_idx]
|
||||
val_data = data[split_idx:]
|
||||
|
||||
with open(train_file, 'w', encoding='utf-8') as f:
|
||||
with open(train_file, "w", encoding="utf-8") as f:
|
||||
for item in train_data:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + '\n')
|
||||
with open(val_file, 'w', encoding='utf-8') as f:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
with open(val_file, "w", encoding="utf-8") as f:
|
||||
for item in val_data:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + '\n')
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Split JSONL file into train and validation sets.')
|
||||
parser.add_argument('--input', required=True, help='Path to input JSONL file')
|
||||
parser.add_argument('--train', required=True, help='Path to output train file')
|
||||
parser.add_argument('--val', required=True, help='Path to output validation file')
|
||||
parser.add_argument('--ratio', type=float, default=0.5, help='Train ratio (default: 0.8)')
|
||||
parser = argparse.ArgumentParser(description="Split JSONL file into train and validation sets.")
|
||||
parser.add_argument("--input", required=True, help="Path to input JSONL file")
|
||||
parser.add_argument("--train", required=True, help="Path to output train file")
|
||||
parser.add_argument("--val", required=True, help="Path to output validation file")
|
||||
parser.add_argument("--ratio", type=float, default=0.5, help="Train ratio (default: 0.8)")
|
||||
|
||||
args = parser.parse_args()
|
||||
split_jsonl(args.input, args.train, args.val, args.ratio)
|
||||
|
|
|
|||
|
|
@ -1,47 +1,47 @@
|
|||
frozenlake_sys_prompt_no_slippery: |
|
||||
You are an AI agent playing FrozenLake game. Your goal is to navigate from Start (S) to Goal (G) while avoiding Holes (H).
|
||||
|
||||
|
||||
Game Rules:
|
||||
- S: Starting position (safe)
|
||||
- F: Frozen surface (safe to walk on)
|
||||
- H: Hole (you fall in and lose)
|
||||
- G: Goal (you win!)
|
||||
- []: Your current position
|
||||
|
||||
|
||||
Actions:
|
||||
- 0: Move LEFT
|
||||
- 1: Move DOWN
|
||||
- 1: Move DOWN
|
||||
- 2: Move RIGHT
|
||||
- 3: Move UP
|
||||
|
||||
Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes.
|
||||
|
||||
Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes.
|
||||
While ensuring a safe arrival at the goal, you should aim to complete the task in as few steps as possible.
|
||||
Think step by step, and respond with your thoughts and then clearly state your action as a number (0-3) in format {"action":"(0-3)"}.
|
||||
|
||||
frozenlake_sys_prompt_slippery: |
|
||||
You are an AI agent playing FrozenLake game. Your goal is to navigate from Start (S) to Goal (G) while avoiding Holes (H).
|
||||
|
||||
|
||||
Game Rules:
|
||||
- S: Starting position (safe)
|
||||
- F: Frozen surface (safe to walk on)
|
||||
- H: Hole (you fall in and lose)
|
||||
- G: Goal (you win!)
|
||||
- []: Your current position
|
||||
|
||||
|
||||
Actions:
|
||||
- 0: Move LEFT
|
||||
- 1: Move DOWN
|
||||
- 1: Move DOWN
|
||||
- 2: Move RIGHT
|
||||
- 3: Move UP
|
||||
|
||||
|
||||
The ice is slippery, so you might not always move in the intended direction!
|
||||
you will move in intended direction with probability of 1/3 else will move in either perpendicular direction with equal probability of 1/3 in both directions.
|
||||
|
||||
|
||||
For example, if action is left, then:
|
||||
- P(move left)=1/3
|
||||
- P(move up)=1/3
|
||||
- P(move down)=1/3
|
||||
|
||||
Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes.
|
||||
|
||||
Your task: Analyze the current state and choose the best action (0-3) to reach the Goal while avoiding Holes.
|
||||
While ensuring a safe arrival at the goal, you should aim to complete the task in as few steps as possible.
|
||||
Think step by step, and respond with your thoughts and then clearly state your action as a number (0-3) in format {{"action":"(0-3)"}}.
|
||||
|
|
|
|||
|
|
@ -1,24 +1,22 @@
|
|||
import os
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Dict, Any
|
||||
|
||||
import gymnasium as gym
|
||||
import ray
|
||||
import requests
|
||||
import random
|
||||
from typing import List, Dict, Any, Optional
|
||||
from dataclasses import dataclass
|
||||
import numpy as np
|
||||
import gymnasium as gym
|
||||
from gymnasium.envs.toy_text.frozen_lake import generate_random_map
|
||||
from openai import OpenAI
|
||||
from loguru import logger
|
||||
import yaml
|
||||
from dotenv import load_dotenv
|
||||
from gymnasium.envs.toy_text.frozen_lake import generate_random_map
|
||||
from loguru import logger
|
||||
from openai import OpenAI
|
||||
from tqdm import tqdm
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
|
||||
@dataclass
|
||||
class GameResult:
|
||||
task_id: str
|
||||
|
|
@ -30,20 +28,23 @@ class GameResult:
|
|||
trajectory: List[Dict]
|
||||
map_config: Dict[str, Any]
|
||||
|
||||
|
||||
@ray.remote
|
||||
class FrozenLakeReactAgent:
|
||||
"""A ReAct Agent for FrozenLake game with task memory learning."""
|
||||
|
||||
def __init__(self,
|
||||
index: int,
|
||||
task_configs: List[Dict],
|
||||
experiment_name: str,
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.7,
|
||||
max_steps: int = 50,
|
||||
num_runs: int = 1,
|
||||
use_task_memory: bool = False,
|
||||
make_task_memory: bool = False):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
task_configs: List[Dict],
|
||||
experiment_name: str,
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.7,
|
||||
max_steps: int = 50,
|
||||
num_runs: int = 1,
|
||||
use_task_memory: bool = False,
|
||||
make_task_memory: bool = False,
|
||||
):
|
||||
|
||||
self.index = index
|
||||
self.task_configs = task_configs
|
||||
|
|
@ -64,11 +65,13 @@ class FrozenLakeReactAgent:
|
|||
def _load_prompts(self) -> Dict[str, str]:
|
||||
"""Load prompts from yaml file"""
|
||||
try:
|
||||
with open("frozenlake_prompts.yaml", 'r', encoding='utf-8') as f:
|
||||
with open("frozenlake_prompts.yaml", "r", encoding="utf-8") as f:
|
||||
return yaml.safe_load(f)
|
||||
except FileNotFoundError:
|
||||
logger.warning("Prompt file not found, using default prompts")
|
||||
raise FileNotFoundError("Prompt file not found. Please check your current path (should be ./cook/frozenlake) and try again.")
|
||||
raise FileNotFoundError(
|
||||
"Prompt file not found. Please check your current path (should be ./cook/frozenlake) and try again.",
|
||||
)
|
||||
|
||||
def call_llm(self, messages: List[Dict]) -> str:
|
||||
"""Call LLM with retry logic"""
|
||||
|
|
@ -79,7 +82,7 @@ class FrozenLakeReactAgent:
|
|||
messages=messages,
|
||||
temperature=self.temperature,
|
||||
extra_body={"enable_thinking": False},
|
||||
seed=0
|
||||
seed=0,
|
||||
)
|
||||
return response.choices[0].message.content
|
||||
except Exception as e:
|
||||
|
|
@ -93,7 +96,7 @@ class FrozenLakeReactAgent:
|
|||
nrow, ncol = desc.shape
|
||||
|
||||
# Convert to string grid
|
||||
grid = [[cell.decode('utf-8') for cell in row] for row in desc]
|
||||
grid = [[cell.decode("utf-8") for cell in row] for row in desc]
|
||||
|
||||
# Get current position
|
||||
row, col = observation // ncol, observation % ncol
|
||||
|
|
@ -134,7 +137,7 @@ class FrozenLakeReactAgent:
|
|||
"workspace_id": workspace_id,
|
||||
"query": query,
|
||||
},
|
||||
timeout=60
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
|
|
@ -155,7 +158,7 @@ class FrozenLakeReactAgent:
|
|||
r'["\']action["\']\s*:\s*["\']([0-3])["\']',
|
||||
r'"action"\s*:\s*"([0-3])"',
|
||||
r"'action'\s*:\s*'([0-3])'",
|
||||
r'\baction["\']?\s*[:=]\s*["\']?([0-3])'
|
||||
r'\baction["\']?\s*[:=]\s*["\']?([0-3])',
|
||||
]
|
||||
|
||||
for pattern in patterns:
|
||||
|
|
@ -190,8 +193,7 @@ class FrozenLakeReactAgent:
|
|||
env = gym.make("FrozenLake-v1", **env_kwargs)
|
||||
|
||||
# Get map description for task memory
|
||||
map_str = '\n'.join([''.join([cell.decode('utf-8') for cell in row])
|
||||
for row in env.unwrapped.desc])
|
||||
map_str = "\n".join(["".join([cell.decode("utf-8") for cell in row]) for row in env.unwrapped.desc])
|
||||
|
||||
# Build messages
|
||||
system_prompt = self.build_system_prompt(is_slippery)
|
||||
|
|
@ -203,7 +205,8 @@ class FrozenLakeReactAgent:
|
|||
memory_content = f"Here are some relevant tips from previous successful games:\n\n{task_memory}\n\nUse these tips to help you succeed."
|
||||
messages.append({"role": "user", "content": memory_content})
|
||||
messages.append(
|
||||
{"role": "assistant", "content": "I'll use these tips to navigate the frozen lake successfully."})
|
||||
{"role": "assistant", "content": "I'll use these tips to navigate the frozen lake successfully."},
|
||||
)
|
||||
|
||||
# Initialize game
|
||||
observation, info = env.reset()
|
||||
|
|
@ -230,16 +233,18 @@ class FrozenLakeReactAgent:
|
|||
done = terminated or truncated
|
||||
|
||||
# Record trajectory step
|
||||
trajectory.append({
|
||||
"step": step,
|
||||
"state": observation,
|
||||
"action": action,
|
||||
"action_name": self.action_map[action],
|
||||
"reward": reward,
|
||||
"next_state": next_observation,
|
||||
"done": done,
|
||||
"llm_response": response
|
||||
})
|
||||
trajectory.append(
|
||||
{
|
||||
"step": step,
|
||||
"state": observation,
|
||||
"action": action,
|
||||
"action_name": self.action_map[action],
|
||||
"reward": reward,
|
||||
"next_state": next_observation,
|
||||
"done": done,
|
||||
"llm_response": response,
|
||||
},
|
||||
)
|
||||
|
||||
if done:
|
||||
if terminated and reward > 0:
|
||||
|
|
@ -275,8 +280,8 @@ class FrozenLakeReactAgent:
|
|||
"map_id": map_id,
|
||||
"is_slippery": is_slippery,
|
||||
"map_size": map_size,
|
||||
"use_task_memory": self.use_task_memory
|
||||
}
|
||||
"use_task_memory": self.use_task_memory,
|
||||
},
|
||||
)
|
||||
|
||||
return result, messages
|
||||
|
|
@ -311,9 +316,9 @@ class FrozenLakeReactAgent:
|
|||
url=base_url + "summary_task_memory",
|
||||
json={
|
||||
"workspace_id": workspace_id,
|
||||
"trajectories": trajectories
|
||||
"trajectories": trajectories,
|
||||
},
|
||||
timeout=300
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
|
|
@ -347,7 +352,7 @@ class FrozenLakeReactAgent:
|
|||
"steps": result.steps,
|
||||
"reward": result.reward,
|
||||
"map_config": result.map_config,
|
||||
"trajectory": result.trajectory
|
||||
"trajectory": result.trajectory,
|
||||
}
|
||||
all_results[-1] = result_dict
|
||||
|
||||
|
|
@ -364,10 +369,10 @@ class FrozenLakeReactAgent:
|
|||
steps=result_dict["steps"],
|
||||
reward=result_dict["reward"],
|
||||
trajectory=result_dict["trajectory"],
|
||||
map_config=result_dict["map_config"]
|
||||
map_config=result_dict["map_config"],
|
||||
)
|
||||
game_results.append(game_result)
|
||||
|
||||
self.save_task_memory(game_results, all_messages)
|
||||
|
||||
return all_results
|
||||
return all_results
|
||||
|
|
|
|||
|
|
@ -4,118 +4,125 @@ Map Management Tool - Pre-generate and manage test maps
|
|||
"""
|
||||
|
||||
import json
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Dict, Any
|
||||
from loguru import logger
|
||||
|
||||
import numpy as np
|
||||
from gymnasium.envs.toy_text.frozen_lake import generate_random_map
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class MapManager:
|
||||
"""Map Manager - pre-generating, storing and loading test maps"""
|
||||
"""Map Manager - pre-generating, storing and loading test maps"""
|
||||
|
||||
def __init__(self, data_dir: str = "./map/"):
|
||||
self.data_dir = Path(data_dir)
|
||||
self.data_dir.mkdir(parents=True, exist_ok=True)
|
||||
def __init__(self, data_dir: str = "./map/"):
|
||||
self.data_dir = Path(data_dir)
|
||||
self.data_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def generate_test_maps(self, num_maps: int, map_size: int = 4,
|
||||
base_seed: int = 10000) -> str:
|
||||
"""
|
||||
Generate test map collection and save
|
||||
def generate_test_maps(
|
||||
self,
|
||||
num_maps: int,
|
||||
map_size: int = 4,
|
||||
base_seed: int = 10000,
|
||||
) -> str:
|
||||
"""
|
||||
Generate test map collection and save
|
||||
|
||||
Args:
|
||||
num_maps: Number of maps to generate
|
||||
map_size: Map size
|
||||
base_seed: Base random seed
|
||||
Args:
|
||||
num_maps: Number of maps to generate
|
||||
map_size: Map size
|
||||
base_seed: Base random seed
|
||||
|
||||
Returns:
|
||||
Path of saved file
|
||||
"""
|
||||
logger.info(f"🗺️ Generating {num_maps} test maps (size={map_size})")
|
||||
Returns:
|
||||
Path of saved file
|
||||
"""
|
||||
logger.info(f"🗺️ Generating {num_maps} test maps (size={map_size})")
|
||||
|
||||
maps_data = []
|
||||
for i in range(num_maps):
|
||||
seed = base_seed + i
|
||||
np.random.seed(seed)
|
||||
map_desc = generate_random_map(size=map_size)
|
||||
maps_data = []
|
||||
for i in range(num_maps):
|
||||
seed = base_seed + i
|
||||
np.random.seed(seed)
|
||||
map_desc = generate_random_map(size=map_size)
|
||||
|
||||
maps_data.append({
|
||||
"map_id": i,
|
||||
"seed": seed,
|
||||
"map_size": map_size,
|
||||
"map_desc": map_desc # Convert to list for JSON serialization
|
||||
})
|
||||
maps_data.append(
|
||||
{
|
||||
"map_id": i,
|
||||
"seed": seed,
|
||||
"map_size": map_size,
|
||||
"map_desc": map_desc, # Convert to list for JSON serialization
|
||||
},
|
||||
)
|
||||
|
||||
# Save to file
|
||||
filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl"
|
||||
filepath = self.data_dir / filename
|
||||
# Save to file
|
||||
filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl"
|
||||
filepath = self.data_dir / filename
|
||||
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
for map_data in maps_data:
|
||||
f.write(json.dumps(map_data, ensure_ascii=False) + "\n")
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
for map_data in maps_data:
|
||||
f.write(json.dumps(map_data, ensure_ascii=False) + "\n")
|
||||
|
||||
logger.info(f"✅ Test maps saved to {filepath}")
|
||||
return str(filepath)
|
||||
logger.info(f"✅ Test maps saved to {filepath}")
|
||||
return str(filepath)
|
||||
|
||||
def load_test_maps(self, filepath: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Load test maps
|
||||
def load_test_maps(self, filepath: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Load test maps
|
||||
|
||||
Args:
|
||||
filepath: Map file path
|
||||
Args:
|
||||
filepath: Map file path
|
||||
|
||||
Returns:
|
||||
Map data list
|
||||
"""
|
||||
if not Path(filepath).exists():
|
||||
raise FileNotFoundError(f"Map file not found: {filepath}")
|
||||
Returns:
|
||||
Map data list
|
||||
"""
|
||||
if not Path(filepath).exists():
|
||||
raise FileNotFoundError(f"Map file not found: {filepath}")
|
||||
|
||||
maps_data = []
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if line.strip():
|
||||
map_data = json.loads(line)
|
||||
# Convert list back to numpy array
|
||||
maps_data.append(map_data)
|
||||
maps_data = []
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if line.strip():
|
||||
map_data = json.loads(line)
|
||||
# Convert list back to numpy array
|
||||
maps_data.append(map_data)
|
||||
|
||||
logger.info(f"📖 Loaded {len(maps_data)} test maps from {filepath}")
|
||||
return maps_data
|
||||
logger.info(f"📖 Loaded {len(maps_data)} test maps from {filepath}")
|
||||
return maps_data
|
||||
|
||||
def get_map_by_index(self, maps_data: List[Dict], index: int) -> Optional[list]:
|
||||
"""Get map by index"""
|
||||
if 0 <= index < len(maps_data):
|
||||
return maps_data[index]["map_desc"]
|
||||
return None
|
||||
def get_map_by_index(self, maps_data: List[Dict], index: int) -> Optional[list]:
|
||||
"""Get map by index"""
|
||||
if 0 <= index < len(maps_data):
|
||||
return maps_data[index]["map_desc"]
|
||||
return None
|
||||
|
||||
def get_or_create_test_maps(self, num_maps: int, map_size: int = 4) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get or create test maps
|
||||
If file exists and has sufficient quantity, load directly; otherwise regenerate
|
||||
"""
|
||||
filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl"
|
||||
filepath = self.data_dir / filename
|
||||
def get_or_create_test_maps(self, num_maps: int, map_size: int = 4) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Get or create test maps
|
||||
If file exists and has sufficient quantity, load directly; otherwise regenerate
|
||||
"""
|
||||
filename = f"test_maps_{num_maps}_{map_size}x{map_size}.jsonl"
|
||||
filepath = self.data_dir / filename
|
||||
|
||||
if filepath.exists():
|
||||
try:
|
||||
maps_data = self.load_test_maps(str(filepath))
|
||||
if len(maps_data) >= num_maps:
|
||||
logger.info(f"✅ Using existing test maps: {filepath}")
|
||||
return maps_data[:num_maps] # Return required number of maps
|
||||
except Exception as e:
|
||||
logger.warning(f"⚠️ Failed to load existing maps: {e}, regenerating...")
|
||||
if filepath.exists():
|
||||
try:
|
||||
maps_data = self.load_test_maps(str(filepath))
|
||||
if len(maps_data) >= num_maps:
|
||||
logger.info(f"✅ Using existing test maps: {filepath}")
|
||||
return maps_data[:num_maps] # Return required number of maps
|
||||
except Exception as e:
|
||||
logger.warning(f"⚠️ Failed to load existing maps: {e}, regenerating...")
|
||||
|
||||
# File doesn't exist or insufficient quantity, regenerate
|
||||
self.generate_test_maps(num_maps, map_size)
|
||||
return self.load_test_maps(str(filepath))
|
||||
# File doesn't exist or insufficient quantity, regenerate
|
||||
self.generate_test_maps(num_maps, map_size)
|
||||
return self.load_test_maps(str(filepath))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Usage example
|
||||
manager = MapManager()
|
||||
# Usage example
|
||||
manager = MapManager()
|
||||
|
||||
# Generate 100 4x4 test maps
|
||||
manager.generate_test_maps(num_maps=100, map_size=4)
|
||||
# Generate 100 4x4 test maps
|
||||
manager.generate_test_maps(num_maps=100, map_size=4)
|
||||
|
||||
# Load and view the first map
|
||||
maps = manager.load_test_maps("./map/test_maps_100_4x4.jsonl")
|
||||
print(f"First map:\n{maps[0]['map_desc']}")
|
||||
# Load and view the first map
|
||||
maps = manager.load_test_maps("./map/test_maps_100_4x4.jsonl")
|
||||
print(f"First map:\n{maps[0]['map_desc']}")
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
import json
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
|
||||
|
||||
|
|
@ -24,7 +25,7 @@ def calculate_best_at_k(scores: List[float], k: int) -> float:
|
|||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
group = scores[i : i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
|
@ -125,7 +126,7 @@ def process_single_result(data: Dict, condition_results: Dict):
|
|||
if part.startswith("map"):
|
||||
try:
|
||||
# Extract number from "mapXX"
|
||||
map_num = ''.join(filter(str.isdigit, part))
|
||||
map_num = "".join(filter(str.isdigit, part))
|
||||
if map_num:
|
||||
map_id = int(map_num)
|
||||
break
|
||||
|
|
@ -189,8 +190,10 @@ def calculate_file_metrics(condition_results: Dict, filename: str) -> Dict:
|
|||
|
||||
file_metrics[f"{condition}_map_details"] = map_success_rates
|
||||
|
||||
logger.info(f"{filename} - {condition}: {overall_success:.3f} success rate, "
|
||||
f"{len(map_results)} maps, {num_runs} runs each")
|
||||
logger.info(
|
||||
f"{filename} - {condition}: {overall_success:.3f} success rate, "
|
||||
f"{len(map_results)} maps, {num_runs} runs each",
|
||||
)
|
||||
|
||||
return file_metrics
|
||||
|
||||
|
|
@ -213,7 +216,7 @@ def generate_analysis_report(all_results: Dict):
|
|||
|
||||
if summary_data:
|
||||
df_summary = pd.DataFrame(summary_data)
|
||||
df_summary = df_summary.set_index('file')
|
||||
df_summary = df_summary.set_index("file")
|
||||
|
||||
print("\n" + "=" * 100)
|
||||
print("FROZENLAKE EXPERIMENT RESULTS SUMMARY")
|
||||
|
|
@ -367,4 +370,4 @@ def main():
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
import os
|
||||
import time
|
||||
import json
|
||||
import ray
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import List, Dict
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
import ray
|
||||
from gymnasium.envs.toy_text.frozen_lake import generate_random_map
|
||||
from loguru import logger
|
||||
|
||||
from frozenlake_react_agent import FrozenLakeReactAgent
|
||||
from map_manager import MapManager
|
||||
|
||||
|
||||
def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery: bool=False) -> List[Dict]:
|
||||
def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery: bool = False) -> List[Dict]:
|
||||
"""Generate random maps for training/task memory generation"""
|
||||
configs = []
|
||||
|
||||
|
|
@ -25,7 +25,7 @@ def generate_training_configs(num_maps: int = 20, map_size: int = 4, is_slippery
|
|||
"map_desc": random_map,
|
||||
"map_size": map_size,
|
||||
"is_slippery": is_slippery,
|
||||
"task_id": f"train_{i}_{is_slippery}"
|
||||
"task_id": f"train_{i}_{is_slippery}",
|
||||
}
|
||||
configs.append(config)
|
||||
|
||||
|
|
@ -43,7 +43,7 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) -
|
|||
configs = []
|
||||
|
||||
for map_data in maps_data:
|
||||
map_desc = np.array([list(row) for row in map_data["map_desc"]], dtype='c')
|
||||
map_desc = np.array([list(row) for row in map_data["map_desc"]], dtype="c")
|
||||
map_id = map_data["map_id"]
|
||||
|
||||
for use_memory in [True, False]:
|
||||
|
|
@ -54,7 +54,7 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) -
|
|||
"is_slippery": is_slippery,
|
||||
"use_task_memory": use_memory,
|
||||
"map_id": map_id,
|
||||
"task_id": f"test_map{map_id}_slip{is_slippery}_mem{use_memory}"
|
||||
"task_id": f"test_map{map_id}_slip{is_slippery}_mem{use_memory}",
|
||||
}
|
||||
configs.append(config)
|
||||
|
||||
|
|
@ -62,7 +62,13 @@ def generate_test_configs(num_test_maps: int = 100, is_slippery: bool = False) -
|
|||
return configs
|
||||
|
||||
|
||||
def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_training_maps= 15, is_slippery: bool= False) -> None:
|
||||
def train(
|
||||
experiment_name: str,
|
||||
max_workers: int = 2,
|
||||
num_runs: int = 3,
|
||||
num_training_maps=15,
|
||||
is_slippery: bool = False,
|
||||
) -> None:
|
||||
"""Phase 1: Generate task memory from random maps"""
|
||||
logger.info("🎯 Starting Training Phase - Generating Task Memory")
|
||||
logger.info("=" * 60)
|
||||
|
|
@ -116,7 +122,7 @@ def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_tra
|
|||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=False,
|
||||
make_task_memory=True
|
||||
make_task_memory=True,
|
||||
)
|
||||
results = agent.execute()
|
||||
dump_results()
|
||||
|
|
@ -130,7 +136,13 @@ def train(experiment_name: str, max_workers: int = 2, num_runs: int = 3, num_tra
|
|||
return results
|
||||
|
||||
|
||||
def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test_maps: int = 100, is_slippery: bool=False) -> None:
|
||||
def test(
|
||||
experiment_name: str,
|
||||
max_workers: int = 2,
|
||||
num_runs: int = 5,
|
||||
num_test_maps: int = 100,
|
||||
is_slippery: bool = False,
|
||||
) -> None:
|
||||
"""Phase 2: Test on fixed maps with/without task memory"""
|
||||
logger.info("🧪 Starting Test Phase - Evaluating Performance")
|
||||
logger.info(f"📊 Testing on {num_test_maps} maps with {num_runs} runs each")
|
||||
|
|
@ -147,8 +159,6 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test
|
|||
logger.info(f"📝 Configs without task memory: {len(no_memory_configs)}")
|
||||
logger.info(f"📝 Configs with task memory: {len(memory_configs)}")
|
||||
|
||||
|
||||
|
||||
def dump_results(suffix: str):
|
||||
output_file = path / f"{experiment_name}_test_{suffix}.jsonl"
|
||||
with open(output_file, "w") as f:
|
||||
|
|
@ -164,7 +174,7 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test
|
|||
experiment_name=experiment_name,
|
||||
max_workers=max_workers,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=False
|
||||
use_task_memory=False,
|
||||
)
|
||||
all_results.extend(results_no_memory)
|
||||
dump_results("no_memory")
|
||||
|
|
@ -177,7 +187,7 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test
|
|||
experiment_name=experiment_name,
|
||||
max_workers=max_workers,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=True
|
||||
use_task_memory=True,
|
||||
)
|
||||
all_results.extend(results_with_memory)
|
||||
dump_results("with_memory")
|
||||
|
|
@ -185,8 +195,13 @@ def test(experiment_name: str, max_workers: int = 2, num_runs: int = 5, num_test
|
|||
return all_results
|
||||
|
||||
|
||||
def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int,
|
||||
num_runs: int, use_task_memory: bool) -> List[Dict]:
|
||||
def run_test_configs(
|
||||
configs: List[Dict],
|
||||
experiment_name: str,
|
||||
max_workers: int,
|
||||
num_runs: int,
|
||||
use_task_memory: bool,
|
||||
) -> List[Dict]:
|
||||
"""Run a set of test configurations"""
|
||||
results = []
|
||||
|
||||
|
|
@ -201,7 +216,7 @@ def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int
|
|||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=False
|
||||
make_task_memory=False,
|
||||
)
|
||||
future = agent.execute.remote()
|
||||
future_list.append(future)
|
||||
|
|
@ -220,7 +235,7 @@ def run_test_configs(configs: List[Dict], experiment_name: str, max_workers: int
|
|||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
use_task_memory=use_task_memory,
|
||||
make_task_memory=False
|
||||
make_task_memory=False,
|
||||
)
|
||||
results = agent.execute()
|
||||
|
||||
|
|
@ -255,21 +270,20 @@ def main():
|
|||
max_workers=max_workers,
|
||||
num_runs=training_runs,
|
||||
num_training_maps=num_training_maps,
|
||||
is_slippery=is_slippery
|
||||
is_slippery=is_slippery,
|
||||
)
|
||||
|
||||
# Wait a bit for task memory service to process
|
||||
logger.info("⏰ Waiting for task memory service to process data...")
|
||||
time.sleep(10)
|
||||
|
||||
|
||||
# Phase 2: Testing (Performance Evaluation)
|
||||
test_results = test(
|
||||
experiment_name=experiment_name,
|
||||
max_workers=max_workers,
|
||||
num_runs=test_runs,
|
||||
num_test_maps=num_test_maps,
|
||||
is_slippery=is_slippery
|
||||
is_slippery=is_slippery,
|
||||
)
|
||||
|
||||
# Summary
|
||||
|
|
@ -293,4 +307,4 @@ def main():
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
|
|
|||
|
|
@ -7,10 +7,11 @@ from reme_ai import ReMeApp
|
|||
# Task Memory Management Examples
|
||||
# ============================================
|
||||
|
||||
|
||||
async def summary_task_memory(app: ReMeApp):
|
||||
"""
|
||||
Experience Summarizer: Learn from execution trajectories
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/summary_task_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -26,11 +27,11 @@ async def summary_task_memory(app: ReMeApp):
|
|||
trajectories=[
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "Help me create a project plan"}
|
||||
{"role": "user", "content": "Help me create a project plan"},
|
||||
],
|
||||
"score": 1.0
|
||||
}
|
||||
]
|
||||
"score": 1.0,
|
||||
},
|
||||
],
|
||||
)
|
||||
print("Summary Task Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -39,7 +40,7 @@ async def summary_task_memory(app: ReMeApp):
|
|||
async def retrieve_task_memory(app: ReMeApp):
|
||||
"""
|
||||
Retriever: Get relevant memories
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/retrieve_task_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -52,7 +53,7 @@ async def retrieve_task_memory(app: ReMeApp):
|
|||
name="retrieve_task_memory",
|
||||
workspace_id="task_workspace",
|
||||
query="How to efficiently manage project progress?",
|
||||
top_k=1
|
||||
top_k=1,
|
||||
)
|
||||
print("Retrieve Task Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -62,10 +63,11 @@ async def retrieve_task_memory(app: ReMeApp):
|
|||
# Personal Memory Management Examples
|
||||
# ============================================
|
||||
|
||||
|
||||
async def summary_personal_memory(app: ReMeApp):
|
||||
"""
|
||||
Memory Integration: Learn from user interactions
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/summary_personal_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -85,11 +87,13 @@ async def summary_personal_memory(app: ReMeApp):
|
|||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "I like to drink coffee while working in the morning"},
|
||||
{"role": "assistant",
|
||||
"content": "I understand, you prefer to start your workday with coffee to stay energized"}
|
||||
]
|
||||
}
|
||||
]
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "I understand, you prefer to start your workday with coffee to stay energized",
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
)
|
||||
print("Summary Personal Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -98,7 +102,7 @@ async def summary_personal_memory(app: ReMeApp):
|
|||
async def retrieve_personal_memory(app: ReMeApp):
|
||||
"""
|
||||
Memory Retrieval: Get personal memory fragments
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/retrieve_personal_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -111,7 +115,7 @@ async def retrieve_personal_memory(app: ReMeApp):
|
|||
name="retrieve_personal_memory",
|
||||
workspace_id="task_workspace",
|
||||
query="What are the user's work habits?",
|
||||
top_k=5
|
||||
top_k=5,
|
||||
)
|
||||
print("Retrieve Personal Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -121,10 +125,11 @@ async def retrieve_personal_memory(app: ReMeApp):
|
|||
# Tool Memory Management Examples
|
||||
# ============================================
|
||||
|
||||
|
||||
async def add_tool_call_result(app: ReMeApp):
|
||||
"""
|
||||
Record tool execution results
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/add_tool_call_result \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -153,9 +158,9 @@ async def add_tool_call_result(app: ReMeApp):
|
|||
"output": "Found 10 relevant results...",
|
||||
"token_cost": 150,
|
||||
"success": True,
|
||||
"time_cost": 2.3
|
||||
}
|
||||
]
|
||||
"time_cost": 2.3,
|
||||
},
|
||||
],
|
||||
)
|
||||
print("Add Tool Call Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -164,7 +169,7 @@ async def add_tool_call_result(app: ReMeApp):
|
|||
async def summary_tool_memory(app: ReMeApp):
|
||||
"""
|
||||
Generate usage guidelines from history
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/summary_tool_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -175,7 +180,7 @@ async def summary_tool_memory(app: ReMeApp):
|
|||
result = await app.async_execute(
|
||||
name="summary_tool_memory",
|
||||
workspace_id="tool_workspace",
|
||||
tool_names="web_search"
|
||||
tool_names="web_search",
|
||||
)
|
||||
print("Summary Tool Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -184,7 +189,7 @@ async def summary_tool_memory(app: ReMeApp):
|
|||
async def retrieve_tool_memory(app: ReMeApp):
|
||||
"""
|
||||
Retrieve tool guidelines before use
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/retrieve_tool_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -195,7 +200,7 @@ async def retrieve_tool_memory(app: ReMeApp):
|
|||
result = await app.async_execute(
|
||||
name="retrieve_tool_memory",
|
||||
workspace_id="tool_workspace",
|
||||
tool_names="web_search"
|
||||
tool_names="web_search",
|
||||
)
|
||||
print("Retrieve Tool Memory Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -205,10 +210,11 @@ async def retrieve_tool_memory(app: ReMeApp):
|
|||
# Vector Store Management Example
|
||||
# ============================================
|
||||
|
||||
|
||||
async def load_vector_store(app: ReMeApp):
|
||||
"""
|
||||
Load pre-built memories
|
||||
|
||||
|
||||
curl -X POST http://localhost:8002/vector_store \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
|
|
@ -221,7 +227,7 @@ async def load_vector_store(app: ReMeApp):
|
|||
name="vector_store",
|
||||
workspace_id="appworld",
|
||||
action="load",
|
||||
path="./docs/library/"
|
||||
path="./docs/library/",
|
||||
)
|
||||
print("Load Vector Store Result:")
|
||||
print(result["answer"])
|
||||
|
|
@ -231,12 +237,13 @@ async def load_vector_store(app: ReMeApp):
|
|||
# Main Execution
|
||||
# ============================================
|
||||
|
||||
|
||||
async def main():
|
||||
"""Run all examples"""
|
||||
async with ReMeApp(
|
||||
"llm.default.model_name=qwen3-30b-a3b-thinking-2507",
|
||||
"embedding_model.default.model_name=text-embedding-v4",
|
||||
"vector_store.default.backend=memory"
|
||||
"llm.default.model_name=qwen3-30b-a3b-thinking-2507",
|
||||
"embedding_model.default.model_name=text-embedding-v4",
|
||||
"vector_store.default.backend=memory",
|
||||
) as app:
|
||||
print("=" * 60)
|
||||
print("Task Memory Examples")
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ import aiohttp
|
|||
# API base URL
|
||||
base_url = "http://0.0.0.0:8002"
|
||||
|
||||
|
||||
async def main():
|
||||
# Create a unique workspace ID
|
||||
workspace_id = "personal_memory_demo"
|
||||
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Step 1: Clear existing memories in the workspace
|
||||
print("Clearing existing memories...")
|
||||
|
|
@ -19,39 +20,57 @@ async def main():
|
|||
"action": "delete",
|
||||
"workspace_id": workspace_id,
|
||||
},
|
||||
headers={"Content-Type": "application/json"}
|
||||
headers={"Content-Type": "application/json"},
|
||||
) as response:
|
||||
result = await response.json()
|
||||
print(json.dumps(result, ensure_ascii=False))
|
||||
|
||||
|
||||
# Step 2: Create a conversation with rich personal information
|
||||
print("\nCreating conversation with personal information...")
|
||||
messages = [
|
||||
{"role": "user", "content": "My name is John Smith, I'm 28 years old, and I work at a tech company in San Francisco"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "My name is John Smith, I'm 28 years old, and I work at a tech company in San Francisco",
|
||||
},
|
||||
{"role": "assistant", "content": "Nice to meet you, John!"},
|
||||
{"role": "user", "content": "I'm a software engineer, mainly doing backend development using Python and Go"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I'm a software engineer, mainly doing backend development using Python and Go",
|
||||
},
|
||||
{"role": "assistant", "content": "I see, you're a backend engineer working with Python and Go."},
|
||||
{"role": "user", "content": "I enjoy playing basketball and watching sci-fi movies. I recently watched Dune Part 2"},
|
||||
{"role": "assistant", "content": "Basketball and sci-fi movies are great hobbies! Dune Part 2 was indeed amazing."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "I enjoy playing basketball and watching sci-fi movies. I recently watched Dune Part 2",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Basketball and sci-fi movies are great hobbies! Dune Part 2 was indeed amazing.",
|
||||
},
|
||||
{"role": "user", "content": "I have a cat named Shadow who is 3 years old"},
|
||||
{"role": "assistant", "content": "Shadow sounds adorable! 3-year-old cats are quite playful."},
|
||||
{"role": "user", "content": "I'm planning a trip to Japan next month, mainly to Tokyo and Kyoto"},
|
||||
{"role": "assistant", "content": "Your Japan trip sounds exciting! Tokyo and Kyoto are both wonderful destinations with their own unique charm."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Your Japan trip sounds exciting! Tokyo and Kyoto are both wonderful destinations with their own unique charm.",
|
||||
},
|
||||
{"role": "user", "content": "I'm really interested in Japanese cuisine, especially sushi and ramen"},
|
||||
{"role": "assistant", "content": "Japanese cuisine is delicious! Sushi and ramen are very popular choices."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Japanese cuisine is delicious! Sushi and ramen are very popular choices.",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
# Step 3: Summarize personal memories from the conversation
|
||||
print("\nSummarizing personal memories...")
|
||||
async with session.post(
|
||||
f"{base_url}/summary_personal_memory",
|
||||
json={
|
||||
"trajectories": [
|
||||
{"messages": messages, "score": 1.0}
|
||||
{"messages": messages, "score": 1.0},
|
||||
],
|
||||
"workspace_id": workspace_id,
|
||||
},
|
||||
headers={"Content-Type": "application/json"}
|
||||
headers={"Content-Type": "application/json"},
|
||||
) as response:
|
||||
result = await response.json()
|
||||
result = json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
|
@ -59,11 +78,11 @@ async def main():
|
|||
|
||||
with open("personal_memory.jsonl", "w") as f:
|
||||
f.write(result)
|
||||
|
||||
|
||||
# Wait for the memories to be processed and stored
|
||||
print("\nWaiting for memories to be processed...")
|
||||
await asyncio.sleep(2)
|
||||
|
||||
|
||||
# Step 4: Retrieve personal memories with different queries
|
||||
queries = [
|
||||
"What's my name and age?",
|
||||
|
|
@ -71,9 +90,9 @@ async def main():
|
|||
"What are my hobbies?",
|
||||
"Do I have any pets?",
|
||||
"What are my travel plans?",
|
||||
"What foods do I like?"
|
||||
"What foods do I like?",
|
||||
]
|
||||
|
||||
|
||||
print("\nRetrieving personal memories...")
|
||||
for query in queries:
|
||||
print(f"\nQuery: {query}")
|
||||
|
|
@ -83,10 +102,11 @@ async def main():
|
|||
"query": query,
|
||||
"workspace_id": workspace_id,
|
||||
},
|
||||
headers={"Content-Type": "application/json"}
|
||||
headers={"Content-Type": "application/json"},
|
||||
) as response:
|
||||
result = await response.json()
|
||||
print(json.dumps(result, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
|
|||
|
|
@ -26,10 +26,10 @@ WORKSPACE_ID = "test_workspace"
|
|||
def handle_api_response(response: requests.Response) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Handle API response with proper error checking
|
||||
|
||||
|
||||
Args:
|
||||
response: Response object from requests
|
||||
|
||||
|
||||
Returns:
|
||||
Response JSON if successful, None otherwise
|
||||
"""
|
||||
|
|
@ -44,7 +44,7 @@ def handle_api_response(response: requests.Response) -> Optional[Dict[str, Any]]
|
|||
def delete_workspace() -> None:
|
||||
"""
|
||||
Delete the current workspace from the vector store
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -53,7 +53,7 @@ def delete_workspace() -> None:
|
|||
json={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "delete",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -64,17 +64,17 @@ def delete_workspace() -> None:
|
|||
def run_agent(query: str, dump_messages: bool = False) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Run the agent with a specific query
|
||||
|
||||
|
||||
Args:
|
||||
query: The query to send to the agent
|
||||
dump_messages: Whether to save messages to a file
|
||||
|
||||
|
||||
Returns:
|
||||
List of message objects from the conversation
|
||||
"""
|
||||
response = requests.post(
|
||||
url=f"{BASE_URL}react",
|
||||
json={"query": query}
|
||||
json={"query": query},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -93,18 +93,18 @@ def run_agent(query: str, dump_messages: bool = False) -> List[Dict[str, Any]]:
|
|||
with open("task_messages.jsonl", "w") as f:
|
||||
f.write(json.dumps(messages, indent=2, ensure_ascii=False))
|
||||
print(f"Messages saved to messages.jsonl")
|
||||
|
||||
|
||||
return messages
|
||||
|
||||
|
||||
def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True) -> None:
|
||||
"""
|
||||
Generate a summary of conversation messages and create task memories
|
||||
|
||||
|
||||
Args:
|
||||
messages: List of message objects from a conversation
|
||||
enable_dump_memory: Whether to save memory list to a file
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -118,9 +118,9 @@ def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True)
|
|||
json={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"trajectories": [
|
||||
{"messages": messages, "score": 1.0}
|
||||
]
|
||||
}
|
||||
{"messages": messages, "score": 1.0},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -141,10 +141,10 @@ def run_summary(messages: List[Dict[str, Any]], enable_dump_memory: bool = True)
|
|||
def run_retrieve(query: str) -> str:
|
||||
"""
|
||||
Retrieve relevant task memories based on a query
|
||||
|
||||
|
||||
Args:
|
||||
query: The query to retrieve relevant memories
|
||||
|
||||
|
||||
Returns:
|
||||
String containing the retrieved memory answer
|
||||
"""
|
||||
|
|
@ -154,7 +154,7 @@ def run_retrieve(query: str) -> str:
|
|||
json={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"query": query,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -170,18 +170,18 @@ def run_retrieve(query: str) -> str:
|
|||
def run_agent_with_memory(query_first: str, query_second: str, enable_dump_memory: bool = True) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Run the agent with memory augmentation
|
||||
|
||||
|
||||
This function demonstrates how to use task memory to enhance agent responses:
|
||||
1. First run the agent with the second query to build memory
|
||||
2. Then summarize the conversation to create memories
|
||||
3. Retrieve relevant memories for the first query
|
||||
4. Run the agent with the first query augmented with retrieved memories
|
||||
|
||||
|
||||
Args:
|
||||
query_first: The query to run with memory augmentation
|
||||
query_second: The query to build initial memories
|
||||
enable_dump_memory: Whether to save memory list to a file
|
||||
|
||||
|
||||
Returns:
|
||||
List of message objects from the final conversation
|
||||
"""
|
||||
|
|
@ -203,17 +203,17 @@ def run_agent_with_memory(query_first: str, query_second: str, enable_dump_memor
|
|||
augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}"
|
||||
print(f"Augmented query: {augmented_query}")
|
||||
messages = run_agent(query=augmented_query)
|
||||
|
||||
|
||||
return messages
|
||||
|
||||
|
||||
def dump_memory(path: str = "./") -> None:
|
||||
"""
|
||||
Dump the vector store memories to disk
|
||||
|
||||
|
||||
Args:
|
||||
path: Directory path to save the memories
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -223,7 +223,7 @@ def dump_memory(path: str = "./") -> None:
|
|||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "dump",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
@ -234,10 +234,10 @@ def dump_memory(path: str = "./") -> None:
|
|||
def load_memory(path: str = "./") -> None:
|
||||
"""
|
||||
Load memories from disk into the vector store
|
||||
|
||||
|
||||
Args:
|
||||
path: Directory path to load the memories from
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -247,7 +247,7 @@ def load_memory(path: str = "./") -> None:
|
|||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "load",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
|
|
|
|||
|
|
@ -26,10 +26,10 @@ WORKSPACE_ID = "test_workspace"
|
|||
async def delete_workspace(client: Client) -> None:
|
||||
"""
|
||||
Delete the current workspace from the vector store
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -38,7 +38,7 @@ async def delete_workspace(client: Client) -> None:
|
|||
arguments={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "delete",
|
||||
}
|
||||
},
|
||||
)
|
||||
print(f"Workspace '{WORKSPACE_ID}' deleted successfully")
|
||||
|
||||
|
|
@ -53,12 +53,12 @@ async def run_agent(client: Client, query: str, dump_messages: bool = False) ->
|
|||
async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dump_memory: bool = True) -> None:
|
||||
"""
|
||||
Generate a summary of conversation messages and create task memories
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
messages: List of message objects from a conversation
|
||||
enable_dump_memory: Whether to save memory list to a file
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -71,9 +71,9 @@ async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dum
|
|||
arguments={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"trajectories": [
|
||||
{"messages": messages, "score": 1.0}
|
||||
]
|
||||
}
|
||||
{"messages": messages, "score": 1.0},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
answer = result.content[0].text
|
||||
|
|
@ -90,11 +90,11 @@ async def run_summary(client: Client, messages: List[Dict[str, Any]], enable_dum
|
|||
async def run_retrieve(client: Client, query: str) -> str:
|
||||
"""
|
||||
Retrieve relevant task memories based on a query
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
query: The query to retrieve relevant memories
|
||||
|
||||
|
||||
Returns:
|
||||
String containing the retrieved memory answer
|
||||
"""
|
||||
|
|
@ -103,7 +103,7 @@ async def run_retrieve(client: Client, query: str) -> str:
|
|||
arguments={
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"query": query,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
answer = result.content[0].text
|
||||
|
|
@ -111,22 +111,27 @@ async def run_retrieve(client: Client, query: str) -> str:
|
|||
return answer
|
||||
|
||||
|
||||
async def run_agent_with_memory(client: Client, query_first: str, query_second: str, enable_dump_memory: bool = True) -> List[Dict[str, Any]]:
|
||||
async def run_agent_with_memory(
|
||||
client: Client,
|
||||
query_first: str,
|
||||
query_second: str,
|
||||
enable_dump_memory: bool = True,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Run the agent with memory augmentation
|
||||
|
||||
|
||||
This function demonstrates how to use task memory to enhance agent responses:
|
||||
1. First run the agent with the second query to build memory
|
||||
2. Then summarize the conversation to create memories
|
||||
3. Retrieve relevant memories for the first query
|
||||
4. Run the agent with the first query augmented with retrieved memories
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
query_first: The query to run with memory augmentation
|
||||
query_second: The query to build initial memories
|
||||
enable_dump_memory: Whether to save memory list to a file
|
||||
|
||||
|
||||
Returns:
|
||||
List of message objects from the final conversation
|
||||
"""
|
||||
|
|
@ -148,18 +153,18 @@ async def run_agent_with_memory(client: Client, query_first: str, query_second:
|
|||
augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}"
|
||||
print(f"Augmented query: {augmented_query}")
|
||||
messages = await run_agent(client, query=augmented_query)
|
||||
|
||||
|
||||
return messages
|
||||
|
||||
|
||||
async def dump_memory(client: Client, path: str = "./") -> None:
|
||||
"""
|
||||
Dump the vector store memories to disk
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
path: Directory path to save the memories
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -169,7 +174,7 @@ async def dump_memory(client: Client, path: str = "./") -> None:
|
|||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "dump",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
print(f"Memory dumped to {path}")
|
||||
|
||||
|
|
@ -177,11 +182,11 @@ async def dump_memory(client: Client, path: str = "./") -> None:
|
|||
async def load_memory(client: Client, path: str = "./") -> None:
|
||||
"""
|
||||
Load memories from disk into the vector store
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
path: Directory path to load the memories from
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -191,7 +196,7 @@ async def load_memory(client: Client, path: str = "./") -> None:
|
|||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "load",
|
||||
"path": path,
|
||||
}
|
||||
},
|
||||
)
|
||||
print(f"Memory loaded from {path}")
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from typing import List, Dict, Any, Optional
|
|||
|
||||
import requests
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from reme_ai.utils.tool_memory_utils import create_mock_tool_call_results
|
||||
|
||||
load_dotenv()
|
||||
|
|
@ -27,7 +28,7 @@ def api_call(endpoint: str, data: dict) -> Optional[Dict[str, Any]]:
|
|||
|
||||
def delete_workspace() -> None:
|
||||
"""删除工作空间
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/vector_store \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -43,7 +44,7 @@ def delete_workspace() -> None:
|
|||
|
||||
def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool:
|
||||
"""添加工具调用结果到记忆库
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/add_tool_call_result \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -62,11 +63,14 @@ def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool:
|
|||
# 统计不同的工具
|
||||
tool_names = set(r.get("tool_name") for r in tool_call_results)
|
||||
print(f"\n[ADD] {len(tool_call_results)} results for {len(tool_names)} tools: {', '.join(sorted(tool_names))}")
|
||||
|
||||
result = api_call("add_tool_call_result", {
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_call_results": tool_call_results
|
||||
})
|
||||
|
||||
result = api_call(
|
||||
"add_tool_call_result",
|
||||
{
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_call_results": tool_call_results,
|
||||
},
|
||||
)
|
||||
if result:
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
print(f"✓ Added successfully, created/updated {len(memory_list)} tool memories")
|
||||
|
|
@ -76,7 +80,7 @@ def add_tool_call_results(tool_call_results: List[Dict[str, Any]]) -> bool:
|
|||
|
||||
def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]:
|
||||
"""总结工具使用模式
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/summary_tool_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -86,11 +90,14 @@ def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]:
|
|||
}'
|
||||
"""
|
||||
print(f"\n[SUMMARIZE] {tool_names}")
|
||||
result = api_call("summary_tool_memory", {
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_names": tool_names
|
||||
})
|
||||
|
||||
result = api_call(
|
||||
"summary_tool_memory",
|
||||
{
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_names": tool_names,
|
||||
},
|
||||
)
|
||||
|
||||
if result:
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
print(f"✓ Summarized {len(memory_list)} tool memories")
|
||||
|
|
@ -98,13 +105,13 @@ def summarize_tool_memory(tool_names: str) -> Optional[Dict[str, Any]]:
|
|||
print(f"\n{'=' * 60}")
|
||||
print(f"Tool: {memory.get('when_to_use', 'N/A')}")
|
||||
print(f"{'=' * 60}")
|
||||
print(memory.get('content', 'No content'))
|
||||
print(memory.get("content", "No content"))
|
||||
return result
|
||||
|
||||
|
||||
def retrieve_tool_memory(tool_names: str, save_to_file: bool = False) -> str:
|
||||
"""检索工具记忆
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/retrieve_tool_memory \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -114,40 +121,45 @@ def retrieve_tool_memory(tool_names: str, save_to_file: bool = False) -> str:
|
|||
}'
|
||||
"""
|
||||
print(f"\n[RETRIEVE] {tool_names}")
|
||||
result = api_call("retrieve_tool_memory", {
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_names": tool_names
|
||||
})
|
||||
|
||||
result = api_call(
|
||||
"retrieve_tool_memory",
|
||||
{
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"tool_names": tool_names,
|
||||
},
|
||||
)
|
||||
|
||||
if not result:
|
||||
return ""
|
||||
|
||||
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
if not memory_list:
|
||||
print("No memories found")
|
||||
return ""
|
||||
|
||||
|
||||
print(f"✓ Retrieved {len(memory_list)} memories")
|
||||
|
||||
|
||||
formatted_memories = []
|
||||
for memory in memory_list:
|
||||
content = f"\nTool: {memory.get('when_to_use', 'N/A')}\n" \
|
||||
f"Calls: {len(memory.get('tool_call_results', []))}\n" \
|
||||
f"{'-' * 60}\n{memory.get('content', 'No content')}\n"
|
||||
content = (
|
||||
f"\nTool: {memory.get('when_to_use', 'N/A')}\n"
|
||||
f"Calls: {len(memory.get('tool_call_results', []))}\n"
|
||||
f"{'-' * 60}\n{memory.get('content', 'No content')}\n"
|
||||
)
|
||||
formatted_memories.append(content)
|
||||
print(content)
|
||||
|
||||
|
||||
if save_to_file:
|
||||
with open("tool_memory.json", "w") as f:
|
||||
json.dump(memory_list, f, indent=2, ensure_ascii=False)
|
||||
print("✓ Saved to tool_memory.json")
|
||||
|
||||
|
||||
return "\n".join(formatted_memories)
|
||||
|
||||
|
||||
def dump_memory(path: str = "./") -> None:
|
||||
"""导出记忆到磁盘
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/vector_store \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -157,18 +169,21 @@ def dump_memory(path: str = "./") -> None:
|
|||
"path": "./"
|
||||
}'
|
||||
"""
|
||||
result = api_call("vector_store", {
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "dump",
|
||||
"path": path
|
||||
})
|
||||
result = api_call(
|
||||
"vector_store",
|
||||
{
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "dump",
|
||||
"path": path,
|
||||
},
|
||||
)
|
||||
if result:
|
||||
print(f"✓ Memory dumped to {path}")
|
||||
|
||||
|
||||
def load_memory(path: str = "./") -> None:
|
||||
"""从磁盘加载记忆
|
||||
|
||||
|
||||
curl example:
|
||||
curl -X POST http://0.0.0.0:8002/vector_store \
|
||||
-H "Content-Type: application/json" \
|
||||
|
|
@ -178,11 +193,14 @@ def load_memory(path: str = "./") -> None:
|
|||
"path": "./"
|
||||
}'
|
||||
"""
|
||||
result = api_call("vector_store", {
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "load",
|
||||
"path": path
|
||||
})
|
||||
result = api_call(
|
||||
"vector_store",
|
||||
{
|
||||
"workspace_id": WORKSPACE_ID,
|
||||
"action": "load",
|
||||
"path": path,
|
||||
},
|
||||
)
|
||||
if result:
|
||||
print(f"✓ Memory loaded from {path}")
|
||||
|
||||
|
|
@ -192,43 +210,43 @@ def main() -> None:
|
|||
print("\n[1] Cleaning workspace...")
|
||||
delete_workspace()
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
# 2. 创建和添加模拟工具调用结果
|
||||
print("\n[2] Adding mock tool call results...")
|
||||
tools_to_test = [
|
||||
("web_search", 30),
|
||||
("database_query", 22),
|
||||
("file_processor", 18)
|
||||
("file_processor", 18),
|
||||
]
|
||||
|
||||
|
||||
# 收集所有工具的结果,然后一次性添加
|
||||
all_mock_results = []
|
||||
for tool_name, count in tools_to_test:
|
||||
mock_results = create_mock_tool_call_results(tool_name, count)
|
||||
all_mock_results.extend(mock_results)
|
||||
|
||||
|
||||
if not add_tool_call_results(all_mock_results):
|
||||
print("✗ Failed to add results")
|
||||
else:
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
# 3. 总结工具记忆
|
||||
print("\n[3] Summarizing tool memories...")
|
||||
all_tool_names = ",".join([tool[0] for tool in tools_to_test])
|
||||
summarize_tool_memory(all_tool_names)
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
# 4. 检索工具记忆
|
||||
print("\n[4] Retrieving tool memories...")
|
||||
for tool_name, _ in tools_to_test:
|
||||
retrieve_tool_memory(tool_name, save_to_file=True)
|
||||
time.sleep(0.5)
|
||||
|
||||
|
||||
# 5. 测试记忆持久化
|
||||
print("\n[5] Testing memory persistence...")
|
||||
dump_memory()
|
||||
load_memory()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("DEMO COMPLETE ✓")
|
||||
print("=" * 60)
|
||||
|
|
|
|||
|
|
@ -38,8 +38,8 @@ class BenchmarkStats:
|
|||
|
||||
def add_result(self, result: Dict[str, Any]):
|
||||
"""添加一个工具调用结果
|
||||
|
||||
Note:
|
||||
|
||||
Note:
|
||||
- score: Quality/relevance of the result (0.0 or 1.0)
|
||||
"""
|
||||
self.total_count += 1
|
||||
|
|
@ -54,13 +54,13 @@ class BenchmarkStats:
|
|||
return {
|
||||
"name": self.name,
|
||||
"total_calls": 0,
|
||||
"avg_score": 0.0
|
||||
"avg_score": 0.0,
|
||||
}
|
||||
|
||||
return {
|
||||
"name": self.name,
|
||||
"total_calls": self.total_count,
|
||||
"avg_score": round(sum(self.scores) / len(self.scores), 3)
|
||||
"avg_score": round(sum(self.scores) / len(self.scores), 3),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -87,18 +87,18 @@ def delete_workspace(workspace_id: str) -> bool:
|
|||
def load_queries(query_file: str = "query.json") -> Dict[str, Any]:
|
||||
"""加载查询数据"""
|
||||
query_path = Path(__file__).parent / query_file
|
||||
with open(query_path, 'r', encoding='utf-8') as f:
|
||||
with open(query_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template: str = "") -> List[ToolCallResult]:
|
||||
"""运行use_mock_search并收集结果(支持并发)
|
||||
|
||||
|
||||
Args:
|
||||
workspace_id: 工作空间ID
|
||||
queries: 查询列表
|
||||
prompt_template: 提示模板
|
||||
|
||||
|
||||
Returns:
|
||||
工具调用结果列表
|
||||
"""
|
||||
|
|
@ -112,10 +112,13 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template:
|
|||
# 提交之前sleep 1秒
|
||||
time.sleep(1)
|
||||
|
||||
result = api_call("use_mock_search", {
|
||||
"workspace_id": workspace_id,
|
||||
"query": prompt_template.format(query=query),
|
||||
})
|
||||
result = api_call(
|
||||
"use_mock_search",
|
||||
{
|
||||
"workspace_id": workspace_id,
|
||||
"query": prompt_template.format(query=query),
|
||||
},
|
||||
)
|
||||
|
||||
if result:
|
||||
tool_call_result = result.get("answer")
|
||||
|
|
@ -128,8 +131,7 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template:
|
|||
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||
# 提交所有任务
|
||||
future_to_query = {
|
||||
executor.submit(process_single_query, idx, query): (idx, query)
|
||||
for idx, query in enumerate(queries)
|
||||
executor.submit(process_single_query, idx, query): (idx, query) for idx, query in enumerate(queries)
|
||||
}
|
||||
|
||||
# 收集结果(按完成顺序)
|
||||
|
|
@ -153,11 +155,11 @@ def run_use_mock_search(workspace_id: str, queries: List[str], prompt_template:
|
|||
|
||||
def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> List[ToolCallResult]:
|
||||
"""批量添加工具调用结果到记忆库,并返回带评分的结果
|
||||
|
||||
|
||||
Args:
|
||||
workspace_id: 工作空间ID
|
||||
results: 工具调用结果列表
|
||||
|
||||
|
||||
Returns:
|
||||
从API返回的memory_list中提取的带评分的ToolCallResult列表
|
||||
"""
|
||||
|
|
@ -171,10 +173,13 @@ def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> L
|
|||
logger.info(f"Adding tool call results to {workspace_id}: {len(tool_call_results)} results")
|
||||
|
||||
# 统一调用API,让后端自动按tool_name分组处理
|
||||
api_result = api_call("add_tool_call_result", {
|
||||
"workspace_id": workspace_id,
|
||||
"tool_call_results": tool_call_results
|
||||
})
|
||||
api_result = api_call(
|
||||
"add_tool_call_result",
|
||||
{
|
||||
"workspace_id": workspace_id,
|
||||
"tool_call_results": tool_call_results,
|
||||
},
|
||||
)
|
||||
|
||||
if not api_result:
|
||||
logger.error("Failed to add results")
|
||||
|
|
@ -206,10 +211,13 @@ def add_tool_call_results(workspace_id: str, results: List[ToolCallResult]) -> L
|
|||
def summarize_tool_memory(workspace_id: str, tool_names: str) -> bool:
|
||||
"""总结工具记忆"""
|
||||
logger.info(f"Summarizing tool memory for {workspace_id}: {tool_names}")
|
||||
result = api_call("summary_tool_memory", {
|
||||
"workspace_id": workspace_id,
|
||||
"tool_names": tool_names
|
||||
})
|
||||
result = api_call(
|
||||
"summary_tool_memory",
|
||||
{
|
||||
"workspace_id": workspace_id,
|
||||
"tool_names": tool_names,
|
||||
},
|
||||
)
|
||||
|
||||
if result:
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
|
|
@ -221,19 +229,22 @@ def summarize_tool_memory(workspace_id: str, tool_names: str) -> bool:
|
|||
|
||||
def retrieve_tool_memory(workspace_id: str, tool_names: str) -> str:
|
||||
"""检索工具记忆并返回格式化的内容
|
||||
|
||||
|
||||
Args:
|
||||
workspace_id: 工作空间ID
|
||||
tool_names: 逗号分隔的工具名称
|
||||
|
||||
|
||||
Returns:
|
||||
格式化的工具记忆内容,每个工具名称作为一级markdown标题
|
||||
"""
|
||||
logger.info(f"Retrieving tool memory for {workspace_id}: {tool_names}")
|
||||
result = api_call("retrieve_tool_memory", {
|
||||
"workspace_id": workspace_id,
|
||||
"tool_names": tool_names
|
||||
})
|
||||
result = api_call(
|
||||
"retrieve_tool_memory",
|
||||
{
|
||||
"workspace_id": workspace_id,
|
||||
"tool_names": tool_names,
|
||||
},
|
||||
)
|
||||
|
||||
if not result:
|
||||
logger.error("Failed to retrieve tool memory")
|
||||
|
|
@ -251,8 +262,9 @@ def retrieve_tool_memory(workspace_id: str, tool_names: str) -> str:
|
|||
tool_name = tool_memory.when_to_use or "Unknown Tool"
|
||||
formatted_section = f"# {tool_name}\n\n{tool_memory.content}"
|
||||
formatted_contents.append(formatted_section)
|
||||
logger.info(f"Retrieved content for tool: {tool_name}, "
|
||||
f"content_length={len(tool_memory.content)}")
|
||||
logger.info(
|
||||
f"Retrieved content for tool: {tool_name}, " f"content_length={len(tool_memory.content)}",
|
||||
)
|
||||
|
||||
# 用两个换行符分隔不同工具的记忆
|
||||
joined_content = "\n\n".join(formatted_contents)
|
||||
|
|
@ -265,7 +277,7 @@ def collect_statistics(results: List[ToolCallResult], stats: BenchmarkStats) ->
|
|||
"""从结果列表中收集统计数据"""
|
||||
for result in results:
|
||||
# 转换为字典用于统计
|
||||
result_dict = result.model_dump() if hasattr(result, 'model_dump') else result
|
||||
result_dict = result.model_dump() if hasattr(result, "model_dump") else result
|
||||
stats.add_result(result_dict)
|
||||
|
||||
|
||||
|
|
@ -276,11 +288,13 @@ def print_comparison_table(stats_list: List[BenchmarkStats]) -> None:
|
|||
|
||||
for stats in stats_list:
|
||||
summary = stats.get_summary()
|
||||
rows.append([
|
||||
summary["name"],
|
||||
summary["total_calls"],
|
||||
summary["avg_score"]
|
||||
])
|
||||
rows.append(
|
||||
[
|
||||
summary["name"],
|
||||
summary["total_calls"],
|
||||
summary["avg_score"],
|
||||
],
|
||||
)
|
||||
|
||||
print("\n" + "=" * 100)
|
||||
print("BENCHMARK RESULTS COMPARISON")
|
||||
|
|
@ -299,8 +313,7 @@ def calculate_improvements(baseline_stats: BenchmarkStats, improved_stats: Bench
|
|||
|
||||
# 平均分数改进(相对提升百分比)
|
||||
if baseline["avg_score"] > 0:
|
||||
improvements["avg_score"] = ((improved["avg_score"] - baseline["avg_score"])
|
||||
/ baseline["avg_score"] * 100)
|
||||
improvements["avg_score"] = (improved["avg_score"] - baseline["avg_score"]) / baseline["avg_score"] * 100
|
||||
else:
|
||||
improvements["avg_score"] = 0.0
|
||||
|
||||
|
|
@ -314,7 +327,7 @@ def print_improvements(improvements: Dict[str, float]) -> None:
|
|||
print("=" * 100)
|
||||
|
||||
metric_labels = {
|
||||
"avg_score": "Average Score"
|
||||
"avg_score": "Average Score",
|
||||
}
|
||||
|
||||
for metric, improvement in improvements.items():
|
||||
|
|
@ -328,19 +341,19 @@ def print_improvements(improvements: Dict[str, float]) -> None:
|
|||
def save_results(results: Dict[str, Any], filename: str = "benchmark_results.json") -> None:
|
||||
"""保存结果到文件"""
|
||||
output_path = Path(__file__).parent / filename
|
||||
with open(output_path, 'w', encoding='utf-8') as f:
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||
logger.info(f"Results saved to {output_path}")
|
||||
|
||||
|
||||
def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: List[str]) -> Dict[str, Any]:
|
||||
"""运行单个epoch的benchmark
|
||||
|
||||
|
||||
Args:
|
||||
epoch_num: epoch编号(从1开始)
|
||||
train_queries: 训练查询列表
|
||||
test_queries: 测试查询列表
|
||||
|
||||
|
||||
Returns:
|
||||
包含该epoch统计结果的字典
|
||||
"""
|
||||
|
|
@ -422,7 +435,7 @@ def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: Lis
|
|||
tool_names_set = set()
|
||||
results_to_use = train_scored_results if train_scored_results else train_results_no_memory
|
||||
for result in results_to_use:
|
||||
tool_name = result.tool_name if hasattr(result, 'tool_name') else None
|
||||
tool_name = result.tool_name if hasattr(result, "tool_name") else None
|
||||
if tool_name:
|
||||
tool_names_set.add(tool_name)
|
||||
|
||||
|
|
@ -507,15 +520,15 @@ def run_single_epoch(epoch_num: int, train_queries: List[str], test_queries: Lis
|
|||
"statistics": {
|
||||
"train_no_memory": train_no_memory_stats.get_summary(),
|
||||
"test_no_memory": test_no_memory_stats.get_summary(),
|
||||
"test_with_memory": test_with_memory_stats.get_summary()
|
||||
"test_with_memory": test_with_memory_stats.get_summary(),
|
||||
},
|
||||
"improvements": improvements
|
||||
"improvements": improvements,
|
||||
}
|
||||
|
||||
|
||||
def main(test_mode: bool = False, run_epoch: int = 3):
|
||||
"""主函数:运行完整的benchmark流程
|
||||
|
||||
|
||||
Args:
|
||||
test_mode: 如果为True,只使用每个难度级别的前3个查询进行快速测试
|
||||
run_epoch: 运行的epoch数量,默认为3
|
||||
|
|
@ -567,11 +580,14 @@ def main(test_mode: bool = False, run_epoch: int = 3):
|
|||
|
||||
# 计算每个场景的平均分数
|
||||
avg_train_no_memory = sum(e["statistics"]["train_no_memory"]["avg_score"] for e in all_epoch_results) / len(
|
||||
all_epoch_results)
|
||||
all_epoch_results,
|
||||
)
|
||||
avg_test_no_memory = sum(e["statistics"]["test_no_memory"]["avg_score"] for e in all_epoch_results) / len(
|
||||
all_epoch_results)
|
||||
all_epoch_results,
|
||||
)
|
||||
avg_test_with_memory = sum(e["statistics"]["test_with_memory"]["avg_score"] for e in all_epoch_results) / len(
|
||||
all_epoch_results)
|
||||
all_epoch_results,
|
||||
)
|
||||
|
||||
# 计算平均改进
|
||||
avg_improvement = sum(e["improvements"]["avg_score"] for e in all_epoch_results) / len(all_epoch_results)
|
||||
|
|
@ -581,7 +597,7 @@ def main(test_mode: bool = False, run_epoch: int = 3):
|
|||
rows = [
|
||||
["Train (No Memory)", f"{avg_train_no_memory:.3f}"],
|
||||
["Test (No Memory)", f"{avg_test_no_memory:.3f}"],
|
||||
["Test (With Memory)", f"{avg_test_with_memory:.3f}"]
|
||||
["Test (With Memory)", f"{avg_test_with_memory:.3f}"],
|
||||
]
|
||||
print(tabulate(rows, headers=headers, tablefmt="grid"))
|
||||
|
||||
|
|
@ -597,13 +613,15 @@ def main(test_mode: bool = False, run_epoch: int = 3):
|
|||
headers = ["Epoch", "Train (No Mem)", "Test (No Mem)", "Test (With Mem)", "Improvement %"]
|
||||
rows = []
|
||||
for e in all_epoch_results:
|
||||
rows.append([
|
||||
f"Epoch {e['epoch']}",
|
||||
f"{e['statistics']['train_no_memory']['avg_score']:.3f}",
|
||||
f"{e['statistics']['test_no_memory']['avg_score']:.3f}",
|
||||
f"{e['statistics']['test_with_memory']['avg_score']:.3f}",
|
||||
f"{e['improvements']['avg_score']:+.2f}%"
|
||||
])
|
||||
rows.append(
|
||||
[
|
||||
f"Epoch {e['epoch']}",
|
||||
f"{e['statistics']['train_no_memory']['avg_score']:.3f}",
|
||||
f"{e['statistics']['test_no_memory']['avg_score']:.3f}",
|
||||
f"{e['statistics']['test_with_memory']['avg_score']:.3f}",
|
||||
f"{e['improvements']['avg_score']:+.2f}%",
|
||||
],
|
||||
)
|
||||
|
||||
print(tabulate(rows, headers=headers, tablefmt="grid"))
|
||||
|
||||
|
|
@ -616,9 +634,9 @@ def main(test_mode: bool = False, run_epoch: int = 3):
|
|||
"train_no_memory": avg_train_no_memory,
|
||||
"test_no_memory": avg_test_no_memory,
|
||||
"test_with_memory": avg_test_with_memory,
|
||||
"improvement": avg_improvement
|
||||
"improvement": avg_improvement,
|
||||
},
|
||||
"per_epoch_results": all_epoch_results
|
||||
"per_epoch_results": all_epoch_results,
|
||||
}
|
||||
|
||||
save_results(benchmark_summary, "tool_memory_benchmark_results.json")
|
||||
|
|
|
|||
|
|
@ -22,6 +22,36 @@ git checkout -b your-feature-branch-name
|
|||
### Making Changes
|
||||
With your new branch checked out, you can now make your changes to the code. Remember to keep your changes as focused as possible. If you're addressing multiple issues or features, it's better to create separate branches and pull requests for each.
|
||||
|
||||
### Set Up Pre-commit Hooks
|
||||
Before committing your changes, you should set up pre-commit hooks to ensure code quality and consistency. Pre-commit hooks will automatically check your code for common issues and format it according to the project's standards.
|
||||
|
||||
**Install pre-commit:**
|
||||
```bash
|
||||
pip install pre-commit
|
||||
```
|
||||
|
||||
**Install the git hooks:**
|
||||
```bash
|
||||
pre-commit install
|
||||
```
|
||||
|
||||
**Run pre-commit manually (optional):**
|
||||
If you want to run pre-commit checks on all files before committing, you can run:
|
||||
```bash
|
||||
pre-commit run --all-files
|
||||
```
|
||||
|
||||
After installation, pre-commit will automatically run on `git commit` to check your code. The hooks will check for:
|
||||
- Code syntax and AST validation
|
||||
- YAML, XML, TOML, and JSON format validation
|
||||
- Trailing whitespace
|
||||
- Code formatting (Black)
|
||||
- Code style (Flake8)
|
||||
- Code quality (Pylint)
|
||||
- Package metadata (Pyroma)
|
||||
|
||||
If any checks fail, please fix the issues before committing.
|
||||
|
||||
### Commit Your Changes
|
||||
Once you've made your changes, it's time to commit them. Write clear and concise commit messages that explain your changes.
|
||||
```bash
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# AppWorld
|
||||
# AppWorld
|
||||
Experiment Quick Start Guide
|
||||
|
||||
This guide helps you quickly set up and run AppWorld experiments with ReMe integration.
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# BFCL
|
||||
# BFCL
|
||||
Experiment Quick Start Guide
|
||||
|
||||
This guide helps you quickly set up and run BFCL experiments with ReMe integration.
|
||||
|
|
@ -40,7 +40,7 @@ Run the main experiment script to collect agent trajectories on training data se
|
|||
python run_bfcl.py
|
||||
```
|
||||
|
||||
**Note**:
|
||||
**Note**:
|
||||
- `max_workers`: Number of parallel workers (default: `4`)
|
||||
- `num_runs`: Number of times each task is repeated (default: `1`)
|
||||
- `model_name`: LLM model name (default: `qwen3-8b`)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ We tested ReMe on Appworld using qwen3-8b:
|
|||
| with ReMe | 0.109 **(+2.6%)** | 0.175 **(+3.5%)** | 0.281 **(+5.3%)** |
|
||||
|
||||
Pass@K measures the probability that at least one of the K generated samples successfully completes the task (
|
||||
score=1).
|
||||
score=1).
|
||||
The current experiment uses an internal AppWorld environment, which may have slight differences.
|
||||
|
||||
You can find more details on reproducing the experiment in [quickstart.md](appworld/quickstart.md).
|
||||
|
|
@ -43,10 +43,10 @@ We tested ReMe on BFCL-V3 multi-turn-base (randomly split 50train/150val) using
|
|||
|
||||
We evaluated Tool Memory effectiveness using a controlled benchmark with three mock search tools using Qwen3-30B-Instruct:
|
||||
|
||||
| Scenario | Avg Score | Improvement |
|
||||
|-----------------------|-----------|--------------------|
|
||||
| Train (No Memory) | 0.650 | - |
|
||||
| Test (No Memory) | 0.672 | Baseline |
|
||||
| Scenario | Avg Score | Improvement |
|
||||
|------------------------|-----------|-------------|
|
||||
| Train (No Memory) | 0.650 | - |
|
||||
| Test (No Memory) | 0.672 | Baseline |
|
||||
| **Test (With Memory)** | **0.772** | **+14.88%** |
|
||||
|
||||
**Key Findings:**
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ kernelspec:
|
|||
name: python3
|
||||
---
|
||||
|
||||
# FrozenLake
|
||||
# FrozenLake
|
||||
Experiment Quick Start Guide
|
||||
|
||||
This guide helps you quickly set up and run FrozenLake experiments with ReMe integration. The FrozenLake experiment demonstrates how task memory can improve an agent's performance in a navigation task.
|
||||
|
|
|
|||
|
|
@ -16,8 +16,8 @@ kernelspec:
|
|||
<em>Remember Me, Refine Me.</em>
|
||||
|
||||
<div class="flex justify-center space-x-3">
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.12+-blue" alt="Python Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-v0.1.10.7-blue?logo=pypi" alt="PyPI Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/python-3.10+-blue" alt="Python Version"></a>
|
||||
<a href="https://pypi.org/project/reme-ai/"><img src="https://img.shields.io/badge/pypi-0.2.0.0-blue?logo=pypi" alt="PyPI Version"></a>
|
||||
<a href="./LICENSE"><img src="https://img.shields.io/badge/license-Apache--2.0-black" alt="License"></a>
|
||||
<a href="https://github.com/agentscope-ai/ReMe"><img src="https://img.shields.io/github/stars/modelscope/ReMe?style=social" alt="GitHub Stars"></a>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -184,11 +184,11 @@ The `summary_task_memory` tool transforms conversation trajectories into valuabl
|
|||
async def run_summary(client, messages):
|
||||
"""
|
||||
Generate a summary of conversation messages and create task memories
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
messages: List of message objects from a conversation
|
||||
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
|
|
@ -227,11 +227,11 @@ The `retrieve_task_memory` tool allows you to retrieve relevant memories based o
|
|||
async def run_retrieve(client, query):
|
||||
"""
|
||||
Retrieve relevant task memories based on a query
|
||||
|
||||
|
||||
Args:
|
||||
client: MCP client instance
|
||||
query: The query to retrieve relevant memories
|
||||
|
||||
|
||||
Returns:
|
||||
String containing the retrieved memory answer
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -88,7 +88,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Retriever: Get relevant memories
|
||||
result = await app.async_execute(
|
||||
name="retrieve_task_memory",
|
||||
|
|
@ -219,7 +219,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Memory Retrieval: Get personal memory fragments
|
||||
result = await app.async_execute(
|
||||
name="retrieve_personal_memory",
|
||||
|
|
@ -366,7 +366,7 @@ async def main():
|
|||
]
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Generate usage guidelines from history
|
||||
result = await app.async_execute(
|
||||
name="summary_tool_memory",
|
||||
|
|
@ -374,7 +374,7 @@ async def main():
|
|||
tool_names="web_search"
|
||||
)
|
||||
print(result)
|
||||
|
||||
|
||||
# Retrieve tool guidelines before use
|
||||
result = await app.async_execute(
|
||||
name="retrieve_tool_memory",
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ Here's a complete example workflow that demonstrates how to use task memory:
|
|||
def run_agent_with_memory(query_first, query_second):
|
||||
# Run agent with second query to build initial memories
|
||||
messages = run_agent(query=query_second)
|
||||
|
||||
|
||||
# Summarize conversation to create memories
|
||||
requests.post(
|
||||
url=f"{BASE_URL}summary_task_memory",
|
||||
|
|
@ -150,7 +150,7 @@ def run_agent_with_memory(query_first, query_second):
|
|||
]
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# Retrieve relevant memories for the first query
|
||||
response = requests.post(
|
||||
url=f"{BASE_URL}retrieve_task_memory",
|
||||
|
|
@ -160,7 +160,7 @@ def run_agent_with_memory(query_first, query_second):
|
|||
}
|
||||
)
|
||||
retrieved_memory = response.json().get("answer", "")
|
||||
|
||||
|
||||
# Run agent with first query augmented with retrieved memories
|
||||
augmented_query = f"{retrieved_memory}\n\nUser Question:\n{query_first}"
|
||||
return run_agent(query=augmented_query)
|
||||
|
|
|
|||
|
|
@ -24,11 +24,11 @@ This benchmark evaluates Tool Memory effectiveness by comparing agent performanc
|
|||
|
||||
Three LLM-based mock search tools with different performance profiles:
|
||||
|
||||
| Tool | Simple Queries | Medium Queries | Complex Queries |
|
||||
|------|---------------|----------------|-----------------|
|
||||
| **SearchToolA** | ⭐⭐⭐ Fast, high success (90%) | ❌ Poor (20% success) | ⚠️ Weak (50% success) |
|
||||
| **SearchToolB** | ⚠️ Over-engineered (30%) | ⭐⭐⭐ Optimal (90% success) | ⚠️ Limited (50% success) |
|
||||
| **SearchToolC** | ⚠️ Overkill (30%) | ⚠️ Excessive (40%) | ⭐⭐⭐ Best (90% success) |
|
||||
| Tool | Simple Queries | Medium Queries | Complex Queries |
|
||||
|-----------------|------------------------------|---------------------------|--------------------------|
|
||||
| **SearchToolA** | ⭐⭐⭐ Fast, high success (90%) | ❌ Poor (20% success) | ⚠️ Weak (50% success) |
|
||||
| **SearchToolB** | ⚠️ Over-engineered (30%) | ⭐⭐⭐ Optimal (90% success) | ⚠️ Limited (50% success) |
|
||||
| **SearchToolC** | ⚠️ Overkill (30%) | ⚠️ Excessive (40%) | ⭐⭐⭐ Best (90% success) |
|
||||
|
||||
**Performance Characteristics:**
|
||||
- `success_rate`: Probability of successful execution (vs "Service busy" error)
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ When an LLM faces numerous MCP tools, it relies heavily on tool descriptions to
|
|||
Imagine an LLM choosing between three search tools:
|
||||
```
|
||||
Tool A: "Search the web for information"
|
||||
Tool B: "Perform web searches with customizable parameters"
|
||||
Tool B: "Perform web searches with customizable parameters"
|
||||
Tool C: "Query search engines and return results"
|
||||
```
|
||||
|
||||
|
|
@ -108,7 +108,7 @@ LLM: "I have 50 search tools, but Tool A has 95% success for technical queries"
|
|||
```
|
||||
Before Tool Memory:
|
||||
- Success rate: 75%
|
||||
- Average time cost: 5.2s
|
||||
- Average time cost: 5.2s
|
||||
- Token cost: 200+ per call
|
||||
- Repeated parameter errors
|
||||
- Random tool selection
|
||||
|
|
@ -636,7 +636,7 @@ op:
|
|||
params:
|
||||
max_history_tool_call_cnt: 200 # Keep more history
|
||||
evaluation_sleep_interval: 0.5 # Faster evaluation
|
||||
|
||||
|
||||
summary_tool_memory_op:
|
||||
params:
|
||||
recent_call_count: 50 # Analyze more calls
|
||||
|
|
@ -649,7 +649,7 @@ op:
|
|||
params:
|
||||
max_history_tool_call_cnt: 50 # Less history needed
|
||||
evaluation_sleep_interval: 1.0 # Standard rate
|
||||
|
||||
|
||||
summary_tool_memory_op:
|
||||
params:
|
||||
recent_call_count: 20 # Analyze fewer calls
|
||||
|
|
@ -661,10 +661,10 @@ op:
|
|||
```{code-cell}
|
||||
memory = retrieve_tool_memory("web_search")['metadata']['memory_list'][0]
|
||||
stats = ToolMemory(**memory).statistic(recent_frequency=30)
|
||||
|
||||
|
||||
print(f"Success Rate: {stats['success_rate']:.2%}")
|
||||
print(f"Avg Score: {stats['avg_score']:.2f}")
|
||||
|
||||
|
||||
if stats['success_rate'] < 0.7:
|
||||
print("⚠️ Low success rate - investigate tool issues")
|
||||
```
|
||||
|
|
|
|||
|
|
@ -12,9 +12,9 @@ kernelspec:
|
|||
name: python3
|
||||
---
|
||||
|
||||
# Vector Store API Guide
|
||||
# Vector Store Configuration Guide
|
||||
|
||||
This guide covers the vector store implementations available in ReMe, their APIs, and how to use them effectively.
|
||||
This guide covers how to configure vector store backends in ReMe using the `default.yaml` configuration file.
|
||||
|
||||
## 📋 Overview
|
||||
|
||||
|
|
@ -43,71 +43,28 @@ All vector stores implement the `BaseVectorStore` interface, providing a consist
|
|||
| **Async Support** | ❌ No | ❌ No | ❌ No | ✅ Native | ❌ No |
|
||||
| **Best For** | Development | Local Apps | Production | Production/Cloud | Testing |
|
||||
|
||||
## 🔄 Common API Methods
|
||||
## ⚙️ Configuration in default.yaml
|
||||
|
||||
All vector store implementations share these core methods:
|
||||
All vector stores are configured in the `vector_store` section of `reme_ai/config/default.yaml`. The configuration structure is:
|
||||
|
||||
### 🔄 Async Support
|
||||
|
||||
All vector stores provide both synchronous and asynchronous versions of every method:
|
||||
|
||||
```python
|
||||
# Synchronous methods
|
||||
store.search(query="example", workspace_id="workspace", top_k=5)
|
||||
store.insert(nodes, workspace_id="workspace")
|
||||
|
||||
# Asynchronous methods (with async_ prefix)
|
||||
await store.async_search(query="example", workspace_id="workspace", top_k=5)
|
||||
await store.async_insert(nodes, workspace_id="workspace")
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: <backend_name> # Required: local, chroma, elasticsearch, qdrant, or memory
|
||||
embedding_model: default # Required: Name of the embedding model configuration
|
||||
params: # Optional: Backend-specific parameters
|
||||
# Backend-specific parameters go here
|
||||
```
|
||||
|
||||
### Workspace Management
|
||||
### Configuration Fields
|
||||
|
||||
```python
|
||||
# Check if workspace exists
|
||||
store.exist_workspace(workspace_id: str) -> bool
|
||||
- **`backend`** (required): The vector store backend to use. Valid values: `local`, `chroma`, `elasticsearch`, `qdrant`, `memory`
|
||||
- **`embedding_model`** (required): The name of the embedding model configuration from the `embedding_model` section
|
||||
- **`params`** (optional): A dictionary of backend-specific parameters that will be passed to the vector store constructor
|
||||
|
||||
# Create a new workspace
|
||||
store.create_workspace(workspace_id: str, **kwargs)
|
||||
## 📁 Vector Store Backend Configurations
|
||||
|
||||
# Delete a workspace
|
||||
store.delete_workspace(workspace_id: str, **kwargs)
|
||||
|
||||
# Copy a workspace
|
||||
store.copy_workspace(src_workspace_id: str, dest_workspace_id: str, **kwargs)
|
||||
```
|
||||
|
||||
### Data Operations
|
||||
|
||||
```python
|
||||
# Insert nodes (single or list)
|
||||
store.insert(nodes: VectorNode | List[VectorNode], workspace_id: str, **kwargs)
|
||||
|
||||
# Delete nodes by ID
|
||||
store.delete(node_ids: str | List[str], workspace_id: str, **kwargs)
|
||||
|
||||
# Search for similar nodes
|
||||
store.search(query: str, workspace_id: str, top_k: int = 1, **kwargs) -> List[VectorNode]
|
||||
|
||||
# Iterate through workspace nodes
|
||||
for node in store.iter_workspace_nodes(workspace_id: str, **kwargs):
|
||||
# Process each node
|
||||
```
|
||||
|
||||
### Import/Export
|
||||
|
||||
```python
|
||||
# Export workspace to file
|
||||
store.dump_workspace(workspace_id: str, path: str | Path = "", callback_fn=None, **kwargs)
|
||||
|
||||
# Import workspace from file
|
||||
store.load_workspace(workspace_id: str, path: str | Path = "", nodes: List[VectorNode] = None,
|
||||
callback_fn=None, **kwargs)
|
||||
```
|
||||
|
||||
## ⚡ Vector Store Implementations
|
||||
|
||||
### 1. 📁 LocalVectorStore (`backend=local`)
|
||||
### 1. LocalVectorStore (`backend=local`)
|
||||
|
||||
A simple file-based vector store that saves data to local JSONL files.
|
||||
|
||||
|
|
@ -118,62 +75,22 @@ A simple file-based vector store that saves data to local JSONL files.
|
|||
|
||||
#### ⚙️ Configuration
|
||||
|
||||
```python
|
||||
from flowllm.storage.vector_store import LocalVectorStore
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from flowllm.utils.common_utils import load_env
|
||||
|
||||
# Load environment variables (for API keys)
|
||||
load_env()
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
|
||||
# Initialize vector store
|
||||
vector_store = LocalVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
store_dir="./file_vector_store", # Directory to store JSONL files
|
||||
batch_size=1024 # Batch size for operations
|
||||
)
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: local
|
||||
embedding_model: default
|
||||
params:
|
||||
store_dir: "./local_vector_store" # Directory to store JSONL files (default: "./local_vector_store")
|
||||
batch_size: 1024 # Batch size for operations (default: 1024)
|
||||
```
|
||||
|
||||
#### 💻 Example Usage
|
||||
#### Configuration Parameters
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
- **`store_dir`** (optional): Directory path where workspace files are stored. Default: `"./local_vector_store"`
|
||||
- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024`
|
||||
|
||||
# Create workspace
|
||||
workspace_id = "my_workspace"
|
||||
vector_store.create_workspace(workspace_id)
|
||||
|
||||
# Create nodes
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Artificial intelligence is revolutionizing technology",
|
||||
metadata={"category": "tech", "source": "article1"}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node2",
|
||||
workspace_id=workspace_id,
|
||||
content="Machine learning enables data-driven insights",
|
||||
metadata={"category": "tech", "source": "article2"}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert nodes
|
||||
vector_store.insert(nodes, workspace_id)
|
||||
|
||||
# Search
|
||||
results = vector_store.search("What is AI?", workspace_id, top_k=2)
|
||||
for result in results:
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Metadata: {result.metadata}")
|
||||
print(f"Score: {result.metadata.get('score', 'N/A')}")
|
||||
```
|
||||
|
||||
### 2. 🔮 ChromaVectorStore (`backend=chroma`)
|
||||
### 2. ChromaVectorStore (`backend=chroma`)
|
||||
|
||||
An embedded vector database that provides persistent storage with advanced features.
|
||||
|
||||
|
|
@ -184,71 +101,22 @@ An embedded vector database that provides persistent storage with advanced featu
|
|||
|
||||
#### ⚙️ Configuration
|
||||
|
||||
```python
|
||||
from flowllm.storage.vector_store import ChromaVectorStore
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from flowllm.utils.common_utils import load_env
|
||||
|
||||
# Load environment variables
|
||||
load_env()
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
|
||||
# Initialize vector store
|
||||
vector_store = ChromaVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
store_dir="./chroma_vector_store", # Directory for Chroma database
|
||||
batch_size=1024 # Batch size for operations
|
||||
)
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: chroma
|
||||
embedding_model: default
|
||||
params:
|
||||
store_dir: "./chroma_vector_store" # Directory for Chroma database (default: "./chroma_vector_store")
|
||||
batch_size: 1024 # Batch size for operations (default: 1024)
|
||||
```
|
||||
|
||||
#### 💻 Example Usage
|
||||
#### Configuration Parameters
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
- **`store_dir`** (optional): Directory path where ChromaDB data is persisted. Default: `"./chroma_vector_store"`
|
||||
- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024`
|
||||
|
||||
workspace_id = "chroma_workspace"
|
||||
|
||||
# Check if workspace exists and create if needed
|
||||
if not vector_store.exist_workspace(workspace_id):
|
||||
vector_store.create_workspace(workspace_id)
|
||||
|
||||
# Create nodes with metadata
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Deep learning models require large datasets",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"difficulty": "advanced",
|
||||
"topic": "deep_learning"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node2",
|
||||
workspace_id=workspace_id,
|
||||
content="Transformer architecture revolutionized NLP",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"difficulty": "intermediate",
|
||||
"topic": "transformers"
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert nodes
|
||||
vector_store.insert(nodes, workspace_id)
|
||||
|
||||
# Search
|
||||
results = vector_store.search("deep learning", workspace_id, top_k=5)
|
||||
for result in results:
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Metadata: {result.metadata}")
|
||||
```
|
||||
|
||||
### 3. 🔍 EsVectorStore (`backend=elasticsearch`)
|
||||
### 3. EsVectorStore (`backend=elasticsearch`)
|
||||
|
||||
Production-grade vector search using Elasticsearch with advanced filtering and scaling capabilities.
|
||||
|
||||
|
|
@ -282,114 +150,24 @@ export FLOW_ES_HOSTS=http://localhost:9200
|
|||
|
||||
#### ⚙️ Configuration
|
||||
|
||||
```python
|
||||
from flowllm.storage.vector_store import EsVectorStore
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from flowllm.utils.common_utils import load_env
|
||||
import os
|
||||
|
||||
# Load environment variables
|
||||
load_env()
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
|
||||
# Initialize vector store
|
||||
vector_store = EsVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
hosts=os.getenv("FLOW_ES_HOSTS", "http://localhost:9200"), # Elasticsearch hosts
|
||||
basic_auth=None, # ("username", "password") for auth
|
||||
batch_size=1024 # Batch size for bulk operations
|
||||
)
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: elasticsearch
|
||||
embedding_model: default
|
||||
params:
|
||||
hosts: "http://localhost:9200" # Elasticsearch host(s) - can be string or list (default: from FLOW_ES_HOSTS env var or "http://localhost:9200")
|
||||
basic_auth: null # Optional: ("username", "password") tuple for authentication
|
||||
batch_size: 1024 # Batch size for bulk operations (default: 1024)
|
||||
```
|
||||
|
||||
#### 🎯 Advanced Filtering
|
||||
#### Configuration Parameters
|
||||
|
||||
EsVectorStore supports advanced filtering capabilities through the `filter_dict` parameter:
|
||||
- **`hosts`** (optional): Elasticsearch host(s) as a string or list of strings. Defaults to the `FLOW_ES_HOSTS` environment variable or `"http://localhost:9200"` if not set
|
||||
- **`basic_auth`** (optional): Tuple of `("username", "password")` for basic authentication. Default: `null` (no authentication)
|
||||
- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024`
|
||||
|
||||
```python
|
||||
# Term filters (exact match)
|
||||
term_filter = {
|
||||
"category": "technology",
|
||||
"author": "research_team"
|
||||
}
|
||||
|
||||
# Range filters (numeric and date ranges)
|
||||
range_filter = {
|
||||
"score": {"gte": 0.8}, # Score >= 0.8
|
||||
"confidence": {"gte": 0.5, "lte": 0.9}, # Between 0.5 and 0.9
|
||||
"timestamp": {"gte": "2024-01-01", "lte": "2024-12-31"}
|
||||
}
|
||||
|
||||
# Combined filters (filters are combined with AND logic)
|
||||
combined_filter = {
|
||||
"category": "AI",
|
||||
"confidence": {"gte": 0.9}
|
||||
}
|
||||
|
||||
# Search with filters applied
|
||||
results = vector_store.search("machine learning", workspace_id, top_k=10, filter_dict=combined_filter)
|
||||
```
|
||||
|
||||
#### ⚡ Performance Optimization
|
||||
|
||||
```python
|
||||
# Refresh index for immediate availability (useful after bulk inserts)
|
||||
vector_store.insert(nodes, workspace_id, refresh=True) # Auto-refresh
|
||||
vector_store.refresh(workspace_id) # Manual refresh
|
||||
|
||||
# Bulk operations with custom batch size
|
||||
vector_store.insert(large_node_list, workspace_id, refresh=False) # Skip refresh for speed
|
||||
vector_store.refresh(workspace_id) # Refresh once after all inserts
|
||||
```
|
||||
|
||||
#### 💻 Example Usage
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
|
||||
# Define workspace
|
||||
workspace_id = "production_workspace"
|
||||
|
||||
# Create workspace if needed
|
||||
if not vector_store.exist_workspace(workspace_id):
|
||||
vector_store.create_workspace(workspace_id)
|
||||
|
||||
# Create nodes with rich metadata
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="doc1",
|
||||
workspace_id=workspace_id,
|
||||
content="Transformer architecture revolutionized NLP",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"subcategory": "NLP",
|
||||
"author": "research_team",
|
||||
"timestamp": "2024-01-15",
|
||||
"confidence": 0.95,
|
||||
"tags": ["transformer", "nlp", "attention"]
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert with refresh for immediate availability
|
||||
vector_store.insert(nodes, workspace_id, refresh=True)
|
||||
|
||||
# Advanced search with filters
|
||||
filter_dict = {
|
||||
"category": "AI",
|
||||
"confidence": {"gte": 0.9}
|
||||
}
|
||||
|
||||
results = vector_store.search("transformer models", workspace_id, top_k=5, filter_dict=filter_dict)
|
||||
|
||||
for result in results:
|
||||
print(f"Score: {result.metadata.get('score', 'N/A')}")
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Metadata: {result.metadata}")
|
||||
```
|
||||
|
||||
### 4. 🎯 QdrantVectorStore (`backend=qdrant`)
|
||||
### 4. QdrantVectorStore (`backend=qdrant`)
|
||||
|
||||
A high-performance vector database designed for production workloads with native async support and advanced filtering.
|
||||
|
||||
|
|
@ -430,284 +208,40 @@ export FLOW_QDRANT_API_KEY=your-api-key
|
|||
|
||||
#### ⚙️ Configuration
|
||||
|
||||
```python
|
||||
from flowllm.storage.vector_store import QdrantVectorStore
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from flowllm.utils.common_utils import load_env
|
||||
import os
|
||||
|
||||
# Load environment variables
|
||||
load_env()
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
|
||||
# Option 1: Use localhost with environment variables
|
||||
vector_store = QdrantVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
host=os.getenv("FLOW_QDRANT_HOST", "localhost"),
|
||||
port=int(os.getenv("FLOW_QDRANT_PORT", "6333")),
|
||||
batch_size=1024
|
||||
)
|
||||
|
||||
# Option 2: Use URL (for Qdrant Cloud or remote servers)
|
||||
vector_store = QdrantVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
url="http://your-qdrant-server:6333",
|
||||
api_key="your-api-key", # Optional, for cloud
|
||||
batch_size=1024
|
||||
)
|
||||
|
||||
# Option 3: Specify custom distance metric
|
||||
from qdrant_client.http.models import Distance
|
||||
|
||||
vector_store = QdrantVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
host="localhost",
|
||||
port=6333,
|
||||
distance=Distance.COSINE, # or Distance.EUCLIDEAN, Distance.DOT
|
||||
batch_size=1024
|
||||
)
|
||||
##### Local Qdrant Instance
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: qdrant
|
||||
embedding_model: default
|
||||
params:
|
||||
host: "localhost" # Qdrant host (default: from FLOW_QDRANT_HOST env var or "localhost")
|
||||
port: 6333 # Qdrant port (default: from FLOW_QDRANT_PORT env var or 6333)
|
||||
batch_size: 1024 # Batch size for operations (default: 1024)
|
||||
distance: "COSINE" # Distance metric: "COSINE", "EUCLIDEAN", or "DOT" (default: "COSINE")
|
||||
```
|
||||
|
||||
#### 🎯 Advanced Filtering
|
||||
|
||||
QdrantVectorStore supports advanced filtering capabilities similar to Elasticsearch:
|
||||
|
||||
```python
|
||||
# Term filters (exact match)
|
||||
term_filter = {
|
||||
"category": "AI",
|
||||
"node_type": "research"
|
||||
}
|
||||
|
||||
# Range filters (numeric)
|
||||
range_filter = {
|
||||
"confidence": {"gte": 0.8, "lte": 1.0}, # Between 0.8 and 1.0
|
||||
"score": {"gt": 0.5} # Greater than 0.5
|
||||
}
|
||||
|
||||
# Combined filters (all conditions must match - AND logic)
|
||||
combined_filter = {
|
||||
"category": "AI",
|
||||
"confidence": {"gte": 0.9},
|
||||
"node_type": "research"
|
||||
}
|
||||
|
||||
# Search with filters
|
||||
results = vector_store.search(
|
||||
query="machine learning",
|
||||
workspace_id=workspace_id,
|
||||
top_k=10,
|
||||
filter_dict=combined_filter
|
||||
)
|
||||
##### Qdrant Cloud or Remote Server
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: qdrant
|
||||
embedding_model: default
|
||||
params:
|
||||
url: "https://your-cluster.qdrant.io:6333" # Qdrant server URL (if provided, host and port are ignored)
|
||||
api_key: "your-api-key" # API key for Qdrant Cloud authentication
|
||||
batch_size: 1024 # Batch size for operations (default: 1024)
|
||||
distance: "COSINE" # Distance metric (default: "COSINE")
|
||||
```
|
||||
|
||||
##### Filter Operations Supported:
|
||||
- **Exact match**: `{"field": "value"}`
|
||||
- **Range queries**:
|
||||
- `gte`: Greater than or equal
|
||||
- `lte`: Less than or equal
|
||||
- `gt`: Greater than
|
||||
- `lt`: Less than
|
||||
#### Configuration Parameters
|
||||
|
||||
#### ⚡ Async Operations
|
||||
|
||||
QdrantVectorStore provides **native async support** for all operations:
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
|
||||
async def main():
|
||||
# All operations have async equivalents
|
||||
|
||||
# Check if workspace exists
|
||||
exists = await vector_store.async_exist_workspace(workspace_id)
|
||||
|
||||
# Create workspace
|
||||
if not exists:
|
||||
await vector_store.async_create_workspace(workspace_id)
|
||||
|
||||
# Insert nodes with async embedding
|
||||
await vector_store.async_insert(nodes, workspace_id)
|
||||
|
||||
# Search with async embedding
|
||||
results = await vector_store.async_search(
|
||||
query="AI research",
|
||||
workspace_id=workspace_id,
|
||||
top_k=5,
|
||||
filter_dict={"category": "AI"}
|
||||
)
|
||||
|
||||
# Delete nodes
|
||||
await vector_store.async_delete(node_ids, workspace_id)
|
||||
|
||||
# Delete workspace
|
||||
await vector_store.async_delete_workspace(workspace_id)
|
||||
|
||||
# Close client
|
||||
await vector_store.async_close()
|
||||
|
||||
# Run async operations
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
#### 💻 Example Usage
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
|
||||
workspace_id = "qdrant_workspace"
|
||||
|
||||
# Check and create workspace
|
||||
if not vector_store.exist_workspace(workspace_id):
|
||||
vector_store.create_workspace(workspace_id)
|
||||
|
||||
# Create nodes with rich metadata
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Artificial intelligence is revolutionizing technology",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"node_type": "research",
|
||||
"confidence": 0.95,
|
||||
"author": "research_team"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node2",
|
||||
workspace_id=workspace_id,
|
||||
content="Machine learning models require large datasets",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"node_type": "tutorial",
|
||||
"confidence": 0.85,
|
||||
"author": "education_team"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="node3",
|
||||
workspace_id=workspace_id,
|
||||
content="Deep learning excels at image recognition",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"node_type": "research",
|
||||
"confidence": 0.92,
|
||||
"author": "research_team"
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert nodes (upsert - creates or updates)
|
||||
vector_store.insert(nodes, workspace_id)
|
||||
|
||||
# Simple search
|
||||
results = vector_store.search("What is AI?", workspace_id, top_k=3)
|
||||
for result in results:
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Score: {result.metadata.get('score', 'N/A')}")
|
||||
print(f"Metadata: {result.metadata}")
|
||||
print("-" * 50)
|
||||
|
||||
# Advanced search with filters
|
||||
filter_dict = {
|
||||
"node_type": "research",
|
||||
"confidence": {"gte": 0.9}
|
||||
}
|
||||
|
||||
filtered_results = vector_store.search(
|
||||
query="AI technology",
|
||||
workspace_id=workspace_id,
|
||||
top_k=5,
|
||||
filter_dict=filter_dict
|
||||
)
|
||||
|
||||
print(f"Found {len(filtered_results)} filtered results")
|
||||
for result in filtered_results:
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Metadata: {result.metadata}")
|
||||
|
||||
# Iterate through all nodes
|
||||
print("\nAll nodes in workspace:")
|
||||
for node in vector_store.iter_workspace_nodes(workspace_id, limit=100):
|
||||
print(f"ID: {node.unique_id}, Content: {node.content[:50]}...")
|
||||
|
||||
# Update a node (delete + insert)
|
||||
updated_node = VectorNode(
|
||||
unique_id="node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Artificial intelligence is transforming industries worldwide",
|
||||
metadata={
|
||||
"category": "AI",
|
||||
"node_type": "research",
|
||||
"confidence": 0.98,
|
||||
"author": "research_team",
|
||||
"updated": True
|
||||
}
|
||||
)
|
||||
vector_store.delete("node1", workspace_id)
|
||||
vector_store.insert(updated_node, workspace_id)
|
||||
|
||||
# Export workspace for backup
|
||||
vector_store.dump_workspace(workspace_id, path="./qdrant_backup")
|
||||
|
||||
# Clean up
|
||||
vector_store.close()
|
||||
```
|
||||
|
||||
#### 🔄 Async Example
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
|
||||
async def async_example():
|
||||
workspace_id = "async_qdrant_workspace"
|
||||
|
||||
# Create workspace
|
||||
if not await vector_store.async_exist_workspace(workspace_id):
|
||||
await vector_store.async_create_workspace(workspace_id)
|
||||
|
||||
# Create nodes
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="async_node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Async operations enable better performance",
|
||||
metadata={"type": "performance", "async": True}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="async_node2",
|
||||
workspace_id=workspace_id,
|
||||
content="Concurrent requests improve throughput",
|
||||
metadata={"type": "performance", "async": True}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert with async embedding
|
||||
await vector_store.async_insert(nodes, workspace_id)
|
||||
|
||||
# Search with async embedding
|
||||
results = await vector_store.async_search(
|
||||
query="performance optimization",
|
||||
workspace_id=workspace_id,
|
||||
top_k=2,
|
||||
filter_dict={"async": True}
|
||||
)
|
||||
|
||||
for result in results:
|
||||
print(f"Score: {result.metadata['score']:.4f}")
|
||||
print(f"Content: {result.content}")
|
||||
|
||||
# Cleanup
|
||||
await vector_store.async_delete_workspace(workspace_id)
|
||||
await vector_store.async_close()
|
||||
|
||||
# Run async example
|
||||
asyncio.run(async_example())
|
||||
```
|
||||
- **`url`** (optional): Complete URL for connecting to Qdrant. If provided, `host` and `port` are ignored. Useful for Qdrant Cloud or custom deployments
|
||||
- **`host`** (optional): Host address of the Qdrant server. Defaults to the `FLOW_QDRANT_HOST` environment variable or `"localhost"` if not set
|
||||
- **`port`** (optional): Port number of the Qdrant server. Defaults to the `FLOW_QDRANT_PORT` environment variable or `6333` if not set
|
||||
- **`api_key`** (optional): API key for authentication (required for Qdrant Cloud). Can also be set via `FLOW_QDRANT_API_KEY` environment variable
|
||||
- **`distance`** (optional): Distance metric for vector similarity. Valid values: `"COSINE"`, `"EUCLIDEAN"`, `"DOT"`. Default: `"COSINE"`
|
||||
- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024`
|
||||
|
||||
#### 🌟 Key Features
|
||||
|
||||
|
|
@ -720,23 +254,7 @@ asyncio.run(async_example())
|
|||
- **Persistent Storage** - Data is automatically persisted to disk
|
||||
- **Efficient Iteration** - Scroll through large collections with pagination
|
||||
|
||||
#### 🚨 Important Notes
|
||||
|
||||
- **Collection = Workspace** - Qdrant uses "collections" which map to workspace_id
|
||||
- **Automatic Embedding** - Nodes without vectors are automatically embedded
|
||||
- **ID-based Upsert** - Using the same unique_id will update existing nodes
|
||||
- **Metadata Indexing** - All metadata fields are automatically indexed for filtering
|
||||
- **Connection Management** - Call `close()` or `async_close()` to cleanup connections
|
||||
|
||||
#### 📊 Performance Tips
|
||||
|
||||
1. **Batch Operations** - Insert multiple nodes at once for better performance
|
||||
2. **Use Async** - For high-concurrency scenarios, use async methods
|
||||
3. **Optimize Filters** - Use indexed metadata fields for faster filtering
|
||||
4. **Pagination** - Use `iter_workspace_nodes()` with appropriate `limit` for large collections
|
||||
5. **Distance Metric** - Choose appropriate distance metric for your use case (COSINE for normalized vectors)
|
||||
|
||||
### 5. ⚡ MemoryVectorStore (`backend=memory`)
|
||||
### 5. MemoryVectorStore (`backend=memory`)
|
||||
|
||||
An ultra-fast in-memory vector store that keeps all data in RAM for maximum performance.
|
||||
|
||||
|
|
@ -748,74 +266,20 @@ An ultra-fast in-memory vector store that keeps all data in RAM for maximum perf
|
|||
|
||||
#### ⚙️ Configuration
|
||||
|
||||
```python
|
||||
from flowllm.storage.vector_store import MemoryVectorStore
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from flowllm.utils.common_utils import load_env
|
||||
|
||||
# Load environment variables
|
||||
load_env()
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
|
||||
# Initialize vector store
|
||||
vector_store = MemoryVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
store_dir="./memory_vector_store", # Directory for backup/restore operations
|
||||
batch_size=1024 # Batch size for operations
|
||||
)
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: memory
|
||||
embedding_model: default
|
||||
params:
|
||||
store_dir: "./memory_vector_store" # Directory for backup/restore operations (default: "./memory_vector_store")
|
||||
batch_size: 1024 # Batch size for operations (default: 1024)
|
||||
```
|
||||
|
||||
#### 💻 Example Usage
|
||||
#### Configuration Parameters
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
|
||||
workspace_id = "memory_workspace"
|
||||
|
||||
# Create workspace in memory
|
||||
vector_store.create_workspace(workspace_id)
|
||||
|
||||
# Create nodes
|
||||
nodes = [
|
||||
VectorNode(
|
||||
unique_id="mem_node1",
|
||||
workspace_id=workspace_id,
|
||||
content="Memory stores provide ultra-fast access to data",
|
||||
metadata={
|
||||
"category": "performance",
|
||||
"type": "memory",
|
||||
"speed": "ultra_fast"
|
||||
}
|
||||
),
|
||||
VectorNode(
|
||||
unique_id="mem_node2",
|
||||
workspace_id=workspace_id,
|
||||
content="In-memory databases excel at low-latency operations",
|
||||
metadata={
|
||||
"category": "performance",
|
||||
"type": "database",
|
||||
"latency": "low"
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
# Insert nodes (stored in memory)
|
||||
vector_store.insert(nodes, workspace_id)
|
||||
|
||||
# Ultra-fast search
|
||||
results = vector_store.search("fast memory access", workspace_id, top_k=2)
|
||||
for result in results:
|
||||
print(f"Content: {result.content}")
|
||||
print(f"Score: {result.metadata.get('score', 'N/A')}")
|
||||
|
||||
# Optional: Save to disk for backup
|
||||
vector_store.dump_workspace(workspace_id, path="./backup")
|
||||
|
||||
# Optional: Load from disk to memory
|
||||
vector_store.load_workspace(workspace_id, path="./backup")
|
||||
```
|
||||
- **`store_dir`** (optional): Directory path for backup/restore operations. Default: `"./memory_vector_store"`
|
||||
- **`batch_size`** (optional): Batch size for bulk operations. Default: `1024`
|
||||
|
||||
#### ⚡ Performance Benefits
|
||||
|
||||
|
|
@ -831,69 +295,80 @@ vector_store.load_workspace(workspace_id, path="./backup")
|
|||
- **No persistence** - Use `dump_workspace()` to save to disk
|
||||
- **Single process** - Not suitable for distributed applications
|
||||
|
||||
## 📝 Working with VectorNode
|
||||
## 📝 Example Configurations
|
||||
|
||||
The `VectorNode` class is the fundamental data unit for all vector stores:
|
||||
|
||||
```python
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
|
||||
# Create a node
|
||||
node = VectorNode(
|
||||
unique_id="unique_identifier", # Unique ID for the node (required)
|
||||
workspace_id="my_workspace", # Workspace ID (required)
|
||||
content="Text content to embed", # Content to be embedded (required)
|
||||
metadata={ # Optional metadata
|
||||
"source": "document1",
|
||||
"category": "technology",
|
||||
"timestamp": "2024-08-29"
|
||||
},
|
||||
vector=None # Vector will be generated automatically if None
|
||||
)
|
||||
### Minimal Configuration (Memory Store)
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: memory
|
||||
embedding_model: default
|
||||
```
|
||||
|
||||
## 🔄 Import/Export Example
|
||||
|
||||
Export and import workspaces for backup or transfer:
|
||||
|
||||
```python
|
||||
# Export workspace to file
|
||||
vector_store.dump_workspace(
|
||||
workspace_id="my_workspace",
|
||||
path="./backup_data" # Directory to store the exported data
|
||||
)
|
||||
|
||||
# Import workspace from file
|
||||
vector_store.load_workspace(
|
||||
workspace_id="new_workspace",
|
||||
path="./backup_data" # Directory containing the exported data
|
||||
)
|
||||
|
||||
# Copy workspace within the same store
|
||||
vector_store.copy_workspace(
|
||||
src_workspace_id="original_workspace",
|
||||
dest_workspace_id="copied_workspace"
|
||||
)
|
||||
### Local File Storage
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: local
|
||||
embedding_model: default
|
||||
params:
|
||||
store_dir: "./my_vector_store"
|
||||
batch_size: 2048
|
||||
```
|
||||
|
||||
### Elasticsearch Production Setup
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: elasticsearch
|
||||
embedding_model: default
|
||||
params:
|
||||
hosts: "http://elasticsearch.example.com:9200"
|
||||
basic_auth: ["username", "password"]
|
||||
batch_size: 2048
|
||||
```
|
||||
|
||||
### Qdrant Cloud Setup
|
||||
```yaml
|
||||
vector_store:
|
||||
default:
|
||||
backend: qdrant
|
||||
embedding_model: default
|
||||
params:
|
||||
url: "https://your-cluster.qdrant.io:6333"
|
||||
api_key: "your-api-key-here"
|
||||
distance: "COSINE"
|
||||
batch_size: 1024
|
||||
```
|
||||
|
||||
## 🔄 Environment Variables
|
||||
|
||||
Some vector store backends support environment variables for configuration:
|
||||
|
||||
- **Elasticsearch**: `FLOW_ES_HOSTS` - Elasticsearch host(s)
|
||||
- **Qdrant**:
|
||||
- `FLOW_QDRANT_HOST` - Qdrant host (default: "localhost")
|
||||
- `FLOW_QDRANT_PORT` - Qdrant port (default: 6333)
|
||||
- `FLOW_QDRANT_API_KEY` - Qdrant API key for authentication
|
||||
|
||||
Environment variables are used as fallbacks when parameters are not explicitly set in the YAML configuration.
|
||||
|
||||
## 🧩 Integration with Embedding Models
|
||||
|
||||
All vector stores require an embedding model to function:
|
||||
All vector stores require an embedding model configuration. The `embedding_model` field in the vector store configuration references a model defined in the `embedding_model` section of `default.yaml`:
|
||||
|
||||
```python
|
||||
from flowllm.embedding_model import OpenAICompatibleEmbeddingModel
|
||||
```yaml
|
||||
embedding_model:
|
||||
default:
|
||||
backend: openai_compatible
|
||||
model_name: text-embedding-v4
|
||||
params:
|
||||
dimensions: 1024
|
||||
|
||||
# Initialize embedding model
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(
|
||||
dimensions=64, # Embedding dimensions
|
||||
model_name="text-embedding-v4", # Model name
|
||||
batch_size=32 # Batch size for embedding generation
|
||||
)
|
||||
|
||||
# Pass to vector store (example with LocalVectorStore)
|
||||
# You can also use: ChromaVectorStore, EsVectorStore, QdrantVectorStore, or MemoryVectorStore
|
||||
vector_store = LocalVectorStore(
|
||||
embedding_model=embedding_model,
|
||||
store_dir="./vector_store"
|
||||
)
|
||||
vector_store:
|
||||
default:
|
||||
backend: memory
|
||||
embedding_model: default # References the embedding_model.default configuration
|
||||
```
|
||||
|
||||
The embedding model configuration provides the model name, backend, and parameters needed for generating vector embeddings.
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ build-backend = "setuptools.build_meta"
|
|||
|
||||
[project]
|
||||
name = "reme_ai"
|
||||
version = "0.1.10.8"
|
||||
description = "Remember me"
|
||||
dynamic = ["version"]
|
||||
description = "Remember Me, Refine Me."
|
||||
authors = [
|
||||
{ name = "jinli.yl", email = "jinli.yl@alibaba-inc.com" },
|
||||
{ name = "dengjiaji.djj", email = "dengjiaji.djj@alibaba-inc.com" },
|
||||
|
|
@ -13,24 +13,32 @@ authors = [
|
|||
]
|
||||
license = { file = "LICENSE" }
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Intended Audience :: Science/Research",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
"Operating System :: OS Independent",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
"Topic :: Software Development :: Libraries :: Python Modules",
|
||||
"Topic :: Software Development :: Libraries :: Application Frameworks",
|
||||
"Typing :: Typed",
|
||||
]
|
||||
|
||||
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"]
|
||||
|
||||
dependencies = [
|
||||
"flowllm[reme]>=0.1.11.6",
|
||||
"flowllm[reme]>=0.2.0.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["jupyter-book", "ghp-import", "myst-nb", "sphinxcontrib-bibtex", "furo", "sphinxcontrib-mermaid"]
|
||||
|
||||
all = ["reme_ai[dev]"]
|
||||
full = ["reme_ai[dev]"]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["."]
|
||||
|
|
@ -44,7 +52,15 @@ reme_ai = [
|
|||
"**/*.json",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
reme = "reme_ai.app:main"
|
||||
[tool.setuptools.dynamic]
|
||||
version = { attr = "reme_ai.__version__" }
|
||||
|
||||
# python -m build && twine upload dist/*
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/agentscope-ai/ReMe"
|
||||
Documentation = "https://reme.agentscope.io/"
|
||||
Repository = "https://github.com/agentscope-ai/ReMe"
|
||||
|
||||
[project.scripts]
|
||||
reme = "reme_ai.main:main"
|
||||
|
||||
# python -m build && twine upload dist/*
|
||||
|
|
|
|||
|
|
@ -1,11 +1,33 @@
|
|||
# pylint: disable=wrong-import-position
|
||||
"""ReMe AI - A memory management framework for AI agents."""
|
||||
|
||||
import os
|
||||
|
||||
os.environ["FLOW_APP_NAME"] = "ReMe"
|
||||
|
||||
__version__ = "0.1.10.8"
|
||||
from . import agent # noqa: E402
|
||||
from . import config # noqa: E402
|
||||
from . import constants # noqa: E402
|
||||
from . import enumeration # noqa: E402
|
||||
from . import retrieve # noqa: E402
|
||||
from . import schema # noqa: E402
|
||||
from . import service # noqa: E402
|
||||
from . import summary # noqa: E402
|
||||
from . import utils # noqa: E402
|
||||
from . import vector_store # noqa: E402
|
||||
from .main import ReMeApp # noqa: E402 F401
|
||||
|
||||
from reme_ai.app import ReMeApp
|
||||
from . import agent
|
||||
from . import retrieve
|
||||
from . import summary
|
||||
from . import vector_store
|
||||
__all__ = [
|
||||
"agent",
|
||||
"config",
|
||||
"constants",
|
||||
"enumeration",
|
||||
"retrieve",
|
||||
"schema",
|
||||
"service",
|
||||
"summary",
|
||||
"utils",
|
||||
"vector_store",
|
||||
]
|
||||
|
||||
__version__ = "0.2.0.0"
|
||||
|
|
|
|||
|
|
@ -1,2 +1,14 @@
|
|||
"""Agent module for ReAct and tool-based agent implementations.
|
||||
|
||||
This module provides submodules for different types of agent operations:
|
||||
- react: ReAct (Reasoning and Acting) agent implementations
|
||||
- tools: Mock search tools for testing and demonstration
|
||||
"""
|
||||
|
||||
from . import react
|
||||
from . import tools
|
||||
|
||||
__all__ = [
|
||||
"react",
|
||||
"tools",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1 +1,11 @@
|
|||
"""ReAct agent operations module.
|
||||
|
||||
This module provides ReAct (Reasoning and Acting) agent implementations for
|
||||
answering user queries through iterative reasoning and search actions.
|
||||
"""
|
||||
|
||||
from .simple_react_op import SimpleReactOp
|
||||
|
||||
__all__ = [
|
||||
"SimpleReactOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,16 +1,44 @@
|
|||
"""Simple ReAct operation module.
|
||||
|
||||
This module provides a simple ReAct (Reasoning and Acting) agent implementation
|
||||
that extends the base ReactSearchOp for answering user queries through iterative
|
||||
reasoning and search actions.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from flowllm import C
|
||||
from flowllm.context.flow_context import FlowContext
|
||||
from flowllm.op.gallery.react_llm_op import ReactLLMOp
|
||||
from flowllm.core.context import C, FlowContext
|
||||
from flowllm.gallery import ReactSearchOp
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class SimpleReactOp(ReactLLMOp):
|
||||
...
|
||||
class SimpleReactOp(ReactSearchOp):
|
||||
"""A simple ReAct (Reasoning and Acting) agent operation.
|
||||
|
||||
This operation extends ReactSearchOp to provide a straightforward implementation
|
||||
of a ReAct agent that answers user queries by reasoning about the problem and
|
||||
taking search actions iteratively until a final answer is reached.
|
||||
|
||||
The agent inherits all functionality from ReactSearchOp, including:
|
||||
- Iterative reasoning and action cycles
|
||||
- Search tool integration
|
||||
- Maximum step limits for preventing infinite loops
|
||||
"""
|
||||
|
||||
|
||||
async def main():
|
||||
"""Main function to demonstrate SimpleReactOp usage.
|
||||
|
||||
This function initializes the FlowLLM context with ReMe configuration,
|
||||
creates a SimpleReactOp instance, and processes a sample query about
|
||||
stock prices for Maotai and Wuliangye.
|
||||
|
||||
Example:
|
||||
Run this module directly to test the SimpleReactOp:
|
||||
```bash
|
||||
python -m reme_ai.agent.react.simple_react_op
|
||||
```
|
||||
"""
|
||||
from reme_ai.config.config_parser import ConfigParser
|
||||
|
||||
C.set_service_config(parser=ConfigParser, config_name="config=default").init_by_service_config()
|
||||
|
|
@ -20,5 +48,6 @@ async def main():
|
|||
await op.async_call(context=context)
|
||||
print(context.response.answer)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
|
|||
|
|
@ -1,3 +1,17 @@
|
|||
"""Mock search tools for testing and demonstration purposes.
|
||||
|
||||
This module provides mock search operations that simulate different search tool behaviors,
|
||||
including LLM-based query classification and result generation.
|
||||
"""
|
||||
|
||||
from .llm_mock_search_op import LLMMockSearchOp
|
||||
from .mock_search_tools import SearchToolA, SearchToolB, SearchToolC
|
||||
from .use_mock_search_op import UseMockSearchOp
|
||||
|
||||
__all__ = [
|
||||
"LLMMockSearchOp",
|
||||
"SearchToolA",
|
||||
"SearchToolB",
|
||||
"SearchToolC",
|
||||
"UseMockSearchOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,13 +1,19 @@
|
|||
"""LLM-based mock search operation for simulating search tool behavior.
|
||||
|
||||
This module provides a mock search operation that uses LLM to classify queries
|
||||
and generate realistic search results based on query complexity levels.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from typing import Dict, Any
|
||||
|
||||
from flowllm.context import FlowContext, C
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.op.base_async_tool_op import BaseAsyncToolOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.schema.tool_call import ToolCall
|
||||
from flowllm.core.context import C, FlowContext
|
||||
from flowllm.core.enumeration.role import Role
|
||||
from flowllm.core.op import BaseAsyncToolOp
|
||||
from flowllm.core.schema import Message
|
||||
from flowllm.core.schema import ToolCall
|
||||
from loguru import logger
|
||||
|
||||
|
||||
|
|
@ -26,15 +32,18 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
- extra_time: Extra sleep time in seconds to simulate latency
|
||||
- relevance_ratio: Probability of returning relevant results (vs random query results)
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
def __init__(self,
|
||||
llm: str = "qwen3_30b_instruct",
|
||||
simple_config: Dict[str, Any] = None,
|
||||
medium_config: Dict[str, Any] = None,
|
||||
complex_config: Dict[str, Any] = None,
|
||||
seed: int = 0,
|
||||
**kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
llm: str = "qwen3_30b_instruct",
|
||||
simple_config: Dict[str, Any] = None,
|
||||
medium_config: Dict[str, Any] = None,
|
||||
complex_config: Dict[str, Any] = None,
|
||||
seed: int = 0,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize the LLM Mock Search Op.
|
||||
|
||||
|
|
@ -55,7 +64,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
seed: Random seed for deterministic behavior, default 0
|
||||
"""
|
||||
super().__init__(llm=llm, **kwargs)
|
||||
|
||||
|
||||
# Set random seed for deterministic behavior
|
||||
self.seed = seed
|
||||
random.seed(self.seed)
|
||||
|
|
@ -65,7 +74,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"success_rate": 0.95,
|
||||
"extra_time": 0.5,
|
||||
"relevance_ratio": 0.98,
|
||||
"content_length": "short"
|
||||
"content_length": "short",
|
||||
}
|
||||
if simple_config:
|
||||
self.simple_config.update(simple_config)
|
||||
|
|
@ -74,7 +83,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"success_rate": 0.85,
|
||||
"extra_time": 1.0,
|
||||
"relevance_ratio": 0.90,
|
||||
"content_length": "medium"
|
||||
"content_length": "medium",
|
||||
}
|
||||
if medium_config:
|
||||
self.medium_config.update(medium_config)
|
||||
|
|
@ -83,22 +92,29 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"success_rate": 0.70,
|
||||
"extra_time": 1.5,
|
||||
"relevance_ratio": 0.80,
|
||||
"content_length": "long"
|
||||
"content_length": "long",
|
||||
}
|
||||
if complex_config:
|
||||
self.complex_config.update(complex_config)
|
||||
|
||||
def build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(**{
|
||||
"description": "Use search keywords to retrieve relevant information from the internet.",
|
||||
"input_schema": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "search keyword or query",
|
||||
"required": True
|
||||
}
|
||||
}
|
||||
})
|
||||
"""Build the tool call schema for the search operation.
|
||||
|
||||
Returns:
|
||||
ToolCall object defining the search tool interface
|
||||
"""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Use search keywords to retrieve relevant information from the internet.",
|
||||
"input_schema": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "search keyword or query",
|
||||
"required": True,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def classify_query(self, query: str) -> str:
|
||||
"""
|
||||
|
|
@ -112,7 +128,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"""
|
||||
classification_prompt = self.prompt_format(
|
||||
prompt_name="classification_prompt",
|
||||
query=query
|
||||
query=query,
|
||||
)
|
||||
|
||||
messages = [Message(role=Role.USER, content=classification_prompt)]
|
||||
|
|
@ -146,7 +162,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
prompt_name="generation_prompt",
|
||||
query=query,
|
||||
complexity=complexity,
|
||||
content_length=content_length
|
||||
content_length=content_length,
|
||||
)
|
||||
|
||||
messages = [Message(role=Role.USER, content=generation_prompt)]
|
||||
|
|
@ -171,7 +187,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"breakthroughs in medical science",
|
||||
"architectural wonders",
|
||||
"wildlife conservation efforts",
|
||||
"developments in artificial intelligence"
|
||||
"developments in artificial intelligence",
|
||||
]
|
||||
|
||||
random_query = random.choice(random_topics)
|
||||
|
|
@ -179,7 +195,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
prompt_name="generation_prompt",
|
||||
query=random_query,
|
||||
complexity="simple",
|
||||
content_length="short"
|
||||
content_length="short",
|
||||
)
|
||||
|
||||
messages = [Message(role=Role.USER, content=generation_prompt)]
|
||||
|
|
@ -188,6 +204,11 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
return f"[Low Relevance Result]\n{response.content}"
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the mock search operation.
|
||||
|
||||
This method classifies the query, applies the appropriate configuration,
|
||||
simulates delays, and generates search results based on success and relevance rates.
|
||||
"""
|
||||
query: str = self.input_dict["query"]
|
||||
logger.info(f"LLMMockSearchOp processing query: {query}")
|
||||
|
||||
|
|
@ -218,9 +239,9 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"success": False,
|
||||
"content": error_message,
|
||||
"query": query,
|
||||
"complexity": complexity
|
||||
"complexity": complexity,
|
||||
}
|
||||
self.set_result(json.dumps(result_dict, ensure_ascii=False))
|
||||
self.set_output(json.dumps(result_dict, ensure_ascii=False))
|
||||
return
|
||||
|
||||
# Step 5: Check relevance ratio
|
||||
|
|
@ -235,7 +256,7 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"content": content,
|
||||
"query": query,
|
||||
"complexity": complexity,
|
||||
"is_relevant": False # Mark as irrelevant for debugging
|
||||
"is_relevant": False, # Mark as irrelevant for debugging
|
||||
}
|
||||
else:
|
||||
# Generate relevant result
|
||||
|
|
@ -246,14 +267,15 @@ class LLMMockSearchOp(BaseAsyncToolOp):
|
|||
"content": content,
|
||||
"query": query,
|
||||
"complexity": complexity,
|
||||
"is_relevant": True # Mark as relevant for debugging
|
||||
"is_relevant": True, # Mark as relevant for debugging
|
||||
}
|
||||
|
||||
self.set_result(json.dumps(result_dict, ensure_ascii=False))
|
||||
self.set_output(json.dumps(result_dict, ensure_ascii=False))
|
||||
|
||||
|
||||
async def async_main():
|
||||
from reme_ai.app import ReMeApp
|
||||
"""Main function for testing the LLMMockSearchOp with various query types."""
|
||||
from reme_ai.main import ReMeApp
|
||||
|
||||
async with ReMeApp():
|
||||
# Test with different query types
|
||||
|
|
@ -267,25 +289,25 @@ async def async_main():
|
|||
custom_simple = {
|
||||
"success_rate": 1,
|
||||
"extra_time": 0,
|
||||
"relevance_ratio": 1
|
||||
"relevance_ratio": 1,
|
||||
}
|
||||
|
||||
custom_medium = {
|
||||
"success_rate": 1,
|
||||
"extra_time": 0,
|
||||
"relevance_ratio": 1
|
||||
"relevance_ratio": 1,
|
||||
}
|
||||
|
||||
custom_complex = {
|
||||
"success_rate": 1,
|
||||
"extra_time": 0,
|
||||
"relevance_ratio": 1
|
||||
"relevance_ratio": 1,
|
||||
}
|
||||
|
||||
op = LLMMockSearchOp(
|
||||
simple_config=custom_simple,
|
||||
medium_config=custom_medium,
|
||||
complex_config=custom_complex
|
||||
complex_config=custom_complex,
|
||||
)
|
||||
|
||||
for query in test_queries:
|
||||
|
|
|
|||
|
|
@ -1,76 +1,76 @@
|
|||
classification_prompt: |
|
||||
You are a query complexity classifier. Analyze the following search query and classify it into one of three categories:
|
||||
|
||||
|
||||
1. **simple** - Simple factual queries that:
|
||||
- Ask for a single, direct fact
|
||||
- Have a clear, unambiguous answer
|
||||
- Require minimal context or explanation
|
||||
- Examples: "What is the capital of France?", "Who invented the telephone?", "When did World War 2 end?"
|
||||
|
||||
|
||||
2. **medium** - Medium complexity queries that:
|
||||
- Require some explanation or context
|
||||
- May involve multiple related facts
|
||||
- Need balanced depth without being exhaustive
|
||||
- Examples: "How does photosynthesis work?", "What are the main causes of climate change?", "Explain blockchain technology"
|
||||
|
||||
|
||||
3. **complex** - Complex research queries that:
|
||||
- Require comprehensive, multi-dimensional analysis
|
||||
- Involve multiple subtopics or perspectives
|
||||
- Need in-depth exploration and connections
|
||||
- Examples: "Analyze the impact of AI on the global economy", "Compare different renewable energy solutions", "What are the geopolitical implications of space exploration?"
|
||||
|
||||
|
||||
Query to classify: {query}
|
||||
|
||||
|
||||
Respond with ONLY one word: simple, medium, or complex.
|
||||
|
||||
classification_prompt_zh: |
|
||||
你是一个查询复杂度分类器。分析以下搜索查询并将其分类为以下三类之一:
|
||||
|
||||
|
||||
1. **simple** - 简单事实查询:
|
||||
- 询问单一、直接的事实
|
||||
- 有明确、无歧义的答案
|
||||
- 需要最少的上下文或解释
|
||||
- 示例:"法国的首都是什么?"、"谁发明了电话?"、"第二次世界大战何时结束?"
|
||||
|
||||
|
||||
2. **medium** - 中等复杂度查询:
|
||||
- 需要一些解释或上下文
|
||||
- 可能涉及多个相关事实
|
||||
- 需要平衡的深度但不需要详尽无遗
|
||||
- 示例:"光合作用如何工作?"、"气候变化的主要原因是什么?"、"解释区块链技术"
|
||||
|
||||
|
||||
3. **complex** - 复杂研究查询:
|
||||
- 需要全面、多维度的分析
|
||||
- 涉及多个子主题或观点
|
||||
- 需要深入探索和联系
|
||||
- 示例:"分析人工智能对全球经济的影响"、"比较不同的可再生能源解决方案"、"太空探索的地缘政治影响是什么?"
|
||||
|
||||
|
||||
要分类的查询:{query}
|
||||
|
||||
|
||||
只用一个词回答:simple、medium 或 complex。
|
||||
|
||||
generation_prompt: |
|
||||
You are a search engine generating mock search results. Generate a {content_length} response for the following query.
|
||||
|
||||
|
||||
Query: {query}
|
||||
Complexity Level: {complexity}
|
||||
|
||||
|
||||
Instructions based on content length:
|
||||
- **short**: Provide a concise answer in 1-3 sentences. Be direct and factual.
|
||||
- **medium**: Provide a balanced answer in 2-4 paragraphs. Include key details and some context.
|
||||
- **long**: Provide a comprehensive answer in 4-6 paragraphs. Include multiple perspectives, detailed explanations, and relevant context.
|
||||
|
||||
|
||||
Generate the search result content now:
|
||||
|
||||
generation_prompt_zh: |
|
||||
你是一个搜索引擎,正在生成模拟搜索结果。为以下查询生成一个 {content_length} 的响应。
|
||||
|
||||
|
||||
查询:{query}
|
||||
复杂度级别:{complexity}
|
||||
|
||||
|
||||
根据内容长度的指示:
|
||||
- **short**(短):提供 1-3 句话的简洁答案。要直接和事实性。
|
||||
- **medium**(中):提供 2-4 段的平衡答案。包括关键细节和一些上下文。
|
||||
- **long**(长):提供 4-6 段的全面答案。包括多个角度、详细解释和相关上下文。
|
||||
|
||||
|
||||
现在生成搜索结果内容:
|
||||
|
||||
|
|
|
|||
|
|
@ -1,77 +1,125 @@
|
|||
from flowllm.context import C
|
||||
from flowllm.schema.tool_call import ToolCall
|
||||
"""Specialized mock search tools with different performance characteristics.
|
||||
|
||||
This module provides three search tools (SearchToolA, SearchToolB, SearchToolC)
|
||||
each optimized for different query complexity levels, allowing for realistic
|
||||
testing of tool selection strategies.
|
||||
"""
|
||||
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.schema import ToolCall
|
||||
|
||||
from reme_ai.agent.tools.llm_mock_search_op import LLMMockSearchOp
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class SearchToolA(LLMMockSearchOp):
|
||||
"""Fast search tool optimized for simple queries.
|
||||
|
||||
This tool is configured for quick responses with high success rates
|
||||
on simple queries, but performs poorly on medium and complex queries.
|
||||
Best suited for simple factual queries.
|
||||
"""
|
||||
|
||||
def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs):
|
||||
"""Initialize SearchToolA with fast, simple-query-optimized configuration.
|
||||
|
||||
Args:
|
||||
llm: LLM model name to use
|
||||
**kwargs: Additional arguments passed to LLMMockSearchOp
|
||||
"""
|
||||
# Configure for fast but shallow performance
|
||||
simple_config = {
|
||||
"success_rate": 0.9, # High success rate for simple queries
|
||||
"extra_time": 0, # Very fast (0.2-0.5s range)
|
||||
"relevance_ratio": 0.9, # High relevance
|
||||
"content_length": "short" # Concise answers
|
||||
"content_length": "short", # Concise answers
|
||||
}
|
||||
|
||||
medium_config = {
|
||||
"success_rate": 0.2, # Lower success for medium queries
|
||||
"extra_time": 0, # Still fast
|
||||
"relevance_ratio": 0.2, # Moderate relevance
|
||||
"content_length": "short" # Limited depth
|
||||
"content_length": "short", # Limited depth
|
||||
}
|
||||
|
||||
complex_config = {
|
||||
"success_rate": 0.5, # Poor success rate for complex queries
|
||||
"extra_time": 0, # Fast but insufficient
|
||||
"relevance_ratio": 0.5, # Low relevance (often misses key aspects)
|
||||
"content_length": "short" # Too shallow for complex topics
|
||||
"content_length": "short", # Too shallow for complex topics
|
||||
}
|
||||
|
||||
super().__init__(llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs)
|
||||
super().__init__(
|
||||
llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def build_tool_call(self) -> ToolCall:
|
||||
"""Build the tool call schema with description indicating simple query optimization.
|
||||
|
||||
Returns:
|
||||
ToolCall object with description indicating best use for simple queries
|
||||
"""
|
||||
tool_call = super().build_tool_call()
|
||||
tool_call.description += " Best suited for simple queries."
|
||||
return tool_call
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class SearchToolB(LLMMockSearchOp):
|
||||
"""Balanced search tool optimized for medium complexity queries.
|
||||
|
||||
This tool provides balanced performance across query types, with
|
||||
excellent results for medium complexity queries. Best suited for
|
||||
queries requiring moderate depth and context.
|
||||
"""
|
||||
|
||||
def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs):
|
||||
"""Initialize SearchToolB with balanced, medium-query-optimized configuration.
|
||||
|
||||
Args:
|
||||
llm: LLM model name to use
|
||||
**kwargs: Additional arguments passed to LLMMockSearchOp
|
||||
"""
|
||||
# Configure for balanced performance
|
||||
simple_config = {
|
||||
"success_rate": 0.3, # Very high success rate
|
||||
"extra_time": 0, # Moderate speed (1.0-1.5s range)
|
||||
"relevance_ratio": 0.3, # High relevance
|
||||
"content_length": "medium" # More detailed than needed for simple
|
||||
"content_length": "medium", # More detailed than needed for simple
|
||||
}
|
||||
|
||||
medium_config = {
|
||||
"success_rate": 0.9, # Excellent success rate
|
||||
"extra_time": 0, # Balanced speed
|
||||
"relevance_ratio": 0.9, # High relevance
|
||||
"content_length": "medium" # Perfect depth for medium queries
|
||||
"content_length": "medium", # Perfect depth for medium queries
|
||||
}
|
||||
|
||||
complex_config = {
|
||||
"success_rate": 0.5, # Good success rate
|
||||
"extra_time": 0, # Still reasonable speed
|
||||
"relevance_ratio": 0.5, # Decent relevance but not exhaustive
|
||||
"content_length": "medium" # Covers main points but lacks depth
|
||||
"content_length": "medium", # Covers main points but lacks depth
|
||||
}
|
||||
|
||||
super().__init__(llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs)
|
||||
super().__init__(
|
||||
llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def build_tool_call(self) -> ToolCall:
|
||||
"""Build the tool call schema with description indicating medium query optimization.
|
||||
|
||||
Returns:
|
||||
ToolCall object with description indicating best use for medium complexity queries
|
||||
"""
|
||||
tool_call = super().build_tool_call()
|
||||
tool_call.description += " Best suited for medium complexity queries."
|
||||
return tool_call
|
||||
|
|
@ -79,37 +127,56 @@ class SearchToolB(LLMMockSearchOp):
|
|||
|
||||
@C.register_op()
|
||||
class SearchToolC(LLMMockSearchOp):
|
||||
"""Comprehensive search tool optimized for complex queries.
|
||||
|
||||
This tool provides thorough, in-depth results with high success rates
|
||||
on complex queries, but may be slower and overly detailed for simple queries.
|
||||
Best suited for complex research queries requiring comprehensive analysis.
|
||||
"""
|
||||
|
||||
def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs):
|
||||
"""Initialize SearchToolC with comprehensive, complex-query-optimized configuration.
|
||||
|
||||
Args:
|
||||
llm: LLM model name to use
|
||||
**kwargs: Additional arguments passed to LLMMockSearchOp
|
||||
"""
|
||||
# Configure for comprehensive but costly performance
|
||||
simple_config = {
|
||||
"success_rate": 0.3, # Good but not optimal (over-processing)
|
||||
"extra_time": 0, # Slow (3.0-4.0s range)
|
||||
"relevance_ratio": 0.3, # High relevance but unnecessary depth
|
||||
"content_length": "long" # Too detailed for simple queries
|
||||
"content_length": "long", # Too detailed for simple queries
|
||||
}
|
||||
|
||||
medium_config = {
|
||||
"success_rate": 0.4, # High success rate
|
||||
"extra_time": 0, # Slow but thorough
|
||||
"relevance_ratio": 0.4, # High relevance with extra context
|
||||
"content_length": "long" # More depth than needed
|
||||
"content_length": "long", # More depth than needed
|
||||
}
|
||||
|
||||
complex_config = {
|
||||
"success_rate": 0.9, # Excellent success rate
|
||||
"extra_time": 0, # Slow but comprehensive (3.5-5.0s range)
|
||||
"relevance_ratio": 0.9, # Very high relevance
|
||||
"content_length": "long" # Perfect depth for complex queries
|
||||
"content_length": "long", # Perfect depth for complex queries
|
||||
}
|
||||
|
||||
super().__init__(llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs)
|
||||
super().__init__(
|
||||
llm=llm,
|
||||
simple_config=simple_config,
|
||||
medium_config=medium_config,
|
||||
complex_config=complex_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def build_tool_call(self) -> ToolCall:
|
||||
"""Build the tool call schema with description indicating complex query optimization.
|
||||
|
||||
Returns:
|
||||
ToolCall object with description indicating best use for complex queries
|
||||
"""
|
||||
tool_call = super().build_tool_call()
|
||||
tool_call.description += " Best suited for complex queries."
|
||||
return tool_call
|
||||
|
|
|
|||
|
|
@ -1,13 +1,19 @@
|
|||
"""Tool selection and execution operation for mock search tools.
|
||||
|
||||
This module provides an operation that intelligently selects and executes
|
||||
the most appropriate mock search tool based on query complexity.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
|
||||
from flowllm.context import C
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.op.base_async_tool_op import BaseAsyncToolOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.schema.tool_call import ToolCall
|
||||
from flowllm.utils.timer import Timer
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncToolOp
|
||||
from flowllm.core.schema import Message
|
||||
from flowllm.core.schema import ToolCall
|
||||
from flowllm.core.utils import Timer
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.agent.tools.mock_search_tools import SearchToolA, SearchToolB, SearchToolC
|
||||
|
|
@ -16,27 +22,61 @@ from reme_ai.schema.memory import ToolCallResult
|
|||
|
||||
@C.register_op()
|
||||
class UseMockSearchOp(BaseAsyncToolOp):
|
||||
"""Operation that selects and executes the most appropriate mock search tool.
|
||||
|
||||
This operation uses LLM to intelligently select from available search tools
|
||||
(SearchToolA, SearchToolB, SearchToolC) based on query characteristics,
|
||||
then executes the selected tool and records performance metrics.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
def __init__(self, llm: str = "qwen3_30b_instruct", **kwargs):
|
||||
"""Initialize the UseMockSearchOp.
|
||||
|
||||
Args:
|
||||
llm: LLM model name to use for tool selection
|
||||
**kwargs: Additional arguments passed to BaseAsyncToolOp
|
||||
"""
|
||||
super().__init__(llm=llm, save_answer=True, **kwargs)
|
||||
|
||||
def build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(**{
|
||||
"description": "Intelligently selects and executes the most appropriate search tool based on query complexity. "
|
||||
"Automatically tracks performance metrics and records tool usage for optimization.",
|
||||
"input_schema": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "query",
|
||||
"required": True
|
||||
}
|
||||
}
|
||||
})
|
||||
"""Build the tool call schema for the search tool selector.
|
||||
|
||||
Returns:
|
||||
ToolCall object defining the search tool selector interface
|
||||
"""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": (
|
||||
"Intelligently selects and executes the most appropriate search tool "
|
||||
"based on query complexity. Automatically tracks performance metrics "
|
||||
"and records tool usage for optimization."
|
||||
),
|
||||
"input_schema": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "query",
|
||||
"required": True,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def select_tool(self, query: str, tool_ops: list[BaseAsyncToolOp]) -> ToolCall | None:
|
||||
assistant_message = await self.llm.achat(messages=[Message(role=Role.USER, content=query)],
|
||||
tools=[x.tool_call for x in tool_ops])
|
||||
"""Select the most appropriate tool for the given query using LLM.
|
||||
|
||||
Args:
|
||||
query: The search query to process
|
||||
tool_ops: List of available tool operations to choose from
|
||||
|
||||
Returns:
|
||||
Selected ToolCall if a tool was chosen, None otherwise
|
||||
"""
|
||||
assistant_message = await self.llm.achat(
|
||||
messages=[Message(role=Role.USER, content=query)],
|
||||
tools=[x.tool_call for x in tool_ops],
|
||||
)
|
||||
logger.info(f"assistant_message={assistant_message.model_dump_json()}")
|
||||
if assistant_message.tool_calls:
|
||||
return assistant_message.tool_calls[0]
|
||||
|
|
@ -44,6 +84,11 @@ class UseMockSearchOp(BaseAsyncToolOp):
|
|||
return None
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the tool selection and execution workflow.
|
||||
|
||||
This method selects an appropriate tool, executes it, measures performance,
|
||||
and creates a ToolCallResult with metrics.
|
||||
"""
|
||||
query: str = self.input_dict["query"]
|
||||
logger.info(f"query={query}")
|
||||
|
||||
|
|
@ -65,9 +110,9 @@ class UseMockSearchOp(BaseAsyncToolOp):
|
|||
output="No appropriate tool was selected for the query",
|
||||
token_cost=0,
|
||||
success=False,
|
||||
time_cost=0.0
|
||||
time_cost=0.0,
|
||||
)
|
||||
self.set_result(error_result.model_dump_json())
|
||||
self.set_output(error_result.model_dump_json())
|
||||
return
|
||||
|
||||
# Step 2: Execute the selected tool
|
||||
|
|
@ -86,9 +131,9 @@ class UseMockSearchOp(BaseAsyncToolOp):
|
|||
output=f"Tool {tool_call.name} not found in available tools",
|
||||
token_cost=0,
|
||||
success=False,
|
||||
time_cost=0.0
|
||||
time_cost=0.0,
|
||||
)
|
||||
self.set_result(error_result.model_dump_json())
|
||||
self.set_output(error_result.model_dump_json())
|
||||
return
|
||||
|
||||
# Step 3: Execute the tool with timer
|
||||
|
|
@ -110,14 +155,15 @@ class UseMockSearchOp(BaseAsyncToolOp):
|
|||
output=content,
|
||||
token_cost=token_cost,
|
||||
success=success,
|
||||
time_cost=round(time_cost, 3)
|
||||
time_cost=round(time_cost, 3),
|
||||
)
|
||||
|
||||
self.set_result(tool_call_result.model_dump_json())
|
||||
self.set_output(tool_call_result.model_dump_json())
|
||||
|
||||
|
||||
async def async_main():
|
||||
from reme_ai.app import ReMeApp
|
||||
"""Main function for testing the UseMockSearchOp with various queries."""
|
||||
from reme_ai.main import ReMeApp
|
||||
|
||||
async with ReMeApp():
|
||||
test_queries = [
|
||||
|
|
@ -125,7 +171,7 @@ async def async_main():
|
|||
"How does quantum computing work?",
|
||||
"Analyze the impact of artificial intelligence on global economy, employment, and society",
|
||||
"When was Python programming language created?",
|
||||
"Compare different types of renewable energy sources",
|
||||
"Compare different types of renewable energy sources",
|
||||
]
|
||||
|
||||
for query in test_queries:
|
||||
|
|
@ -136,4 +182,3 @@ async def async_main():
|
|||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(async_main())
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,14 @@
|
|||
"""Configuration module for ReMe.
|
||||
|
||||
This module provides configuration parsing capabilities for the ReMe framework.
|
||||
It includes:
|
||||
|
||||
- ConfigParser: Configuration parser class that extends PydanticConfigParser
|
||||
to provide configuration parsing with awareness of the current module location
|
||||
"""
|
||||
|
||||
from .config_parser import ConfigParser
|
||||
|
||||
__all__ = [
|
||||
"ConfigParser",
|
||||
]
|
||||
|
|
@ -1,6 +1,24 @@
|
|||
from flowllm.config.pydantic_config_parser import PydanticConfigParser
|
||||
"""Configuration parser module for ReMe.
|
||||
|
||||
This module provides configuration parsing capabilities for the ReMe framework.
|
||||
It extends the PydanticConfigParser from FlowLLM to provide configuration parsing
|
||||
with awareness of the current module location.
|
||||
"""
|
||||
|
||||
from flowllm.core.utils import PydanticConfigParser
|
||||
|
||||
|
||||
class ConfigParser(PydanticConfigParser):
|
||||
"""Configuration parser for ReMe framework.
|
||||
|
||||
Extends PydanticConfigParser to provide configuration parsing capabilities
|
||||
with awareness of the current module location. Uses the default.yaml
|
||||
configuration file as the default configuration source.
|
||||
|
||||
Attributes:
|
||||
current_file: Path to the current file, used for relative config file resolution.
|
||||
default_config_name: Default configuration file name (without .yaml extension).
|
||||
"""
|
||||
|
||||
current_file: str = __file__
|
||||
default_config_name: str = "default"
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ http:
|
|||
|
||||
flow:
|
||||
retrieve_task_memory:
|
||||
flow_content: build_query_op >> recall_vector_store_op >> rerank_memory_op >> rewrite_memory_op
|
||||
flow_content: BuildQueryOp() >> RecallVectorStoreOp() >> RerankMemoryOp(enable_llm_rerank=True, enable_score_filter=False, top_k=5) >> RewriteMemoryOp(enable_llm_rewrite=True)
|
||||
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -23,7 +23,7 @@ flow:
|
|||
required: true
|
||||
|
||||
summary_task_memory:
|
||||
flow_content: trajectory_preprocess_op >> (success_extraction_op|failure_extraction_op|comparative_extraction_op) >> memory_validation_op >> update_vector_store_op
|
||||
flow_content: TrajectoryPreprocessOp(success_threshold=1.0) >> (SuccessExtractionOp()|FailureExtractionOp()|ComparativeExtractionOp(enable_soft_comparison=True)) >> MemoryValidationOp(validation_threshold=0.5) >> UpdateVectorStoreOp()
|
||||
description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -32,7 +32,7 @@ flow:
|
|||
required: false
|
||||
|
||||
retrieve_task_memory_simple:
|
||||
flow_content: build_query_op >> recall_vector_store_op >> merge_memory_op
|
||||
flow_content: BuildQueryOp() >> RecallVectorStoreOp() >> MergeMemoryOp()
|
||||
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -41,7 +41,7 @@ flow:
|
|||
required: true
|
||||
|
||||
summary_task_memory_simple:
|
||||
flow_content: simple_summary_op >> update_vector_store_op
|
||||
flow_content: SimpleSummaryOp() >> UpdateVectorStoreOp()
|
||||
description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -50,7 +50,7 @@ flow:
|
|||
required: false
|
||||
|
||||
retrieve_personal_memory:
|
||||
flow_content: set_query_op >> (extract_time_op | (retrieve_memory_op >> semantic_rank_op)) >> fuse_rerank_op
|
||||
flow_content: SetQueryOp() >> (ExtractTimeOp() | (RetrieveMemoryOp() >> SemanticRankOp())) >> FuseRerankOp()
|
||||
description: "Retrieves the most relevant personal memories from historical data based on the query to enhance response quality"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -59,7 +59,7 @@ flow:
|
|||
required: true
|
||||
|
||||
summary_personal_memory:
|
||||
flow_content: info_filter_op >> (get_observation_op | get_observation_with_time_op | load_today_memory_op) >> contra_repeat_op >> update_vector_store_op
|
||||
flow_content: InfoFilterOp() >> (GetObservationOp() | GetObservationWithTimeOp() | LoadTodayMemoryOp()) >> ContraRepeatOp() >> UpdateVectorStoreOp()
|
||||
description: "Consolidates user observations and memories by filtering information and removing redundancies for efficient storage"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -68,7 +68,7 @@ flow:
|
|||
required: false
|
||||
|
||||
retrieve_tool_memory:
|
||||
flow_content: retrieve_tool_memory_op
|
||||
flow_content: RetrieveToolMemoryOp()
|
||||
description: "Retrieves tool memories from the vector database based on tool names to provide tool usage patterns and best practices"
|
||||
input_schema:
|
||||
tool_names:
|
||||
|
|
@ -77,7 +77,7 @@ flow:
|
|||
required: true
|
||||
|
||||
add_tool_call_result:
|
||||
flow_content: parse_tool_call_result_op >> update_vector_store_op
|
||||
flow_content: ParseToolCallResultOp(max_history_tool_call_cnt=100, evaluation_sleep_interval=1.0) >> UpdateVectorStoreOp()
|
||||
description: "Evaluates and adds tool call results to the tool memory database, creating new memory or updating existing memory for the specified tool"
|
||||
input_schema:
|
||||
tool_call_results:
|
||||
|
|
@ -86,7 +86,7 @@ flow:
|
|||
required: true
|
||||
|
||||
summary_tool_memory:
|
||||
flow_content: summary_tool_memory_op >> update_vector_store_op
|
||||
flow_content: SummaryToolMemoryOp(recent_call_count=20, summary_sleep_interval=1.0) >> UpdateVectorStoreOp()
|
||||
description: "Analyzes tool call history and generates comprehensive usage patterns, best practices, and recommendations for the specified tools"
|
||||
input_schema:
|
||||
tool_names:
|
||||
|
|
@ -95,7 +95,7 @@ flow:
|
|||
required: true
|
||||
|
||||
use_mock_search:
|
||||
flow_content: use_mock_search_op
|
||||
flow_content: UseMockSearchOp()
|
||||
description: "Simulates intelligent search tool selection and execution based on query complexity, with automatic tool memory recording"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -104,7 +104,7 @@ flow:
|
|||
required: true
|
||||
|
||||
vector_store:
|
||||
flow_content: vector_store_action_op
|
||||
flow_content: VectorStoreActionOp()
|
||||
description: "Directly operates on the vector store with various management actions"
|
||||
input_schema:
|
||||
action:
|
||||
|
|
@ -114,7 +114,7 @@ flow:
|
|||
enum: [ copy, delete, delete_ids, dump, load, list]
|
||||
|
||||
record_task_memory:
|
||||
flow_content: update_memory_freq_op >> update_memory_utility_op >> update_vector_store_op
|
||||
flow_content: UpdateMemoryFreqOp() >> UpdateMemoryUtilityOp() >> UpdateVectorStoreOp()
|
||||
description: "Update the freq & utility attributes of retrieved task memories"
|
||||
input_schema:
|
||||
workspace_id:
|
||||
|
|
@ -131,7 +131,7 @@ flow:
|
|||
required: true
|
||||
|
||||
delete_task_memory:
|
||||
flow_content: delete_memory_op >> update_vector_store_op
|
||||
flow_content: DeleteMemoryOp() >> UpdateVectorStoreOp()
|
||||
description: "Delete task memories when utility/freq < utility_threshold and freq >= freq_threshold"
|
||||
input_schema:
|
||||
workspace_id:
|
||||
|
|
@ -148,7 +148,7 @@ flow:
|
|||
required: true
|
||||
|
||||
react:
|
||||
flow_content: simple_react_op
|
||||
flow_content: SimpleReactOp()
|
||||
description: "React to the current task with an agent"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -156,69 +156,6 @@ flow:
|
|||
description: "user query"
|
||||
required: true
|
||||
|
||||
op:
|
||||
# retriever ops
|
||||
rerank_memory_op:
|
||||
backend: rerank_memory_op
|
||||
llm: default
|
||||
params:
|
||||
enable_llm_rerank: true
|
||||
enable_score_filter: false
|
||||
top_k: 5
|
||||
|
||||
rewrite_memory_op:
|
||||
backend: rewrite_memory_op
|
||||
llm: default
|
||||
params:
|
||||
enable_llm_rewrite: true
|
||||
|
||||
#summarizer ops
|
||||
trajectory_preprocess_op:
|
||||
backend: trajectory_preprocess_op
|
||||
params:
|
||||
success_threshold: 1.0
|
||||
|
||||
success_extraction_op:
|
||||
backend: success_extraction_op
|
||||
llm: default
|
||||
|
||||
failure_extraction_op:
|
||||
backend: failure_extraction_op
|
||||
llm: default
|
||||
|
||||
comparative_extraction_op:
|
||||
backend: comparative_extraction_op
|
||||
llm: default
|
||||
params:
|
||||
enable_soft_comparison: true
|
||||
|
||||
memory_validation_op:
|
||||
backend: memory_validation_op
|
||||
llm: default
|
||||
params:
|
||||
validation_threshold: 0.5
|
||||
|
||||
memory_deduplication_op:
|
||||
backend: memory_deduplication_op
|
||||
vector_store: default
|
||||
params:
|
||||
similarity_threshold: 0.5
|
||||
|
||||
# tool memory ops
|
||||
parse_tool_call_result_op:
|
||||
backend: parse_tool_call_result_op
|
||||
llm: default
|
||||
params:
|
||||
max_history_tool_call_cnt: 100
|
||||
evaluation_sleep_interval: 1.0
|
||||
|
||||
summary_tool_memory_op:
|
||||
backend: summary_tool_memory_op
|
||||
llm: default
|
||||
params:
|
||||
recent_call_count: 20
|
||||
summary_sleep_interval: 1.0
|
||||
|
||||
llm:
|
||||
default:
|
||||
backend: openai_compatible
|
||||
|
|
|
|||
|
|
@ -1,7 +1,92 @@
|
|||
"""Constants module for ReMe AI.
|
||||
|
||||
This module provides access to all constants used throughout the application,
|
||||
including common workflow keys and language-specific constants.
|
||||
"""
|
||||
|
||||
from . import common_constants
|
||||
from . import language_constants
|
||||
|
||||
# Export all constants from common_constants
|
||||
from .common_constants import (
|
||||
WORKFLOW_NAME,
|
||||
RESULT,
|
||||
MEMORIES,
|
||||
CHAT_MESSAGES,
|
||||
CHAT_MESSAGES_SCATTER,
|
||||
CHAT_KWARGS,
|
||||
USER_NAME,
|
||||
TARGET_NAME,
|
||||
MEMORY_MANAGER,
|
||||
QUERY_WITH_TS,
|
||||
RETRIEVE_MEMORY_NODES,
|
||||
RANKED_MEMORY_NODES,
|
||||
NOT_REFLECTED_NODES,
|
||||
NOT_UPDATED_NODES,
|
||||
EXTRACT_TIME_DICT,
|
||||
NEW_OBS_NODES,
|
||||
NEW_OBS_WITH_TIME_NODES,
|
||||
INSIGHT_NODES,
|
||||
TODAY_NODES,
|
||||
MERGE_OBS_NODES,
|
||||
TIME_INFER,
|
||||
)
|
||||
|
||||
# Export all constants from language_constants
|
||||
from .language_constants import (
|
||||
DATATIME_WORD_LIST,
|
||||
WEEKDAYS,
|
||||
MONTH_DICT,
|
||||
NONE_WORD,
|
||||
REPEATED_WORD,
|
||||
CONTRADICTORY_WORD,
|
||||
CONTAINED_WORD,
|
||||
COLON_WORD,
|
||||
COMMA_WORD,
|
||||
DEFAULT_HUMAN_NAME,
|
||||
DATATIME_KEY_MAP,
|
||||
TIME_INFER_WORD,
|
||||
USER_NAME_EXPRESSION,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Module exports
|
||||
"common_constants",
|
||||
"language_constants"
|
||||
"language_constants",
|
||||
# Common constants
|
||||
"WORKFLOW_NAME",
|
||||
"RESULT",
|
||||
"MEMORIES",
|
||||
"CHAT_MESSAGES",
|
||||
"CHAT_MESSAGES_SCATTER",
|
||||
"CHAT_KWARGS",
|
||||
"USER_NAME",
|
||||
"TARGET_NAME",
|
||||
"MEMORY_MANAGER",
|
||||
"QUERY_WITH_TS",
|
||||
"RETRIEVE_MEMORY_NODES",
|
||||
"RANKED_MEMORY_NODES",
|
||||
"NOT_REFLECTED_NODES",
|
||||
"NOT_UPDATED_NODES",
|
||||
"EXTRACT_TIME_DICT",
|
||||
"NEW_OBS_NODES",
|
||||
"NEW_OBS_WITH_TIME_NODES",
|
||||
"INSIGHT_NODES",
|
||||
"TODAY_NODES",
|
||||
"MERGE_OBS_NODES",
|
||||
"TIME_INFER",
|
||||
# Language constants
|
||||
"DATATIME_WORD_LIST",
|
||||
"WEEKDAYS",
|
||||
"MONTH_DICT",
|
||||
"NONE_WORD",
|
||||
"REPEATED_WORD",
|
||||
"CONTRADICTORY_WORD",
|
||||
"CONTAINED_WORD",
|
||||
"COLON_WORD",
|
||||
"COMMA_WORD",
|
||||
"DEFAULT_HUMAN_NAME",
|
||||
"DATATIME_KEY_MAP",
|
||||
"TIME_INFER_WORD",
|
||||
"USER_NAME_EXPRESSION",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,7 +1,10 @@
|
|||
# common_constants.py
|
||||
# This module defines constants used as keys throughout the application to maintain a consistent reference
|
||||
# for data structures related to workflow management, chat interactions, context storage, memory operations,
|
||||
# node processing, and temporal inference functionalities.
|
||||
"""Common constants module.
|
||||
|
||||
This module defines constants used as keys throughout the application to maintain
|
||||
a consistent reference for data structures related to workflow management, chat
|
||||
interactions, context storage, memory operations, node processing, and temporal
|
||||
inference functionalities.
|
||||
"""
|
||||
|
||||
WORKFLOW_NAME = "workflow_name"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,11 @@
|
|||
"""Language constants module.
|
||||
|
||||
This module provides language-specific constants and mappings for datetime
|
||||
expressions, weekdays, months, and other linguistic elements used throughout
|
||||
the application. It supports multiple languages (currently Chinese and English)
|
||||
and facilitates internationalization of temporal and linguistic processing.
|
||||
"""
|
||||
|
||||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
# This dictionary maps languages to lists of words related to datetime expressions.
|
||||
|
|
@ -29,64 +37,101 @@ DATATIME_WORD_LIST = {
|
|||
],
|
||||
LanguageEnum.EN: [
|
||||
# Units of Time
|
||||
"year", "yr",
|
||||
"month", "mo",
|
||||
"week", "wk",
|
||||
"day", "d",
|
||||
"hour", "hr",
|
||||
"minute", "min",
|
||||
"second", "sec",
|
||||
|
||||
"year",
|
||||
"yr",
|
||||
"month",
|
||||
"mo",
|
||||
"week",
|
||||
"wk",
|
||||
"day",
|
||||
"d",
|
||||
"hour",
|
||||
"hr",
|
||||
"minute",
|
||||
"min",
|
||||
"second",
|
||||
"sec",
|
||||
# Days of the Week
|
||||
"Monday", "Mon",
|
||||
"Tuesday", "Tue", "Tues",
|
||||
"Wednesday", "Wed",
|
||||
"Thursday", "Thu", "Thur", "Thurs",
|
||||
"Friday", "Fri",
|
||||
"Saturday", "Sat",
|
||||
"Sunday", "Sun",
|
||||
|
||||
"Monday",
|
||||
"Mon",
|
||||
"Tuesday",
|
||||
"Tue",
|
||||
"Tues",
|
||||
"Wednesday",
|
||||
"Wed",
|
||||
"Thursday",
|
||||
"Thu",
|
||||
"Thur",
|
||||
"Thurs",
|
||||
"Friday",
|
||||
"Fri",
|
||||
"Saturday",
|
||||
"Sat",
|
||||
"Sunday",
|
||||
"Sun",
|
||||
# Months of the Year
|
||||
"January", "Jan",
|
||||
"February", "Feb",
|
||||
"March", "Mar",
|
||||
"April", "Apr",
|
||||
"May", "May",
|
||||
"June", "Jun",
|
||||
"July", "Jul",
|
||||
"August", "Aug",
|
||||
"September", "Sep", "Sept",
|
||||
"October", "Oct",
|
||||
"November", "Nov",
|
||||
"December", "Dec",
|
||||
|
||||
"January",
|
||||
"Jan",
|
||||
"February",
|
||||
"Feb",
|
||||
"March",
|
||||
"Mar",
|
||||
"April",
|
||||
"Apr",
|
||||
"May",
|
||||
"May",
|
||||
"June",
|
||||
"Jun",
|
||||
"July",
|
||||
"Jul",
|
||||
"August",
|
||||
"Aug",
|
||||
"September",
|
||||
"Sep",
|
||||
"Sept",
|
||||
"October",
|
||||
"Oct",
|
||||
"November",
|
||||
"Nov",
|
||||
"December",
|
||||
"Dec",
|
||||
# Relative Time References
|
||||
"Today",
|
||||
"Tomorrow", "Tmrw",
|
||||
"Yesterday", "Yday",
|
||||
"Tomorrow",
|
||||
"Tmrw",
|
||||
"Yesterday",
|
||||
"Yday",
|
||||
"Now",
|
||||
"Morning", "AM", "a.m.",
|
||||
"Afternoon", "PM", "p.m.",
|
||||
"Morning",
|
||||
"AM",
|
||||
"a.m.",
|
||||
"Afternoon",
|
||||
"PM",
|
||||
"p.m.",
|
||||
"Evening",
|
||||
"Night",
|
||||
"Midnight",
|
||||
"Noon",
|
||||
|
||||
# Seasonal References
|
||||
"Spring",
|
||||
"Summer",
|
||||
"Autumn", "Fall",
|
||||
"Autumn",
|
||||
"Fall",
|
||||
"Winter",
|
||||
|
||||
# General Time References
|
||||
"Century", "cent.",
|
||||
"Century",
|
||||
"cent.",
|
||||
"Decade",
|
||||
"Millennium",
|
||||
"Quarter", "Q1", "Q2", "Q3", "Q4",
|
||||
"Quarter",
|
||||
"Q1",
|
||||
"Q2",
|
||||
"Q3",
|
||||
"Q4",
|
||||
"Semester",
|
||||
"Fortnight",
|
||||
"Weekend"
|
||||
]
|
||||
"Weekend",
|
||||
],
|
||||
}
|
||||
|
||||
# A mapping of weekdays for each supported language, facilitating calendar-related operations and understanding
|
||||
|
|
@ -99,7 +144,7 @@ WEEKDAYS = {
|
|||
"周四",
|
||||
"周五",
|
||||
"周六",
|
||||
"周日"
|
||||
"周日",
|
||||
],
|
||||
LanguageEnum.EN: [
|
||||
"Monday",
|
||||
|
|
@ -109,7 +154,7 @@ WEEKDAYS = {
|
|||
"Friday",
|
||||
"Saturday",
|
||||
"Sunday",
|
||||
]
|
||||
],
|
||||
}
|
||||
|
||||
MONTH_DICT = {
|
||||
|
|
@ -140,49 +185,49 @@ MONTH_DICT = {
|
|||
"October",
|
||||
"November",
|
||||
"December",
|
||||
]
|
||||
],
|
||||
}
|
||||
|
||||
# Constants for the word 'none' in different languages
|
||||
NONE_WORD = {
|
||||
LanguageEnum.CN: "无",
|
||||
LanguageEnum.EN: "none"
|
||||
LanguageEnum.EN: "none",
|
||||
}
|
||||
|
||||
# Constants for the word 'repeated' in different languages
|
||||
REPEATED_WORD = {
|
||||
LanguageEnum.CN: "重复",
|
||||
LanguageEnum.EN: "repeated"
|
||||
LanguageEnum.EN: "repeated",
|
||||
}
|
||||
|
||||
# Constants for the word 'contradictory' in different languages
|
||||
CONTRADICTORY_WORD = {
|
||||
LanguageEnum.CN: "矛盾",
|
||||
LanguageEnum.EN: "contradiction"
|
||||
LanguageEnum.EN: "contradiction",
|
||||
}
|
||||
|
||||
# Constants for the phrase 'included' in different languages
|
||||
CONTAINED_WORD = {
|
||||
LanguageEnum.CN: "被包含",
|
||||
LanguageEnum.EN: "contained"
|
||||
LanguageEnum.EN: "contained",
|
||||
}
|
||||
|
||||
# Constants for the symbol ':' in different languages' representations
|
||||
COLON_WORD = {
|
||||
LanguageEnum.CN: ":",
|
||||
LanguageEnum.EN: ":"
|
||||
LanguageEnum.EN: ":",
|
||||
}
|
||||
|
||||
# Constants for the symbol ',' in different languages' representations
|
||||
COMMA_WORD = {
|
||||
LanguageEnum.CN: ",",
|
||||
LanguageEnum.EN: ","
|
||||
LanguageEnum.EN: ",",
|
||||
}
|
||||
|
||||
# Default human name placeholders for different languages
|
||||
DEFAULT_HUMAN_NAME = {
|
||||
LanguageEnum.CN: "用户",
|
||||
LanguageEnum.EN: "user"
|
||||
LanguageEnum.EN: "user",
|
||||
}
|
||||
|
||||
# Mapping of datetime terms from natural language to standardized keys for each supported language
|
||||
|
|
@ -200,16 +245,16 @@ DATATIME_KEY_MAP = {
|
|||
"Day": "day",
|
||||
"Week": "week",
|
||||
"Weekday": "weekday",
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
# Phrase for indicating inferred time in different languages
|
||||
TIME_INFER_WORD = {
|
||||
LanguageEnum.CN: "推断时间",
|
||||
LanguageEnum.EN: "Inference time"
|
||||
LanguageEnum.EN: "Inference time",
|
||||
}
|
||||
|
||||
USER_NAME_EXPRESSION = {
|
||||
LanguageEnum.CN: "用户姓名是{name}。",
|
||||
LanguageEnum.EN: "User's name is {name}."
|
||||
LanguageEnum.EN: "User's name is {name}.",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,11 @@
|
|||
"""Enumeration module for ReMe.
|
||||
|
||||
This module provides enumerations used throughout the ReMe system,
|
||||
including language enumerations and other type definitions.
|
||||
"""
|
||||
|
||||
from reme_ai.enumeration.language_enum import LanguageEnum
|
||||
|
||||
__all__ = [
|
||||
"LanguageEnum",
|
||||
]
|
||||
|
|
@ -1,3 +1,8 @@
|
|||
"""Language enumeration module.
|
||||
|
||||
This module provides enumerations for supported languages in the ReMe system.
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
|
|
@ -9,6 +14,7 @@ class LanguageEnum(str, Enum):
|
|||
- CN: Represents the Chinese language.
|
||||
- EN: Represents the English language.
|
||||
"""
|
||||
|
||||
CN = "cn"
|
||||
|
||||
EN = "en"
|
||||
|
|
|
|||
|
|
@ -12,50 +12,57 @@ which extends FlowLLM with specialized memory management capabilities including:
|
|||
import asyncio
|
||||
import sys
|
||||
|
||||
from flowllm import FlowLLMApp, C
|
||||
from flowllm.schema.flow_response import FlowResponse
|
||||
from flowllm.core.application import Application
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.schema import FlowResponse
|
||||
|
||||
from reme_ai.config.config_parser import ConfigParser
|
||||
|
||||
|
||||
class ReMeApp(FlowLLMApp):
|
||||
class ReMeApp(Application):
|
||||
"""
|
||||
ReMeApp - Main application class for Reflexive Memory system.
|
||||
|
||||
|
||||
ReMeApp extends FlowLLMApp to provide enhanced memory capabilities for AI agents.
|
||||
It manages multiple types of memories and provides both synchronous and asynchronous
|
||||
execution interfaces for memory-enhanced workflows.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
*args,
|
||||
llm_api_key: str = None,
|
||||
llm_api_base: str = None,
|
||||
embedding_api_key: str = None,
|
||||
embedding_api_base: str = None,
|
||||
config_path: str = None,
|
||||
**kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
llm_api_key: str = None,
|
||||
llm_api_base: str = None,
|
||||
embedding_api_key: str = None,
|
||||
embedding_api_base: str = None,
|
||||
config_path: str = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize ReMeApp with configuration for LLM, embeddings, and vector stores.
|
||||
|
||||
⚠️ IMPORTANT: The initialization parameters here are consistent with the command-line
|
||||
|
||||
⚠️ IMPORTANT: The initialization parameters here are consistent with the command-line
|
||||
startup parameters shown in README.md. You can use the same configuration in both ways:
|
||||
|
||||
Command-line startup (from README):
|
||||
$ reme \\
|
||||
backend=http \\
|
||||
http.port=8002 \\
|
||||
llm.default.model_name=qwen3-30b-a3b-thinking-2507 \\
|
||||
embedding_model.default.model_name=text-embedding-v4 \\
|
||||
|
||||
Command-line startup:
|
||||
```bash
|
||||
reme \
|
||||
backend=http \
|
||||
http.port=8002 \
|
||||
llm.default.model_name=qwen3-30b-a3b-thinking-2507 \
|
||||
embedding_model.default.model_name=text-embedding-v4 \
|
||||
vector_store.default.backend=memory
|
||||
|
||||
```
|
||||
|
||||
Python API equivalent:
|
||||
>>> app = ReMeApp(
|
||||
... "llm.default.model_name=qwen3-30b-a3b-thinking-2507",
|
||||
... "embedding_model.default.model_name=text-embedding-v4",
|
||||
... "vector_store.default.backend=memory"
|
||||
... )
|
||||
|
||||
```python
|
||||
app = ReMeApp(
|
||||
"llm.default.model_name=qwen3-30b-a3b-thinking-2507",
|
||||
"embedding_model.default.model_name=text-embedding-v4",
|
||||
"vector_store.default.backend=memory"
|
||||
)
|
||||
```
|
||||
|
||||
Both approaches accept the same configuration parameters and produce identical results.
|
||||
|
||||
Args:
|
||||
|
|
@ -101,40 +108,42 @@ class ReMeApp(FlowLLMApp):
|
|||
This overrides the default configuration with your custom settings.
|
||||
**kwargs: Additional keyword arguments passed to parser. Same format as args but as key-value pairs.
|
||||
Example: model_name="gpt-4", temperature=0.7
|
||||
|
||||
|
||||
Raises:
|
||||
AssertionError: If required configurations are missing or invalid.
|
||||
|
||||
|
||||
Note:
|
||||
- Parameters here mirror the command-line options in README.md exactly
|
||||
- API keys can be provided via arguments or environment variables (see example.env)
|
||||
- The parser (ConfigParser) handles merging default configs with custom overrides
|
||||
- Vector store configuration determines where memories are persisted
|
||||
- For detailed startup examples and all available parameters, refer to README.md Quick Start section
|
||||
|
||||
|
||||
See Also:
|
||||
- README.md "Quick Start" section for command-line startup examples
|
||||
- README.md "Environment Configuration" for environment variable setup
|
||||
- example.env for all available environment variables
|
||||
"""
|
||||
super().__init__(*args,
|
||||
llm_api_key=llm_api_key,
|
||||
llm_api_base=llm_api_base,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_api_base=embedding_api_base,
|
||||
service_config=None,
|
||||
parser=ConfigParser,
|
||||
config_path=config_path,
|
||||
load_default_config=True,
|
||||
**kwargs)
|
||||
super().__init__(
|
||||
*args,
|
||||
llm_api_key=llm_api_key,
|
||||
llm_api_base=llm_api_base,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_api_base=embedding_api_base,
|
||||
service_config=None,
|
||||
parser=ConfigParser,
|
||||
config_path=config_path,
|
||||
load_default_config=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def async_execute(self, name: str, **kwargs) -> dict:
|
||||
"""
|
||||
Asynchronously execute a named flow with given parameters.
|
||||
|
||||
|
||||
This method executes a registered flow (workflow) by name and returns the result.
|
||||
Flows are defined in the configuration and registered during app initialization.
|
||||
|
||||
|
||||
Args:
|
||||
name: Name of the flow to execute. Must be registered in C.flow_dict.
|
||||
Common flows in ReMe:
|
||||
|
|
@ -148,23 +157,25 @@ class ReMeApp(FlowLLMApp):
|
|||
- context (dict): Additional context for the flow
|
||||
- max_results (int): Maximum number of results to return
|
||||
- threshold (float): Similarity threshold for retrieval
|
||||
|
||||
|
||||
Returns:
|
||||
dict: Flow execution result as a dictionary containing:
|
||||
- status: Execution status (success/failure)
|
||||
- result: Flow output data
|
||||
- metadata: Additional execution metadata
|
||||
|
||||
|
||||
Raises:
|
||||
AssertionError: If the flow name is not registered in C.flow_dict.
|
||||
|
||||
|
||||
Example:
|
||||
>>> result = await app.async_execute(
|
||||
... "task_memory_flow",
|
||||
... query="Show me all Python debugging tasks",
|
||||
... max_results=10
|
||||
... )
|
||||
>>> print(result['result'])
|
||||
```python
|
||||
result = await app.async_execute(
|
||||
"task_memory_flow",
|
||||
query="Show me all Python debugging tasks",
|
||||
max_results=10
|
||||
)
|
||||
print(result['result'])
|
||||
```
|
||||
"""
|
||||
assert name in C.flow_dict, f"Invalid flow_name={name} !"
|
||||
result: FlowResponse = await self.async_execute_flow(name=name, **kwargs)
|
||||
|
|
@ -173,28 +184,30 @@ class ReMeApp(FlowLLMApp):
|
|||
def execute(self, name: str, **kwargs) -> dict:
|
||||
"""
|
||||
Synchronously execute a named flow with given parameters.
|
||||
|
||||
|
||||
This is a convenience wrapper around async_execute() for synchronous contexts.
|
||||
It internally uses asyncio.run() to execute the async flow.
|
||||
|
||||
|
||||
Args:
|
||||
name: Name of the flow to execute. See async_execute() for available flows.
|
||||
**kwargs: Keyword arguments passed to the flow. See async_execute() for details.
|
||||
|
||||
|
||||
Returns:
|
||||
dict: Flow execution result. Same format as async_execute().
|
||||
|
||||
|
||||
Raises:
|
||||
AssertionError: If the flow name is not registered.
|
||||
|
||||
|
||||
Example:
|
||||
>>> app = ReMeApp()
|
||||
>>> result = app.execute(
|
||||
... "tool_memory_flow",
|
||||
... query="How to use the search tool effectively?"
|
||||
... )
|
||||
>>> print(result)
|
||||
|
||||
```python
|
||||
app = ReMeApp()
|
||||
result = app.execute(
|
||||
"tool_memory_flow",
|
||||
query="How to use the search tool effectively?"
|
||||
)
|
||||
print(result)
|
||||
```
|
||||
|
||||
Note:
|
||||
For better performance in async contexts, prefer using async_execute() directly.
|
||||
This method creates a new event loop for each call, which has overhead.
|
||||
|
|
@ -205,24 +218,25 @@ class ReMeApp(FlowLLMApp):
|
|||
def main():
|
||||
"""
|
||||
Entry point for running ReMeApp as a service.
|
||||
|
||||
|
||||
This function initializes ReMeApp with command-line arguments and starts the service.
|
||||
It's typically called when running the module directly (python -m reme_ai.app).
|
||||
|
||||
|
||||
Command-line arguments are passed directly to ReMeApp.__init__(), allowing
|
||||
configuration via command line:
|
||||
|
||||
|
||||
Example:
|
||||
$ python -m reme_ai.app --llm_api_key=sk-xxx --config_path=config.yaml
|
||||
|
||||
python -m reme_ai.app --llm_api_key=sk-xxx --config_path=config.yaml
|
||||
|
||||
The app runs as a context manager, ensuring proper cleanup of resources
|
||||
(database connections, API clients, etc.) on shutdown.
|
||||
|
||||
|
||||
Note:
|
||||
Press Ctrl+C to gracefully shutdown the service.
|
||||
"""
|
||||
with ReMeApp(*sys.argv[1:]) as app:
|
||||
app.run_service()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,3 +1,17 @@
|
|||
"""Retrieval module for memory operations.
|
||||
|
||||
This module provides submodules for different types of memory retrieval:
|
||||
- personal: Personal memory retrieval operations
|
||||
- task: Task memory retrieval operations
|
||||
- tool: Tool memory retrieval operations
|
||||
"""
|
||||
|
||||
from . import personal
|
||||
from . import task
|
||||
from . import tool
|
||||
|
||||
__all__ = [
|
||||
"personal",
|
||||
"task",
|
||||
"tool",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,9 @@
|
|||
"""Personal memory retrieval operations module.
|
||||
|
||||
This module provides operations for retrieving, ranking, and formatting personal memories
|
||||
from a vector store, including time extraction, semantic ranking, and memory formatting.
|
||||
"""
|
||||
|
||||
from .extract_time_op import ExtractTimeOp
|
||||
from .fuse_rerank_op import FuseRerankOp
|
||||
from .print_memory_op import PrintMemoryOp
|
||||
|
|
@ -13,5 +19,5 @@ __all__ = [
|
|||
"ReadMessageOp",
|
||||
"RetrieveMemoryOp",
|
||||
"SemanticRankOp",
|
||||
"SetQueryOp"
|
||||
"SetQueryOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,9 +1,16 @@
|
|||
"""Time extraction operation for personal memories.
|
||||
|
||||
This module provides functionality to extract time-related information from queries
|
||||
using LLM-based extraction and pattern matching.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import Dict
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT
|
||||
|
|
@ -13,16 +20,25 @@ from reme_ai.utils.datetime_handler import DatetimeHandler
|
|||
|
||||
@C.register_op()
|
||||
class ExtractTimeOp(BaseAsyncOp):
|
||||
file_path: str = __file__
|
||||
EXTRACT_TIME_PATTERN = r"-\s*(\S+)[::]\s*(\S+)"
|
||||
|
||||
"""
|
||||
A specialized worker class designed to identify and extract time-related information
|
||||
from text generated by an LLM, translating date-time keywords based on the set language,
|
||||
and storing this extracted data within a shared context.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
EXTRACT_TIME_PATTERN = r"-\s*(\S+)[::]\s*(\S+)"
|
||||
|
||||
def get_language_value(self, value_dict: dict):
|
||||
"""
|
||||
Get value from dictionary based on current language setting.
|
||||
|
||||
Args:
|
||||
value_dict: Dictionary with language keys
|
||||
|
||||
Returns:
|
||||
Value for current language or English fallback
|
||||
"""
|
||||
return value_dict.get(self.language, value_dict.get("en"))
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -51,8 +67,11 @@ class ExtractTimeOp(BaseAsyncOp):
|
|||
# Create message with system and few-shot examples
|
||||
system_prompt = self.prompt_format(prompt_name="extract_time_system")
|
||||
few_shot = self.prompt_format(prompt_name="extract_time_few_shot")
|
||||
user_prompt = self.prompt_format(prompt_name="extract_time_user_query",
|
||||
query=query, query_time_str=query_time_str)
|
||||
user_prompt = self.prompt_format(
|
||||
prompt_name="extract_time_user_query",
|
||||
query=query,
|
||||
query_time_str=query_time_str,
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_prompt}"
|
||||
logger.info(f"Extracting time from query: {query[:100]}...")
|
||||
|
|
@ -77,10 +96,10 @@ class ExtractTimeOp(BaseAsyncOp):
|
|||
def _parse_time_from_response(self, response_text: str) -> Dict[str, str]:
|
||||
"""
|
||||
Parse time information from LLM response using regex.
|
||||
|
||||
|
||||
Args:
|
||||
response_text: Raw LLM response content
|
||||
|
||||
|
||||
Returns:
|
||||
Dictionary of extracted time information
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
"""Fuse reranking operation for personal memories.
|
||||
|
||||
This module provides functionality to rerank memory nodes by combining scores,
|
||||
memory types, and temporal relevance to improve retrieval quality.
|
||||
"""
|
||||
|
||||
from typing import Dict, List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.constants.common_constants import EXTRACT_TIME_DICT
|
||||
|
|
@ -12,12 +19,20 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
"""
|
||||
Reranks the memory nodes by scores, types, and temporal relevance. Formats the top-K reranked nodes to print.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
@staticmethod
|
||||
def match_memory_time(extract_time_dict: Dict[str, str], memory: BaseMemory):
|
||||
"""
|
||||
Determines whether the memory is relevant based on time matching.
|
||||
|
||||
Args:
|
||||
extract_time_dict: Dictionary containing extracted time information
|
||||
memory: Memory object to check for time relevance
|
||||
|
||||
Returns:
|
||||
Tuple of (match_event_flag, match_msg_flag) indicating temporal matches
|
||||
"""
|
||||
if extract_time_dict:
|
||||
match_event_flag = True
|
||||
|
|
@ -25,18 +40,16 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
event_value = memory.metadata.get(f"event_{k}", "")
|
||||
if event_value in ["-1", v]:
|
||||
continue
|
||||
else:
|
||||
match_event_flag = False
|
||||
break
|
||||
match_event_flag = False
|
||||
break
|
||||
|
||||
match_msg_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
msg_value = memory.metadata.get(f"msg_{k}", "")
|
||||
if msg_value == v:
|
||||
continue
|
||||
else:
|
||||
match_msg_flag = False
|
||||
break
|
||||
match_msg_flag = False
|
||||
break
|
||||
else:
|
||||
match_event_flag = False
|
||||
match_msg_flag = False
|
||||
|
|
@ -59,12 +72,15 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
"""
|
||||
# Get operation parameters
|
||||
fuse_score_threshold = self.op_params.get("fuse_score_threshold", 0.1)
|
||||
fuse_ratio_dict = self.op_params.get("fuse_ratio_dict", {
|
||||
"conversation": 0.5,
|
||||
"observation": 1,
|
||||
"obs_customized": 1.2,
|
||||
"insight": 2.0
|
||||
})
|
||||
fuse_ratio_dict = self.op_params.get(
|
||||
"fuse_ratio_dict",
|
||||
{
|
||||
"conversation": 0.5,
|
||||
"observation": 1,
|
||||
"obs_customized": 1.2,
|
||||
"insight": 2.0,
|
||||
},
|
||||
)
|
||||
fuse_time_ratio = self.op_params.get("fuse_time_ratio", 2.0)
|
||||
output_memory_max_count = self.op_params.get("output_memory_max_count", 5)
|
||||
|
||||
|
|
@ -83,14 +99,19 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
|
||||
# Perform reranking based on score, type, and time relevance
|
||||
reranked_memories = self._apply_fuse_reranking(
|
||||
memory_list, extract_time_dict, fuse_score_threshold,
|
||||
fuse_ratio_dict, fuse_time_ratio
|
||||
memory_list,
|
||||
extract_time_dict,
|
||||
fuse_score_threshold,
|
||||
fuse_ratio_dict,
|
||||
fuse_time_ratio,
|
||||
)
|
||||
|
||||
# Sort and select top-k memories
|
||||
reranked_memories = sorted(reranked_memories,
|
||||
key=lambda x: x.score or 0.0,
|
||||
reverse=True)[:output_memory_max_count]
|
||||
reranked_memories = sorted(
|
||||
reranked_memories,
|
||||
key=lambda x: x.score or 0.0,
|
||||
reverse=True,
|
||||
)[:output_memory_max_count]
|
||||
|
||||
logger.info(f"Final reranked memories: {len(reranked_memories)}")
|
||||
|
||||
|
|
@ -101,13 +122,27 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
self.context.response.metadata["memory_list"] = reranked_memories
|
||||
self.context.response.answer = "\n".join(formatted_memories)
|
||||
|
||||
def _apply_fuse_reranking(self,
|
||||
memory_list: List[BaseMemory],
|
||||
extract_time_dict: Dict[str, str],
|
||||
fuse_score_threshold: float,
|
||||
fuse_ratio_dict: Dict[str, float],
|
||||
fuse_time_ratio: float) -> List[BaseMemory]:
|
||||
"""Apply fuse reranking logic to memories"""
|
||||
def _apply_fuse_reranking(
|
||||
self,
|
||||
memory_list: List[BaseMemory],
|
||||
extract_time_dict: Dict[str, str],
|
||||
fuse_score_threshold: float,
|
||||
fuse_ratio_dict: Dict[str, float],
|
||||
fuse_time_ratio: float,
|
||||
) -> List[BaseMemory]:
|
||||
"""
|
||||
Apply fuse reranking logic to memories.
|
||||
|
||||
Args:
|
||||
memory_list: List of memories to rerank
|
||||
extract_time_dict: Dictionary containing extracted time information
|
||||
fuse_score_threshold: Minimum score threshold for memories
|
||||
fuse_ratio_dict: Dictionary mapping memory types to score multipliers
|
||||
fuse_time_ratio: Multiplier for time-relevant memories
|
||||
|
||||
Returns:
|
||||
List of reranked memories with updated scores
|
||||
"""
|
||||
reranked_memories = []
|
||||
|
||||
for memory in memory_list:
|
||||
|
|
@ -130,23 +165,35 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
original_score = memory_score
|
||||
memory.score = memory_score * type_ratio * time_ratio
|
||||
|
||||
logger.debug(f"Memory reranked: {original_score:.3f} -> {memory.score:.3f} "
|
||||
f"(type={type_ratio}, time={time_ratio})")
|
||||
logger.debug(
|
||||
f"Memory reranked: {original_score:.3f} -> {memory.score:.3f} "
|
||||
f"(type={type_ratio}, time={time_ratio})",
|
||||
)
|
||||
|
||||
reranked_memories.append(memory)
|
||||
|
||||
return reranked_memories
|
||||
|
||||
def _format_memories_for_output(self, memories: List[BaseMemory]) -> List[str]:
|
||||
"""Format memories for final output"""
|
||||
"""
|
||||
Format memories for final output.
|
||||
|
||||
Args:
|
||||
memories: List of memories to format
|
||||
|
||||
Returns:
|
||||
List of formatted memory strings
|
||||
"""
|
||||
formatted_memories = []
|
||||
|
||||
for memory in memories:
|
||||
# Log reranking details
|
||||
logger.info(f"Final memory: Score={memory.score:.3f}, "
|
||||
f"Event={memory.metadata.get('match_event_flag', '0')}, "
|
||||
f"Msg={memory.metadata.get('match_msg_flag', '0')}, "
|
||||
f"Content={memory.content[:50]}...")
|
||||
logger.info(
|
||||
f"Final memory: Score={memory.score:.3f}, "
|
||||
f"Event={memory.metadata.get('match_event_flag', '0')}, "
|
||||
f"Msg={memory.metadata.get('match_msg_flag', '0')}, "
|
||||
f"Content={memory.content[:50]}...",
|
||||
)
|
||||
|
||||
# Format memory with timestamp if available
|
||||
formatted_content = self._format_memory_with_timestamp(memory, self.language)
|
||||
|
|
@ -167,8 +214,9 @@ class FuseRerankOp(BaseAsyncOp):
|
|||
Formatted memory content string
|
||||
"""
|
||||
try:
|
||||
if hasattr(memory, 'timestamp') and memory.timestamp:
|
||||
if hasattr(memory, "timestamp") and memory.timestamp:
|
||||
from reme_ai.utils.datetime_handler import DatetimeHandler
|
||||
|
||||
dt_handler = DatetimeHandler(memory.timestamp)
|
||||
datetime_str = dt_handler.datetime_format("%Y-%m-%d %H:%M:%S")
|
||||
weekday = dt_handler.get_dt_info_dict(language)["weekday"]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
"""Memory printing operation for personal memories.
|
||||
|
||||
This module provides functionality to format and print memories in various formats
|
||||
for display or output purposes.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -11,6 +18,7 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
"""
|
||||
Formats the memories to print.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -39,7 +47,15 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _format_memories_for_print(memories: List[BaseMemory]) -> str:
|
||||
"""Format memories for printing"""
|
||||
"""
|
||||
Format memories for printing.
|
||||
|
||||
Args:
|
||||
memories: List of memory objects to format
|
||||
|
||||
Returns:
|
||||
Formatted string representation of memories
|
||||
"""
|
||||
if not memories:
|
||||
return "No memories available."
|
||||
|
||||
|
|
@ -51,10 +67,10 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
memory_text += f" Content: {memory.content}\n"
|
||||
|
||||
# Add additional metadata if available
|
||||
if hasattr(memory, 'metadata') and memory.metadata:
|
||||
if hasattr(memory, "metadata") and memory.metadata:
|
||||
metadata_items = []
|
||||
for key, value in memory.metadata.items():
|
||||
if key not in ['when_to_use', 'content']:
|
||||
if key not in ["when_to_use", "content"]:
|
||||
metadata_items.append(f"{key}: {value}")
|
||||
if metadata_items:
|
||||
memory_text += f" Metadata: {', '.join(metadata_items)}\n"
|
||||
|
|
@ -67,10 +83,10 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
def format_memories_for_output(memories: List) -> str:
|
||||
"""
|
||||
Format memory list for output string.
|
||||
|
||||
|
||||
Args:
|
||||
memories: List of memory objects
|
||||
|
||||
|
||||
Returns:
|
||||
Formatted string
|
||||
"""
|
||||
|
|
@ -79,8 +95,8 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
|
||||
formatted_parts = []
|
||||
for i, memory in enumerate(memories, 1):
|
||||
when_to_use = getattr(memory, 'when_to_use', '') or memory.get('when_to_use', '')
|
||||
content = getattr(memory, 'content', '') or memory.get('content', '')
|
||||
when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "")
|
||||
content = getattr(memory, "content", "") or memory.get("content", "")
|
||||
|
||||
part = f"Memory {i}:\n"
|
||||
if when_to_use:
|
||||
|
|
@ -96,10 +112,10 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
def format_memories_for_simple_output(memories: List) -> str:
|
||||
"""
|
||||
Format memory list for simple flow output.
|
||||
|
||||
|
||||
Args:
|
||||
memories: List of memory objects
|
||||
|
||||
|
||||
Returns:
|
||||
Formatted string suitable for response.answer
|
||||
"""
|
||||
|
|
@ -110,8 +126,8 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
|
||||
for memory in memories:
|
||||
# Safely get field values
|
||||
when_to_use = getattr(memory, 'when_to_use', '') or memory.get('when_to_use', '')
|
||||
content = getattr(memory, 'content', '') or memory.get('content', '')
|
||||
when_to_use = getattr(memory, "when_to_use", "") or memory.get("when_to_use", "")
|
||||
content = getattr(memory, "content", "") or memory.get("content", "")
|
||||
|
||||
# Skip memories with empty content
|
||||
if not content:
|
||||
|
|
@ -125,7 +141,9 @@ class PrintMemoryOp(BaseAsyncOp):
|
|||
if len(content_parts) == 1: # Only title
|
||||
return "No relevant memories with valid content found."
|
||||
|
||||
content_parts.append("\nPlease consider the helpful parts from these in answering the question, "
|
||||
"to make the response more comprehensive and substantial.")
|
||||
content_parts.append(
|
||||
"\nPlease consider the helpful parts from these in answering the question, "
|
||||
"to make the response more comprehensive and substantial.",
|
||||
)
|
||||
|
||||
return "\n".join(content_parts)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,14 @@
|
|||
"""Message reading operation for personal memories.
|
||||
|
||||
This module provides functionality to read and filter unmemorized chat messages
|
||||
from the context for processing.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
|
||||
|
|
@ -10,6 +17,7 @@ class ReadMessageOp(BaseAsyncOp):
|
|||
"""
|
||||
Fetches unmemorized chat messages.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -19,20 +27,20 @@ class ReadMessageOp(BaseAsyncOp):
|
|||
# Get chat messages from context
|
||||
chat_messages = self.context.chat_messages
|
||||
target_name = self.context.target_name
|
||||
contextual_msg_max_count = self.op_params.get('contextual_msg_max_count', 10)
|
||||
contextual_msg_max_count = self.op_params.get("contextual_msg_max_count", 10)
|
||||
|
||||
chat_messages_not_memorized: List[List[Message]] = []
|
||||
for messages in chat_messages:
|
||||
if not messages:
|
||||
continue
|
||||
|
||||
if hasattr(messages[0], 'memorized') and messages[0].memorized:
|
||||
if hasattr(messages[0], "memorized") and messages[0].memorized:
|
||||
continue
|
||||
|
||||
contain_flag = False
|
||||
|
||||
for msg in messages:
|
||||
if hasattr(msg, 'role_name') and msg.role_name == target_name:
|
||||
if hasattr(msg, "role_name") and msg.role_name == target_name:
|
||||
contain_flag = True
|
||||
break
|
||||
|
||||
|
|
@ -44,7 +52,7 @@ class ReadMessageOp(BaseAsyncOp):
|
|||
chat_message_scatter.extend(messages)
|
||||
|
||||
# Sort by time_created if available
|
||||
if chat_message_scatter and hasattr(chat_message_scatter[0], 'time_created'):
|
||||
if chat_message_scatter and hasattr(chat_message_scatter[0], "time_created"):
|
||||
chat_message_scatter.sort(key=lambda _: _.time_created)
|
||||
|
||||
# Store result in context
|
||||
|
|
|
|||
|
|
@ -1,7 +1,14 @@
|
|||
"""Memory retrieval operation for personal memories.
|
||||
|
||||
This module provides functionality to retrieve memories from a vector store
|
||||
based on query similarity and score thresholds.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import VectorNode
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, vector_node_to_memory
|
||||
|
|
@ -16,6 +23,15 @@ class RetrieveMemoryOp(BaseAsyncOp):
|
|||
"""
|
||||
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the memory retrieval operation.
|
||||
|
||||
This method:
|
||||
1. Retrieves memories from vector store based on query similarity
|
||||
2. Removes duplicate memories based on content
|
||||
3. Filters memories by score threshold if specified
|
||||
4. Stores the retrieved memories in context metadata
|
||||
"""
|
||||
recall_key: str = self.op_params.get("recall_key", "query")
|
||||
top_k: int = self.context.get("top_k", 3)
|
||||
|
||||
|
|
@ -23,9 +39,11 @@ class RetrieveMemoryOp(BaseAsyncOp):
|
|||
assert query, "query should be not empty!"
|
||||
|
||||
workspace_id: str = self.context.workspace_id
|
||||
nodes: List[VectorNode] = await self.vector_store.async_search(query=query,
|
||||
workspace_id=workspace_id,
|
||||
top_k=top_k)
|
||||
nodes: List[VectorNode] = await self.vector_store.async_search(
|
||||
query=query,
|
||||
workspace_id=workspace_id,
|
||||
top_k=top_k,
|
||||
)
|
||||
memory_list: List[BaseMemory] = []
|
||||
memory_content_list: List[str] = []
|
||||
for node in nodes:
|
||||
|
|
|
|||
|
|
@ -1,11 +1,19 @@
|
|||
"""Semantic ranking operation for personal memories.
|
||||
|
||||
This module provides functionality to rank memories semantically using LLM-based
|
||||
relevance scoring to improve retrieval quality.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Role
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
||||
|
||||
|
|
@ -17,6 +25,7 @@ class SemanticRankOp(BaseAsyncOp):
|
|||
assigning scores, sorting the nodes, and storing the ranked nodes back,
|
||||
while logging relevant information.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -73,7 +82,14 @@ class SemanticRankOp(BaseAsyncOp):
|
|||
|
||||
async def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]:
|
||||
"""
|
||||
Use LLM to semantically rank memories based on relevance to the query
|
||||
Use LLM to semantically rank memories based on relevance to the query.
|
||||
|
||||
Args:
|
||||
query: User query to rank memories against
|
||||
memories: List of memories to rank
|
||||
|
||||
Returns:
|
||||
List of memories with updated semantic scores
|
||||
"""
|
||||
if not memories:
|
||||
return memories
|
||||
|
|
@ -84,7 +100,7 @@ class SemanticRankOp(BaseAsyncOp):
|
|||
# Create prompt for semantic ranking
|
||||
prompt = f"""Given the query: "{query}"
|
||||
|
||||
Please rank the following memories by their semantic relevance to the query.
|
||||
Please rank the following memories by their semantic relevance to the query.
|
||||
Rate each memory on a scale of 0.0 to 1.0 where 1.0 is most relevant.
|
||||
|
||||
Memories:
|
||||
|
|
@ -112,10 +128,18 @@ Please respond in JSON format:
|
|||
|
||||
@staticmethod
|
||||
def parse_llm_ranking_response(response: str) -> List[dict]:
|
||||
"""Parse LLM ranking response to extract rankings."""
|
||||
"""
|
||||
Parse LLM ranking response to extract rankings.
|
||||
|
||||
Args:
|
||||
response: Raw LLM response string containing ranking JSON
|
||||
|
||||
Returns:
|
||||
List of ranking dictionaries with index and score
|
||||
"""
|
||||
try:
|
||||
# Try to extract JSON blocks
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_pattern = r"```json\s*([\s\S]*?)\s*```"
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
|
|
@ -135,7 +159,16 @@ Please respond in JSON format:
|
|||
|
||||
@staticmethod
|
||||
def apply_semantic_scores_to_memories(memories: List, rankings: List[dict]) -> int:
|
||||
"""Apply semantic ranking scores to memory objects."""
|
||||
"""
|
||||
Apply semantic ranking scores to memory objects.
|
||||
|
||||
Args:
|
||||
memories: List of memory objects to update
|
||||
rankings: List of ranking dictionaries with index and score
|
||||
|
||||
Returns:
|
||||
Number of memories successfully updated with scores
|
||||
"""
|
||||
applied_count = 0
|
||||
|
||||
for ranking in rankings:
|
||||
|
|
@ -144,21 +177,29 @@ Please respond in JSON format:
|
|||
|
||||
if 0 <= idx < len(memories):
|
||||
# Set score on memory object
|
||||
if hasattr(memories[idx], 'score'):
|
||||
if hasattr(memories[idx], "score"):
|
||||
memories[idx].score = score
|
||||
applied_count += 1
|
||||
else:
|
||||
# Add score as metadata if score attribute doesn't exist
|
||||
if not hasattr(memories[idx], 'metadata'):
|
||||
if not hasattr(memories[idx], "metadata"):
|
||||
memories[idx].metadata = {}
|
||||
memories[idx].metadata['semantic_score'] = score
|
||||
memories[idx].metadata["semantic_score"] = score
|
||||
applied_count += 1
|
||||
|
||||
return applied_count
|
||||
|
||||
@staticmethod
|
||||
def format_memories_for_llm_ranking(memories: List) -> str:
|
||||
"""Format memories for LLM ranking input."""
|
||||
"""
|
||||
Format memories for LLM ranking input.
|
||||
|
||||
Args:
|
||||
memories: List of memory objects to format
|
||||
|
||||
Returns:
|
||||
Formatted string representation of memories for LLM input
|
||||
"""
|
||||
formatted_memories = []
|
||||
|
||||
for i, memory in enumerate(memories):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,14 @@
|
|||
"""Query setting operation for personal memories.
|
||||
|
||||
This module provides functionality to set query and timestamp in the context
|
||||
for downstream memory retrieval operations.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
from typing import Tuple
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.constants.common_constants import QUERY_WITH_TS
|
||||
|
|
|
|||
|
|
@ -1,4 +1,17 @@
|
|||
"""Task memory retrieval operations module.
|
||||
|
||||
This module provides operations for building queries, reranking memories,
|
||||
rewriting memory context, and merging memories for task-related retrieval.
|
||||
"""
|
||||
|
||||
from .build_query_op import BuildQueryOp
|
||||
from .merge_memory_op import MergeMemoryOp
|
||||
from .rerank_memory_op import RerankMemoryOp
|
||||
from .rewrite_memory_op import RewriteMemoryOp
|
||||
|
||||
__all__ = [
|
||||
"BuildQueryOp",
|
||||
"MergeMemoryOp",
|
||||
"RerankMemoryOp",
|
||||
"RewriteMemoryOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,15 +1,38 @@
|
|||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.utils.llm_utils import merge_messages_content
|
||||
from loguru import logger
|
||||
"""Query building operation module.
|
||||
|
||||
from reme_ai.schema import Message, Role
|
||||
This module provides functionality to build retrieval queries from either
|
||||
explicit query strings or conversation messages, optionally using LLM to
|
||||
generate optimized queries.
|
||||
"""
|
||||
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from flowllm.core.utils import merge_messages_content
|
||||
from loguru import logger
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class BuildQueryOp(BaseAsyncOp):
|
||||
"""Build retrieval query from context or messages.
|
||||
|
||||
This operation constructs a query string for memory retrieval. It can use
|
||||
an explicit query from context, or generate one from conversation messages
|
||||
using either LLM-based generation or simple message concatenation.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the query building operation.
|
||||
|
||||
Builds a query string from either:
|
||||
1. An explicit query in the context
|
||||
2. Conversation messages (using LLM or simple concatenation)
|
||||
|
||||
Stores the built query in context.query.
|
||||
"""
|
||||
if "query" in self.context:
|
||||
query = self.context.query
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
query_build: |
|
||||
# Execution Process
|
||||
{execution_process}
|
||||
|
||||
Read through the entire execution process to understand which part is currently being executed.
|
||||
|
||||
Read through the entire execution process to understand which part is currently being executed.
|
||||
Generate a `query` that reflects the current state, which will later be used to search for similar problems in the database and help resolve the issue at hand.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
"""Memory merging operation module.
|
||||
|
||||
This module provides functionality to merge multiple retrieved memories
|
||||
into a single formatted context string for use in LLM responses.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -8,8 +15,19 @@ from reme_ai.schema.memory import BaseMemory
|
|||
|
||||
@C.register_op()
|
||||
class MergeMemoryOp(BaseAsyncOp):
|
||||
"""Merge multiple memories into a single formatted context.
|
||||
|
||||
This operation takes a list of retrieved memories and formats them into
|
||||
a single context string that can be used to guide LLM responses. It includes
|
||||
instructions for the LLM to consider the helpful parts from these memories.
|
||||
"""
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the memory merging operation.
|
||||
|
||||
Merges memories from context metadata into a formatted string with
|
||||
instructions for the LLM. Stores the merged result in response.answer.
|
||||
"""
|
||||
memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"]
|
||||
|
||||
if not memory_list:
|
||||
|
|
@ -21,7 +39,9 @@ class MergeMemoryOp(BaseAsyncOp):
|
|||
continue
|
||||
|
||||
content_collector.append(f"- {memory.when_to_use} {memory.content}\n")
|
||||
content_collector.append("Please consider the helpful parts from these in answering the question, "
|
||||
"to make the response more comprehensive and substantial.")
|
||||
content_collector.append(
|
||||
"Please consider the helpful parts from these in answering the question, "
|
||||
"to make the response more comprehensive and substantial.",
|
||||
)
|
||||
self.context.response.answer = "\n".join(content_collector)
|
||||
logger.info(f"response.answer={self.context.response.answer}")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,18 @@
|
|||
"""Memory reranking operation module.
|
||||
|
||||
This module provides functionality to rerank and filter retrieved memories
|
||||
using LLM-based reranking and score-based filtering to select the most relevant
|
||||
memories for the current task.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -12,13 +20,22 @@ from reme_ai.schema.memory import BaseMemory
|
|||
|
||||
@C.register_op()
|
||||
class RerankMemoryOp(BaseAsyncOp):
|
||||
"""Rerank and filter recalled experiences using LLM and score-based filtering.
|
||||
|
||||
This operation takes recalled memories and applies multiple filtering and
|
||||
ranking strategies to select the most relevant memories for the current task.
|
||||
It supports LLM-based reranking and score-based filtering.
|
||||
"""
|
||||
Rerank and filter recalled experiences using LLM and score-based filtering
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute rerank operation"""
|
||||
"""Execute the memory reranking operation.
|
||||
|
||||
Applies LLM-based reranking (optional) and score-based filtering (optional)
|
||||
to select the top-k most relevant memories. Stores the reranked results
|
||||
in the context response metadata.
|
||||
"""
|
||||
memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"]
|
||||
retrieval_query: str = self.context.query
|
||||
enable_llm_rerank = self.op_params.get("enable_llm_rerank", True)
|
||||
|
|
@ -52,7 +69,15 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
self.context.response.metadata["memory_list"] = reranked_memories
|
||||
|
||||
async def _llm_rerank(self, query: str, candidates: List[BaseMemory]) -> List[BaseMemory]:
|
||||
"""LLM-based reranking of candidate experiences"""
|
||||
"""LLM-based reranking of candidate experiences.
|
||||
|
||||
Args:
|
||||
query: The retrieval query used to rank candidates.
|
||||
candidates: List of memory candidates to rerank.
|
||||
|
||||
Returns:
|
||||
List of memories reranked by relevance to the query.
|
||||
"""
|
||||
if not candidates:
|
||||
return candidates
|
||||
|
||||
|
|
@ -63,7 +88,8 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
prompt_name="memory_rerank_prompt",
|
||||
query=query,
|
||||
candidates=candidates_text,
|
||||
num_candidates=len(candidates))
|
||||
num_candidates=len(candidates),
|
||||
)
|
||||
|
||||
response = await self.llm.achat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
|
|
@ -89,7 +115,15 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _score_based_filter(memories: List[BaseMemory], min_score: float) -> List[BaseMemory]:
|
||||
"""Filter memories based on quality scores"""
|
||||
"""Filter memories based on quality scores.
|
||||
|
||||
Args:
|
||||
memories: List of memories to filter.
|
||||
min_score: Minimum combined score threshold for filtering.
|
||||
|
||||
Returns:
|
||||
List of memories that meet the minimum score threshold.
|
||||
"""
|
||||
filtered_memories = []
|
||||
|
||||
for memory in memories:
|
||||
|
|
@ -110,7 +144,14 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _format_candidates_for_rerank(candidates: List[BaseMemory]) -> str:
|
||||
"""Format candidates for LLM reranking"""
|
||||
"""Format candidates for LLM reranking.
|
||||
|
||||
Args:
|
||||
candidates: List of memory candidates to format.
|
||||
|
||||
Returns:
|
||||
Formatted string representation of candidates for LLM evaluation.
|
||||
"""
|
||||
formatted_candidates = []
|
||||
|
||||
for i, candidate in enumerate(candidates):
|
||||
|
|
@ -127,10 +168,17 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _parse_rerank_response(response: str) -> List[int]:
|
||||
"""Parse LLM reranking response to extract ranked indices"""
|
||||
"""Parse LLM reranking response to extract ranked indices.
|
||||
|
||||
Args:
|
||||
response: The LLM response containing ranked indices.
|
||||
|
||||
Returns:
|
||||
List of indices representing the reranked order.
|
||||
"""
|
||||
try:
|
||||
# Try to extract JSON format
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_pattern = r"```json\s*([\s\S]*?)\s*```"
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
|
|
@ -141,7 +189,7 @@ class RerankMemoryOp(BaseAsyncOp):
|
|||
return parsed
|
||||
|
||||
# Try to extract numbers from text
|
||||
numbers = re.findall(r'\b\d+\b', response)
|
||||
numbers = re.findall(r"\b\d+\b", response)
|
||||
return [int(num) for num in numbers if int(num) < 100] # Reasonable upper bound
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
memory_rerank_prompt: |
|
||||
You are an expert AI analyst tasked with reranking retrieved experiences based on their relevance to a specific query.
|
||||
|
||||
|
||||
Your task is to analyze the candidates and rank them by relevance, considering:
|
||||
● DIRECT RELEVANCE: How directly applicable the experience is to the current query
|
||||
● SITUATION SIMILARITY: How similar the experience context is to the current situation
|
||||
● ACTIONABILITY: How actionable and specific the experience is
|
||||
● QUALITY: The overall quality and clarity of the experience
|
||||
|
||||
|
||||
# Current Query
|
||||
{query}
|
||||
|
||||
|
||||
# Candidate Experiences (Total: {num_candidates})
|
||||
{candidates}
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Provide a ranked list of candidate indices (0-based) from most relevant to least relevant:
|
||||
```json
|
||||
|
|
@ -21,5 +21,5 @@ memory_rerank_prompt: |
|
|||
"reasoning": "Brief explanation of ranking rationale"
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
Note: Include ALL candidate indices in the ranking, even if some are less relevant.
|
||||
|
|
@ -1,10 +1,17 @@
|
|||
"""Memory rewriting operation module.
|
||||
|
||||
This module provides functionality to rewrite and format retrieved memories
|
||||
into context messages that can be used by LLMs for task completion.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -12,17 +19,25 @@ from reme_ai.schema.memory import BaseMemory
|
|||
|
||||
@C.register_op()
|
||||
class RewriteMemoryOp(BaseAsyncOp):
|
||||
"""Generate and rewrite context messages from reranked experiences.
|
||||
|
||||
This operation takes reranked memories and formats them into context messages
|
||||
that can be used by LLMs. It optionally uses LLM-based rewriting to make
|
||||
the context more relevant and actionable for the current task.
|
||||
"""
|
||||
Generate and rewrite context messages from reranked experiences
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute rewrite operation"""
|
||||
"""Execute the memory rewrite operation.
|
||||
|
||||
Retrieves memories from context metadata, formats them, and optionally
|
||||
rewrites them using LLM to make them more relevant for the current query.
|
||||
Stores the rewritten context in the response answer field.
|
||||
"""
|
||||
memory_list: List[BaseMemory] = self.context.response.metadata["memory_list"]
|
||||
query: str = self.context.query
|
||||
messages: List[Message] = \
|
||||
[Message(**x) if isinstance(x, dict) else x for x in self.context.get('messages', [])]
|
||||
messages: List[Message] = [Message(**x) if isinstance(x, dict) else x for x in self.context.get("messages", [])]
|
||||
|
||||
if not memory_list:
|
||||
logger.info("No reranked memories to rewrite")
|
||||
|
|
@ -39,7 +54,16 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
self.context.response.metadata["memory_list"] = [memory.model_dump() for memory in memory_list]
|
||||
|
||||
async def _generate_context_message(self, query: str, messages: List[Message], memories: List[BaseMemory]) -> str:
|
||||
"""Generate context message from retrieved memories"""
|
||||
"""Generate context message from retrieved memories.
|
||||
|
||||
Args:
|
||||
query: The current query string.
|
||||
messages: List of conversation messages for context.
|
||||
memories: List of retrieved memories to format.
|
||||
|
||||
Returns:
|
||||
Formatted context string, optionally rewritten by LLM.
|
||||
"""
|
||||
if not memories:
|
||||
return ""
|
||||
|
||||
|
|
@ -60,7 +84,16 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
return self._format_memories_for_context(memories)
|
||||
|
||||
async def _rewrite_context(self, query: str, context_content: str, messages: List[Message]) -> str:
|
||||
"""LLM-based context rewriting to make experiences more relevant and actionable"""
|
||||
"""LLM-based context rewriting to make experiences more relevant and actionable.
|
||||
|
||||
Args:
|
||||
query: The current query string.
|
||||
context_content: The formatted context content to rewrite.
|
||||
messages: List of conversation messages for additional context.
|
||||
|
||||
Returns:
|
||||
Rewritten context string optimized for the current task.
|
||||
"""
|
||||
if not context_content:
|
||||
return context_content
|
||||
|
||||
|
|
@ -72,7 +105,8 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
prompt_name="memory_rewrite_prompt",
|
||||
current_query=query,
|
||||
current_context=current_context,
|
||||
original_context=context_content)
|
||||
original_context=context_content,
|
||||
)
|
||||
|
||||
response = await self.llm.achat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
|
|
@ -91,7 +125,14 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _format_memories_for_context(memories: List[BaseMemory]) -> str:
|
||||
"""Format memories for context generation"""
|
||||
"""Format memories for context generation.
|
||||
|
||||
Args:
|
||||
memories: List of memories to format.
|
||||
|
||||
Returns:
|
||||
Formatted string containing all memories with their conditions and content.
|
||||
"""
|
||||
formatted_memories = []
|
||||
|
||||
for i, memory in enumerate(memories, 1):
|
||||
|
|
@ -105,7 +146,14 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _extract_context(messages: List[Message]) -> str:
|
||||
"""Extract relevant context from messages"""
|
||||
"""Extract relevant context from messages.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages.
|
||||
|
||||
Returns:
|
||||
Formatted string containing recent conversation context.
|
||||
"""
|
||||
if not messages:
|
||||
return ""
|
||||
|
||||
|
|
@ -125,10 +173,18 @@ class RewriteMemoryOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _parse_json_response(response: str, key: str) -> str:
|
||||
"""Parse JSON response to extract specific key"""
|
||||
"""Parse JSON response to extract specific key.
|
||||
|
||||
Args:
|
||||
response: The response string that may contain JSON.
|
||||
key: The key to extract from the JSON object.
|
||||
|
||||
Returns:
|
||||
The value associated with the key, or the response string if parsing fails.
|
||||
"""
|
||||
try:
|
||||
# Try to extract JSON blocks
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_pattern = r"```json\s*([\s\S]*?)\s*```"
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
|
|
|
|||
|
|
@ -1,23 +1,23 @@
|
|||
memory_rewrite_prompt: |
|
||||
You are an expert AI assistant tasked with rewriting and reorganizing context content to make it more relevant and actionable for the current task.
|
||||
|
||||
|
||||
Your task is to take the original context (containing multiple experiences) and rewrite it as a cohesive, task-specific guidance that directly addresses the current situation.
|
||||
|
||||
|
||||
REWRITING GUIDELINES:
|
||||
● RELEVANCE FOCUS: Emphasize the most relevant aspects of each experience. Prioritize the most relevant experiences. Use clear, direct language.
|
||||
● ACTIONABLE INSIGHTS: Extract specific, actionable guidance. Make the context immediately actionable
|
||||
● COHERENT NARRATIVE: Create a flowing narrative rather than disconnected tips
|
||||
● SITUATIONAL AWARENESS: Adapt the guidance to the current situation
|
||||
|
||||
|
||||
# Current Task/Query
|
||||
{current_query}
|
||||
|
||||
|
||||
# Current Trajectory
|
||||
{current_context}
|
||||
|
||||
|
||||
# Original Context Content (Multiple Experiences)
|
||||
{original_context}
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Provide the rewritten context:
|
||||
```json
|
||||
|
|
@ -25,7 +25,7 @@ memory_rewrite_prompt: |
|
|||
"rewritten_context": "A cohesive, task-specific context message that reorganizes and adapts the original experiences for the current task. This should be written as a unified guidance rather than separate experience items.",
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
Guidelines:
|
||||
- Rewrite as a unified, flowing guidance
|
||||
- Adapt terminology and examples to match the current task domain
|
||||
|
|
|
|||
|
|
@ -1 +1,11 @@
|
|||
"""Tool memory retrieval operations module.
|
||||
|
||||
This module provides operations for retrieving tool memories from a vector store
|
||||
based on tool names, including formatting and matching tool memory results.
|
||||
"""
|
||||
|
||||
from .retrieve_tool_memory_op import RetrieveToolMemoryOp
|
||||
|
||||
__all__ = [
|
||||
"RetrieveToolMemoryOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,7 +1,15 @@
|
|||
"""Tool memory retrieval operation module.
|
||||
|
||||
This module provides functionality to retrieve tool memories from a vector store
|
||||
based on tool names, format them into structured documents, and match them with
|
||||
the requested tools.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import VectorNode
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import ToolMemory, vector_node_to_memory
|
||||
|
|
@ -9,26 +17,51 @@ from reme_ai.schema.memory import ToolMemory, vector_node_to_memory
|
|||
|
||||
@C.register_op()
|
||||
class RetrieveToolMemoryOp(BaseAsyncOp):
|
||||
"""Retrieves tool memories from vector store based on tool names.
|
||||
|
||||
This operation searches for tool memories in the vector store using tool names,
|
||||
validates that the retrieved memories match the requested tools, and formats
|
||||
them into a structured document format for use in the context.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
@staticmethod
|
||||
def _format_tool_memories(memories: List[ToolMemory]) -> str:
|
||||
"""Format tool memories into a structured document format.
|
||||
|
||||
def _format_tool_memories(self, memories: List[ToolMemory]) -> str:
|
||||
"""Format tool memories into a structured document format"""
|
||||
lines = []
|
||||
lines.append(f"Retrieved {len(memories)} tool memory(ies):\n")
|
||||
Args:
|
||||
memories: List of ToolMemory objects to format.
|
||||
|
||||
Returns:
|
||||
A formatted string containing all tool memories with separators.
|
||||
"""
|
||||
lines = [f"Retrieved {len(memories)} tool memory(ies):\n"]
|
||||
|
||||
for idx, memory in enumerate(memories, 1):
|
||||
lines.append(f"Tool: {memory.when_to_use}")
|
||||
lines.append(memory.content)
|
||||
|
||||
|
||||
if idx < len(memories):
|
||||
lines.append("\n---\n")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the tool memory retrieval operation.
|
||||
|
||||
This method:
|
||||
1. Extracts tool names from context
|
||||
2. Searches for each tool in the vector store
|
||||
3. Validates that retrieved memories match the requested tools
|
||||
4. Formats the memories into a structured document
|
||||
5. Stores the results in context response
|
||||
|
||||
The operation expects 'tool_names' in the context, which should be a
|
||||
comma-separated string of tool names. For each tool name, it retrieves
|
||||
the top matching memory from the vector store and validates that it
|
||||
matches exactly.
|
||||
"""
|
||||
tool_names: str = self.context.get("tool_names", "")
|
||||
workspace_id: str = self.context.workspace_id
|
||||
|
||||
|
|
@ -49,7 +82,7 @@ class RetrieveToolMemoryOp(BaseAsyncOp):
|
|||
nodes: List[VectorNode] = await self.vector_store.async_search(
|
||||
query=tool_name,
|
||||
workspace_id=workspace_id,
|
||||
top_k=1
|
||||
top_k=1,
|
||||
)
|
||||
|
||||
if nodes:
|
||||
|
|
@ -59,9 +92,11 @@ class RetrieveToolMemoryOp(BaseAsyncOp):
|
|||
# Ensure it's a ToolMemory and when_to_use matches
|
||||
if isinstance(memory, ToolMemory) and memory.when_to_use == tool_name:
|
||||
matched_tool_memories.append(memory)
|
||||
logger.info(f"Found tool_memory for tool_name={tool_name}, "
|
||||
f"memory_id={memory.memory_id}, "
|
||||
f"total_calls={len(memory.tool_call_results)}")
|
||||
logger.info(
|
||||
f"Found tool_memory for tool_name={tool_name}, "
|
||||
f"memory_id={memory.memory_id}, "
|
||||
f"total_calls={len(memory.tool_call_results)}",
|
||||
)
|
||||
else:
|
||||
logger.warning(f"No exact match found for tool_name={tool_name}")
|
||||
else:
|
||||
|
|
@ -83,6 +118,8 @@ class RetrieveToolMemoryOp(BaseAsyncOp):
|
|||
|
||||
# Log retrieval results
|
||||
for memory in matched_tool_memories:
|
||||
logger.info(f"Retrieved tool: {memory.when_to_use}, "
|
||||
f"total_calls={len(memory.tool_call_results)}, "
|
||||
f"content_length={len(memory.content)}")
|
||||
logger.info(
|
||||
f"Retrieved tool: {memory.when_to_use}, "
|
||||
f"total_calls={len(memory.tool_call_results)}, "
|
||||
f"content_length={len(memory.content)}",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1 +1,34 @@
|
|||
from flowllm.schema.message import Message, Role, Trajectory # noqa
|
||||
"""Schema module for ReMe.
|
||||
|
||||
This module provides data structures and schemas for memory management,
|
||||
including memory types, tool call results, and conversion utilities.
|
||||
"""
|
||||
|
||||
from flowllm.core.enumeration import Role # noqa
|
||||
from flowllm.core.schema import Message, Trajectory # noqa
|
||||
|
||||
from reme_ai.schema.memory import (
|
||||
BaseMemory,
|
||||
PersonalMemory,
|
||||
TaskMemory,
|
||||
ToolCallResult,
|
||||
ToolMemory,
|
||||
dict_to_memory,
|
||||
vector_node_to_memory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# FlowLLM schema imports
|
||||
"Message",
|
||||
"Role",
|
||||
"Trajectory",
|
||||
# Memory classes
|
||||
"BaseMemory",
|
||||
"TaskMemory",
|
||||
"PersonalMemory",
|
||||
"ToolMemory",
|
||||
"ToolCallResult",
|
||||
# Utility functions
|
||||
"vector_node_to_memory",
|
||||
"dict_to_memory",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,10 @@
|
|||
"""Memory schema definitions for ReMe.
|
||||
|
||||
This module defines the core memory data structures used in the ReMe system,
|
||||
including base memory classes and specialized memory types for tasks, personal
|
||||
information, and tool call results.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
|
|
@ -5,12 +12,31 @@ from abc import ABC
|
|||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
from flowllm.core.schema import VectorNode
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class BaseMemory(BaseModel, ABC):
|
||||
"""Base class for all memory types in the ReMe system.
|
||||
|
||||
This abstract base class provides common fields and methods for all memory
|
||||
types, including workspace identification, content storage, timestamps,
|
||||
and conversion to/from vector nodes for storage and retrieval.
|
||||
|
||||
Attributes:
|
||||
workspace_id: Identifier for the workspace this memory belongs to.
|
||||
memory_id: Unique identifier for this memory instance.
|
||||
memory_type: Type of memory (task, personal, tool, etc.).
|
||||
when_to_use: Description of when this memory should be retrieved.
|
||||
content: The actual content of the memory (string or bytes).
|
||||
score: Relevance score for this memory (0.0 to 1.0).
|
||||
time_created: Timestamp when the memory was created.
|
||||
time_modified: Timestamp when the memory was last modified.
|
||||
author: Identifier of the entity that created this memory.
|
||||
metadata: Additional metadata dictionary for extensibility.
|
||||
"""
|
||||
|
||||
workspace_id: str = Field(default="")
|
||||
memory_id: str = Field(default_factory=lambda: uuid4().hex)
|
||||
memory_type: str = Field(default=...)
|
||||
|
|
@ -26,90 +52,195 @@ class BaseMemory(BaseModel, ABC):
|
|||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
def update_modified_time(self):
|
||||
"""Update the time_modified field to the current timestamp."""
|
||||
self.time_modified = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
def update_metadata(self, new_metadata):
|
||||
"""Update the metadata dictionary with new values.
|
||||
|
||||
Args:
|
||||
new_metadata: Dictionary containing new metadata to replace existing metadata.
|
||||
"""
|
||||
self.metadata = new_metadata
|
||||
|
||||
def to_vector_node(self) -> VectorNode:
|
||||
"""Convert this memory instance to a VectorNode for storage.
|
||||
|
||||
Returns:
|
||||
VectorNode: A vector node representation of this memory.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def from_vector_node(cls, node: VectorNode):
|
||||
"""Create a memory instance from a VectorNode.
|
||||
|
||||
Args:
|
||||
node: VectorNode containing memory data.
|
||||
|
||||
Returns:
|
||||
BaseMemory: A memory instance reconstructed from the vector node.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TaskMemory(BaseMemory):
|
||||
"""Memory type for storing task-related information.
|
||||
|
||||
TaskMemory is used to store information about tasks, including when to use
|
||||
the memory and the task content itself. It extends BaseMemory with
|
||||
task-specific behavior.
|
||||
|
||||
Attributes:
|
||||
memory_type: Always set to "task" for task memories.
|
||||
"""
|
||||
|
||||
memory_type: str = Field(default="task")
|
||||
|
||||
def to_vector_node(self) -> VectorNode:
|
||||
return VectorNode(unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"metadata": self.metadata,
|
||||
})
|
||||
"""Convert this TaskMemory to a VectorNode.
|
||||
|
||||
Returns:
|
||||
VectorNode: Vector node representation with when_to_use as content
|
||||
and all other fields stored in metadata.
|
||||
"""
|
||||
return VectorNode(
|
||||
unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"metadata": self.metadata,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_vector_node(cls, node: VectorNode) -> "TaskMemory":
|
||||
"""Create a TaskMemory instance from a VectorNode.
|
||||
|
||||
Args:
|
||||
node: VectorNode containing task memory data.
|
||||
|
||||
Returns:
|
||||
TaskMemory: Reconstructed TaskMemory instance.
|
||||
"""
|
||||
metadata = node.metadata.copy()
|
||||
return cls(workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
when_to_use=node.content,
|
||||
content=metadata.pop("content"),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
metadata=metadata.pop("metadata", {}))
|
||||
return cls(
|
||||
workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
when_to_use=node.content,
|
||||
content=metadata.pop("content"),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
metadata=metadata.pop("metadata", {}),
|
||||
)
|
||||
|
||||
|
||||
class PersonalMemory(BaseMemory):
|
||||
"""Memory type for storing personal information and user preferences.
|
||||
|
||||
PersonalMemory extends BaseMemory with fields specific to personal data,
|
||||
including target information and reflection subject attributes. This is
|
||||
used for storing user preferences, personal insights, and reflection data.
|
||||
|
||||
Attributes:
|
||||
memory_type: Always set to "personal" for personal memories.
|
||||
target: Target identifier or category for this personal memory.
|
||||
reflection_subject: Subject of reflection for storing reflection attributes.
|
||||
"""
|
||||
|
||||
memory_type: str = Field(default="personal")
|
||||
target: str = Field(default="")
|
||||
reflection_subject: str = Field(default="") # For storing reflection subject attributes
|
||||
|
||||
def to_vector_node(self) -> VectorNode:
|
||||
return VectorNode(unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"target": self.target,
|
||||
"reflection_subject": self.reflection_subject,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"metadata": self.metadata,
|
||||
})
|
||||
"""Convert this PersonalMemory to a VectorNode.
|
||||
|
||||
Returns:
|
||||
VectorNode: Vector node representation with when_to_use as content
|
||||
and all other fields including target and reflection_subject
|
||||
stored in metadata.
|
||||
"""
|
||||
return VectorNode(
|
||||
unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"target": self.target,
|
||||
"reflection_subject": self.reflection_subject,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"metadata": self.metadata,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_vector_node(cls, node: VectorNode) -> "PersonalMemory":
|
||||
"""Create a PersonalMemory instance from a VectorNode.
|
||||
|
||||
Args:
|
||||
node: VectorNode containing personal memory data.
|
||||
|
||||
Returns:
|
||||
PersonalMemory: Reconstructed PersonalMemory instance.
|
||||
"""
|
||||
metadata = node.metadata.copy()
|
||||
return cls(workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
when_to_use=node.content,
|
||||
content=metadata.pop("content"),
|
||||
target=metadata.pop("target", ""),
|
||||
reflection_subject=metadata.pop("reflection_subject", ""),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
metadata=metadata.pop("metadata", {}))
|
||||
return cls(
|
||||
workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
when_to_use=node.content,
|
||||
content=metadata.pop("content"),
|
||||
target=metadata.pop("target", ""),
|
||||
reflection_subject=metadata.pop("reflection_subject", ""),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
metadata=metadata.pop("metadata", {}),
|
||||
)
|
||||
|
||||
|
||||
class ToolCallResult(BaseModel):
|
||||
"""Represents the result of a tool invocation.
|
||||
|
||||
This class stores comprehensive information about a tool call, including
|
||||
inputs, outputs, performance metrics, evaluation, and deduplication hash.
|
||||
|
||||
Attributes:
|
||||
create_time: Timestamp when the tool was invoked.
|
||||
tool_name: Name of the tool that was called.
|
||||
input: Input parameters passed to the tool (dict or string).
|
||||
output: Output result from the tool execution.
|
||||
token_cost: Number of tokens consumed by the tool call (-1 if unknown).
|
||||
success: Whether the tool invocation completed successfully.
|
||||
time_cost: Time taken for the tool invocation in seconds.
|
||||
summary: Brief summary of the tool call result.
|
||||
evaluation: Detailed evaluation of the tool invocation.
|
||||
score: Quality score from 0.0 (failure) to 1.0 (complete success).
|
||||
is_summarized: Whether this tool call has been included in a summary.
|
||||
call_hash: MD5 hash of input and output for deduplication.
|
||||
metadata: Additional metadata dictionary.
|
||||
"""
|
||||
|
||||
create_time: str = Field(default="", description="Time of tool invocation")
|
||||
tool_name: str = Field(default=..., description="Name of the tool")
|
||||
input: dict | str = Field(default="", description="Tool input")
|
||||
|
|
@ -126,24 +257,38 @@ class ToolCallResult(BaseModel):
|
|||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
def generate_hash(self) -> str:
|
||||
"""Generate hash value from tool input and output for deduplication"""
|
||||
"""Generate hash value from tool input and output for deduplication.
|
||||
|
||||
Creates an MD5 hash from the combined input and output strings.
|
||||
This hash is used to identify duplicate tool calls.
|
||||
|
||||
Returns:
|
||||
str: MD5 hash hexdigest of the combined input and output.
|
||||
"""
|
||||
# Convert input to string if it's a dict
|
||||
input_str = json.dumps(self.input, sort_keys=True) if isinstance(self.input, dict) else str(self.input)
|
||||
|
||||
|
||||
# Combine input and output
|
||||
combined = f"{input_str}|{self.output}"
|
||||
|
||||
|
||||
# Generate MD5 hash
|
||||
hash_value = hashlib.md5(combined.encode('utf-8')).hexdigest()
|
||||
|
||||
hash_value = hashlib.md5(combined.encode("utf-8")).hexdigest()
|
||||
|
||||
return hash_value
|
||||
|
||||
|
||||
def ensure_hash(self):
|
||||
"""Ensure call_hash is set, generate if empty"""
|
||||
"""Ensure call_hash is set, generate if empty."""
|
||||
if not self.call_hash:
|
||||
self.call_hash = self.generate_hash()
|
||||
|
||||
def from_mcp_tool_result(self, tool_result: CallToolResult, max_char_len: int = None):
|
||||
"""Populate this instance from an MCP CallToolResult.
|
||||
|
||||
Args:
|
||||
tool_result: MCP CallToolResult to extract data from.
|
||||
max_char_len: Optional maximum character length for output content.
|
||||
If provided, output will be truncated to this length.
|
||||
"""
|
||||
text_list = []
|
||||
for content in tool_result.content:
|
||||
if isinstance(content, TextContent):
|
||||
|
|
@ -162,28 +307,60 @@ class ToolCallResult(BaseModel):
|
|||
|
||||
|
||||
class ToolMemory(BaseMemory):
|
||||
"""Memory type for storing tool call execution history.
|
||||
|
||||
ToolMemory extends BaseMemory to store a collection of tool call results,
|
||||
allowing tracking of tool usage patterns, performance metrics, and
|
||||
execution history for analysis and summarization.
|
||||
|
||||
Attributes:
|
||||
memory_type: Always set to "tool" for tool memories.
|
||||
tool_call_results: List of ToolCallResult instances representing
|
||||
historical tool invocations.
|
||||
"""
|
||||
|
||||
memory_type: str = Field(default="tool")
|
||||
tool_call_results: List[ToolCallResult] = Field(default_factory=list)
|
||||
|
||||
def to_vector_node(self) -> VectorNode:
|
||||
return VectorNode(unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"tool_call_results": [x.model_dump() for x in self.tool_call_results],
|
||||
"metadata": self.metadata,
|
||||
})
|
||||
"""Convert this ToolMemory to a VectorNode.
|
||||
|
||||
Returns:
|
||||
VectorNode: Vector node representation with when_to_use as content
|
||||
and all tool_call_results serialized in metadata.
|
||||
"""
|
||||
return VectorNode(
|
||||
unique_id=self.memory_id,
|
||||
workspace_id=self.workspace_id,
|
||||
content=self.when_to_use,
|
||||
metadata={
|
||||
"memory_type": self.memory_type,
|
||||
"content": self.content,
|
||||
"score": self.score,
|
||||
"time_created": self.time_created,
|
||||
"time_modified": self.time_modified,
|
||||
"author": self.author,
|
||||
"tool_call_results": [x.model_dump() for x in self.tool_call_results],
|
||||
"metadata": self.metadata,
|
||||
},
|
||||
)
|
||||
|
||||
def statistic(self, recent_frequency: int = 20) -> dict:
|
||||
"""
|
||||
Calculate statistical information for the most recent N tool calls.
|
||||
Returns avg token_cost, success rate, avg time_cost, and avg score.
|
||||
"""Calculate statistical information for the most recent N tool calls.
|
||||
|
||||
Analyzes the most recent tool calls and computes average metrics including
|
||||
token cost, success rate, time cost, and quality scores.
|
||||
|
||||
Args:
|
||||
recent_frequency: Number of most recent tool calls to analyze.
|
||||
Defaults to 20.
|
||||
|
||||
Returns:
|
||||
dict: Dictionary containing:
|
||||
- avg_token_cost: Average token consumption (rounded to 2 decimals)
|
||||
- avg_time_cost: Average execution time in seconds (rounded to 3 decimals)
|
||||
- success_rate: Ratio of successful calls (rounded to 4 decimals)
|
||||
- avg_score: Average quality score (rounded to 3 decimals)
|
||||
"""
|
||||
if not self.tool_call_results:
|
||||
return {
|
||||
|
|
@ -192,54 +369,80 @@ class ToolMemory(BaseMemory):
|
|||
"avg_token_cost": 0.0,
|
||||
"success_rate": 0.0,
|
||||
"avg_time_cost": 0.0,
|
||||
"avg_score": 0.0
|
||||
"avg_score": 0.0,
|
||||
}
|
||||
|
||||
|
||||
# Get the most recent N tool calls (or all if less than N)
|
||||
recent_calls = self.tool_call_results[-recent_frequency:]
|
||||
total_calls = len(self.tool_call_results)
|
||||
# total_calls = len(self.tool_call_results)
|
||||
recent_calls_count = len(recent_calls)
|
||||
|
||||
|
||||
# Calculate statistics
|
||||
total_token_cost = sum(call.token_cost for call in recent_calls if call.token_cost >= 0)
|
||||
valid_token_calls = [call for call in recent_calls if call.token_cost >= 0]
|
||||
avg_token_cost = total_token_cost / len(valid_token_calls) if valid_token_calls else 0.0
|
||||
|
||||
|
||||
successful_calls = sum(1 for call in recent_calls if call.success)
|
||||
success_rate = successful_calls / recent_calls_count if recent_calls_count > 0 else 0.0
|
||||
|
||||
|
||||
total_time_cost = sum(call.time_cost for call in recent_calls)
|
||||
avg_time_cost = total_time_cost / recent_calls_count if recent_calls_count > 0 else 0.0
|
||||
|
||||
|
||||
total_score = sum(call.score for call in recent_calls)
|
||||
avg_score = total_score / recent_calls_count if recent_calls_count > 0 else 0.0
|
||||
|
||||
|
||||
return {
|
||||
"avg_token_cost": round(avg_token_cost, 2),
|
||||
"avg_time_cost": round(avg_time_cost, 3),
|
||||
"success_rate": round(success_rate, 4),
|
||||
"avg_score": round(avg_score, 3)
|
||||
"avg_score": round(avg_score, 3),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_vector_node(cls, node: VectorNode) -> "ToolMemory":
|
||||
"""Create a ToolMemory instance from a VectorNode.
|
||||
|
||||
Args:
|
||||
node: VectorNode containing tool memory data.
|
||||
|
||||
Returns:
|
||||
ToolMemory: Reconstructed ToolMemory instance with tool_call_results
|
||||
deserialized from metadata.
|
||||
"""
|
||||
metadata = node.metadata.copy()
|
||||
tool_call_results = [ToolCallResult(**result) for result in metadata.pop("tool_call_results", [])]
|
||||
return cls(workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
when_to_use=node.content,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
content=metadata.pop("content"),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
tool_call_results=tool_call_results,
|
||||
metadata=metadata.pop("metadata", {}))
|
||||
|
||||
return cls(
|
||||
workspace_id=node.workspace_id,
|
||||
memory_id=node.unique_id,
|
||||
when_to_use=node.content,
|
||||
memory_type=metadata.pop("memory_type"),
|
||||
content=metadata.pop("content"),
|
||||
score=metadata.pop("score"),
|
||||
time_created=metadata.pop("time_created"),
|
||||
time_modified=metadata.pop("time_modified"),
|
||||
author=metadata.pop("author"),
|
||||
tool_call_results=tool_call_results,
|
||||
metadata=metadata.pop("metadata", {}),
|
||||
)
|
||||
|
||||
|
||||
def vector_node_to_memory(node: VectorNode):
|
||||
"""Convert a VectorNode to the appropriate memory type.
|
||||
|
||||
This function inspects the memory_type in the node's metadata and
|
||||
reconstructs the appropriate memory subclass (TaskMemory, PersonalMemory,
|
||||
or ToolMemory).
|
||||
|
||||
Args:
|
||||
node: VectorNode containing memory data with memory_type in metadata.
|
||||
|
||||
Returns:
|
||||
BaseMemory: Instance of the appropriate memory subclass based on
|
||||
memory_type.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If memory_type is not recognized or not present.
|
||||
"""
|
||||
memory_type = node.metadata.get("memory_type")
|
||||
if memory_type == "task":
|
||||
return TaskMemory.from_vector_node(node)
|
||||
|
|
@ -255,6 +458,23 @@ def vector_node_to_memory(node: VectorNode):
|
|||
|
||||
|
||||
def dict_to_memory(memory_dict: dict):
|
||||
"""Create a memory instance from a dictionary.
|
||||
|
||||
This function creates the appropriate memory subclass based on the
|
||||
memory_type field in the dictionary. Defaults to TaskMemory if
|
||||
memory_type is not specified.
|
||||
|
||||
Args:
|
||||
memory_dict: Dictionary containing memory data with optional
|
||||
memory_type field.
|
||||
|
||||
Returns:
|
||||
BaseMemory: Instance of the appropriate memory subclass based on
|
||||
memory_type.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If memory_type is not recognized.
|
||||
"""
|
||||
memory_type = memory_dict.get("memory_type", "task")
|
||||
if memory_type == "task":
|
||||
return TaskMemory(**memory_dict)
|
||||
|
|
@ -270,13 +490,15 @@ def dict_to_memory(memory_dict: dict):
|
|||
|
||||
|
||||
def task_main():
|
||||
"""Test function for TaskMemory serialization and deserialization."""
|
||||
e1 = TaskMemory(
|
||||
workspace_id="w_1024",
|
||||
memory_id="123",
|
||||
when_to_use="test case use",
|
||||
content="test content",
|
||||
score=0.99,
|
||||
metadata={})
|
||||
metadata={},
|
||||
)
|
||||
print(e1.model_dump_json(indent=2))
|
||||
v1 = e1.to_vector_node()
|
||||
print(v1.model_dump_json(indent=2))
|
||||
|
|
@ -285,6 +507,7 @@ def task_main():
|
|||
|
||||
|
||||
def personal_main():
|
||||
"""Test function for PersonalMemory serialization and deserialization."""
|
||||
p1 = PersonalMemory(
|
||||
workspace_id="w_2048",
|
||||
memory_id="456",
|
||||
|
|
@ -293,7 +516,8 @@ def personal_main():
|
|||
target="user_preferences",
|
||||
reflection_subject="learning_style",
|
||||
score=0.85,
|
||||
metadata={"category": "user_profile"})
|
||||
metadata={"category": "user_profile"},
|
||||
)
|
||||
print("PersonalMemory test:")
|
||||
print(p1.model_dump_json(indent=2))
|
||||
v1 = p1.to_vector_node()
|
||||
|
|
@ -305,6 +529,7 @@ def personal_main():
|
|||
|
||||
|
||||
def tool_main():
|
||||
"""Test function for ToolMemory serialization and deserialization."""
|
||||
# Create sample tool call results
|
||||
tool_result1 = ToolCallResult(
|
||||
create_time="2025-10-15 10:30:00",
|
||||
|
|
@ -315,9 +540,9 @@ def tool_main():
|
|||
success=True,
|
||||
time_cost=0.5,
|
||||
evaluation="Successfully executed",
|
||||
score=0.95
|
||||
score=0.95,
|
||||
)
|
||||
|
||||
|
||||
tool_result2 = ToolCallResult(
|
||||
create_time="2025-10-15 10:31:00",
|
||||
tool_name="data_processor",
|
||||
|
|
@ -327,9 +552,9 @@ def tool_main():
|
|||
success=True,
|
||||
time_cost=1.2,
|
||||
evaluation="Good performance",
|
||||
score=0.88
|
||||
score=0.88,
|
||||
)
|
||||
|
||||
|
||||
t1 = ToolMemory(
|
||||
workspace_id="w_4096",
|
||||
memory_id="789",
|
||||
|
|
@ -338,8 +563,9 @@ def tool_main():
|
|||
content="tool execution test content",
|
||||
score=0.92,
|
||||
tool_call_results=[tool_result1, tool_result2],
|
||||
metadata={"execution_context": "test_environment"})
|
||||
|
||||
metadata={"execution_context": "test_environment"},
|
||||
)
|
||||
|
||||
print("ToolMemory test:")
|
||||
print(t1.model_dump_json(indent=2))
|
||||
v1 = t1.to_vector_node()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,15 @@
|
|||
"""Memory service modules for ReMe.
|
||||
|
||||
This package provides memory service implementations for managing different
|
||||
types of memories including task memories and personal memories.
|
||||
"""
|
||||
|
||||
from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeMemoryService
|
||||
from reme_ai.service.personal_memory_service import PersonalMemoryService
|
||||
from reme_ai.service.task_memory_service import TaskMemoryService
|
||||
|
||||
__all__ = [
|
||||
"AgentscopeRuntimeMemoryService",
|
||||
"PersonalMemoryService",
|
||||
"TaskMemoryService",
|
||||
]
|
||||
|
|
@ -1,18 +1,48 @@
|
|||
"""Base memory service for Agentscope runtime integration.
|
||||
|
||||
This module provides the abstract base class AgentscopeRuntimeMemoryService
|
||||
which defines the interface for memory services that integrate with
|
||||
Agentscope runtime. Concrete implementations should inherit from this class
|
||||
and implement the abstract methods.
|
||||
"""
|
||||
|
||||
from abc import abstractmethod, ABC
|
||||
from typing import Optional, Dict, Any
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from reme_ai.app import ReMeApp
|
||||
from reme_ai.main import ReMeApp
|
||||
|
||||
|
||||
class AgentscopeRuntimeMemoryService(ABC):
|
||||
"""Abstract base class for memory services integrated with Agentscope runtime.
|
||||
|
||||
This class provides a common interface for memory services and manages
|
||||
the underlying ReMeApp instance and session-to-memory-id mappings.
|
||||
Subclasses must implement the abstract methods to provide specific
|
||||
memory management functionality.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the memory service.
|
||||
|
||||
Creates a new ReMeApp instance and initializes the session-to-memory-id
|
||||
mapping dictionary.
|
||||
"""
|
||||
self.app = ReMeApp()
|
||||
self.session_id_dict: dict = {}
|
||||
|
||||
def add_session_memory_id(self, session_id: str, memory_id):
|
||||
"""Add a memory ID to a session's memory list.
|
||||
|
||||
Associates a memory_id with a session_id by adding it to the
|
||||
session's memory list. If the session doesn't exist, it will
|
||||
be created.
|
||||
|
||||
Args:
|
||||
session_id: The session identifier.
|
||||
memory_id: The memory identifier to associate with the session.
|
||||
"""
|
||||
if session_id not in self.session_id_dict:
|
||||
self.session_id_dict[session_id] = []
|
||||
|
||||
|
|
@ -37,21 +67,38 @@ class AgentscopeRuntimeMemoryService(ABC):
|
|||
"""
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry."""
|
||||
"""Async context manager entry.
|
||||
|
||||
Starts the service when entering an async context.
|
||||
|
||||
Returns:
|
||||
The service instance.
|
||||
"""
|
||||
await self.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Async context manager exit."""
|
||||
"""Async context manager exit.
|
||||
|
||||
Stops the service when exiting an async context.
|
||||
|
||||
Args:
|
||||
exc_type: Exception type if an exception occurred.
|
||||
exc_val: Exception value if an exception occurred.
|
||||
exc_tb: Exception traceback if an exception occurred.
|
||||
|
||||
Returns:
|
||||
False to propagate exceptions, True to suppress them.
|
||||
"""
|
||||
await self.stop()
|
||||
return False
|
||||
|
||||
@abstractmethod
|
||||
async def add_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
session_id: Optional[str] = None,
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
session_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Adds messages to the memory service.
|
||||
|
|
@ -64,14 +111,13 @@ class AgentscopeRuntimeMemoryService(ABC):
|
|||
|
||||
@abstractmethod
|
||||
async def search_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""
|
||||
Searches messages from the memory service.
|
||||
|
|
@ -86,13 +132,12 @@ class AgentscopeRuntimeMemoryService(ABC):
|
|||
|
||||
@abstractmethod
|
||||
async def list_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
self,
|
||||
user_id: str,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""
|
||||
Lists the memory items for a given user with filters, such as
|
||||
|
|
@ -105,9 +150,9 @@ class AgentscopeRuntimeMemoryService(ABC):
|
|||
|
||||
@abstractmethod
|
||||
async def delete_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
session_id: Optional[str] = None,
|
||||
self,
|
||||
user_id: str,
|
||||
session_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Deletes the memory items for a given user with certain session id,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,15 @@
|
|||
"""Personal memory service for managing personalized memories.
|
||||
|
||||
This module provides the PersonalMemoryService class which extends the base
|
||||
AgentscopeRuntimeMemoryService to handle personal memory operations.
|
||||
It supports creating, retrieving, listing, and deleting personal memories
|
||||
using flow-based execution.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional, Dict, Any, List
|
||||
|
||||
from flowllm.schema.flow_response import FlowResponse
|
||||
from flowllm.core.schema import FlowResponse
|
||||
from loguru import logger
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
|
@ -10,17 +18,49 @@ from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeM
|
|||
|
||||
|
||||
class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
||||
"""Service for managing personalized memories.
|
||||
|
||||
PersonalMemoryService empowers you to generate, retrieve, and share
|
||||
customized memories. Leveraging advanced LLM, embedding, and vector store
|
||||
technologies, it builds a comprehensive memory system with intelligent,
|
||||
context- and time-aware memory management.
|
||||
"""
|
||||
|
||||
async def start(self):
|
||||
"""Start the personal memory service.
|
||||
|
||||
Returns:
|
||||
The result of starting the underlying application.
|
||||
"""
|
||||
return await self.app.async_start()
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the personal memory service.
|
||||
|
||||
Releases resources and stops the underlying application.
|
||||
"""
|
||||
return await self.app.async_stop()
|
||||
|
||||
async def health(self) -> bool:
|
||||
"""Check the health status of the service.
|
||||
|
||||
Returns:
|
||||
True if the service is healthy, False otherwise.
|
||||
"""
|
||||
return True
|
||||
|
||||
async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None:
|
||||
"""Add personal memory from messages.
|
||||
|
||||
Processes the provided messages and creates personal memories using
|
||||
the summary_personal_memory flow. The created memories are associated
|
||||
with the given session_id.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
messages: List of messages (dict or BaseModel instances) to process.
|
||||
session_id: Optional session identifier to associate with the memory.
|
||||
"""
|
||||
new_messages: List[dict] = []
|
||||
for message in messages:
|
||||
if isinstance(message, dict):
|
||||
|
|
@ -33,8 +73,8 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
|||
kwargs = {
|
||||
"workspace_id": user_id,
|
||||
"trajectories": [
|
||||
{"messages": new_messages, "score": 1.0}
|
||||
]
|
||||
{"messages": new_messages, "score": 1.0},
|
||||
],
|
||||
}
|
||||
|
||||
result: FlowResponse = await self.app.async_execute_flow(name="summary_personal_memory", **kwargs)
|
||||
|
|
@ -44,11 +84,29 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
|||
self.add_session_memory_id(session_id, memory_id)
|
||||
logger.info(f"[personal_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}")
|
||||
|
||||
async def search_memory(self, user_id: str, messages: list, filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
)) -> list:
|
||||
async def search_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""Search for personal memories matching the given messages.
|
||||
|
||||
Searches the memory store for personal memories relevant to the provided
|
||||
messages using the retrieve_personal_memory flow. The query is extracted
|
||||
from the last message in the messages list.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
messages: List of messages (dict or BaseModel instances) to search with.
|
||||
filters: Optional filters including top_k for controlling search results.
|
||||
|
||||
Returns:
|
||||
List containing the search result answer.
|
||||
"""
|
||||
new_messages: List[dict] = []
|
||||
for message in messages:
|
||||
if isinstance(message, dict):
|
||||
|
|
@ -64,7 +122,7 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
|||
kwargs = {
|
||||
"workspace_id": user_id,
|
||||
"query": query,
|
||||
"top_k": filters.get("top_k", 1) if filters else 1
|
||||
"top_k": filters.get("top_k", 1) if filters else 1,
|
||||
}
|
||||
|
||||
result: FlowResponse = await self.app.async_execute_flow(name="retrieve_personal_memory", **kwargs)
|
||||
|
|
@ -72,11 +130,26 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
|||
|
||||
return [result.answer]
|
||||
|
||||
async def list_memory(self, user_id: str, filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
)) -> list:
|
||||
async def list_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""List all personal memories for a user.
|
||||
|
||||
Retrieves all personal memories associated with the given user_id
|
||||
from the vector store.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
filters: Optional filters (currently not used but kept for API consistency).
|
||||
|
||||
Returns:
|
||||
List of memory items for the user.
|
||||
"""
|
||||
result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list")
|
||||
logger.info(f"[personal_memory_service] list_memory result: {result}")
|
||||
|
||||
|
|
@ -86,29 +159,51 @@ class PersonalMemoryService(AgentscopeRuntimeMemoryService):
|
|||
return result
|
||||
|
||||
async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None:
|
||||
"""Delete personal memories for a user session.
|
||||
|
||||
Deletes all memories associated with the given session_id for the user.
|
||||
If no session_id is provided or no memories exist for the session,
|
||||
no deletion is performed.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
session_id: Optional session identifier. If provided, only memories
|
||||
associated with this session will be deleted.
|
||||
"""
|
||||
delete_ids = self.session_id_dict.get(session_id, [])
|
||||
if not delete_ids:
|
||||
return
|
||||
|
||||
result = await self.app.async_execute_flow(name="vector_store",
|
||||
workspace_id=user_id,
|
||||
action="delete_ids",
|
||||
memory_ids=delete_ids)
|
||||
result = await self.app.async_execute_flow(
|
||||
name="vector_store",
|
||||
workspace_id=user_id,
|
||||
action="delete_ids",
|
||||
memory_ids=delete_ids,
|
||||
)
|
||||
result = result.metadata["action_result"]
|
||||
logger.info(f"[personal_memory_service] delete memory result={result}")
|
||||
|
||||
|
||||
async def main():
|
||||
"""Main function for testing the PersonalMemoryService.
|
||||
|
||||
Demonstrates the usage of PersonalMemoryService by adding, searching,
|
||||
listing, and deleting personal memories.
|
||||
"""
|
||||
async with PersonalMemoryService() as service:
|
||||
logger.info("========== start personal memory service ==========")
|
||||
|
||||
await service.add_memory(user_id="u_12345",
|
||||
messages=[{"content": "I really enjoy playing tennis on weekends"}],
|
||||
session_id="s_123456")
|
||||
await service.add_memory(
|
||||
user_id="u_12345",
|
||||
messages=[{"content": "I really enjoy playing tennis on weekends"}],
|
||||
session_id="s_123456",
|
||||
)
|
||||
|
||||
await service.search_memory(user_id="u_12345",
|
||||
messages=[{"content": "What do I like to do for fun?"}],
|
||||
filters={"top_k": 1})
|
||||
await service.search_memory(
|
||||
user_id="u_12345",
|
||||
messages=[{"content": "What do I like to do for fun?"}],
|
||||
filters={"top_k": 1},
|
||||
)
|
||||
|
||||
await service.list_memory(user_id="u_12345")
|
||||
await service.delete_memory(user_id="u_12345", session_id="s_123456")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,15 @@
|
|||
"""Task memory service for managing task-oriented memories.
|
||||
|
||||
This module provides the TaskMemoryService class which extends the base
|
||||
AgentscopeRuntimeMemoryService to handle task-related memory operations.
|
||||
It supports creating, retrieving, listing, and deleting task memories
|
||||
using flow-based execution.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional, Dict, Any, List
|
||||
|
||||
from flowllm.schema.flow_response import FlowResponse
|
||||
from flowllm.core.schema import FlowResponse
|
||||
from loguru import logger
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
|
@ -10,17 +18,49 @@ from reme_ai.service.agentscope_runtime_memory_service import AgentscopeRuntimeM
|
|||
|
||||
|
||||
class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
||||
"""Service for managing task-oriented memories.
|
||||
|
||||
TaskMemoryService helps efficiently manage and schedule task-related memories,
|
||||
enhancing both the accuracy and efficiency of task execution. Powered by LLM
|
||||
capabilities, it supports flexible creation, retrieval, update, and deletion
|
||||
of memories across diverse task scenarios.
|
||||
"""
|
||||
|
||||
async def start(self):
|
||||
"""Start the task memory service.
|
||||
|
||||
Returns:
|
||||
The result of starting the underlying application.
|
||||
"""
|
||||
return await self.app.async_start()
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the task memory service.
|
||||
|
||||
Releases resources and stops the underlying application.
|
||||
"""
|
||||
return await self.app.async_stop()
|
||||
|
||||
async def health(self) -> bool:
|
||||
"""Check the health status of the service.
|
||||
|
||||
Returns:
|
||||
True if the service is healthy, False otherwise.
|
||||
"""
|
||||
return True
|
||||
|
||||
async def add_memory(self, user_id: str, messages: list, session_id: Optional[str] = None) -> None:
|
||||
"""Add task memory from messages.
|
||||
|
||||
Processes the provided messages and creates task memories using
|
||||
the summary_task_memory flow. The created memories are associated
|
||||
with the given session_id.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
messages: List of messages (dict or BaseModel instances) to process.
|
||||
session_id: Optional session identifier to associate with the memory.
|
||||
"""
|
||||
new_messages: List[dict] = []
|
||||
for message in messages:
|
||||
if isinstance(message, dict):
|
||||
|
|
@ -33,8 +73,8 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
|||
kwargs = {
|
||||
"workspace_id": user_id,
|
||||
"trajectories": [
|
||||
{"messages": new_messages, "score": 1.0}
|
||||
]
|
||||
{"messages": new_messages, "score": 1.0},
|
||||
],
|
||||
}
|
||||
|
||||
result: FlowResponse = await self.app.async_execute_flow(name="summary_task_memory", **kwargs)
|
||||
|
|
@ -44,11 +84,28 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
|||
self.add_session_memory_id(session_id, memory_id)
|
||||
logger.info(f"[task_memory_service] user_id={user_id} session_id={session_id} add memory: {memory}")
|
||||
|
||||
async def search_memory(self, user_id: str, messages: list, filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
)) -> list:
|
||||
async def search_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""Search for task memories matching the given messages.
|
||||
|
||||
Searches the memory store for task memories relevant to the provided
|
||||
messages using the retrieve_task_memory flow.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
messages: List of messages (dict or BaseModel instances) to search with.
|
||||
filters: Optional filters including top_k for controlling search results.
|
||||
|
||||
Returns:
|
||||
List containing the search result answer.
|
||||
"""
|
||||
new_messages: List[dict] = []
|
||||
for message in messages:
|
||||
if isinstance(message, dict):
|
||||
|
|
@ -61,7 +118,7 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
|||
kwargs = {
|
||||
"workspace_id": user_id,
|
||||
"messages": new_messages,
|
||||
"top_k": filters.get("top_k", 1) if filters else 1
|
||||
"top_k": filters.get("top_k", 1) if filters else 1,
|
||||
}
|
||||
|
||||
result: FlowResponse = await self.app.async_execute_flow(name="retrieve_task_memory", **kwargs)
|
||||
|
|
@ -69,11 +126,26 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
|||
|
||||
return [result.answer]
|
||||
|
||||
async def list_memory(self, user_id: str, filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, "
|
||||
"such as top_k, score etc.",
|
||||
default=None,
|
||||
)) -> list:
|
||||
async def list_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
filters: Optional[Dict[str, Any]] = Field(
|
||||
description="Associated filters for the messages, " "such as top_k, score etc.",
|
||||
default=None,
|
||||
),
|
||||
) -> list:
|
||||
"""List all task memories for a user.
|
||||
|
||||
Retrieves all task memories associated with the given user_id
|
||||
from the vector store.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
filters: Optional filters (currently not used but kept for API consistency).
|
||||
|
||||
Returns:
|
||||
List of memory items for the user.
|
||||
"""
|
||||
result = await self.app.async_execute_flow(name="vector_store", workspace_id=user_id, action="list")
|
||||
print("list_memory result:", result)
|
||||
|
||||
|
|
@ -83,29 +155,51 @@ class TaskMemoryService(AgentscopeRuntimeMemoryService):
|
|||
return result
|
||||
|
||||
async def delete_memory(self, user_id: str, session_id: Optional[str] = None) -> None:
|
||||
"""Delete task memories for a user session.
|
||||
|
||||
Deletes all memories associated with the given session_id for the user.
|
||||
If no session_id is provided or no memories exist for the session,
|
||||
no deletion is performed.
|
||||
|
||||
Args:
|
||||
user_id: The user identifier.
|
||||
session_id: Optional session identifier. If provided, only memories
|
||||
associated with this session will be deleted.
|
||||
"""
|
||||
delete_ids = self.session_id_dict.get(session_id, [])
|
||||
if not delete_ids:
|
||||
return
|
||||
|
||||
result = await self.app.async_execute_flow(name="vector_store",
|
||||
workspace_id=user_id,
|
||||
action="delete_ids",
|
||||
memory_ids=delete_ids)
|
||||
result = await self.app.async_execute_flow(
|
||||
name="vector_store",
|
||||
workspace_id=user_id,
|
||||
action="delete_ids",
|
||||
memory_ids=delete_ids,
|
||||
)
|
||||
result = result.metadata["action_result"]
|
||||
logger.info(f"[task_memory_service] delete memory result={result}")
|
||||
|
||||
|
||||
async def main():
|
||||
"""Main function for testing the TaskMemoryService.
|
||||
|
||||
Demonstrates the usage of TaskMemoryService by adding, searching,
|
||||
listing, and deleting task memories.
|
||||
"""
|
||||
async with TaskMemoryService() as service:
|
||||
logger.info("========== start task memory service ==========")
|
||||
|
||||
await service.add_memory(user_id="u_123456",
|
||||
messages=[{"content": "please use web search tool to search financial news:"}],
|
||||
session_id="s_123456")
|
||||
await service.add_memory(
|
||||
user_id="u_123456",
|
||||
messages=[{"content": "please use web search tool to search financial news:"}],
|
||||
session_id="s_123456",
|
||||
)
|
||||
|
||||
await service.search_memory(user_id="u_123456",
|
||||
messages=[{"content": "please use web search tool to search financial news"}],
|
||||
filters={"top_k": 1})
|
||||
await service.search_memory(
|
||||
user_id="u_123456",
|
||||
messages=[{"content": "please use web search tool to search financial news"}],
|
||||
filters={"top_k": 1},
|
||||
)
|
||||
|
||||
await service.list_memory(user_id="u_123456")
|
||||
await service.delete_memory(user_id="u_123456", session_id="s_123456")
|
||||
|
|
|
|||
|
|
@ -1,3 +1,17 @@
|
|||
"""Summary operations module.
|
||||
|
||||
This module provides summary operations for different types of memories:
|
||||
- Personal memory summary operations
|
||||
- Task memory summary operations
|
||||
- Tool memory summary operations
|
||||
"""
|
||||
|
||||
from . import personal
|
||||
from . import task
|
||||
from . import tool
|
||||
|
||||
__all__ = [
|
||||
"personal",
|
||||
"task",
|
||||
"tool",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,8 +1,30 @@
|
|||
from .contra_repeat_op import ContraRepeatOp
|
||||
from .get_observation_op import GetObservationOp
|
||||
from .get_observation_with_time_op import GetObservationWithTimeOp
|
||||
from .get_reflection_subject_op import GetReflectionSubjectOp
|
||||
from .info_filter_op import InfoFilterOp
|
||||
from .load_today_memory_op import LoadTodayMemoryOp
|
||||
from .long_contra_repeat_op import LongContraRepeatOp
|
||||
from .update_insight_op import UpdateInsightOp
|
||||
"""Personal memory summary operations module.
|
||||
|
||||
This module provides operations for processing personal memories, including:
|
||||
- Filtering messages based on information content
|
||||
- Extracting observations from chat messages
|
||||
- Generating reflection subjects
|
||||
- Updating insights based on new observations
|
||||
- Detecting and handling contradictions and redundancies
|
||||
- Loading today's memories for deduplication
|
||||
"""
|
||||
|
||||
from reme_ai.summary.personal.contra_repeat_op import ContraRepeatOp
|
||||
from reme_ai.summary.personal.get_observation_op import GetObservationOp
|
||||
from reme_ai.summary.personal.get_observation_with_time_op import GetObservationWithTimeOp
|
||||
from reme_ai.summary.personal.get_reflection_subject_op import GetReflectionSubjectOp
|
||||
from reme_ai.summary.personal.info_filter_op import InfoFilterOp
|
||||
from reme_ai.summary.personal.load_today_memory_op import LoadTodayMemoryOp
|
||||
from reme_ai.summary.personal.long_contra_repeat_op import LongContraRepeatOp
|
||||
from reme_ai.summary.personal.update_insight_op import UpdateInsightOp
|
||||
|
||||
__all__ = [
|
||||
"ContraRepeatOp",
|
||||
"GetObservationOp",
|
||||
"GetObservationWithTimeOp",
|
||||
"GetReflectionSubjectOp",
|
||||
"InfoFilterOp",
|
||||
"LoadTodayMemoryOp",
|
||||
"LongContraRepeatOp",
|
||||
"UpdateInsightOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,10 +1,20 @@
|
|||
"""Module for detecting and handling contradictory and repetitive memories.
|
||||
|
||||
This module provides the ContraRepeatOp class which processes memory nodes
|
||||
to identify and handle contradictory and repetitive information. It collects
|
||||
observation memories from context, constructs prompts for language model
|
||||
analysis, parses responses to detect contradictions or redundancies, and
|
||||
filters the processed memories accordingly.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import List, Tuple
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -22,6 +32,7 @@ class ContraRepeatOp(BaseAsyncOp):
|
|||
- Parses the model's response to detect contradictions or redundancies.
|
||||
- Filters and returns the processed memories.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -71,12 +82,16 @@ class ContraRepeatOp(BaseAsyncOp):
|
|||
user_name = self.context.get("user_name", "user")
|
||||
|
||||
# Create prompt using the new pattern
|
||||
system_prompt = self.prompt_format(prompt_name="contra_repeat_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name)
|
||||
system_prompt = self.prompt_format(
|
||||
prompt_name="contra_repeat_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name,
|
||||
)
|
||||
few_shot = self.prompt_format(prompt_name="contra_repeat_few_shot", user_name=user_name)
|
||||
user_query = self.prompt_format(prompt_name="contra_repeat_user_query",
|
||||
user_query="\n".join(user_query_list))
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="contra_repeat_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
logger.info(f"contra_repeat_prompt={full_prompt}")
|
||||
|
|
@ -105,7 +120,9 @@ class ContraRepeatOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _parse_and_filter_memories(response_text: str, memories: List[BaseMemory]) -> Tuple[
|
||||
List[BaseMemory], List[str]]:
|
||||
List[BaseMemory],
|
||||
List[str],
|
||||
]:
|
||||
"""Parse LLM response and filter memories based on contradiction/containment analysis"""
|
||||
|
||||
# Parse the response to extract judgments
|
||||
|
|
@ -128,7 +145,7 @@ class ContraRepeatOp(BaseAsyncOp):
|
|||
continue
|
||||
|
||||
judgment_lower = judgment.lower()
|
||||
if judgment_lower in ['矛盾', 'contradiction', '被包含', 'contained']:
|
||||
if judgment_lower in ["矛盾", "contradiction", "被包含", "contained"]:
|
||||
indices_to_remove.add(idx)
|
||||
deleted_memory_ids.append(memories[idx].memory_id)
|
||||
logger.info(f"Marking memory {idx + 1} for removal: {judgment} - {memories[idx].content[:100]}...")
|
||||
|
|
|
|||
|
|
@ -1,8 +1,17 @@
|
|||
"""Module for generating observations from chat messages.
|
||||
|
||||
This module provides the GetObservationOp class which extracts personal
|
||||
observations from chat messages. It filters messages to exclude those with
|
||||
time-related keywords and uses LLM-based extraction to generate structured
|
||||
observation memories from the filtered messages.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, PersonalMemory
|
||||
|
|
@ -14,6 +23,7 @@ class GetObservationOp(BaseAsyncOp):
|
|||
"""
|
||||
A specialized operation class to generate observations from chat messages using BaseAsyncOp.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -68,13 +78,17 @@ class GetObservationOp(BaseAsyncOp):
|
|||
user_query_list.append(f"{i + 1} {user_name}: {msg.content}")
|
||||
|
||||
# Create prompt using the prompt format method
|
||||
system_prompt = self.prompt_format(prompt_name="get_observation_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name)
|
||||
system_prompt = self.prompt_format(
|
||||
prompt_name="get_observation_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name,
|
||||
)
|
||||
few_shot = self.prompt_format(prompt_name="get_observation_few_shot", user_name=user_name)
|
||||
user_query = self.prompt_format(prompt_name="get_observation_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
user_name=user_name)
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="get_observation_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
user_name=user_name,
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
logger.info(f"get_observation_prompt={full_prompt}")
|
||||
|
|
@ -104,8 +118,8 @@ class GetObservationOp(BaseAsyncOp):
|
|||
metadata={
|
||||
"keywords": obs["keywords"],
|
||||
"source_message": filtered_messages[idx].content,
|
||||
"observation_type": "personal_info"
|
||||
}
|
||||
"observation_type": "personal_info",
|
||||
},
|
||||
)
|
||||
observation_memories.append(observation)
|
||||
logger.info(f"Created observation: {obs['content'][:50]}...")
|
||||
|
|
@ -127,19 +141,21 @@ class GetObservationOp(BaseAsyncOp):
|
|||
# Handle both Chinese and English patterns
|
||||
if match[0]: # Chinese pattern
|
||||
idx_str, content, keywords = match[0], match[1], match[2]
|
||||
else: # English pattern
|
||||
else: # English pattern
|
||||
idx_str, content, keywords = match[3], match[4], match[5]
|
||||
|
||||
try:
|
||||
idx = int(idx_str)
|
||||
# Skip if content indicates no meaningful observation
|
||||
content_lower = content.lower().strip()
|
||||
if content_lower not in ['无', 'none', '', 'repeat']:
|
||||
observations.append({
|
||||
"index": idx,
|
||||
"content": content.strip(),
|
||||
"keywords": keywords.strip() if keywords else ""
|
||||
})
|
||||
if content_lower not in ["无", "none", "", "repeat"]:
|
||||
observations.append(
|
||||
{
|
||||
"index": idx,
|
||||
"content": content.strip(),
|
||||
"keywords": keywords.strip() if keywords else "",
|
||||
},
|
||||
)
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid index format: {idx_str}")
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -1,8 +1,18 @@
|
|||
"""Module for extracting observations with time information from chat messages.
|
||||
|
||||
This module provides the GetObservationWithTimeOp class which extracts
|
||||
personal observations with time information from chat messages. It filters
|
||||
messages to only include those with time-related keywords and uses LLM-based
|
||||
extraction to generate structured observation memories with time information
|
||||
from the filtered messages.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, PersonalMemory
|
||||
|
|
@ -14,6 +24,7 @@ class GetObservationWithTimeOp(BaseAsyncOp):
|
|||
"""
|
||||
A specialized operation class to extract observations with time information from chat messages using BaseAsyncOp.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -77,13 +88,17 @@ class GetObservationWithTimeOp(BaseAsyncOp):
|
|||
user_query_list.append(f"{i + 1} {dt} {user_name}{colon}{msg.content}")
|
||||
|
||||
# Create prompt using the prompt format method
|
||||
system_prompt = self.prompt_format(prompt_name="get_observation_with_time_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name)
|
||||
system_prompt = self.prompt_format(
|
||||
prompt_name="get_observation_with_time_system",
|
||||
num_obs=len(user_query_list),
|
||||
user_name=user_name,
|
||||
)
|
||||
few_shot = self.prompt_format(prompt_name="get_observation_with_time_few_shot", user_name=user_name)
|
||||
user_query = self.prompt_format(prompt_name="get_observation_with_time_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
user_name=user_name)
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="get_observation_with_time_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
user_name=user_name,
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
logger.info(f"get_observation_with_time_prompt={full_prompt}")
|
||||
|
|
@ -114,8 +129,8 @@ class GetObservationWithTimeOp(BaseAsyncOp):
|
|||
"keywords": obs["keywords"],
|
||||
"time_info": obs.get("time_info", ""),
|
||||
"source_message": filtered_messages[idx].content,
|
||||
"observation_type": "personal_info_with_time"
|
||||
}
|
||||
"observation_type": "personal_info_with_time",
|
||||
},
|
||||
)
|
||||
observation_memories.append(observation)
|
||||
logger.info(f"Created observation with time: {obs['content'][:50]}...")
|
||||
|
|
@ -135,8 +150,12 @@ class GetObservationWithTimeOp(BaseAsyncOp):
|
|||
"""Parse observation with time response to extract structured data"""
|
||||
# Pattern to match both Chinese and English observation formats with time information
|
||||
# Chinese: 信息:<1> <时间信息或不输出> <明确的重要信息或"无"> <关键词>
|
||||
# English: Information: <1> <Time information or do not output> <Clear important information or "None"> <Keywords>
|
||||
pattern = r"信息:<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>|Information:\s*<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>"
|
||||
# English: Information: <1> <Time information or do not output>
|
||||
# <Clear important information or "None"> <Keywords>
|
||||
pattern = (
|
||||
r"信息:<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>|"
|
||||
r"Information:\s*<(\d+)>\s*<([^<>]*)>\s*<([^<>]+)>\s*<([^<>]*)>"
|
||||
)
|
||||
matches = re.findall(pattern, response_text, re.IGNORECASE | re.MULTILINE)
|
||||
|
||||
observations = []
|
||||
|
|
@ -144,20 +163,22 @@ class GetObservationWithTimeOp(BaseAsyncOp):
|
|||
# Handle both Chinese and English patterns
|
||||
if match[0]: # Chinese pattern
|
||||
idx_str, time_info, content, keywords = match[0], match[1], match[2], match[3]
|
||||
else: # English pattern
|
||||
else: # English pattern
|
||||
idx_str, time_info, content, keywords = match[4], match[5], match[6], match[7]
|
||||
|
||||
try:
|
||||
idx = int(idx_str)
|
||||
# Skip if content indicates no meaningful observation
|
||||
content_lower = content.lower().strip()
|
||||
if content_lower not in ['无', 'none', '', 'repeat']:
|
||||
observations.append({
|
||||
"index": idx,
|
||||
"time_info": time_info.strip() if time_info else "",
|
||||
"content": content.strip(),
|
||||
"keywords": keywords.strip() if keywords else ""
|
||||
})
|
||||
if content_lower not in ["无", "none", "", "repeat"]:
|
||||
observations.append(
|
||||
{
|
||||
"index": idx,
|
||||
"time_info": time_info.strip() if time_info else "",
|
||||
"content": content.strip(),
|
||||
"keywords": keywords.strip() if keywords else "",
|
||||
},
|
||||
)
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid index format: {idx_str}")
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -1,7 +1,16 @@
|
|||
"""Module for generating reflection subjects from personal memories.
|
||||
|
||||
This module provides the GetReflectionSubjectOp class which retrieves
|
||||
unreflected memory nodes, generates reflection prompts with current insights,
|
||||
invokes an LLM for fresh insights, parses the LLM responses, forms new
|
||||
insight nodes, and updates memory statuses accordingly.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, PersonalMemory
|
||||
|
|
@ -14,6 +23,7 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
generating reflection prompts with current insights, invoking an LLM for fresh insights,
|
||||
parsing the LLM responses, forming new insight nodes, and updating memory statuses accordingly.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
def new_insight_memory(self, insight_content: str, target: str) -> PersonalMemory:
|
||||
|
|
@ -35,8 +45,8 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
author=getattr(self.llm, "model_name", "system"),
|
||||
metadata={
|
||||
"insight_type": "reflection_subject",
|
||||
"memory_type": "personal_topic"
|
||||
}
|
||||
"memory_type": "personal_topic",
|
||||
},
|
||||
)
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -68,15 +78,16 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
# Extract existing insight subjects to avoid duplication
|
||||
existing_subjects = []
|
||||
if existing_insights:
|
||||
existing_subjects = [memory.content for memory in existing_insights if
|
||||
hasattr(memory, 'content') and memory.content]
|
||||
existing_subjects = [
|
||||
memory.content for memory in existing_insights if hasattr(memory, "content") and memory.content
|
||||
]
|
||||
|
||||
logger.info(f"Found {len(existing_subjects)} existing insight subjects")
|
||||
|
||||
# Prepare memory content for LLM analysis
|
||||
memory_contents = []
|
||||
for memory in personal_memories:
|
||||
if hasattr(memory, 'content') and memory.content.strip():
|
||||
if hasattr(memory, "content") and memory.content.strip():
|
||||
memory_contents.append(memory.content.strip())
|
||||
|
||||
if not memory_contents:
|
||||
|
|
@ -86,24 +97,32 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
|
||||
# Generate reflection subjects using LLM
|
||||
insight_memories = await self._generate_reflection_subjects(
|
||||
memory_contents, existing_subjects, user_name, reflect_num_questions
|
||||
memory_contents,
|
||||
existing_subjects,
|
||||
user_name,
|
||||
reflect_num_questions,
|
||||
)
|
||||
|
||||
# Store results in context
|
||||
self.context.response.metadata["insight_memories"] = insight_memories
|
||||
logger.info(f"Generated {len(insight_memories)} new reflection subject memories")
|
||||
|
||||
async def _generate_reflection_subjects(self, memory_contents: List[str], existing_subjects: List[str],
|
||||
user_name: str, num_questions: int) -> List[BaseMemory]:
|
||||
async def _generate_reflection_subjects(
|
||||
self,
|
||||
memory_contents: List[str],
|
||||
existing_subjects: List[str],
|
||||
user_name: str,
|
||||
num_questions: int,
|
||||
) -> List[BaseMemory]:
|
||||
"""
|
||||
Generate new reflection subjects using LLM analysis of memory contents.
|
||||
|
||||
|
||||
Args:
|
||||
memory_contents: List of memory content strings
|
||||
existing_subjects: List of already existing subject strings
|
||||
user_name: Target username
|
||||
num_questions: Maximum number of new subjects to generate
|
||||
|
||||
|
||||
Returns:
|
||||
List of PersonalMemory objects representing new reflection subjects
|
||||
"""
|
||||
|
|
@ -111,17 +130,17 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
system_prompt = self.prompt_format(
|
||||
prompt_name="get_reflection_subject_system",
|
||||
user_name=user_name,
|
||||
num_questions=num_questions
|
||||
num_questions=num_questions,
|
||||
)
|
||||
few_shot = self.prompt_format(
|
||||
prompt_name="get_reflection_subject_few_shot",
|
||||
user_name=user_name
|
||||
user_name=user_name,
|
||||
)
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="get_reflection_subject_user_query",
|
||||
user_name=user_name,
|
||||
exist_keys=", ".join(existing_subjects) if existing_subjects else "None",
|
||||
user_query="\n".join(memory_contents)
|
||||
user_query="\n".join(memory_contents),
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
|
|
@ -140,7 +159,7 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
for subject in new_subjects:
|
||||
insight_memory = self.new_insight_memory(
|
||||
insight_content=subject,
|
||||
target=user_name
|
||||
target=user_name,
|
||||
)
|
||||
insight_memories.append(insight_memory)
|
||||
logger.info(f"Created reflection subject: {subject}")
|
||||
|
|
@ -151,7 +170,15 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
return await self.llm.achat(messages=[Message(content=full_prompt)], callback_fn=parse_reflection_response)
|
||||
|
||||
def get_language_value(self, value_dict: dict):
|
||||
"""Get language-specific value from dictionary"""
|
||||
"""Get language-specific value from dictionary.
|
||||
|
||||
Args:
|
||||
value_dict: Dictionary mapping language codes to values
|
||||
|
||||
Returns:
|
||||
The value corresponding to the current language, or the English
|
||||
value as fallback if the current language is not found
|
||||
"""
|
||||
return value_dict.get(self.language, value_dict.get("en"))
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -161,19 +188,20 @@ class GetReflectionSubjectOp(BaseAsyncOp):
|
|||
existing_subjects = []
|
||||
|
||||
# Split response into lines and clean up
|
||||
lines = response_text.strip().split('\n')
|
||||
lines = response_text.strip().split("\n")
|
||||
subjects = []
|
||||
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
# Skip empty lines, "None" responses, and existing subjects
|
||||
if (line and
|
||||
line not in ['无', 'None', ''] and
|
||||
line not in existing_subjects and
|
||||
not line.startswith('新增') and # Skip Chinese header
|
||||
not line.startswith('New ') and # Skip English header
|
||||
len(line) > 1): # Skip single character responses
|
||||
subjects.append(line)
|
||||
# Check basic validity first
|
||||
if not line or line in ["无", "None", ""] or len(line) <= 1:
|
||||
continue
|
||||
# Check if it's a header or duplicate
|
||||
is_header = line.startswith("新增") or line.startswith("New ")
|
||||
if is_header or line in existing_subjects:
|
||||
continue
|
||||
subjects.append(line)
|
||||
|
||||
logger.info(f"Parsed {len(subjects)} new reflection subjects from response")
|
||||
return subjects
|
||||
|
|
|
|||
|
|
@ -1,8 +1,17 @@
|
|||
"""Module for filtering messages based on information content scores.
|
||||
|
||||
This module provides the InfoFilterOp class which filters chat messages by
|
||||
retaining only those that include significant information about the user.
|
||||
It uses LLM-based scoring to evaluate the information content of messages
|
||||
and filters them based on configurable score thresholds.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message, Trajectory
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message, Trajectory
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import PersonalMemory
|
||||
|
|
@ -14,6 +23,7 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
A specialized operation class to filter messages based on information content scores using BaseAsyncOp.
|
||||
This filters chat messages by retaining only those that include significant information about the user.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -36,7 +46,7 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
user_name = self.context.get("user_name", "user")
|
||||
|
||||
# Filter and process messages
|
||||
info_messages = self._filter_and_process_messages(messages, user_name, info_filter_msg_max_size)
|
||||
info_messages = self._filter_and_process_messages(messages, info_filter_msg_max_size)
|
||||
if not info_messages:
|
||||
logger.warning("No messages left after filtering")
|
||||
self.context.messages = []
|
||||
|
|
@ -52,7 +62,7 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
logger.info(f"Filtered to {len(filtered_memories)} high-information messages")
|
||||
|
||||
@staticmethod
|
||||
def _filter_and_process_messages(messages: List[Message], user_name: str, max_size: int) -> List[Message]:
|
||||
def _filter_and_process_messages(messages: List[Message], max_size: int) -> List[Message]:
|
||||
"""Filter and process messages for information filtering"""
|
||||
info_messages = []
|
||||
|
||||
|
|
@ -60,7 +70,7 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
# Ensure metadata exists
|
||||
|
||||
# Skip memorized messages
|
||||
if msg.metadata.get('memorized', False):
|
||||
if msg.metadata.get("memorized", False):
|
||||
continue
|
||||
|
||||
# Only process messages from the target user
|
||||
|
|
@ -68,7 +78,7 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
# if role_name and role_name != user_name:
|
||||
# continue
|
||||
|
||||
elif msg.role.value != "user":
|
||||
if msg.role.value != "user":
|
||||
continue
|
||||
|
||||
# Truncate long messages
|
||||
|
|
@ -81,8 +91,12 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
logger.info(f"Filtered messages from {len(messages)} to {len(info_messages)}")
|
||||
return info_messages
|
||||
|
||||
async def _filter_messages_with_llm(self, info_messages: List[Message], user_name: str, preserved_scores: str) -> List[
|
||||
PersonalMemory]:
|
||||
async def _filter_messages_with_llm(
|
||||
self,
|
||||
info_messages: List[Message],
|
||||
user_name: str,
|
||||
preserved_scores: str,
|
||||
) -> List[PersonalMemory]:
|
||||
"""Filter messages using LLM to score information content"""
|
||||
|
||||
# Build prompt for information filtering
|
||||
|
|
@ -92,12 +106,16 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
user_query_list.append(f"{i + 1} {user_name}{colon} {msg.content}")
|
||||
|
||||
# Create prompt using the prompt format method
|
||||
system_prompt = self.prompt_format(prompt_name="info_filter_system",
|
||||
batch_size=len(info_messages),
|
||||
user_name=user_name)
|
||||
system_prompt = self.prompt_format(
|
||||
prompt_name="info_filter_system",
|
||||
batch_size=len(info_messages),
|
||||
user_name=user_name,
|
||||
)
|
||||
few_shot = self.prompt_format(prompt_name="info_filter_few_shot", user_name=user_name)
|
||||
user_query = self.prompt_format(prompt_name="info_filter_user_query",
|
||||
user_query="\n".join(user_query_list))
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="info_filter_user_query",
|
||||
user_query="\n".join(user_query_list),
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
logger.info(f"info_filter_prompt={full_prompt}")
|
||||
|
|
@ -134,11 +152,11 @@ class InfoFilterOp(BaseAsyncOp):
|
|||
metadata={
|
||||
"info_score": score,
|
||||
"filter_type": "info_content",
|
||||
"original_message_time": getattr(message, 'time_created', None),
|
||||
"original_message_time": getattr(message, "time_created", None),
|
||||
"role_name": message.metadata.pop("role_name", user_name),
|
||||
"memorized": True,
|
||||
**message.metadata # Include all original metadata
|
||||
}
|
||||
**message.metadata, # Include all original metadata
|
||||
},
|
||||
)
|
||||
filtered_memories.append(memory)
|
||||
logger.info(f"Info filter: kept message with score {score}: {message.content[:50]}...")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,16 @@
|
|||
"""Module for loading today's memories from vector store.
|
||||
|
||||
This module provides the LoadTodayMemoryOp class which loads memories from
|
||||
the current date for deduplication purposes. It focuses specifically on
|
||||
retrieving and deduplicating memories from the current date using vector
|
||||
store search with date filtering.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import VectorNode
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, vector_node_to_memory
|
||||
|
|
@ -14,12 +23,13 @@ class LoadTodayMemoryOp(BaseAsyncOp):
|
|||
Operation to load today's memories from vector store for deduplication.
|
||||
Focuses specifically on retrieving and deduplicating memories from the current date.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Load today's memories from vector store and perform deduplication.
|
||||
|
||||
|
||||
This operation:
|
||||
1. Retrieves memories from today using vector store search
|
||||
2. Converts vector nodes to memory objects
|
||||
|
|
@ -50,12 +60,12 @@ class LoadTodayMemoryOp(BaseAsyncOp):
|
|||
async def _retrieve_today_memories(self, workspace_id: str, user_name: str, top_k: int) -> List[BaseMemory]:
|
||||
"""
|
||||
Retrieve memories from today using vector store with date filtering.
|
||||
|
||||
|
||||
Args:
|
||||
workspace_id: Workspace identifier
|
||||
user_name: Target username
|
||||
top_k: Maximum number of memories to retrieve
|
||||
|
||||
|
||||
Returns:
|
||||
List of today's memories
|
||||
"""
|
||||
|
|
@ -70,7 +80,7 @@ class LoadTodayMemoryOp(BaseAsyncOp):
|
|||
filter_dict = {
|
||||
"memory_type": "personal",
|
||||
"target": user_name,
|
||||
"created_date": today_date
|
||||
"created_date": today_date,
|
||||
}
|
||||
|
||||
# Search vector store with date filter
|
||||
|
|
@ -78,7 +88,8 @@ class LoadTodayMemoryOp(BaseAsyncOp):
|
|||
query=" ",
|
||||
workspace_id=workspace_id,
|
||||
top_k=top_k,
|
||||
filter_dict=filter_dict)
|
||||
filter_dict=filter_dict,
|
||||
)
|
||||
|
||||
logger.info(f"Vector store returned {len(nodes)} nodes for today")
|
||||
|
||||
|
|
@ -96,10 +107,10 @@ class LoadTodayMemoryOp(BaseAsyncOp):
|
|||
def _convert_nodes_to_memories(nodes: List[VectorNode]) -> List[BaseMemory]:
|
||||
"""
|
||||
Convert vector nodes to memory objects.
|
||||
|
||||
|
||||
Args:
|
||||
nodes: List of vector nodes from vector store
|
||||
|
||||
|
||||
Returns:
|
||||
List of converted memory objects
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,19 @@
|
|||
"""Module for handling contradictions and redundancies in long conversations.
|
||||
|
||||
This module provides the LongContraRepeatOp class which manages and updates
|
||||
memory entries within a conversation scope by identifying and handling
|
||||
contradictions or redundancies. It extends BaseAsyncOp to provide specialized
|
||||
functionality for long conversations with potential contradictory or
|
||||
repetitive statements.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, PersonalMemory
|
||||
|
|
@ -17,12 +27,13 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
specialized functionality for long conversations with potential contradictory
|
||||
or repetitive statements.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Analyze memories for contradictions and redundancies, resolving conflicts.
|
||||
|
||||
|
||||
Process:
|
||||
1. Get updated insight memories from previous operation
|
||||
2. Check for contradictions and redundancies among memories
|
||||
|
|
@ -51,7 +62,7 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
sorted_memories = sorted(
|
||||
updated_insights,
|
||||
key=lambda x: x.time_created,
|
||||
reverse=True
|
||||
reverse=True,
|
||||
)[:max_memories_to_process]
|
||||
|
||||
if len(sorted_memories) <= 1:
|
||||
|
|
@ -71,10 +82,10 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
async def _analyze_and_resolve_conflicts(self, memories: List[BaseMemory]) -> List[BaseMemory]:
|
||||
"""
|
||||
Analyze memories for contradictions and redundancies using LLM.
|
||||
|
||||
|
||||
Args:
|
||||
memories: List of memories to analyze
|
||||
|
||||
|
||||
Returns:
|
||||
List of filtered memories with conflicts resolved
|
||||
"""
|
||||
|
|
@ -89,15 +100,15 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
system_prompt = self.prompt_format(
|
||||
prompt_name="long_contra_repeat_system",
|
||||
num_obs=len(memory_texts),
|
||||
user_name=user_name
|
||||
user_name=user_name,
|
||||
)
|
||||
few_shot = self.prompt_format(
|
||||
prompt_name="long_contra_repeat_few_shot",
|
||||
user_name=user_name
|
||||
user_name=user_name,
|
||||
)
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="long_contra_repeat_user_query",
|
||||
user_query="\n".join(memory_texts)
|
||||
user_query="\n".join(memory_texts),
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
|
|
@ -139,7 +150,7 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
memory = memories[memory_idx]
|
||||
judgment_lower = judgment.lower()
|
||||
|
||||
if judgment_lower in ['矛盾', 'contradiction']:
|
||||
if judgment_lower in ["矛盾", "contradiction"]:
|
||||
# For contradictory memories, either modify content or mark for removal
|
||||
if modified_content.strip():
|
||||
# Create new memory with modified content
|
||||
|
|
@ -147,9 +158,9 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
workspace_id=memory.workspace_id,
|
||||
memory_id=memory.memory_id,
|
||||
content=modified_content.strip(),
|
||||
target=memory.target if hasattr(memory, 'target') else user_name,
|
||||
target=memory.target if hasattr(memory, "target") else user_name,
|
||||
author=memory.author,
|
||||
metadata={**memory.metadata, 'modified_by': 'long_contra_repeat'}
|
||||
metadata={**memory.metadata, "modified_by": "long_contra_repeat"},
|
||||
)
|
||||
modified_memory.update_time_modified()
|
||||
filtered_memories.append(modified_memory)
|
||||
|
|
@ -158,7 +169,7 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
# Remove contradictory memory without modification
|
||||
logger.info(f"Removing contradictory memory {idx}: {memory.content[:50]}...")
|
||||
|
||||
elif judgment_lower in ['被包含', 'contained']:
|
||||
elif judgment_lower in ["被包含", "contained"]:
|
||||
# Remove contained/redundant memories
|
||||
logger.info(f"Removing contained memory {idx}: {memory.content[:50]}...")
|
||||
|
||||
|
|
@ -179,7 +190,15 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
return filtered_memories
|
||||
|
||||
def get_language_value(self, value_dict: dict):
|
||||
"""Get language-specific value from dictionary"""
|
||||
"""Get language-specific value from dictionary.
|
||||
|
||||
Args:
|
||||
value_dict: Dictionary mapping language codes to values
|
||||
|
||||
Returns:
|
||||
The value corresponding to the current language, or the English
|
||||
value as fallback if the current language is not found
|
||||
"""
|
||||
return value_dict.get(self.language, value_dict.get("en"))
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -188,7 +207,10 @@ class LongContraRepeatOp(BaseAsyncOp):
|
|||
# Pattern to match both Chinese and English judgment formats
|
||||
# Chinese: 判断:<序号> <矛盾|被包含|无> <修改后的内容>
|
||||
# English: Judgment: <Index> <Contradiction|Contained|None> <Modified content>
|
||||
pattern = r"判断:<(\d+)>\s*<(矛盾|被包含|无)>\s*<([^<>]*)>|Judgment:\s*<(\d+)>\s*<(Contradiction|Contained|None)>\s*<([^<>]*)>"
|
||||
pattern = (
|
||||
r"判断:<(\d+)>\s*<(矛盾|被包含|无)>\s*<([^<>]*)>|"
|
||||
r"Judgment:\s*<(\d+)>\s*<(Contradiction|Contained|None)>\s*<([^<>]*)>"
|
||||
)
|
||||
matches = re.findall(pattern, response_text, re.IGNORECASE | re.MULTILINE)
|
||||
|
||||
judgments = []
|
||||
|
|
|
|||
|
|
@ -1,8 +1,17 @@
|
|||
"""Module for updating personal insight memories based on new observations.
|
||||
|
||||
This module provides the UpdateInsightOp class which updates insight values
|
||||
in a memory system by filtering insight nodes based on their association with
|
||||
observed nodes, utilizing a ranking model to prioritize them, generating
|
||||
refreshed insights via an LLM, and managing node statuses and content updates.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.schema.message import Message
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import PersonalMemory
|
||||
|
|
@ -15,12 +24,13 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
based on their association with observed nodes, utilizes a ranking model to prioritize them,
|
||||
generates refreshed insights via an LLM, and manages node statuses and content updates.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Update insight values based on new observation memories.
|
||||
|
||||
|
||||
Process:
|
||||
1. Get insight subjects and personal memories from context
|
||||
2. Find relevant observations for each insight subject
|
||||
|
|
@ -50,7 +60,9 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
|
||||
# Score and filter insights based on relevance to observations
|
||||
scored_insights = self._score_insights_by_relevance(
|
||||
insight_memories, personal_memories, update_insight_threshold
|
||||
insight_memories,
|
||||
personal_memories,
|
||||
update_insight_threshold,
|
||||
)
|
||||
|
||||
if not scored_insights:
|
||||
|
|
@ -64,9 +76,11 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
|
||||
# Update each selected insight
|
||||
updated_insights = []
|
||||
for insight_memory, relevance_score, relevant_observations in top_insights:
|
||||
for insight_memory, _relevance_score, relevant_observations in top_insights:
|
||||
updated_insight = await self._update_insight_with_observations(
|
||||
insight_memory, relevant_observations, user_name
|
||||
insight_memory,
|
||||
relevant_observations,
|
||||
user_name,
|
||||
)
|
||||
if updated_insight:
|
||||
updated_insights.append(updated_insight)
|
||||
|
|
@ -75,17 +89,20 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
self.context.response.metadata["updated_insight_memories"] = updated_insights
|
||||
logger.info(f"Successfully updated {len(updated_insights)} insight memories")
|
||||
|
||||
def _score_insights_by_relevance(self, insight_memories: List[PersonalMemory],
|
||||
observation_memories: List[PersonalMemory],
|
||||
threshold: float) -> List[tuple]:
|
||||
def _score_insights_by_relevance(
|
||||
self,
|
||||
insight_memories: List[PersonalMemory],
|
||||
observation_memories: List[PersonalMemory],
|
||||
threshold: float,
|
||||
) -> List[tuple]:
|
||||
"""
|
||||
Score insight memories based on relevance to observation memories.
|
||||
|
||||
|
||||
Args:
|
||||
insight_memories: List of insight memories to score
|
||||
observation_memories: List of observation memories for comparison
|
||||
threshold: Minimum relevance score threshold
|
||||
|
||||
|
||||
Returns:
|
||||
List[tuple]: List of (insight_memory, relevance_score, relevant_observations)
|
||||
"""
|
||||
|
|
@ -95,13 +112,15 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
relevant_observations = []
|
||||
max_relevance = 0.0
|
||||
|
||||
insight_subject = getattr(insight_memory, 'reflection_subject', '') or insight_memory.content
|
||||
insight_subject = getattr(insight_memory, "reflection_subject", "") or insight_memory.content
|
||||
insight_keywords = set(insight_memory.content.lower().split())
|
||||
|
||||
# Find observations relevant to this insight
|
||||
for obs_memory in observation_memories:
|
||||
relevance_score = self._calculate_relevance_score(
|
||||
insight_memory, obs_memory, insight_keywords
|
||||
insight_memory,
|
||||
obs_memory,
|
||||
insight_keywords,
|
||||
)
|
||||
|
||||
if relevance_score >= threshold:
|
||||
|
|
@ -112,18 +131,36 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
if relevant_observations:
|
||||
scored_insights.append((insight_memory, max_relevance, relevant_observations))
|
||||
logger.info(
|
||||
f"Insight '{insight_subject[:40]}...' scored {max_relevance:.3f} with {len(relevant_observations)} observations"
|
||||
f"Insight '{insight_subject[:40]}...' scored {max_relevance:.3f} "
|
||||
f"with {len(relevant_observations)} observations",
|
||||
)
|
||||
|
||||
return scored_insights
|
||||
|
||||
@staticmethod
|
||||
def _calculate_relevance_score(insight_memory: PersonalMemory,
|
||||
obs_memory: PersonalMemory, insight_keywords: set) -> float:
|
||||
"""Calculate relevance score between insight and observation memory"""
|
||||
def _calculate_relevance_score(
|
||||
insight_memory: PersonalMemory,
|
||||
obs_memory: PersonalMemory,
|
||||
insight_keywords: set,
|
||||
) -> float:
|
||||
"""Calculate relevance score between insight and observation memory.
|
||||
|
||||
The relevance score is calculated based on:
|
||||
- High relevance (0.9) if both memories share the same reflection subject
|
||||
- Medium relevance based on keyword overlap using Jaccard similarity
|
||||
|
||||
Args:
|
||||
insight_memory: The insight memory to compare
|
||||
obs_memory: The observation memory to compare against
|
||||
insight_keywords: Set of keywords extracted from insight memory content
|
||||
|
||||
Returns:
|
||||
float: Relevance score between 0.0 and 1.0, where higher values
|
||||
indicate greater relevance
|
||||
"""
|
||||
# High relevance for same reflection subject
|
||||
insight_subject = getattr(insight_memory, 'reflection_subject', '')
|
||||
obs_subject = getattr(obs_memory, 'reflection_subject', '')
|
||||
insight_subject = getattr(insight_memory, "reflection_subject", "")
|
||||
obs_subject = getattr(obs_memory, "reflection_subject", "")
|
||||
|
||||
if insight_subject and obs_subject and insight_subject == obs_subject:
|
||||
return 0.9
|
||||
|
|
@ -135,22 +172,26 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
|
||||
return intersection / union if union > 0 else 0.0
|
||||
|
||||
async def _update_insight_with_observations(self, insight_memory: PersonalMemory,
|
||||
relevant_observations: List[PersonalMemory],
|
||||
user_name: str) -> PersonalMemory:
|
||||
async def _update_insight_with_observations(
|
||||
self,
|
||||
insight_memory: PersonalMemory,
|
||||
relevant_observations: List[PersonalMemory],
|
||||
user_name: str,
|
||||
) -> PersonalMemory:
|
||||
"""
|
||||
Update a single insight memory based on relevant observations using LLM.
|
||||
|
||||
|
||||
Args:
|
||||
insight_memory: The insight memory to update
|
||||
relevant_observations: List of relevant observation memories
|
||||
user_name: The target username
|
||||
|
||||
|
||||
Returns:
|
||||
PersonalMemory: Updated insight memory or None if update failed
|
||||
"""
|
||||
logger.info(
|
||||
f"Updating insight: {insight_memory.content[:50]}... with {len(relevant_observations)} observations")
|
||||
f"Updating insight: {insight_memory.content[:50]}... with {len(relevant_observations)} observations",
|
||||
)
|
||||
|
||||
# Build observation context
|
||||
observation_texts = [obs.content for obs in relevant_observations]
|
||||
|
|
@ -161,10 +202,12 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
|
||||
system_prompt = self.prompt_format(prompt_name="update_insight_system", user_name=user_name)
|
||||
few_shot = self.prompt_format(prompt_name="update_insight_few_shot", user_name=user_name)
|
||||
user_query = self.prompt_format(prompt_name="update_insight_user_query",
|
||||
user_query="\n".join(observation_texts),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key_value)
|
||||
user_query = self.prompt_format(
|
||||
prompt_name="update_insight_user_query",
|
||||
user_query="\n".join(observation_texts),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key_value,
|
||||
)
|
||||
|
||||
full_prompt = f"{system_prompt}\n\n{few_shot}\n\n{user_query}"
|
||||
logger.info(f"update_insight_prompt={full_prompt}")
|
||||
|
|
@ -177,7 +220,7 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
# Parse the response to extract updated insight
|
||||
updated_content = UpdateInsightOp.parse_update_insight_response(response_text, self.language)
|
||||
|
||||
if not updated_content or updated_content.lower() in ['无', 'none', '']:
|
||||
if not updated_content or updated_content.lower() in ["无", "none", ""]:
|
||||
logger.info(f"No update needed for insight: {insight_memory.content[:50]}...")
|
||||
return insight_memory
|
||||
|
||||
|
|
@ -198,8 +241,8 @@ class UpdateInsightOp(BaseAsyncOp):
|
|||
**insight_memory.metadata,
|
||||
"updated_by": "update_insight_op",
|
||||
"original_content": insight_memory.content,
|
||||
"update_reason": "integrated_new_observations"
|
||||
}
|
||||
"update_reason": "integrated_new_observations",
|
||||
},
|
||||
)
|
||||
updated_insight.update_time_modified()
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,10 @@
|
|||
"""Task memory operations module.
|
||||
|
||||
This module provides various operations for extracting, processing, and managing
|
||||
task memories from trajectories, including success/failure extraction, comparative
|
||||
analysis, deduplication, and validation.
|
||||
"""
|
||||
|
||||
from .comparative_extraction_op import ComparativeExtractionOp
|
||||
from .failure_extraction_op import FailureExtractionOp
|
||||
from .memory_deduplication_op import MemoryDeduplicationOp
|
||||
|
|
@ -7,3 +14,15 @@ from .simple_summary_op import SimpleSummaryOp
|
|||
from .success_extraction_op import SuccessExtractionOp
|
||||
from .trajectory_preprocess_op import TrajectoryPreprocessOp
|
||||
from .trajectory_segmentation_op import TrajectorySegmentationOp
|
||||
|
||||
__all__ = [
|
||||
"ComparativeExtractionOp",
|
||||
"FailureExtractionOp",
|
||||
"MemoryDeduplicationOp",
|
||||
"MemoryValidationOp",
|
||||
"SimpleComparativeSummaryOp",
|
||||
"SimpleSummaryOp",
|
||||
"SuccessExtractionOp",
|
||||
"TrajectoryPreprocessOp",
|
||||
"TrajectorySegmentationOp",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,8 +1,15 @@
|
|||
"""Comparative extraction operation for task memory generation.
|
||||
|
||||
This module provides operations to extract comparative task memories by comparing
|
||||
different trajectories with varying scores or success/failure outcomes.
|
||||
"""
|
||||
|
||||
from typing import List, Tuple, Optional
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,6 +19,16 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
|
|||
|
||||
@C.register_op()
|
||||
class ComparativeExtractionOp(BaseAsyncOp):
|
||||
"""Extract comparative task memories by comparing different scoring trajectories.
|
||||
|
||||
This operation performs two types of comparisons:
|
||||
1. Soft comparison: Compares highest vs lowest scoring trajectories
|
||||
2. Hard comparison: Compares similar success vs failure step sequences
|
||||
|
||||
The extracted memories help identify what makes some trajectories more successful
|
||||
than others.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -27,20 +44,24 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
highest_traj, lowest_traj = self._find_highest_lowest_scoring_trajectories(all_trajectories)
|
||||
if highest_traj and lowest_traj and highest_traj.score > lowest_traj.score:
|
||||
logger.info(
|
||||
f"Extracting soft comparative task memories: highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})")
|
||||
f"Extracting soft comparative task memories: "
|
||||
f"highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})",
|
||||
)
|
||||
soft_task_memories = await self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
|
||||
comparative_task_memories.extend(soft_task_memories)
|
||||
|
||||
# Hard comparison: success vs failure (if similarity search is enabled)
|
||||
if (success_trajectories and failure_trajectories and
|
||||
self.op_params.get("enable_similarity_comparison", False)):
|
||||
if success_trajectories and failure_trajectories and self.op_params.get("enable_similarity_comparison", False):
|
||||
|
||||
similar_pairs = self._find_similar_step_sequences(success_trajectories, failure_trajectories)
|
||||
logger.info(f"Found {len(similar_pairs)} similar pairs for hard comparison")
|
||||
|
||||
for success_steps, failure_steps, similarity_score in similar_pairs:
|
||||
hard_task_memories = await self._extract_hard_comparative_task_memory(success_steps, failure_steps,
|
||||
similarity_score)
|
||||
hard_task_memories = await self._extract_hard_comparative_task_memory(
|
||||
success_steps,
|
||||
failure_steps,
|
||||
similarity_score,
|
||||
)
|
||||
comparative_task_memories.extend(hard_task_memories)
|
||||
|
||||
logger.info(f"Extracted {len(comparative_task_memories)} comparative task memories")
|
||||
|
|
@ -50,7 +71,9 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
|
||||
@staticmethod
|
||||
def _find_highest_lowest_scoring_trajectories(trajectories: List[Trajectory]) -> Tuple[
|
||||
Optional[Trajectory], Optional[Trajectory]]:
|
||||
Optional[Trajectory],
|
||||
Optional[Trajectory],
|
||||
]:
|
||||
"""Find the highest and lowest scoring trajectories"""
|
||||
if len(trajectories) < 2:
|
||||
return None, None
|
||||
|
|
@ -75,8 +98,11 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
"""Get trajectory score"""
|
||||
return trajectory.score
|
||||
|
||||
async def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
|
||||
BaseMemory]:
|
||||
async def _extract_soft_comparative_task_memory(
|
||||
self,
|
||||
higher_traj: Trajectory,
|
||||
lower_traj: Trajectory,
|
||||
) -> List[BaseMemory]:
|
||||
"""Extract soft comparative task memory (high score vs low score)"""
|
||||
higher_steps = self._get_trajectory_steps(higher_traj)
|
||||
lower_steps = self._get_trajectory_steps(lower_traj)
|
||||
|
|
@ -88,7 +114,7 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
higher_steps=merge_messages_content(higher_steps),
|
||||
lower_steps=merge_messages_content(lower_steps),
|
||||
higher_score=f"{higher_score:.2f}",
|
||||
lower_score=f"{lower_score:.2f}"
|
||||
lower_score=f"{lower_score:.2f}",
|
||||
)
|
||||
|
||||
def parse_task_memories(message: Message) -> List[BaseMemory]:
|
||||
|
|
@ -100,24 +126,30 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
workspace_id=self.context.get("workspace_id", ""),
|
||||
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
|
||||
content=tm_data.get("experience", ""),
|
||||
author=getattr(self.llm, 'model_name', 'system'),
|
||||
metadata=tm_data
|
||||
author=getattr(self.llm, "model_name", "system"),
|
||||
metadata=tm_data,
|
||||
)
|
||||
task_memories.append(task_memory)
|
||||
|
||||
return task_memories
|
||||
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(
|
||||
messages=[FlowMessage(role=Role.USER, content=prompt)],
|
||||
callback_fn=parse_task_memories,
|
||||
)
|
||||
|
||||
async def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
|
||||
failure_steps: List[Message], similarity_score: float) -> List[
|
||||
BaseMemory]:
|
||||
async def _extract_hard_comparative_task_memory(
|
||||
self,
|
||||
success_steps: List[Message],
|
||||
failure_steps: List[Message],
|
||||
similarity_score: float,
|
||||
) -> List[BaseMemory]:
|
||||
"""Extract hard comparative task memory (success vs failure)"""
|
||||
prompt = self.prompt_format(
|
||||
prompt_name="hard_comparative_step_task_memory_prompt",
|
||||
success_steps=merge_messages_content(success_steps),
|
||||
failure_steps=merge_messages_content(failure_steps),
|
||||
similarity_score=similarity_score
|
||||
similarity_score=similarity_score,
|
||||
)
|
||||
|
||||
def parse_task_memories(message: Message) -> List[BaseMemory]:
|
||||
|
|
@ -129,19 +161,22 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
workspace_id=self.context.get("workspace_id", ""),
|
||||
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
|
||||
content=tm_data.get("experience", ""),
|
||||
author=getattr(self.llm, 'model_name', 'system'),
|
||||
metadata=tm_data
|
||||
author=getattr(self.llm, "model_name", "system"),
|
||||
metadata=tm_data,
|
||||
)
|
||||
task_memories.append(task_memory)
|
||||
|
||||
return task_memories
|
||||
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(
|
||||
messages=[FlowMessage(role=Role.USER, content=prompt)],
|
||||
callback_fn=parse_task_memories,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_trajectory_steps(trajectory: Trajectory) -> List[Message]:
|
||||
"""Get trajectory steps, prioritizing segmented steps"""
|
||||
if hasattr(trajectory, 'segments') and trajectory.segments:
|
||||
if hasattr(trajectory, "segments") and trajectory.segments:
|
||||
# If there are segments, merge all segments
|
||||
all_steps = []
|
||||
for segment in trajectory.segments:
|
||||
|
|
@ -150,9 +185,11 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
else:
|
||||
return trajectory.messages
|
||||
|
||||
def _find_similar_step_sequences(self, success_trajectories: List[Trajectory],
|
||||
failure_trajectories: List[Trajectory]) -> List[
|
||||
Tuple[List[Message], List[Message], float]]:
|
||||
def _find_similar_step_sequences(
|
||||
self,
|
||||
success_trajectories: List[Trajectory],
|
||||
failure_trajectories: List[Trajectory],
|
||||
) -> List[Tuple[List[Message], List[Message], float]]:
|
||||
"""Find similar step sequences for comparison"""
|
||||
if not self.op_params.get("enable_similarity_comparison", False):
|
||||
return []
|
||||
|
|
@ -163,14 +200,14 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
# Get step sequences
|
||||
success_step_sequences = []
|
||||
for traj in success_trajectories:
|
||||
if hasattr(traj.metadata, 'segments') and traj.metadata["segments"]:
|
||||
if hasattr(traj.metadata, "segments") and traj.metadata["segments"]:
|
||||
success_step_sequences.extend(traj.metadata["segments"])
|
||||
else:
|
||||
success_step_sequences.append(traj.messages)
|
||||
|
||||
failure_step_sequences = []
|
||||
for traj in failure_trajectories:
|
||||
if hasattr(traj.metadata, 'segments') and traj.metadata["segments"]:
|
||||
if hasattr(traj.metadata, "segments") and traj.metadata["segments"]:
|
||||
failure_step_sequences.extend(traj.metadata["segments"])
|
||||
else:
|
||||
failure_step_sequences.append(traj.messages)
|
||||
|
|
@ -188,8 +225,14 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
failure_texts = [merge_messages_content(seq) for seq in failure_step_sequences]
|
||||
|
||||
# Get embedding vectors
|
||||
if hasattr(self, 'vector_store') and self.vector_store and hasattr(
|
||||
self.vector_store, 'embedding_model'):
|
||||
if (
|
||||
hasattr(self, "vector_store")
|
||||
and self.vector_store
|
||||
and hasattr(
|
||||
self.vector_store,
|
||||
"embedding_model",
|
||||
)
|
||||
):
|
||||
success_embeddings = self.vector_store.embedding_model.get_embeddings(success_texts)
|
||||
failure_embeddings = self.vector_store.embedding_model.get_embeddings(failure_texts)
|
||||
|
||||
|
|
@ -201,11 +244,13 @@ class ComparativeExtractionOp(BaseAsyncOp):
|
|||
similarity = self._calculate_cosine_similarity(s_emb, f_emb)
|
||||
|
||||
if similarity > similarity_threshold:
|
||||
similar_pairs.append((
|
||||
success_step_sequences[i],
|
||||
failure_step_sequences[j],
|
||||
similarity
|
||||
))
|
||||
similar_pairs.append(
|
||||
(
|
||||
success_step_sequences[i],
|
||||
failure_step_sequences[j],
|
||||
similarity,
|
||||
),
|
||||
)
|
||||
|
||||
# Return top most similar pairs
|
||||
max_pairs = self.op_params.get("max_similarity_pairs", 3)
|
||||
|
|
|
|||
|
|
@ -1,28 +1,28 @@
|
|||
soft_comparative_step_task_memory_prompt: |
|
||||
You are an expert AI analyst comparing higher-scoring and lower-scoring step sequences to extract performance insights.
|
||||
|
||||
|
||||
Your task is to identify the key differences between higher and lower performing approaches at the step level.
|
||||
Focus on what made the higher-scoring approach more effective, even when both approaches may have had partial success.
|
||||
|
||||
|
||||
SOFT COMPARATIVE ANALYSIS FRAMEWORK:
|
||||
● PERFORMANCE FACTORS: Identify what specifically contributed to the higher score
|
||||
● APPROACH DIFFERENCES: Compare methodologies and execution strategies
|
||||
● EFFICIENCY ANALYSIS: Analyze why one approach was more efficient or effective
|
||||
● OPTIMIZATION INSIGHTS: Extract lessons for improving performance
|
||||
|
||||
|
||||
EXTRACTION PRINCIPLES:
|
||||
● Focus on INCREMENTAL IMPROVEMENTS and performance optimization
|
||||
● Extract QUALITY INDICATORS that differentiate better vs good approaches
|
||||
● Identify REFINEMENT STRATEGIES that lead to higher scores
|
||||
● Frame insights as PERFORMANCE ENHANCEMENT guidelines
|
||||
|
||||
|
||||
# Higher-Scoring Step Sequence (Score: {higher_score})
|
||||
{higher_steps}
|
||||
|
||||
|
||||
# Lower-Scoring Step Sequence (Score: {lower_score})
|
||||
{lower_steps}
|
||||
|
||||
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Generate 1-2 performance improvement insights as JSON objects:
|
||||
```json
|
||||
|
|
@ -40,30 +40,30 @@ soft_comparative_step_task_memory_prompt: |
|
|||
|
||||
hard_comparative_step_task_memory_prompt: |
|
||||
You are an expert AI analyst comparing successful and failed step sequences to extract differential insights.
|
||||
|
||||
|
||||
Your task is to identify the key differences between success and failure patterns at the step level.
|
||||
Focus on critical decision points, technique variations, and approach differences.
|
||||
|
||||
|
||||
COMPARATIVE ANALYSIS FRAMEWORK:
|
||||
● DECISION CONTRAST: Compare critical decisions made in success vs failure cases
|
||||
● TECHNIQUE VARIATIONS: Identify different approaches and their outcomes
|
||||
● TIMING DIFFERENCES: Analyze when certain actions were taken and their impact
|
||||
● SUCCESS FACTORS: Extract what specifically made the difference
|
||||
|
||||
|
||||
EXTRACTION PRINCIPLES:
|
||||
● Frame comparisons as PRINCIPLES as well as case-specific SOLUTIONS
|
||||
● Identify PATTERNS that differentiate effective vs ineffective approaches
|
||||
● Extract RULES that can guide future similar situations
|
||||
● Focus on UNDERLYING MECHANISMS rather than surface-level differences
|
||||
|
||||
|
||||
# Successful Step Sequence
|
||||
{success_steps}
|
||||
|
||||
|
||||
# Failed Step Sequence
|
||||
{failure_steps}
|
||||
|
||||
|
||||
# Similarity Score: {similarity_score}
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Generate 1-2 comparative insights as JSON objects:
|
||||
```json
|
||||
|
|
|
|||
|
|
@ -1,8 +1,15 @@
|
|||
"""Failure extraction operation for task memory generation.
|
||||
|
||||
This module provides operations to extract task memories from failed trajectories,
|
||||
identifying mistakes, pitfalls, and lessons learned from failures.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,6 +19,13 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
|
|||
|
||||
@C.register_op()
|
||||
class FailureExtractionOp(BaseAsyncOp):
|
||||
"""Extract task memories from failed trajectories.
|
||||
|
||||
This operation analyzes failed trajectories (or their segments) to extract
|
||||
lessons learned, common mistakes, and anti-patterns that should be avoided
|
||||
in similar future tasks.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -28,7 +42,7 @@ class FailureExtractionOp(BaseAsyncOp):
|
|||
|
||||
# Process trajectories
|
||||
for trajectory in failure_trajectories:
|
||||
if hasattr(trajectory, 'segments') and trajectory.segments:
|
||||
if hasattr(trajectory, "segments") and trajectory.segments:
|
||||
# Process segmented step sequences
|
||||
for segment in trajectory.segments:
|
||||
task_memories = await self._extract_failure_task_memory_from_steps(segment, trajectory)
|
||||
|
|
@ -43,17 +57,21 @@ class FailureExtractionOp(BaseAsyncOp):
|
|||
# Add task memories to context
|
||||
self.context.failure_task_memories = failure_task_memories
|
||||
|
||||
async def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
|
||||
async def _extract_failure_task_memory_from_steps(
|
||||
self,
|
||||
steps: List[Message],
|
||||
trajectory: Trajectory,
|
||||
) -> List[BaseMemory]:
|
||||
"""Extract task memory from failed step sequences"""
|
||||
step_content = merge_messages_content(steps)
|
||||
context = get_trajectory_context(trajectory, steps)
|
||||
|
||||
prompt = self.prompt_format(
|
||||
prompt_name="failure_step_task_memory_prompt",
|
||||
query=trajectory.metadata.get('query', ''),
|
||||
query=trajectory.metadata.get("query", ""),
|
||||
step_sequence=step_content,
|
||||
context=context,
|
||||
outcome="failed"
|
||||
outcome="failed",
|
||||
)
|
||||
|
||||
def parse_task_memories(message: Message) -> List[BaseMemory]:
|
||||
|
|
@ -65,11 +83,14 @@ class FailureExtractionOp(BaseAsyncOp):
|
|||
workspace_id=self.context.get("workspace_id", ""),
|
||||
when_to_use=tm_data.get("when_to_use", tm_data.get("condition", "")),
|
||||
content=tm_data.get("experience", ""),
|
||||
author=getattr(self.llm, 'model_name', 'system'),
|
||||
metadata=tm_data
|
||||
author=getattr(self.llm, "model_name", "system"),
|
||||
metadata=tm_data,
|
||||
)
|
||||
task_memories.append(task_memory)
|
||||
|
||||
return task_memories
|
||||
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(
|
||||
messages=[FlowMessage(role=Role.USER, content=prompt)],
|
||||
callback_fn=parse_task_memories,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,31 +1,31 @@
|
|||
failure_step_task_memory_prompt: |
|
||||
You are an expert AI analyst reviewing failed step sequences from an AI agent execution.
|
||||
|
||||
|
||||
Your task is to extract learning task memories from failures to prevent similar mistakes in future executions.
|
||||
Focus on identifying error patterns, missed opportunities, and alternative approaches.
|
||||
|
||||
|
||||
ANALYSIS FRAMEWORK:
|
||||
● FAILURE POINT IDENTIFICATION: Pinpoint where and why the steps went wrong
|
||||
● ERROR PATTERN ANALYSIS: Identify recurring mistakes or problematic approaches
|
||||
● ALTERNATIVE APPROACHES: Suggest what could have been done differently
|
||||
● PREVENTION STRATEGIES: Extract actionable insights to avoid similar failures
|
||||
|
||||
|
||||
EXTRACTION PRINCIPLES:
|
||||
● Extract GENERAL PRINCIPLES as well as SPECIFIC INSTRUCTIONS
|
||||
● Focus on PATTERNS and RULES as well as particular instances
|
||||
|
||||
|
||||
# Original Query
|
||||
{query}
|
||||
|
||||
|
||||
# Step Sequence Analysis
|
||||
{step_sequence}
|
||||
|
||||
|
||||
# Context Information
|
||||
{context}
|
||||
|
||||
|
||||
# Outcome
|
||||
This step sequence was part of a {outcome} trajectory.
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Generate 1-3 step-level failure prevention insights as JSON objects:
|
||||
```json
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
"""Memory deduplication operation for task memory management.
|
||||
|
||||
This module provides operations to remove duplicate or highly similar task
|
||||
memories by comparing embeddings and calculating similarity scores.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
|
@ -8,6 +15,13 @@ from reme_ai.schema.memory import BaseMemory
|
|||
|
||||
@C.register_op()
|
||||
class MemoryDeduplicationOp(BaseAsyncOp):
|
||||
"""Remove duplicate task memories using embedding similarity.
|
||||
|
||||
This operation identifies and removes duplicate or highly similar task
|
||||
memories by comparing their embeddings against both existing memories
|
||||
in the vector store and other memories in the current batch.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -25,7 +39,9 @@ class MemoryDeduplicationOp(BaseAsyncOp):
|
|||
deduplicated_task_memories = await self._deduplicate_task_memories(task_memories)
|
||||
|
||||
logger.info(
|
||||
f"Deduplication complete: {len(deduplicated_task_memories)} deduplicated task memories out of {len(task_memories)}")
|
||||
f"Deduplication complete: {len(deduplicated_task_memories)} deduplicated "
|
||||
f"task memories out of {len(task_memories)}",
|
||||
)
|
||||
|
||||
# Update context
|
||||
self.context.response.metadata["memory_list"] = deduplicated_task_memories
|
||||
|
|
@ -70,24 +86,25 @@ class MemoryDeduplicationOp(BaseAsyncOp):
|
|||
async def _get_existing_task_memory_embeddings(self, workspace_id: str) -> List[List[float]]:
|
||||
"""Get embeddings of existing task memories"""
|
||||
try:
|
||||
if not hasattr(self, 'vector_store') or not self.vector_store or not workspace_id:
|
||||
if not hasattr(self, "vector_store") or not self.vector_store or not workspace_id:
|
||||
return []
|
||||
|
||||
# Query existing task memory nodes
|
||||
existing_nodes = await self.vector_store.async_search(
|
||||
query="...", # Empty query to get all
|
||||
workspace_id=workspace_id,
|
||||
top_k=self.op_params.get("max_existing_task_memories", 1000)
|
||||
top_k=self.op_params.get("max_existing_task_memories", 1000),
|
||||
)
|
||||
|
||||
# Extract embeddings
|
||||
existing_embeddings = []
|
||||
for node in existing_nodes:
|
||||
if hasattr(node, 'embedding') and node.embedding:
|
||||
if hasattr(node, "embedding") and node.embedding:
|
||||
existing_embeddings.append(node.embedding)
|
||||
|
||||
logger.debug(
|
||||
f"Retrieved {len(existing_embeddings)} existing task memory embeddings from workspace {workspace_id}")
|
||||
f"Retrieved {len(existing_embeddings)} existing task memory embeddings from workspace {workspace_id}",
|
||||
)
|
||||
return existing_embeddings
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -112,9 +129,12 @@ class MemoryDeduplicationOp(BaseAsyncOp):
|
|||
logger.error(f"Error generating embedding for task memory: {e}")
|
||||
return None
|
||||
|
||||
def _is_similar_to_existing_task_memories(self, current_embedding: List[float],
|
||||
existing_embeddings: List[List[float]],
|
||||
threshold: float) -> bool:
|
||||
def _is_similar_to_existing_task_memories(
|
||||
self,
|
||||
current_embedding: List[float],
|
||||
existing_embeddings: List[List[float]],
|
||||
threshold: float,
|
||||
) -> bool:
|
||||
"""Check if current embedding is similar to existing embeddings"""
|
||||
for existing_embedding in existing_embeddings:
|
||||
similarity = self._calculate_cosine_similarity(current_embedding, existing_embedding)
|
||||
|
|
@ -123,9 +143,13 @@ class MemoryDeduplicationOp(BaseAsyncOp):
|
|||
return True
|
||||
return False
|
||||
|
||||
def _is_similar_to_current_task_memories(self, current_embedding: List[float],
|
||||
current_task_memories: List[BaseMemory],
|
||||
threshold: float) -> bool:
|
||||
def _is_similar_to_current_task_memories(
|
||||
self,
|
||||
current_embedding: List[float],
|
||||
current_task_memories: List[BaseMemory],
|
||||
threshold: float,
|
||||
) -> bool:
|
||||
"""Check if current embedding is similar to other memories in current batch."""
|
||||
for existing_task_memory in current_task_memories:
|
||||
existing_embedding = self._get_task_memory_embedding(existing_task_memory)
|
||||
if existing_embedding is None:
|
||||
|
|
|
|||
|
|
@ -1,18 +1,32 @@
|
|||
"""Memory validation operation for task memory quality control.
|
||||
|
||||
This module provides operations to validate the quality of extracted task
|
||||
memories using LLM-based evaluation, ensuring only high-quality memories
|
||||
are stored.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message
|
||||
from reme_ai.schema.memory import BaseMemory
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class MemoryValidationOp(BaseAsyncOp):
|
||||
"""Validate quality of extracted task memories.
|
||||
|
||||
This operation uses LLM-based evaluation to assess the quality of extracted
|
||||
task memories, filtering out low-quality or invalid memories based on
|
||||
validation scores and criteria.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
|
|
@ -59,15 +73,16 @@ class MemoryValidationOp(BaseAsyncOp):
|
|||
prompt = self.prompt_format(
|
||||
prompt_name="task_memory_validation_prompt",
|
||||
condition=task_memory.when_to_use,
|
||||
task_memory_content=task_memory.content)
|
||||
task_memory_content=task_memory.content,
|
||||
)
|
||||
|
||||
def parse_validation(message: Message) -> Dict[str, Any]:
|
||||
def parse_validation(message: FlowMessage) -> Dict[str, Any]:
|
||||
try:
|
||||
response_content = message.content
|
||||
|
||||
# Parse validation result
|
||||
# Extract JSON blocks
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_pattern = r"```json\s*([\s\S]*?)\s*```"
|
||||
json_blocks = re.findall(json_pattern, response_content)
|
||||
|
||||
if json_blocks:
|
||||
|
|
@ -85,8 +100,11 @@ class MemoryValidationOp(BaseAsyncOp):
|
|||
"is_valid": is_valid and score >= validation_threshold,
|
||||
"score": score,
|
||||
"feedback": response_content,
|
||||
"reason": "" if (
|
||||
is_valid and score >= validation_threshold) else f"Low validation score ({score:.2f}) or marked as invalid"
|
||||
"reason": (
|
||||
""
|
||||
if (is_valid and score >= validation_threshold)
|
||||
else f"Low validation score ({score:.2f}) or marked as invalid"
|
||||
),
|
||||
}
|
||||
|
||||
except Exception as e_inner:
|
||||
|
|
@ -95,10 +113,13 @@ class MemoryValidationOp(BaseAsyncOp):
|
|||
"is_valid": False,
|
||||
"score": 0.0,
|
||||
"feedback": "",
|
||||
"reason": f"Parse error: {str(e_inner)}"
|
||||
"reason": f"Parse error: {str(e_inner)}",
|
||||
}
|
||||
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_validation)
|
||||
return await self.llm.achat(
|
||||
messages=[FlowMessage(role=Role.USER, content=prompt)],
|
||||
callback_fn=parse_validation,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"LLM validation failed: {e}")
|
||||
|
|
@ -106,5 +127,5 @@ class MemoryValidationOp(BaseAsyncOp):
|
|||
"is_valid": False,
|
||||
"score": 0.0,
|
||||
"feedback": "",
|
||||
"reason": f"LLM validation error: {str(e)}"
|
||||
"reason": f"LLM validation error: {str(e)}",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
task_memory_validation_prompt: |
|
||||
You are an expert AI analyst tasked with validating the quality and usefulness of extracted step-level task memories.
|
||||
|
||||
|
||||
Your task is to assess whether the extracted task memory is actionable, accurate, and valuable for future agent executions.
|
||||
|
||||
|
||||
VALIDATION CRITERIA:
|
||||
● ACTIONABILITY: Is the task memory specific enough to guide future actions?
|
||||
● ACCURACY: Does the task memory correctly reflect the patterns observed?
|
||||
● RELEVANCE: Is the task memory applicable to similar future scenarios?
|
||||
● CLARITY: Is the task memory clearly articulated and understandable?
|
||||
● UNIQUENESS: Does the task memory provide novel insights or common knowledge?
|
||||
|
||||
|
||||
# Task Memory to Validate
|
||||
Condition: {condition}
|
||||
Task Memory Content: {task_memory_content}
|
||||
|
||||
|
||||
OUTPUT FORMAT:
|
||||
Provide validation assessment:
|
||||
```json
|
||||
|
|
@ -24,6 +24,6 @@ task_memory_validation_prompt: |
|
|||
"recommendations": "Suggestions for improvement if applicable"
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
Score should be between 0.0 (poor quality) and 1.0 (excellent quality).
|
||||
Mark as invalid if score is below 0.3 or if there are fundamental issues with the task memory.
|
||||
|
|
|
|||
|
|
@ -1,9 +1,16 @@
|
|||
"""Simple comparative summary operation for task memory generation.
|
||||
|
||||
This module provides a simplified operation to extract task memories by
|
||||
comparing trajectories with different scores for the same task.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import List, Dict
|
||||
|
||||
from flowllm import C, BaseAsyncOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from flowllm.core.context import C
|
||||
from flowllm.core.enumeration import Role
|
||||
from flowllm.core.op import BaseAsyncOp
|
||||
from flowllm.core.schema import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -13,13 +20,30 @@ from reme_ai.utils.op_utils import merge_messages_content
|
|||
|
||||
@C.register_op()
|
||||
class SimpleComparativeSummaryOp(BaseAsyncOp):
|
||||
"""Extract task memories by comparing trajectories with different scores.
|
||||
|
||||
This operation compares the highest and lowest scoring trajectories for
|
||||
each task to extract comparative insights and best practices.
|
||||
"""
|
||||
|
||||
file_path: str = __file__
|
||||
|
||||
async def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
|
||||
summary_prompt = self.prompt_format(prompt_name="summary_prompt",
|
||||
execution_process_a=merge_messages_content(trajectory_a.messages),
|
||||
execution_process_b=merge_messages_content(trajectory_b.messages),
|
||||
summary_example=self.get_prompt("summary_example"))
|
||||
"""Compare two trajectories and extract comparative task memories.
|
||||
|
||||
Args:
|
||||
trajectory_a: First trajectory to compare (typically higher scoring)
|
||||
trajectory_b: Second trajectory to compare (typically lower scoring)
|
||||
|
||||
Returns:
|
||||
List of extracted task memories from the comparison
|
||||
"""
|
||||
summary_prompt = self.prompt_format(
|
||||
prompt_name="summary_prompt",
|
||||
execution_process_a=merge_messages_content(trajectory_a.messages),
|
||||
execution_process_b=merge_messages_content(trajectory_b.messages),
|
||||
summary_example=self.get_prompt("summary_example"),
|
||||
)
|
||||
|
||||
def parse_content(message: Message):
|
||||
content = message.content
|
||||
|
|
@ -33,10 +57,14 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
|
|||
when_to_use = tm_dict.get("when_to_use", "").strip()
|
||||
task_memory_content = tm_dict.get("experience", "").strip()
|
||||
if when_to_use and task_memory_content:
|
||||
task_memory_list.append(TaskMemory(workspace_id=self.context.get("workspace_id", ""),
|
||||
when_to_use=when_to_use,
|
||||
content=task_memory_content,
|
||||
author=getattr(self.llm, 'model_name', 'system')))
|
||||
task_memory_list.append(
|
||||
TaskMemory(
|
||||
workspace_id=self.context.get("workspace_id", ""),
|
||||
when_to_use=when_to_use,
|
||||
content=task_memory_content,
|
||||
author=getattr(self.llm, "model_name", "system"),
|
||||
),
|
||||
)
|
||||
|
||||
return task_memory_list
|
||||
|
||||
|
|
@ -44,9 +72,17 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
|
|||
logger.exception(f"parse content failed!\n{content}")
|
||||
raise e
|
||||
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
|
||||
return await self.llm.achat(
|
||||
messages=[FlowMessage(role=Role.USER, content=summary_prompt)],
|
||||
callback_fn=parse_content,
|
||||
)
|
||||
|
||||
async def async_execute(self):
|
||||
"""Execute the comparative summary operation.
|
||||
|
||||
Groups trajectories by task_id, compares the highest and lowest scoring
|
||||
trajectories for each task, and extracts task memories from the comparison.
|
||||
"""
|
||||
trajectories: list = self.context.get("trajectories", [])
|
||||
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
|
||||
|
||||
|
|
@ -57,14 +93,16 @@ class SimpleComparativeSummaryOp(BaseAsyncOp):
|
|||
task_id_dict[trajectory.task_id].append(trajectory)
|
||||
|
||||
memory_list = []
|
||||
for task_id, task_trajectories in task_id_dict.items():
|
||||
for _, task_trajectories in task_id_dict.items():
|
||||
task_trajectories: List[Trajectory] = sorted(task_trajectories, key=lambda x: x.score, reverse=True)
|
||||
if len(task_trajectories) < 2:
|
||||
continue
|
||||
|
||||
if task_trajectories[0].score > task_trajectories[-1].score:
|
||||
task_memories = await self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
|
||||
trajectory_b=task_trajectories[-1])
|
||||
task_memories = await self.compare_summary_trajectory(
|
||||
trajectory_a=task_trajectories[0],
|
||||
trajectory_b=task_trajectories[-1],
|
||||
)
|
||||
memory_list.extend(task_memories)
|
||||
|
||||
self.context.response.answer = json.dumps([x.model_dump() for x in memory_list])
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue