mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
rename reme4->reme
This commit is contained in:
parent
26cb5ca62f
commit
195f97857a
602 changed files with 194 additions and 81331 deletions
|
|
@ -1,360 +0,0 @@
|
|||
# flake8: noqa: E402, E501
|
||||
# pylint: disable=E0611
|
||||
"""A minimal ReAct Agent for AppWorld tasks."""
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
import json
|
||||
import datetime
|
||||
from typing import List, Any
|
||||
|
||||
|
||||
import ray
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from loguru import logger
|
||||
from openai import OpenAI
|
||||
from jinja2 import Template
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from prompt import NEW_PROMPT_TEMPLATE
|
||||
from appworld import AppWorld, load_task_ids
|
||||
|
||||
os.environ["APPWORLD_ROOT"] = "."
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
|
||||
@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 = 129024,
|
||||
num_trials: int = 1,
|
||||
use_memory: bool = False,
|
||||
memory_base_url: str = "http://0.0.0.0:8002/",
|
||||
use_memory_addition: bool = False,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
):
|
||||
|
||||
self.index: int = index
|
||||
self.task_ids: List[str] = task_ids
|
||||
self.experiment_name: str = experiment_name
|
||||
self.model_name: str = model_name
|
||||
self.temperature: float = temperature
|
||||
self.max_interactions: int = max_interactions
|
||||
self.max_response_size: int = max_response_size
|
||||
self.num_trials: int = num_trials
|
||||
self.use_memory: bool = use_memory
|
||||
self.use_memory_addition: bool = use_memory_addition if use_memory else False
|
||||
self.use_memory_deletion: bool = use_memory_deletion if use_memory else False
|
||||
self.delete_freq: int = delete_freq
|
||||
self.freq_threshold: int = freq_threshold
|
||||
self.utility_threshold: float = utility_threshold
|
||||
|
||||
self.llm_client = OpenAI()
|
||||
self.memory_base_url: str = memory_base_url
|
||||
|
||||
self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)]
|
||||
self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)]
|
||||
|
||||
for run_id in range(num_trials):
|
||||
for _ in range(len(task_ids)):
|
||||
self.retrieved_memory_list[run_id].append([])
|
||||
self.history[run_id].append([])
|
||||
|
||||
def call_llm(self, messages: list) -> str:
|
||||
"""Call the LLM to generate a response to the messages."""
|
||||
for i in range(100):
|
||||
try:
|
||||
response = self.llm_client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=messages,
|
||||
temperature=self.temperature,
|
||||
extra_body={"enable_thinking": False},
|
||||
seed=0,
|
||||
)
|
||||
|
||||
return response.choices[0].message.content
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"encounter error with {e.args}")
|
||||
time.sleep(1 + i * 10)
|
||||
|
||||
return "call llm error"
|
||||
|
||||
def prompt_messages(self, run_id, task_index, previous_memories: None, world: AppWorld):
|
||||
"""Prompt the messages to the LLM."""
|
||||
app_descriptions = json.dumps(
|
||||
[{"name": k, "description": v} for (k, v) in world.task.app_descriptions.items()],
|
||||
indent=1,
|
||||
)
|
||||
dictionary = {"supervisor": world.task.supervisor, "app_descriptions": app_descriptions}
|
||||
sys_prompt = Template(NEW_PROMPT_TEMPLATE.lstrip()).render(dictionary)
|
||||
query = world.task.instruction
|
||||
if self.use_memory:
|
||||
if len(previous_memories) == 0:
|
||||
response = self.get_memory(world.task.instruction)
|
||||
if response and "memory_list" in response["metadata"]:
|
||||
self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"]
|
||||
task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1:", response["answer"])
|
||||
logger.info(f"loaded task_memory: {task_memory}")
|
||||
query = (
|
||||
"Task:\n"
|
||||
+ query
|
||||
+ "\n\nSome Related Experience to help you to complete the task:\n"
|
||||
+ task_memory
|
||||
)
|
||||
else:
|
||||
formatted_memories = []
|
||||
for i, memory in enumerate(previous_memories, 1):
|
||||
condition = memory["when_to_use"]
|
||||
memory_content = memory["content"]
|
||||
memory_text = f"Experience {i}:\n When to use: {condition}\n Content: {memory_content}\n"
|
||||
formatted_memories.append(memory_text)
|
||||
query = (
|
||||
"Task:\n"
|
||||
+ query
|
||||
+ "\n\nSome Related Experience to help you to complete the task:\n"
|
||||
+ "\n".join(formatted_memories)
|
||||
)
|
||||
messages = [
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": query},
|
||||
]
|
||||
self.history[run_id][task_index] = messages
|
||||
|
||||
@staticmethod
|
||||
def get_reward(world) -> float:
|
||||
"""Get the reward for the Appworld world."""
|
||||
tracker = world.evaluate()
|
||||
num_passes = len(tracker.passes)
|
||||
num_failures = len(tracker.failures)
|
||||
return num_passes / (num_passes + num_failures)
|
||||
|
||||
def extract_code_and_fix_content(
|
||||
self,
|
||||
text: str,
|
||||
ignore_multiple_calls=True,
|
||||
) -> tuple[str, str]:
|
||||
"""Extract the code and fix the content."""
|
||||
full_code_regex = r"```python\n(.*?)```"
|
||||
partial_code_regex = r".*```python\n(.*)"
|
||||
|
||||
original_text = text
|
||||
output_code = ""
|
||||
match_end = 0
|
||||
# Handle multiple calls
|
||||
for re_match in re.finditer(full_code_regex, original_text, flags=re.DOTALL):
|
||||
code = re_match.group(1).strip()
|
||||
if ignore_multiple_calls:
|
||||
text = original_text[: re_match.end()]
|
||||
return code, text
|
||||
output_code += code + "\n"
|
||||
match_end = re_match.end()
|
||||
# check for partial code match at end (no terminating ```) following the last match
|
||||
partial_match = re.match(
|
||||
partial_code_regex,
|
||||
original_text[match_end:],
|
||||
flags=re.DOTALL,
|
||||
)
|
||||
if partial_match:
|
||||
output_code += partial_match.group(1).strip()
|
||||
# terminated due to stop condition. Add stop condition to output.
|
||||
if not text.endswith("\n"):
|
||||
text = text + "\n"
|
||||
text = text + "```"
|
||||
if len(output_code) == 0:
|
||||
return text, text
|
||||
else:
|
||||
return output_code, text
|
||||
|
||||
def execute(self):
|
||||
"""Execute the Appworld tasks."""
|
||||
result = []
|
||||
counter = 0
|
||||
for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"run_index={self.index}")):
|
||||
t_result = None
|
||||
previous_memories = []
|
||||
# Run each task num_trials times
|
||||
for run_id in range(self.num_trials):
|
||||
start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
with AppWorld(task_id=task_id, experiment_name=f"{self.experiment_name}_run_{run_id}") as world:
|
||||
before_score = self.get_reward(world)
|
||||
for i in range(self.max_interactions):
|
||||
if i == 0:
|
||||
self.prompt_messages(
|
||||
run_id=run_id,
|
||||
task_index=task_index,
|
||||
previous_memories=previous_memories,
|
||||
world=world,
|
||||
)
|
||||
code_msg = self.call_llm(self.history[run_id][task_index])
|
||||
code, _ = self.extract_code_and_fix_content(code_msg)
|
||||
self.history[run_id][task_index].append({"role": "assistant", "content": code})
|
||||
|
||||
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]
|
||||
self.history[run_id][task_index].append(
|
||||
{"role": "user", "content": "Output:\n```\n" + output + "```\n\n"},
|
||||
)
|
||||
|
||||
if world.task_completed():
|
||||
break
|
||||
|
||||
after_score = self.get_reward(world)
|
||||
uplift_score = after_score - before_score
|
||||
|
||||
if self.use_memory:
|
||||
if self.use_memory_addition:
|
||||
new_traj_list = [
|
||||
self.get_traj_from_task_history(task_id, self.history[run_id][task_index], after_score),
|
||||
]
|
||||
previous_memories = self.summary_memory(new_traj_list)
|
||||
if after_score == 1:
|
||||
self.add_memory(previous_memories)
|
||||
|
||||
# update the freq & utility attributes of retrieved memories
|
||||
update_utility: bool = after_score == 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()
|
||||
|
||||
t_result = {
|
||||
"task_id": world.task_id,
|
||||
"run_id": run_id,
|
||||
"experiment_name": self.experiment_name,
|
||||
"task_completed": world.task_completed(),
|
||||
"before_score": before_score,
|
||||
"after_score": after_score,
|
||||
"uplift_score": uplift_score,
|
||||
"task_history": self.history[run_id][task_index],
|
||||
"task_start_time": start_time,
|
||||
}
|
||||
if after_score == 1:
|
||||
break
|
||||
result.append(t_result)
|
||||
|
||||
return result
|
||||
|
||||
def handle_api_response(self, response: requests.Response):
|
||||
"""Handle API response with proper error checking"""
|
||||
if response.status_code != 200:
|
||||
print(f"Error: {response.status_code}")
|
||||
print(response.text)
|
||||
return None
|
||||
|
||||
return response.json()
|
||||
|
||||
def get_memory(self, query: str):
|
||||
"""Retrieve relevant task memories based on a query"""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}retrieve_task_memory",
|
||||
json={
|
||||
"query": query,
|
||||
"enable_llm_rerank": False,
|
||||
"enable_score_filter": False,
|
||||
"top_k": 5,
|
||||
"enable_llm_rewrite": False,
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
logger.info(f"query: {query}, response: {result}")
|
||||
return result
|
||||
|
||||
def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float):
|
||||
"""Get the trajectory from the task history."""
|
||||
pattern = r"\n\nSome Related Experience to help you to complete the task:.*"
|
||||
task_history[1]["content"] = re.sub(pattern, "", task_history[1]["content"], flags=re.DOTALL)
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"messages": task_history,
|
||||
"score": reward,
|
||||
}
|
||||
|
||||
def summary_memory(self, trajectories):
|
||||
"""Generate a summary of conversation messages and create task memories"""
|
||||
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}summary_task_memory",
|
||||
json={
|
||||
"trajectories": trajectories,
|
||||
"success_threshold": 1.0,
|
||||
"enable_soft_comparison": True,
|
||||
"validation_threshold": 0.5,
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
if not result:
|
||||
return []
|
||||
|
||||
# Extract memory list from response
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
print(f"Task memory list created: {len(memory_list)} memories")
|
||||
return memory_list
|
||||
|
||||
def add_memory(self, memory_list):
|
||||
"""Add the memory to the memory pool."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}add_task_memory",
|
||||
json={
|
||||
"memory_list": memory_list,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
def update_memory_information(self, memory_list, update_utility: bool = False):
|
||||
"""Update the memory information."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}record_task_memory",
|
||||
json={
|
||||
"memory_list": memory_list,
|
||||
"update_utility": update_utility,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(response.json())
|
||||
|
||||
def delete_memory(self):
|
||||
"""Delete the memory from the memory pool."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}delete_task_memory",
|
||||
json={
|
||||
"freq_threshold": self.freq_threshold,
|
||||
"utility_threshold": self.utility_threshold,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function to run the Appworld React Agent."""
|
||||
dataset_name = "train"
|
||||
task_ids = load_task_ids(dataset_name)
|
||||
agent = AppworldReactAgent(index=0, task_ids=task_ids[0:1], experiment_name=dataset_name, num_trials=1)
|
||||
result = agent.execute()
|
||||
logger.info(f"result={json.dumps(result)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,660 +0,0 @@
|
|||
# flake8: noqa: E402, E501
|
||||
# pylint: disable=C0114,C0301
|
||||
# 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.
|
||||
PROMPT_TEMPLATE = """
|
||||
USER:
|
||||
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
|
||||
|
||||
To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
|
||||
|
||||
Here are three key APIs that you need to know to get more information
|
||||
|
||||
# To get a list of apps that are available to you.
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
|
||||
# To get the list of apis under any app listed above, e.g. supervisor
|
||||
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
|
||||
|
||||
# To get the specification of a particular api, e.g. supervisor app's show_account_passwords
|
||||
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
|
||||
|
||||
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task.
|
||||
|
||||
For example, consider the task:
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
|
||||
Task:
|
||||
|
||||
What is the password for my Spotify account?
|
||||
|
||||
ASSISTANT:
|
||||
# Okay. Lets first find which apps are available to get the password by looking at the app descriptions.
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
|
||||
USER:
|
||||
[
|
||||
{
|
||||
"name": "api_docs",
|
||||
"description": "An app to search and explore API documentation."
|
||||
},
|
||||
{
|
||||
"name": "supervisor",
|
||||
"description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task."
|
||||
},
|
||||
...
|
||||
{
|
||||
"name": "spotify",
|
||||
"description": "A music streaming app to stream songs and manage song, album and playlist libraries."
|
||||
},
|
||||
{
|
||||
"name": "venmo",
|
||||
"description": "A social payment app to send, receive and request money to and from others."
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# Looks like the supervisor app could help me with that. Lets see what apis are available under this app.
|
||||
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
|
||||
|
||||
|
||||
USER:
|
||||
[
|
||||
...
|
||||
"show_account_passwords : Show your supervisor's account passwords."
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure.
|
||||
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
|
||||
|
||||
USER:
|
||||
{
|
||||
'app_name': 'supervisor',
|
||||
'api_name': 'show_account_passwords',
|
||||
'path': '/account_passwords',
|
||||
'method': 'GET',
|
||||
'description': "Show your supervisor's app account passwords.",
|
||||
'parameters': [],
|
||||
'response_schemas': {
|
||||
'success': [{'account_name': 'string', 'password': 'string'}],
|
||||
'failure': {'message': 'string'}
|
||||
}
|
||||
}
|
||||
|
||||
ASSISTANT:
|
||||
# Okay, it requires no arguments. So I can just call it directly.
|
||||
print(apis.supervisor.show_account_passwords())
|
||||
|
||||
USER:
|
||||
[
|
||||
{
|
||||
"account_name": "spotify",
|
||||
"password": "dummy_spotify_pass"
|
||||
},
|
||||
{
|
||||
"account_name": "file_system",
|
||||
"password": "dummy_fs_pass"
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# So the Spotify password is an entry in the `passwords` list with the account_name=spotify.
|
||||
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
|
||||
print(spotify_password)
|
||||
|
||||
|
||||
USER:
|
||||
dummy_spotify_pass
|
||||
|
||||
ASSISTANT:
|
||||
# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer.
|
||||
apis.supervisor.complete_task(answer=spotify_password)
|
||||
|
||||
|
||||
USER:
|
||||
Marked the active task complete.
|
||||
|
||||
|
||||
----------------------------------------------
|
||||
|
||||
USER:
|
||||
**Key instructions and disclaimers**:
|
||||
|
||||
1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself.
|
||||
2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments.
|
||||
3. You can use the variables from the previous code blocks in the subsequent code blocks.
|
||||
4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
|
||||
5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them.
|
||||
6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions.
|
||||
7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library.
|
||||
8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation.
|
||||
9. For APIs that return results in "pages", make sure to consider all pages.
|
||||
10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is.
|
||||
11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone.
|
||||
12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list.
|
||||
13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app.
|
||||
14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=<answer>)`. For tasks that do not require an answer, just skip the answer argument or pass it as None.
|
||||
15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten".
|
||||
16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit.
|
||||
17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else.
|
||||
|
||||
USER:
|
||||
Using these APIs, now generate code to solve the actual task:
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
|
||||
Task:
|
||||
|
||||
{{ instruction }}
|
||||
"""
|
||||
|
||||
PROMPT_TEMPLATE_WITH_EXPERIENCE = """
|
||||
USER:
|
||||
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
|
||||
|
||||
To do this, you will need to interact with app/s (e.g., spotify, venmo, etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
|
||||
|
||||
Here are three key APIs that you need to know to get more information
|
||||
|
||||
# To get a list of apps that are available to you.
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
|
||||
# To get the list of apis under any app listed above, e.g. supervisor
|
||||
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
|
||||
|
||||
# To get the specification of a particular api, e.g. supervisor app's show_account_passwords
|
||||
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
|
||||
|
||||
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that the environment will execute, to solve the task.
|
||||
|
||||
For example, consider the task:
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
|
||||
Task:
|
||||
|
||||
What is the password for my Spotify account?
|
||||
|
||||
ASSISTANT:
|
||||
# Okay. Lets first find which apps are available to get the password by looking at the app descriptions.
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
|
||||
USER:
|
||||
[
|
||||
{
|
||||
"name": "api_docs",
|
||||
"description": "An app to search and explore API documentation."
|
||||
},
|
||||
{
|
||||
"name": "supervisor",
|
||||
"description": "An app to access supervisor's personal information, account credentials, addresses, payment cards, and manage the assigned task."
|
||||
},
|
||||
...
|
||||
{
|
||||
"name": "spotify",
|
||||
"description": "A music streaming app to stream songs and manage song, album and playlist libraries."
|
||||
},
|
||||
{
|
||||
"name": "venmo",
|
||||
"description": "A social payment app to send, receive and request money to and from others."
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# Looks like the supervisor app could help me with that. Lets see what apis are available under this app.
|
||||
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
|
||||
|
||||
|
||||
USER:
|
||||
[
|
||||
...
|
||||
"show_account_passwords : Show your supervisor's account passwords."
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# I can use `show_account_passwords` to get the passwords. Let me see its detailed specification to understand its arguments and output structure.
|
||||
print(apis.api_docs.show_api_doc(app_name='supervisor', api_name='show_account_passwords'))
|
||||
|
||||
USER:
|
||||
{
|
||||
'app_name': 'supervisor',
|
||||
'api_name': 'show_account_passwords',
|
||||
'path': '/account_passwords',
|
||||
'method': 'GET',
|
||||
'description': "Show your supervisor's app account passwords.",
|
||||
'parameters': [],
|
||||
'response_schemas': {
|
||||
'success': [{'account_name': 'string', 'password': 'string'}],
|
||||
'failure': {'message': 'string'}
|
||||
}
|
||||
}
|
||||
|
||||
ASSISTANT:
|
||||
# Okay, it requires no arguments. So I can just call it directly.
|
||||
print(apis.supervisor.show_account_passwords())
|
||||
|
||||
USER:
|
||||
[
|
||||
{
|
||||
"account_name": "spotify",
|
||||
"password": "dummy_spotify_pass"
|
||||
},
|
||||
{
|
||||
"account_name": "file_system",
|
||||
"password": "dummy_fs_pass"
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
# So the Spotify password is an entry in the `passwords` list with the account_name=spotify.
|
||||
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
|
||||
print(spotify_password)
|
||||
|
||||
|
||||
USER:
|
||||
dummy_spotify_pass
|
||||
|
||||
ASSISTANT:
|
||||
# When the task is completed, I need to call apis.supervisor.complete_task(). If there is an answer, I need to pass it as an argument `answer`. I will pass the spotify_password as an answer.
|
||||
apis.supervisor.complete_task(answer=spotify_password)
|
||||
|
||||
|
||||
USER:
|
||||
Marked the active task complete.
|
||||
|
||||
|
||||
----------------------------------------------
|
||||
|
||||
USER:
|
||||
**Key instructions and disclaimers**:
|
||||
|
||||
1. The email addresses, access tokens and variables (e.g. spotify_password) in the example above were only for demonstration. Obtain the correct information by calling relevant APIs yourself.
|
||||
2. Only generate valid code blocks, i.e., do not put them in ```...``` or add any extra formatting. Any thoughts should be put as code comments.
|
||||
3. You can use the variables from the previous code blocks in the subsequent code blocks.
|
||||
4. Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
|
||||
5. The provided Python environment has access to its standard library. But modules and functions that have a risk of affecting the underlying OS, file system or process are disabled. You will get an error if do call them.
|
||||
6. Any reference to a file system in the task instructions means the file system *app*, operable via given APIs, and not the actual file system the code is running on. So do not write code making calls to os-level modules and functions.
|
||||
7. To interact with apps, only use the provided APIs, and not the corresponding Python packages. E.g., do NOT use `spotipy` for Spotify. Remember, the environment only has the standard library.
|
||||
8. The provided API documentation has both the input arguments and the output JSON schemas. All calls to APIs and parsing its outputs must be as per this documentation.
|
||||
9. For APIs that return results in "pages", make sure to consider all pages.
|
||||
10. To obtain current date or time, use Python functions like `datetime.now()` or obtain it from the phone app. Do not rely on your existing knowledge of what the current date or time is.
|
||||
11. For all temporal requests, use proper time boundaries, e.g., if I ask for something that happened yesterday, make sure to consider the time between 00:00:00 and 23:59:59. All requests are concerning a single, default (no) time zone.
|
||||
12. Any reference to my friends, family or any other person or relation refers to the people in my phone's contacts list.
|
||||
13. All my personal information, and information about my app account credentials, physical addresses and owned payment cards are stored in the "supervisor" app. You can access them via the APIs provided by the supervisor app.
|
||||
14. Once you have completed the task, call `apis.supervisor.complete_task()`. If the task asks for some information, return it as the answer argument, i.e. call `apis.supervisor.complete_task(answer=<answer>)`. For tasks that do not require an answer, just skip the answer argument or pass it as None.
|
||||
15. The answers, when given, should be just entity or number, not full sentences, e.g., `answer=10` for "How many songs are in the Spotify queue?". When an answer is a number, it should be in numbers, not in words, e.g., "10" and not "ten".
|
||||
16. You can also pass `status="fail"` in the complete_task API if you are sure you cannot solve it and want to exit.
|
||||
17. You must make all decisions completely autonomously and not ask for any clarifications or confirmations from me or anyone else.
|
||||
18. Some Related Experience to help you to complete the task:
|
||||
{{experience}}
|
||||
|
||||
USER:
|
||||
Using these APIs, now generate code to solve the actual task:
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
|
||||
Task:
|
||||
|
||||
{{ instruction }}
|
||||
"""
|
||||
|
||||
NEW_PROMPT_TEMPLATE = """
|
||||
USER:
|
||||
I am your supervisor and you are a super intelligent AI Assistant whose job is to achieve my day-to-day tasks completely autonomously.
|
||||
|
||||
To do this, you will need to interact with app/s (e.g., spotify, venmo etc) using their associated APIs on my behalf. For this you will undertake a *multi-step conversation* using a python REPL environment. That is, you will write the python code and the environment will execute it and show you the result, based on which, you will write python code for the next step and so on, until you've achieved the goal. This environment will let you interact with app/s using their associated APIs on my behalf.
|
||||
|
||||
Here are three key APIs that you need to know to get more information
|
||||
|
||||
# To get a list of apps that are available to you.
|
||||
|
||||
```python
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
```
|
||||
|
||||
# To get the list of apis under any app listed above, e.g. spotify
|
||||
|
||||
```python
|
||||
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
|
||||
```
|
||||
|
||||
# To get the specification of a particular api, e.g. spotify app's login api
|
||||
|
||||
```python
|
||||
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login'))
|
||||
```
|
||||
|
||||
Each code execution will produce an output that you can use in subsequent calls. Using these APIs, you can now generate code, that I will execute, to solve the task. Let's start with the task
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
Task: How many playlists do I have in Spotify?
|
||||
|
||||
ASSISTANT:
|
||||
Okay. Lets first find which APIs are available to use in Spotify.
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
|
||||
```
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
[
|
||||
...
|
||||
"login : Login to your account.",
|
||||
"logout : Logout from your account.",
|
||||
...
|
||||
]
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
Okay. Looks like I can use the `login` api. Lets find its specifications.
|
||||
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='login'))
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
{
|
||||
"app_name": "spotify",
|
||||
"api_name": "login",
|
||||
"path": "/auth/token",
|
||||
"method": "POST",
|
||||
"description": "Login to your account.",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "username",
|
||||
"type": "string",
|
||||
"required": true,
|
||||
"description": "Your account email.",
|
||||
"default": null,
|
||||
"constraints": []
|
||||
},
|
||||
{
|
||||
"name": "password",
|
||||
"type": "string",
|
||||
"required": true,
|
||||
"description": "Your account password.",
|
||||
"default": null,
|
||||
"constraints": []
|
||||
}
|
||||
],
|
||||
"response_schemas": {
|
||||
"success": {
|
||||
"token_type": "string",
|
||||
"access_token": "string"
|
||||
},
|
||||
"failure": {
|
||||
"message": "string"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
I need the supervisor's username and password. Lets see if any app can help me get that.
|
||||
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_app_descriptions())
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
{{ app_descriptions }}
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
Looks like the supervisor app could help me with that.
|
||||
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_api_descriptions(app_name='supervisor'))
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
[
|
||||
...
|
||||
"show_account_passwords : Show your supervisor's account passwords."
|
||||
...
|
||||
]
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
I can use `show_account_passwords` to get the passwords. So I will call that and save it in a variable for future reference.
|
||||
Code:
|
||||
```python
|
||||
passwords=apis.supervisor.show_account_passwords()
|
||||
print(passwords)
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
[
|
||||
{
|
||||
"account_name": "spotify",
|
||||
"password": "dummy_spotify_pass"
|
||||
},
|
||||
{
|
||||
"account_name": "file_system",
|
||||
"password": "dummy_fs_pass"
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
So the Spotify password is an entry in the `passwords` list with the account_name=spotify. I can use it to login now with the provided email address: {{ supervisor.email }}.
|
||||
|
||||
Code:
|
||||
```python
|
||||
spotify_password = [account_password["account_name"] == "spotify" for account_password in passwords][0]["password"]
|
||||
login_result = apis.spotify.login(username='{{ supervisor.email }}', password=spotify_password)
|
||||
print(login_result)
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
{
|
||||
"token_type": "Bearer",
|
||||
"access_token": "fake_access_token"
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
First, let's check the available APIs for the Spotify app.
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_api_descriptions(app_name='spotify'))
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
[
|
||||
...
|
||||
{
|
||||
"name": "show_playlist_library",
|
||||
"description": "Get a list of playlists in the user's playlist library."
|
||||
},
|
||||
{
|
||||
"name": "show_playlist",
|
||||
"description": "Get detailed information about a specific playlist. You can view your own playlists or others' playlists if they are public."
|
||||
},
|
||||
...
|
||||
]
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
The `show_playlist_library` API seems to be the one we need. Let's check its specifications.
|
||||
Code:
|
||||
```python
|
||||
print(apis.api_docs.show_api_doc(app_name='spotify', api_name='show_playlist_library'))
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
{
|
||||
"app_name": "spotify",
|
||||
"api_name": "show_playlist_library",
|
||||
"path": "/private_playlists",
|
||||
"method": "GET",
|
||||
"description": "Get a list of playlists in the user's playlist library.",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "access_token",
|
||||
"type": "string",
|
||||
"required": true,
|
||||
"description": "Access token obtained from spotify app login.",
|
||||
"default": null,
|
||||
"constraints": []
|
||||
},
|
||||
{
|
||||
"name": "page_index",
|
||||
"type": "integer",
|
||||
"required": false,
|
||||
"description": "The index of the page to retrieve.",
|
||||
"default": 0,
|
||||
"constraints": [
|
||||
"value >= 0.0"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "page_limit",
|
||||
"type": "integer",
|
||||
"required": false,
|
||||
"description": "The maximum number of results to return per page.",
|
||||
"default": 5,
|
||||
"constraints": [
|
||||
"value >= 1.0, <= 20.0"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "is_public",
|
||||
"type": "boolean",
|
||||
"required": false,
|
||||
"description": "Whether to show public playlists or private playlists.",
|
||||
"default": null,
|
||||
"constraints": []
|
||||
}
|
||||
],
|
||||
"response_schema": [
|
||||
{
|
||||
"title": "string",
|
||||
"created_at": "2019-01-01T00:00:00",
|
||||
"is_public": true,
|
||||
"rating": 0.0,
|
||||
"like_count": 1,
|
||||
"owner_email": "user@example.com",
|
||||
"playlist_id": 1,
|
||||
"song_ids": [
|
||||
1
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
I need to page through all the playlists to get the list of playlists and save it in `playlists`.
|
||||
Code:
|
||||
```python
|
||||
page_index = 0
|
||||
playlists = []
|
||||
while page_index < 10:
|
||||
playlist_page = apis.spotify.show_playlist_library(access_token=spotify_access_token, page_index=page_index)
|
||||
if playlist_page:
|
||||
playlists.extend(playlist_page)
|
||||
page_index += 1
|
||||
else:
|
||||
break
|
||||
num_playlists = len(playlists)
|
||||
print(num_playlists)
|
||||
|
||||
```
|
||||
|
||||
USER:
|
||||
Output:
|
||||
```
|
||||
23
|
||||
```
|
||||
|
||||
|
||||
ASSISTANT:
|
||||
Now that the task is completed, I can call apis.supervisor.complete_task(). Since this task has an answer to be returned, I will pass that as an argument.
|
||||
|
||||
Code:
|
||||
```python
|
||||
apis.supervisor.complete_task(answer=num_playlists)
|
||||
```
|
||||
|
||||
|
||||
USER:
|
||||
Output:
|
||||
Marked the active task complete.
|
||||
|
||||
|
||||
----------------------------------------------
|
||||
|
||||
USER:
|
||||
**Key instructions**:
|
||||
(1) Make sure to end code blocks with ``` followed by a newline(\n).
|
||||
|
||||
(2) Remember you can use the variables in your code in subsequent code blocks.
|
||||
|
||||
(3) Remember that the email addresses, access tokens and variables (e.g. spotify_password) in the example above are not valid anymore.
|
||||
|
||||
(4) You can use the "supervisor" app to get information about my accounts and use the "phone" app to get information about friends and family.
|
||||
|
||||
(5) Always look at API specifications (using apis.api_docs.show_api_doc) before calling an API.
|
||||
|
||||
(6) Write small chunks of code and only one chunk of code in every step. Make sure everything is working correctly before making any irreversible change.
|
||||
|
||||
(7) Many APIs return items in "pages". Make sure to run through all the pages by looping over `page_index`.
|
||||
|
||||
(8) Once you have completed the task, make sure to call apis.supervisor.complete_task(). If the task asked for some information, return it as the answer argument, i.e. call apis.supervisor.complete_task(answer=<answer>). Many tasks do not require an answer, so in those cases, just call apis.supervisor.complete_task() i.e. do not pass any argument.
|
||||
|
||||
USER:
|
||||
Using these APIs, now generate code to solve the actual task:
|
||||
|
||||
My name is: {{ supervisor.first_name }} {{ supervisor.last_name }}. My personal email is {{ supervisor.email }} and phone number is {{ supervisor.phone_number }}.
|
||||
|
||||
"""
|
||||
|
|
@ -1,197 +0,0 @@
|
|||
# pylint: disable=E0611
|
||||
"""Run the Appworld React Agent."""
|
||||
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import ray
|
||||
import requests
|
||||
from loguru import logger
|
||||
from dotenv import load_dotenv
|
||||
from appworld import load_task_ids
|
||||
from appworld_react_agent import AppworldReactAgent
|
||||
|
||||
os.environ["APPWORLD_ROOT"] = "."
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
|
||||
def run_agent(
|
||||
run_index: int,
|
||||
max_workers: int,
|
||||
model_name: str,
|
||||
dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
num_trials: int = 1,
|
||||
use_memory: bool = False,
|
||||
memory_base_url: str = "http://0.0.0.0:8002/",
|
||||
use_memory_addition: bool = False,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
batch_size: int = 4,
|
||||
):
|
||||
"""Run the Appworld React Agent."""
|
||||
experiment_name = dataset_name + "_" + experiment_suffix
|
||||
path: Path = Path(f"./exp_result/{model_name}")
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
task_ids = load_task_ids(dataset_name)
|
||||
|
||||
result: list = []
|
||||
|
||||
def dump_file():
|
||||
with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f:
|
||||
for x in result:
|
||||
f.write(json.dumps(x) + "\n")
|
||||
|
||||
if max_workers > 1:
|
||||
# Process tasks in batches
|
||||
total_tasks = len(task_ids)
|
||||
num_batches = (total_tasks + batch_size - 1) // batch_size # Ceiling division
|
||||
|
||||
logger.info(f"Total tasks: {total_tasks}, Batch size: {batch_size}, Number of batches: {num_batches}")
|
||||
|
||||
for batch_idx in range(num_batches):
|
||||
# Initialize Ray for this batch
|
||||
start_idx = batch_idx * batch_size
|
||||
end_idx = min(start_idx + batch_size, total_tasks)
|
||||
batch_task_ids = task_ids[start_idx:end_idx]
|
||||
|
||||
logger.info(f"Starting batch {batch_idx + 1}/{num_batches} with {len(batch_task_ids)} tasks")
|
||||
|
||||
# Initialize Ray with the number of CPUs needed for this batch
|
||||
ray.init(num_cpus=len(batch_task_ids))
|
||||
|
||||
future_list: list = []
|
||||
for i, task_id in enumerate(batch_task_ids):
|
||||
actor = AppworldReactAgent.remote(
|
||||
index=start_idx + i,
|
||||
model_name=model_name,
|
||||
task_ids=[task_id],
|
||||
experiment_name=experiment_name,
|
||||
num_trials=num_trials,
|
||||
use_memory=use_memory,
|
||||
memory_base_url=memory_base_url,
|
||||
use_memory_addition=use_memory_addition,
|
||||
use_memory_deletion=use_memory_deletion,
|
||||
delete_freq=delete_freq,
|
||||
freq_threshold=freq_threshold,
|
||||
utility_threshold=utility_threshold,
|
||||
)
|
||||
future = actor.execute.remote()
|
||||
future_list.append(future)
|
||||
time.sleep(1)
|
||||
|
||||
logger.info(f"Batch {batch_idx + 1} submit complete, waiting for results...")
|
||||
|
||||
# Collect results from this batch
|
||||
for i, (task_id, future) in enumerate(zip(batch_task_ids, future_list)):
|
||||
try:
|
||||
t_result = ray.get(future)
|
||||
if t_result:
|
||||
if isinstance(t_result, list):
|
||||
result.extend(t_result)
|
||||
else:
|
||||
result.append(t_result)
|
||||
except Exception:
|
||||
logger.exception(f"run ray error with task_id={task_id}")
|
||||
|
||||
logger.info(f"Batch {batch_idx + 1}: task {i + 1}/{len(batch_task_ids)} complete")
|
||||
|
||||
# Shutdown Ray to free resources before next batch
|
||||
ray.shutdown()
|
||||
logger.info(f"Batch {batch_idx + 1}/{num_batches} complete, Ray resources released")
|
||||
|
||||
# Optional: small delay between batches
|
||||
if batch_idx < num_batches - 1:
|
||||
time.sleep(2)
|
||||
|
||||
dump_file()
|
||||
|
||||
else:
|
||||
agent = AppworldReactAgent(
|
||||
index=run_index,
|
||||
model_name=model_name,
|
||||
task_ids=task_ids,
|
||||
experiment_name=experiment_name,
|
||||
num_trials=num_trials,
|
||||
use_memory=use_memory,
|
||||
memory_base_url=memory_base_url,
|
||||
use_memory_addition=use_memory_addition,
|
||||
use_memory_deletion=use_memory_deletion,
|
||||
delete_freq=delete_freq,
|
||||
freq_threshold=freq_threshold,
|
||||
utility_threshold=utility_threshold,
|
||||
)
|
||||
result = agent.execute()
|
||||
|
||||
dump_file()
|
||||
|
||||
|
||||
def handle_api_response(response: requests.Response):
|
||||
"""Handle API response with proper error checking"""
|
||||
if response.status_code != 200:
|
||||
print(f"Error: {response.status_code}")
|
||||
print(response.text)
|
||||
return None
|
||||
|
||||
return response.json()
|
||||
|
||||
|
||||
def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"):
|
||||
"""Load memories from disk into the vector store"""
|
||||
response = requests.post(
|
||||
url=f"{api_url}load_memory",
|
||||
json={
|
||||
"load_file_path": path,
|
||||
"clear_existing": True,
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
if result:
|
||||
print(f"Memory loaded from {path}")
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function to run the Appworld React Agent."""
|
||||
max_workers = 16
|
||||
batch_size = 8
|
||||
|
||||
num_runs = 4 # Number of runs
|
||||
num_trials = 1 # for self-reflection
|
||||
model_name = "qwen3-8b"
|
||||
use_memory = True
|
||||
use_memory_addition = False
|
||||
use_memory_deletion = False
|
||||
memory_base_url = "http://0.0.0.0:8002/"
|
||||
|
||||
if use_memory:
|
||||
load_file_path = "docs/library/paper_data/task/appworld_qwen3_8b.jsonl"
|
||||
load_memory(load_file_path, memory_base_url)
|
||||
|
||||
for i in range(num_runs):
|
||||
run_agent(
|
||||
run_index=i,
|
||||
max_workers=max_workers,
|
||||
model_name=model_name,
|
||||
dataset_name="test_normal",
|
||||
experiment_suffix="with-fixed-memory",
|
||||
num_trials=num_trials,
|
||||
use_memory=use_memory,
|
||||
memory_base_url=memory_base_url,
|
||||
use_memory_addition=use_memory_addition,
|
||||
use_memory_deletion=use_memory_deletion,
|
||||
delete_freq=5,
|
||||
freq_threshold=5,
|
||||
utility_threshold=0.5,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,164 +0,0 @@
|
|||
"""Run the experiment statistic."""
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
|
||||
|
||||
def calculate_best_at_k(scores: list, k: int) -> float:
|
||||
"""
|
||||
Calculate best@k
|
||||
Divide scores into groups of size k, take the maximum value in each group,
|
||||
then average these maximum values
|
||||
|
||||
Args:
|
||||
scores: List of after_score values for all runs of a task
|
||||
k: Group size
|
||||
|
||||
Returns:
|
||||
best@k value
|
||||
"""
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i : i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def calculate_pass_at_k(scores: list, k: int) -> float:
|
||||
"""Calculate pass@k."""
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
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_maxs.append(is_pass)
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def get_possible_k_values(total_runs: int) -> list:
|
||||
"""
|
||||
Get all possible k values (factors of total_runs)
|
||||
|
||||
Args:
|
||||
total_runs: Total number of runs
|
||||
|
||||
Returns:
|
||||
List of k values in descending order
|
||||
"""
|
||||
k_values = []
|
||||
for k in range(1, total_runs + 1):
|
||||
if total_runs % k == 0:
|
||||
k_values.append(k)
|
||||
return sorted(k_values, reverse=True) # Sort from large to small
|
||||
|
||||
|
||||
def run_exp_statistic():
|
||||
"""Run the experiment statistic."""
|
||||
path: Path = Path("./exp_result/qwen3-8b")
|
||||
|
||||
# Store results for all experiments
|
||||
all_results = {}
|
||||
|
||||
for file in path.glob("*.jsonl"): # [f for f in path.glob("*.jsonl") if not f.stem[-1].isdigit()]
|
||||
# Group results by task_id
|
||||
task_results = defaultdict(list)
|
||||
|
||||
with open(file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
data = json.loads(line)
|
||||
|
||||
if isinstance(data, list):
|
||||
for part_data in data:
|
||||
task_id = part_data["task_id"]
|
||||
after_score = part_data["after_score"]
|
||||
task_results[task_id].append(after_score)
|
||||
else:
|
||||
task_id = data["task_id"]
|
||||
after_score = data["after_score"]
|
||||
task_results[task_id].append(after_score)
|
||||
|
||||
if not task_results:
|
||||
logger.warning(f"No valid data found in file {file}")
|
||||
continue
|
||||
|
||||
# Check if each task has consistent number of runs
|
||||
run_counts = [len(scores) for scores in task_results.values()]
|
||||
if len(set(run_counts)) > 1:
|
||||
logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}")
|
||||
continue
|
||||
|
||||
num_runs = run_counts[0]
|
||||
logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task")
|
||||
|
||||
# Get all possible k values
|
||||
k_values = get_possible_k_values(num_runs)
|
||||
logger.info(f"Calculable best@k values: {k_values}")
|
||||
|
||||
# Calculate various best@k values
|
||||
file_results = {"file": file.name}
|
||||
|
||||
for k in k_values:
|
||||
best_at_k_scores = []
|
||||
pass_at_k_scores = []
|
||||
for task_id, scores in task_results.items():
|
||||
try:
|
||||
best_k_score = calculate_best_at_k(scores, k)
|
||||
pass_at_k_score = calculate_pass_at_k(scores, k)
|
||||
pass_at_k_scores.append(pass_at_k_score)
|
||||
best_at_k_scores.append(best_k_score)
|
||||
except ValueError as e:
|
||||
logger.error(f"Error calculating best@{k} for task {task_id}: {e}")
|
||||
continue
|
||||
|
||||
if best_at_k_scores:
|
||||
avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores)
|
||||
file_results[f"best@{k}"] = avg_best_at_k
|
||||
logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}")
|
||||
|
||||
if pass_at_k_scores:
|
||||
avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores)
|
||||
file_results[f"pass@{k}"] = avg_pass_at_k
|
||||
logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}")
|
||||
|
||||
all_results[file.name] = file_results
|
||||
|
||||
# Create and display table
|
||||
if all_results:
|
||||
df = pd.DataFrame(list(all_results.values()))
|
||||
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@")]
|
||||
# best_columns = [col for col in df.columns]
|
||||
pass_columns.sort(key=lambda x: x, reverse=False)
|
||||
df = df[pass_columns]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Experiment Results Summary Table")
|
||||
print("=" * 80)
|
||||
print(df.round(4))
|
||||
print("=" * 80)
|
||||
|
||||
# Save table to CSV
|
||||
output_path = path / "experiment_summary.csv"
|
||||
df.to_csv(output_path)
|
||||
logger.info(f"Results table saved to: {output_path}")
|
||||
else:
|
||||
logger.warning("No valid experiment results found")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_exp_statistic()
|
||||
|
|
@ -1,726 +0,0 @@
|
|||
# flake8: noqa: E402
|
||||
# pylint: disable=too-many-return-statements
|
||||
"""A minimal ReAct Agent for BFCL-v3(multi-turn) tasks."""
|
||||
|
||||
import re
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
import warnings
|
||||
import tempfile
|
||||
import datetime
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Any
|
||||
|
||||
import ray
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from loguru import logger
|
||||
from openai import OpenAI
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from bfcl_utils import (
|
||||
load_test_case,
|
||||
handle_user_turn,
|
||||
handle_tool_calls,
|
||||
extract_tool_schema,
|
||||
extract_single_turn_response,
|
||||
extract_multi_turn_responses,
|
||||
capture_and_print_score_files,
|
||||
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 (
|
||||
is_empty_execute_response,
|
||||
)
|
||||
from bfcl_eval.eval_checker.eval_runner import (
|
||||
multi_turn_runner,
|
||||
ast_file_runner,
|
||||
)
|
||||
from bfcl_eval.eval_checker.eval_runner_helper import record_cost_latency
|
||||
from bfcl_eval.utils import (
|
||||
is_multi_turn,
|
||||
is_relevance_or_irrelevance,
|
||||
find_file_with_suffix,
|
||||
load_file,
|
||||
)
|
||||
|
||||
os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl"
|
||||
os.environ["BFCL_ANSWER_PATH"] = "data/possible_answer"
|
||||
load_dotenv("../../.env")
|
||||
|
||||
|
||||
@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_trials: 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:8002/",
|
||||
):
|
||||
|
||||
self.index: int = index
|
||||
self.task_ids: List[str] = task_ids
|
||||
self.categories: List[str] = [task_id.rsplit("_", 1)[0] if "_" in task_id else task_id for task_id in task_ids]
|
||||
self.experiment_name: str = experiment_name
|
||||
self.data_path: str = data_path
|
||||
self.answer_path: Path = answer_path
|
||||
self.model_name: str = model_name
|
||||
self.temperature: float = temperature
|
||||
self.max_interactions: int = max_interactions
|
||||
self.max_response_size: int = max_response_size
|
||||
self.num_trials: int = num_trials
|
||||
self.enable_thinking: bool = enable_thinking
|
||||
self.use_memory: bool = use_memory
|
||||
self.use_memory_addition: bool = use_memory_addition if use_memory else False
|
||||
self.use_memory_deletion: bool = use_memory_deletion if use_memory else False
|
||||
self.delete_freq: int = delete_freq
|
||||
self.freq_threshold: int = freq_threshold
|
||||
self.utility_threshold: float = utility_threshold
|
||||
self.memory_base_url: str = memory_base_url
|
||||
|
||||
self.llm_client = OpenAI()
|
||||
|
||||
self.history: List[List[List[dict]]] = [[] for _ in range(num_trials)]
|
||||
self.retrieved_memory_list: List[List[List[Any]]] = [[] for _ in range(num_trials)]
|
||||
self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)]
|
||||
self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_trials)]
|
||||
self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_trials)]
|
||||
self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_trials)]
|
||||
|
||||
for run_id in range(num_trials):
|
||||
for task_index in range(len(task_ids)):
|
||||
self.init_state(run_id, task_index)
|
||||
|
||||
def init_state(self, run_id, i) -> Dict[str, Any]:
|
||||
"""Initialize the state of the agent."""
|
||||
self.test_entry[run_id].append(load_test_case(self.data_path, self.task_ids[i]))
|
||||
self.original_test_entry[run_id].append(self.test_entry[run_id][i].get("extra", {}))
|
||||
self.tool_schema[run_id].append(extract_tool_schema(self.test_entry[run_id][i].get("tools", [{}])))
|
||||
|
||||
msg = self.test_entry[run_id][i].get("messages", [])
|
||||
self.history[run_id].append(msg)
|
||||
self.retrieved_memory_list[run_id].append([])
|
||||
self.current_turn[run_id][i] = 1
|
||||
|
||||
def update_task_history_with_memory(self, run_id, task_index, previous_memories: None):
|
||||
"""Update the task history with memory."""
|
||||
query = self.history[run_id][task_index][0]["content"]
|
||||
if len(previous_memories) == 0:
|
||||
response = self.get_memory(query)
|
||||
if response and "memory_list" in response["metadata"]:
|
||||
self.retrieved_memory_list[run_id][task_index] = response["metadata"]["memory_list"]
|
||||
task_memory = re.sub(r"\bMemory\s*(\d+)\s*[:]", r"Experience \1 :", response["answer"])
|
||||
logger.info(f"loaded task_memory: {task_memory}")
|
||||
self.history[run_id][task_index][0] = self.get_query_with_memory(query, task_memory)
|
||||
else:
|
||||
formatted_memories = []
|
||||
for i, memory in enumerate(previous_memories, 1):
|
||||
condition = memory["when_to_use"]
|
||||
memory_content = memory["content"]
|
||||
memory_text = f"Experience {i} :\n When to use: {condition}\n Content: {memory_content}\n"
|
||||
formatted_memories.append(memory_text)
|
||||
self.history[run_id][task_index][0] = self.get_query_with_memory(query, "\n".join(formatted_memories))
|
||||
|
||||
def get_query_with_memory(self, query: str, memory: str):
|
||||
"""Get the query with memory."""
|
||||
return {
|
||||
"role": "user",
|
||||
"content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + memory,
|
||||
}
|
||||
|
||||
def get_query_without_experience(self, query: str):
|
||||
"""Get the query without experience."""
|
||||
if "\n\nSome Related Experience" in query:
|
||||
query = query.split("\n\nSome Related Experience")[0].split("Task:\n")[-1]
|
||||
return query
|
||||
|
||||
def get_traj_from_task_history(self, task_id: str, task_history: list, reward: float):
|
||||
"""Get the trajectory from the task history."""
|
||||
return {
|
||||
"task_id": task_id,
|
||||
"messages": task_history,
|
||||
"score": reward,
|
||||
}
|
||||
|
||||
def handle_api_response(self, response: requests.Response):
|
||||
"""Handle API response with proper error checking"""
|
||||
if response.status_code != 200:
|
||||
print(f"Error: {response.status_code}")
|
||||
print(response.text)
|
||||
return None
|
||||
|
||||
return response.json()
|
||||
|
||||
def get_memory(self, query: str):
|
||||
"""Retrieve relevant task memories based on a query"""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}retrieve_task_memory",
|
||||
json={
|
||||
"query": query,
|
||||
"enable_llm_rerank": False,
|
||||
"enable_score_filter": False,
|
||||
"top_k": 5,
|
||||
"enable_llm_rewrite": False,
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
logger.info(f"query: {query}, response: {result}")
|
||||
return result
|
||||
|
||||
def summary_memory(self, trajectories):
|
||||
"""Generate a summary of conversation messages and create task memories"""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}summary_task_memory",
|
||||
json={
|
||||
"trajectories": trajectories,
|
||||
"success_threshold": 1.0,
|
||||
"enable_soft_comparison": True,
|
||||
"validation_threshold": 0.5,
|
||||
},
|
||||
)
|
||||
|
||||
result = self.handle_api_response(response)
|
||||
if not result:
|
||||
return []
|
||||
|
||||
# Extract memory list from response
|
||||
memory_list = result.get("metadata", {}).get("memory_list", [])
|
||||
logger.info(f"add new memories: {memory_list}")
|
||||
return memory_list
|
||||
|
||||
def add_memory(self, memory_list):
|
||||
"""Add the memory to the memory pool."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}add_task_memory",
|
||||
json={
|
||||
"memory_list": memory_list,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
def update_memory_information(self, memory_list, update_utility: bool = False):
|
||||
"""Update the memory information."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}record_task_memory",
|
||||
json={
|
||||
"memory_list": memory_list,
|
||||
"update_utility": update_utility,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info(response.json())
|
||||
|
||||
def delete_memory(self):
|
||||
"""Delete the memory from the memory pool."""
|
||||
response = requests.post(
|
||||
url=f"{self.memory_base_url}delete_task_memory",
|
||||
json={
|
||||
"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:
|
||||
"""Call the LLM."""
|
||||
for i in range(100):
|
||||
try:
|
||||
response = self.llm_client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=messages,
|
||||
tools=tool_schemas,
|
||||
temperature=self.temperature,
|
||||
seed=0,
|
||||
extra_body={"enable_thinking": self.enable_thinking},
|
||||
stream=self.enable_thinking,
|
||||
parallel_tool_calls=True,
|
||||
)
|
||||
if not self.enable_thinking:
|
||||
out_msg = response.choices[0].message
|
||||
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
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
if not chunk.choices:
|
||||
# Handle usage information
|
||||
continue
|
||||
|
||||
delta = chunk.choices[0].delta
|
||||
# Handle AI's thought process (chain reasoning)
|
||||
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": ""},
|
||||
},
|
||||
)
|
||||
|
||||
# Collect tool call ID (used for subsequent function calls)
|
||||
if 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
|
||||
|
||||
# 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
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": answer_content,
|
||||
"reasoning_content": reasoning_content,
|
||||
}
|
||||
if tool_info:
|
||||
msg["tool_calls"] = tool_info
|
||||
return msg
|
||||
except Exception as e:
|
||||
logger.exception(f"encounter error with {e.args}")
|
||||
time.sleep(1 + i * 10)
|
||||
|
||||
return "call llm error"
|
||||
|
||||
def env_step(self, run_id: int, index: int, messages: str) -> str:
|
||||
"""
|
||||
Process one step in the conversation.
|
||||
Both single turn and multi turn are supported.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages, with the last one being assistant response
|
||||
test_entry: Test entry containing initial_config, involved_classes, question etc.
|
||||
**kwargs: Additional arguments for compatibility
|
||||
|
||||
Returns:
|
||||
Dict containing next message and tools if applicable
|
||||
"""
|
||||
try:
|
||||
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",
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
# 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)}",
|
||||
)
|
||||
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],
|
||||
)
|
||||
except Exception as e:
|
||||
warnings.warn(f"Errors during tool invocation: {str(e)}")
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
else:
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
|
||||
except Exception as e:
|
||||
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]],
|
||||
) -> List[str]:
|
||||
"""
|
||||
Convert OpenAI format tool calls to execution format.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls in OpenAI format
|
||||
|
||||
Returns:
|
||||
List of function calls in string format
|
||||
"""
|
||||
execution_list = []
|
||||
|
||||
for tool_call in tool_calls:
|
||||
function = tool_call.get("function", {})
|
||||
function_name = function.get("name", "")
|
||||
|
||||
try:
|
||||
arguments = function.get("arguments", "{}")
|
||||
if isinstance(arguments, str):
|
||||
args_dict = json.loads(arguments)
|
||||
else:
|
||||
args_dict = arguments
|
||||
|
||||
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
|
||||
execution_list.append(f"{function_name}({args_str})")
|
||||
|
||||
except Exception:
|
||||
execution_list.append(f"{function_name}()")
|
||||
|
||||
return execution_list
|
||||
|
||||
def get_reward(self, run_id, index) -> float:
|
||||
"""Get the reward."""
|
||||
try:
|
||||
if not self.history[run_id][index] or not self.original_test_entry[run_id][index]:
|
||||
return 0.0
|
||||
|
||||
model_name = "env_handler"
|
||||
handler = QwenAPIHandler(
|
||||
model_name,
|
||||
temperature=1.0,
|
||||
) # FIXME: magic number
|
||||
|
||||
model_result_data = self._convert_conversation_to_eval_format(run_id, index)
|
||||
|
||||
prompt_data = [self.original_test_entry[run_id][index]]
|
||||
|
||||
state = {"leaderboard_table": {}}
|
||||
record_cost_latency(
|
||||
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,
|
||||
)
|
||||
else:
|
||||
# Find the corresponding possible answer file
|
||||
|
||||
possible_answer_file = find_file_with_suffix(
|
||||
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]]
|
||||
if is_multi_turn(self.categories[index]):
|
||||
accuracy, _ = self._eval_multi_turn_test(
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
self.categories[index],
|
||||
)
|
||||
else:
|
||||
accuracy, _ = self._eval_single_turn_test(
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
self.categories[index],
|
||||
)
|
||||
print(f"model_result_data: {model_result_data}")
|
||||
if possible_answer:
|
||||
print(f"possible_answer: {possible_answer}")
|
||||
else:
|
||||
print("possible_answer: None")
|
||||
|
||||
return accuracy
|
||||
|
||||
except Exception:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
return 0
|
||||
|
||||
def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert conversation history to evaluation format.
|
||||
|
||||
Args:
|
||||
conversation_result: Result from run_conversation
|
||||
original_test_entry: Original test entry data
|
||||
|
||||
Returns:
|
||||
Data in format expected by multi_turn_runner or other runners
|
||||
"""
|
||||
if is_multi_turn(self.categories[index]):
|
||||
turns_data = extract_multi_turn_responses(self.history[run_id][index])
|
||||
else:
|
||||
turns_data = extract_single_turn_response(self.history[run_id][index])
|
||||
|
||||
model_result_data = {
|
||||
"id": self.task_ids[index],
|
||||
"result": turns_data,
|
||||
"latency": 0,
|
||||
"input_token_count": 0,
|
||||
"output_token_count": 0,
|
||||
}
|
||||
|
||||
return model_result_data
|
||||
|
||||
def _eval_multi_turn_test(
|
||||
self,
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
test_category,
|
||||
):
|
||||
"""
|
||||
Evaluate multi-turn test.
|
||||
|
||||
Args:
|
||||
handler: Model handler instance
|
||||
model_result_data: Model result data
|
||||
prompt_data: Prompt data
|
||||
possible_answer: Possible answer data
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
|
||||
Returns:
|
||||
Tuple of (accuracy, total_count)
|
||||
"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
score_dir = Path(temp_dir)
|
||||
accuracy, total_count = multi_turn_runner(
|
||||
handler=handler,
|
||||
model_result=[model_result_data],
|
||||
prompt=prompt_data,
|
||||
possible_answer=possible_answer,
|
||||
model_name=model_name,
|
||||
test_category=test_category,
|
||||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir,
|
||||
model_name,
|
||||
test_category,
|
||||
"multi_turn",
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
def _eval_single_turn_test(
|
||||
self,
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
test_category,
|
||||
):
|
||||
"""
|
||||
Evaluate single-turn AST test.
|
||||
|
||||
Args:
|
||||
handler: Model handler instance
|
||||
model_result_data: Model result data
|
||||
prompt_data: Prompt data
|
||||
possible_answer: Possible answer data
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
|
||||
Returns:
|
||||
Tuple of (accuracy, total_count)
|
||||
"""
|
||||
language = "Python"
|
||||
if "java" in test_category.lower():
|
||||
language = "Java"
|
||||
elif "js" in test_category.lower() or "javascript" in test_category.lower():
|
||||
language = "JavaScript"
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
score_dir = Path(temp_dir)
|
||||
accuracy, total_count = ast_file_runner(
|
||||
handler=handler,
|
||||
model_result=[model_result_data],
|
||||
prompt=prompt_data,
|
||||
possible_answer=possible_answer,
|
||||
language=language,
|
||||
test_category=test_category,
|
||||
model_name=model_name,
|
||||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir,
|
||||
model_name,
|
||||
test_category,
|
||||
"single_turn",
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
def execute(self):
|
||||
"""Execute the agent."""
|
||||
result = []
|
||||
counter = 0
|
||||
for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"ray_index={self.index}")):
|
||||
t_result = None
|
||||
previous_memories = []
|
||||
for run_id in range(self.num_trials):
|
||||
try:
|
||||
start_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
for i in range(self.max_interactions):
|
||||
if self.use_memory and i == 0:
|
||||
self.update_task_history_with_memory(run_id, task_index, previous_memories)
|
||||
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}
|
||||
# 2. Returns tool invocation result: {"messages":
|
||||
# [{"role": "tool", "content": {<exec_results>}, 'tool_call_id': 'chatcmpl-tool-xxx'}]}
|
||||
# <exec_results>: when success, returns result dicts, e.g., {"travel_cost_list": [x]},
|
||||
# when error, returns error message,
|
||||
# e.g., {"error": "cd: temporary: No such directory. You cannot use path ..."}
|
||||
# 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 = []
|
||||
next_user_msg = ""
|
||||
for idx, msg in enumerate(env_output.get("messages", [])):
|
||||
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
|
||||
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]},
|
||||
)
|
||||
else:
|
||||
self.history[run_id][task_index].append({"role": "user", "content": next_user_msg})
|
||||
|
||||
logger.info(f"index={self.index} task_id={task_id} iteration={i}")
|
||||
|
||||
if self.task_completed(run_id, task_index):
|
||||
break
|
||||
|
||||
reward = self.get_reward(run_id, task_index)
|
||||
if self.use_memory:
|
||||
if self.use_memory_addition:
|
||||
new_traj_list = [
|
||||
self.get_traj_from_task_history(task_id, self.history[run_id][task_index], reward),
|
||||
]
|
||||
previous_memories = self.summary_memory(new_traj_list)
|
||||
if reward == 1:
|
||||
self.add_memory(previous_memories)
|
||||
|
||||
# update the freq & utility attributes of retrieved memories
|
||||
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()
|
||||
|
||||
t_result = {
|
||||
"run_id": run_id,
|
||||
"task_id": self.task_ids[task_index],
|
||||
"experiment_name": self.experiment_name,
|
||||
"task_completed": self.task_completed(run_id, task_index),
|
||||
"reward": reward,
|
||||
"task_history": self.history[run_id][task_index],
|
||||
"task_start_time": start_time,
|
||||
}
|
||||
if reward == 1:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"encounter error with {e.args}")
|
||||
result.append(t_result)
|
||||
return result
|
||||
|
||||
def task_completed(self, run_id, index):
|
||||
"""
|
||||
Check if task is completed.
|
||||
|
||||
Returns:
|
||||
True if task is completed, False otherwise
|
||||
"""
|
||||
return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]"
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function to run the BFCLAgent."""
|
||||
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_ids=[task_ids[0]],
|
||||
experiment_name=f"qwen3_8b_{dataset_name}",
|
||||
)
|
||||
result = agent.execute()
|
||||
logger.info(f"result={json.dumps(result)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,399 +0,0 @@
|
|||
"""Utils for evaluation on BFCL tasks"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Any
|
||||
|
||||
from bfcl_eval.constants.default_prompts import (
|
||||
DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
|
||||
)
|
||||
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]:
|
||||
"""
|
||||
load test cases by id
|
||||
"""
|
||||
if not Path(data_path).exists():
|
||||
raise FileNotFoundError(f"BFCL data file '{data_path}' not found")
|
||||
|
||||
if test_id is None:
|
||||
raise ValueError("task_id is required")
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
if str(test_id).isdigit(): # pylint: disable=R1720
|
||||
idx = int(test_id)
|
||||
for line_no, line in enumerate(f):
|
||||
if line_no == idx:
|
||||
return json.loads(line)
|
||||
raise ValueError(f"Test case index {idx} not found in {data_path}")
|
||||
else:
|
||||
for line in f:
|
||||
data = json.loads(line)
|
||||
if data.get("id") == test_id:
|
||||
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,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle user turn by returning appropriate content from test_entry["question"].
|
||||
For non-first turns, processes user query and tools.
|
||||
|
||||
Args:
|
||||
test_entry: Test entry containing conversation data
|
||||
current_turn: Current turn number
|
||||
|
||||
Returns:
|
||||
Response containing next user message and tools
|
||||
"""
|
||||
try:
|
||||
current_turn_message = []
|
||||
tools = compile_tools(test_entry)
|
||||
questions = test_entry.get("question", [])
|
||||
holdout_function = test_entry.get("holdout_function", {})
|
||||
|
||||
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."
|
||||
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)
|
||||
|
||||
except Exception as e:
|
||||
return create_error_response(f"Failed to process user message: {str(e)}")
|
||||
|
||||
|
||||
def handle_tool_calls( # pylint: disable=W0613
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
decoded_calls: list[str],
|
||||
test_entry: Dict[str, Any],
|
||||
current_turn: int,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle tool calls from assistant.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls in OpenAI format
|
||||
decoded_calls: List of decoded function calls
|
||||
test_entry: Test entry containing environment data
|
||||
current_turn: Current turn number
|
||||
|
||||
Returns:
|
||||
Response containing tool execution results
|
||||
"""
|
||||
execution_results, _ = execute_multi_turn_func_call(
|
||||
func_call_list=decoded_calls,
|
||||
initial_config=test_entry["initial_config"],
|
||||
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"]),
|
||||
is_evaL_run=False,
|
||||
)
|
||||
# print('execution_results in handler_tool_calls:', execution_results)
|
||||
|
||||
return create_tool_response(tool_calls, execution_results)
|
||||
|
||||
|
||||
def compile_tools(test_entry: dict) -> list:
|
||||
"""
|
||||
Compile functions into tools format.
|
||||
|
||||
Args:
|
||||
test_entry: Test entry containing functions
|
||||
|
||||
Returns:
|
||||
List of tools in OpenAI format
|
||||
"""
|
||||
functions: list = test_entry["function"]
|
||||
test_category: str = test_entry["id"].rsplit("_", 1)[0]
|
||||
|
||||
functions = func_doc_language_specific_pre_processing(functions, test_category)
|
||||
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],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for tool calls.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls
|
||||
execution_results: List of execution results
|
||||
|
||||
Returns:
|
||||
Response containing tool execution results
|
||||
"""
|
||||
tool_messages = []
|
||||
for i, (tool_call, result) in enumerate(zip(tool_calls, execution_results)):
|
||||
tool_messages.append(
|
||||
{
|
||||
"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]],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response containing user message.
|
||||
|
||||
Args:
|
||||
question_turn: List of messages for current turn
|
||||
tools: List of available tools
|
||||
|
||||
Returns:
|
||||
Response containing user message and tools
|
||||
"""
|
||||
user_content = ""
|
||||
for msg in question_turn:
|
||||
if msg["role"] == "user":
|
||||
user_content = msg["content"]
|
||||
break
|
||||
|
||||
return {"messages": [{"role": "user", "content": user_content}], "tools": tools}
|
||||
|
||||
|
||||
def create_completion_response() -> Dict[str, Any]:
|
||||
"""
|
||||
Create response indicating conversation completion.
|
||||
|
||||
Returns:
|
||||
Response with completion message
|
||||
"""
|
||||
return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
|
||||
|
||||
|
||||
def create_error_response(error_message: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for error conditions.
|
||||
|
||||
Args:
|
||||
error_message: Error message to include
|
||||
|
||||
Returns:
|
||||
Response containing error message
|
||||
"""
|
||||
return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
|
||||
|
||||
|
||||
def decode_execute(result):
|
||||
"""
|
||||
Decode execute results for compatibility with evaluation framework.
|
||||
|
||||
Args:
|
||||
result: Result to decode
|
||||
|
||||
Returns:
|
||||
List of decoded function calls
|
||||
"""
|
||||
return default_decode_execute_prompting(result)
|
||||
|
||||
|
||||
def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
|
||||
"""
|
||||
Extract single-turn response from conversation messages.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages
|
||||
|
||||
Returns:
|
||||
String representation of the response
|
||||
"""
|
||||
for message in reversed(messages):
|
||||
if message["role"] == "assistant":
|
||||
if "tool_calls" in message and message["tool_calls"]:
|
||||
formatted_calls = []
|
||||
for tool_call in message["tool_calls"]:
|
||||
formatted_call = format_single_tool_call_for_eval(
|
||||
tool_call,
|
||||
)
|
||||
if formatted_call:
|
||||
formatted_calls.append(formatted_call)
|
||||
return "\n".join(formatted_calls) if formatted_calls else ""
|
||||
elif message.get("content"):
|
||||
return message["content"]
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
def extract_multi_turn_responses(
|
||||
messages: List[Dict[str, Any]],
|
||||
) -> List[List[str]]:
|
||||
"""
|
||||
Extract multi-turn responses from conversation messages.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages
|
||||
|
||||
Returns:
|
||||
List of turns, each turn is a list of function call strings
|
||||
"""
|
||||
turns_data = []
|
||||
current_turn_responses = []
|
||||
|
||||
i = 0
|
||||
while i < len(messages):
|
||||
message = messages[i]
|
||||
|
||||
if message["role"] == "user":
|
||||
if current_turn_responses:
|
||||
turns_data.append(current_turn_responses)
|
||||
current_turn_responses = []
|
||||
|
||||
i += 1
|
||||
while i < len(messages) and messages[i]["role"] == "assistant":
|
||||
assistant_msg = messages[i]
|
||||
|
||||
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,
|
||||
)
|
||||
if formatted_call:
|
||||
current_turn_responses.append(formatted_call)
|
||||
|
||||
i += 1
|
||||
|
||||
while i < len(messages) and messages[i]["role"] == "tool":
|
||||
i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
if current_turn_responses:
|
||||
turns_data.append(current_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
|
||||
|
||||
Returns:
|
||||
Formatted string representation
|
||||
"""
|
||||
function = tool_call.get("function", {})
|
||||
function_name = function.get("name", "")
|
||||
|
||||
try:
|
||||
arguments = function.get("arguments", "{}")
|
||||
if isinstance(arguments, str):
|
||||
args_dict = json.loads(arguments)
|
||||
else:
|
||||
args_dict = arguments
|
||||
|
||||
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
|
||||
return f"{function_name}({args_str})"
|
||||
|
||||
except Exception:
|
||||
return f"{function_name}()"
|
||||
|
||||
|
||||
def capture_and_print_score_files(
|
||||
score_dir: Path,
|
||||
model_name: str,
|
||||
test_category: str,
|
||||
eval_type: str,
|
||||
):
|
||||
"""
|
||||
Capture and print contents of score files written to score_dir.
|
||||
|
||||
Args:
|
||||
score_dir: Directory containing score files
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
eval_type: Type of evaluation (relevance/multi_turn/single_turn)
|
||||
"""
|
||||
try:
|
||||
print(f"\n=== {eval_type.upper()} Evaluation Result Files ===")
|
||||
print(f"Model: {model_name}")
|
||||
print(f"Test Category: {test_category}")
|
||||
print(f"Evaluation Type: {eval_type}")
|
||||
|
||||
for file_path in score_dir.rglob("*"):
|
||||
if file_path.is_file():
|
||||
relative_path = file_path.relative_to(score_dir)
|
||||
print(f"\n--- File: {relative_path} ---")
|
||||
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
if (
|
||||
file_path.suffix == ".json"
|
||||
or content.strip().startswith("{")
|
||||
or content.strip().startswith("[")
|
||||
):
|
||||
try:
|
||||
lines = content.strip().split("\n")
|
||||
formatted_lines = []
|
||||
for line in lines:
|
||||
if line.strip():
|
||||
parsed = json.loads(line)
|
||||
formatted_lines.append(
|
||||
json.dumps(
|
||||
parsed,
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
),
|
||||
)
|
||||
content = "\n".join(formatted_lines)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
print(content)
|
||||
|
||||
except UnicodeDecodeError:
|
||||
print(f"[Binary file, size: {file_path.stat().st_size} bytes]")
|
||||
except Exception as e:
|
||||
print(f"[Error reading file: {str(e)}]")
|
||||
|
||||
print(f"=== {eval_type.upper()} Evaluation Result Files End ===\n")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error capturing evaluation result files: {str(e)}")
|
||||
|
||||
|
||||
def extract_tool_schema(tools):
|
||||
"""Reformat tool schema"""
|
||||
for i in range(len(tools)): # pylint: disable=C0200
|
||||
tools[i]["function"].pop("response")
|
||||
return tools
|
||||
|
|
@ -1,206 +0,0 @@
|
|||
# pylint: disable=C0114
|
||||
DEFAULT_TRAIN_IDS: set[str] = {
|
||||
"multi_turn_base_102",
|
||||
"multi_turn_base_107",
|
||||
"multi_turn_base_110",
|
||||
"multi_turn_base_114",
|
||||
"multi_turn_base_115",
|
||||
"multi_turn_base_118",
|
||||
"multi_turn_base_122",
|
||||
"multi_turn_base_123",
|
||||
"multi_turn_base_128",
|
||||
"multi_turn_base_13",
|
||||
"multi_turn_base_130",
|
||||
"multi_turn_base_132",
|
||||
"multi_turn_base_133",
|
||||
"multi_turn_base_143",
|
||||
"multi_turn_base_144",
|
||||
"multi_turn_base_146",
|
||||
"multi_turn_base_15",
|
||||
"multi_turn_base_158",
|
||||
"multi_turn_base_169",
|
||||
"multi_turn_base_17",
|
||||
"multi_turn_base_172",
|
||||
"multi_turn_base_176",
|
||||
"multi_turn_base_182",
|
||||
"multi_turn_base_187",
|
||||
"multi_turn_base_197",
|
||||
"multi_turn_base_199",
|
||||
"multi_turn_base_22",
|
||||
"multi_turn_base_23",
|
||||
"multi_turn_base_24",
|
||||
"multi_turn_base_36",
|
||||
"multi_turn_base_40",
|
||||
"multi_turn_base_44",
|
||||
"multi_turn_base_47",
|
||||
"multi_turn_base_48",
|
||||
"multi_turn_base_5",
|
||||
"multi_turn_base_51",
|
||||
"multi_turn_base_59",
|
||||
"multi_turn_base_63",
|
||||
"multi_turn_base_65",
|
||||
"multi_turn_base_66",
|
||||
"multi_turn_base_67",
|
||||
"multi_turn_base_68",
|
||||
"multi_turn_base_70",
|
||||
"multi_turn_base_75",
|
||||
"multi_turn_base_77",
|
||||
"multi_turn_base_78",
|
||||
"multi_turn_base_79",
|
||||
"multi_turn_base_81",
|
||||
"multi_turn_base_83",
|
||||
"multi_turn_base_93",
|
||||
}
|
||||
|
||||
DEFAULT_VAL_IDS: set[str] = {
|
||||
"multi_turn_base_0",
|
||||
"multi_turn_base_1",
|
||||
"multi_turn_base_10",
|
||||
"multi_turn_base_100",
|
||||
"multi_turn_base_101",
|
||||
"multi_turn_base_103",
|
||||
"multi_turn_base_104",
|
||||
"multi_turn_base_105",
|
||||
"multi_turn_base_106",
|
||||
"multi_turn_base_108",
|
||||
"multi_turn_base_109",
|
||||
"multi_turn_base_11",
|
||||
"multi_turn_base_111",
|
||||
"multi_turn_base_112",
|
||||
"multi_turn_base_113",
|
||||
"multi_turn_base_116",
|
||||
"multi_turn_base_117",
|
||||
"multi_turn_base_119",
|
||||
"multi_turn_base_12",
|
||||
"multi_turn_base_120",
|
||||
"multi_turn_base_121",
|
||||
"multi_turn_base_124",
|
||||
"multi_turn_base_125",
|
||||
"multi_turn_base_126",
|
||||
"multi_turn_base_127",
|
||||
"multi_turn_base_129",
|
||||
"multi_turn_base_131",
|
||||
"multi_turn_base_134",
|
||||
"multi_turn_base_135",
|
||||
"multi_turn_base_136",
|
||||
"multi_turn_base_137",
|
||||
"multi_turn_base_138",
|
||||
"multi_turn_base_139",
|
||||
"multi_turn_base_14",
|
||||
"multi_turn_base_140",
|
||||
"multi_turn_base_141",
|
||||
"multi_turn_base_142",
|
||||
"multi_turn_base_145",
|
||||
"multi_turn_base_147",
|
||||
"multi_turn_base_148",
|
||||
"multi_turn_base_149",
|
||||
"multi_turn_base_150",
|
||||
"multi_turn_base_151",
|
||||
"multi_turn_base_152",
|
||||
"multi_turn_base_153",
|
||||
"multi_turn_base_154",
|
||||
"multi_turn_base_155",
|
||||
"multi_turn_base_156",
|
||||
"multi_turn_base_157",
|
||||
"multi_turn_base_159",
|
||||
"multi_turn_base_16",
|
||||
"multi_turn_base_160",
|
||||
"multi_turn_base_161",
|
||||
"multi_turn_base_162",
|
||||
"multi_turn_base_163",
|
||||
"multi_turn_base_164",
|
||||
"multi_turn_base_165",
|
||||
"multi_turn_base_166",
|
||||
"multi_turn_base_167",
|
||||
"multi_turn_base_168",
|
||||
"multi_turn_base_170",
|
||||
"multi_turn_base_171",
|
||||
"multi_turn_base_173",
|
||||
"multi_turn_base_174",
|
||||
"multi_turn_base_175",
|
||||
"multi_turn_base_177",
|
||||
"multi_turn_base_178",
|
||||
"multi_turn_base_179",
|
||||
"multi_turn_base_18",
|
||||
"multi_turn_base_180",
|
||||
"multi_turn_base_181",
|
||||
"multi_turn_base_183",
|
||||
"multi_turn_base_184",
|
||||
"multi_turn_base_185",
|
||||
"multi_turn_base_186",
|
||||
"multi_turn_base_188",
|
||||
"multi_turn_base_189",
|
||||
"multi_turn_base_19",
|
||||
"multi_turn_base_190",
|
||||
"multi_turn_base_191",
|
||||
"multi_turn_base_192",
|
||||
"multi_turn_base_193",
|
||||
"multi_turn_base_194",
|
||||
"multi_turn_base_195",
|
||||
"multi_turn_base_196",
|
||||
"multi_turn_base_198",
|
||||
"multi_turn_base_2",
|
||||
"multi_turn_base_20",
|
||||
"multi_turn_base_21",
|
||||
"multi_turn_base_25",
|
||||
"multi_turn_base_26",
|
||||
"multi_turn_base_27",
|
||||
"multi_turn_base_28",
|
||||
"multi_turn_base_29",
|
||||
"multi_turn_base_3",
|
||||
"multi_turn_base_30",
|
||||
"multi_turn_base_31",
|
||||
"multi_turn_base_32",
|
||||
"multi_turn_base_33",
|
||||
"multi_turn_base_34",
|
||||
"multi_turn_base_35",
|
||||
"multi_turn_base_37",
|
||||
"multi_turn_base_38",
|
||||
"multi_turn_base_39",
|
||||
"multi_turn_base_4",
|
||||
"multi_turn_base_41",
|
||||
"multi_turn_base_42",
|
||||
"multi_turn_base_43",
|
||||
"multi_turn_base_45",
|
||||
"multi_turn_base_46",
|
||||
"multi_turn_base_49",
|
||||
"multi_turn_base_50",
|
||||
"multi_turn_base_52",
|
||||
"multi_turn_base_53",
|
||||
"multi_turn_base_54",
|
||||
"multi_turn_base_55",
|
||||
"multi_turn_base_56",
|
||||
"multi_turn_base_57",
|
||||
"multi_turn_base_58",
|
||||
"multi_turn_base_6",
|
||||
"multi_turn_base_60",
|
||||
"multi_turn_base_61",
|
||||
"multi_turn_base_62",
|
||||
"multi_turn_base_64",
|
||||
"multi_turn_base_69",
|
||||
"multi_turn_base_7",
|
||||
"multi_turn_base_71",
|
||||
"multi_turn_base_72",
|
||||
"multi_turn_base_73",
|
||||
"multi_turn_base_74",
|
||||
"multi_turn_base_76",
|
||||
"multi_turn_base_8",
|
||||
"multi_turn_base_80",
|
||||
"multi_turn_base_82",
|
||||
"multi_turn_base_84",
|
||||
"multi_turn_base_85",
|
||||
"multi_turn_base_86",
|
||||
"multi_turn_base_87",
|
||||
"multi_turn_base_88",
|
||||
"multi_turn_base_89",
|
||||
"multi_turn_base_9",
|
||||
"multi_turn_base_90",
|
||||
"multi_turn_base_91",
|
||||
"multi_turn_base_92",
|
||||
"multi_turn_base_94",
|
||||
"multi_turn_base_95",
|
||||
"multi_turn_base_96",
|
||||
"multi_turn_base_97",
|
||||
"multi_turn_base_98",
|
||||
"multi_turn_base_99",
|
||||
}
|
||||
|
|
@ -1,235 +0,0 @@
|
|||
# pylint: disable=W0621,W1514
|
||||
"""Init task memory pool"""
|
||||
import argparse
|
||||
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]:
|
||||
"""
|
||||
load training cases by id
|
||||
"""
|
||||
if not Path(data_path).exists():
|
||||
raise FileNotFoundError(f"BFCL data file '{data_path}' not found")
|
||||
|
||||
if task_id is None:
|
||||
raise ValueError("task_id is required")
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
if str(task_id).isdigit(): # pylint: disable=R1720
|
||||
idx = int(task_id)
|
||||
for line_no, line in enumerate(f):
|
||||
if line_no == idx:
|
||||
return json.loads(line)
|
||||
raise ValueError(f"Task case index {idx} not found in {data_path}")
|
||||
else:
|
||||
for line in f:
|
||||
data = json.loads(line)
|
||||
if data.get("id") == task_id:
|
||||
return data
|
||||
raise ValueError(f"Task case id '{task_id}' not found in {data_path}")
|
||||
|
||||
|
||||
def get_tool_prompt(tools):
|
||||
"""Construct prompt with provided 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>'
|
||||
)
|
||||
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 _, trajectories in grouped.items():
|
||||
if len(trajectories) == 1:
|
||||
# when only one trajectory, retain it
|
||||
filtered_groups.append(trajectories)
|
||||
elif len(trajectories) == 2:
|
||||
# when there are two trajectories, retain them
|
||||
filtered_groups.append(trajectories)
|
||||
else:
|
||||
# when there are more than two trajectories, choose the two with the highest and lowest rewards
|
||||
trajectories.sort(key=lambda t: t["reward"])
|
||||
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) -> Dict[str, Any]:
|
||||
"""
|
||||
post trajectories to summarizer service
|
||||
|
||||
Args:
|
||||
trajectories: trajectory list
|
||||
service_url: summarizer service URL
|
||||
|
||||
Returns:
|
||||
response json
|
||||
"""
|
||||
trajectory_dicts = [
|
||||
{
|
||||
"task_id": traj["task_id"],
|
||||
"messages": traj["task_history"],
|
||||
"score": traj["reward"],
|
||||
}
|
||||
for traj in trajectories
|
||||
]
|
||||
|
||||
request_data = {
|
||||
"trajectories": trajectory_dicts,
|
||||
"success_threshold": 1.0,
|
||||
"enable_soft_comparison": True,
|
||||
"validation_threshold": 0.5,
|
||||
}
|
||||
|
||||
try:
|
||||
response = requests.post(f"{service_url}/summary_task_memory", json=request_data)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as e:
|
||||
return {"error": str(e), "trajectories_count": len(trajectories)}
|
||||
|
||||
|
||||
def process_trajectories_with_threads(
|
||||
grouped_trajectories: List[List[Any]],
|
||||
service_url: 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
|
||||
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): i for i, group in enumerate(grouped_trajectories)
|
||||
}
|
||||
|
||||
for future in as_completed(future_to_group):
|
||||
group_index = future_to_group[future]
|
||||
try:
|
||||
result = future.result()
|
||||
result["group_index"] = group_index
|
||||
result["group_size"] = len(grouped_trajectories[group_index])
|
||||
results.append(result)
|
||||
if "memory_list" in result["metadata"]:
|
||||
print(f'✅ Group {group_index} processed: {result["metadata"].get("memory_list", 0)}')
|
||||
memory_list = result["metadata"].get("memory_list", [])
|
||||
response = requests.post(url=f"{service_url}/add_task_memory", json={"memory_list": memory_list})
|
||||
response.raise_for_status()
|
||||
else:
|
||||
print(f"❌ Group {group_index} processed: error")
|
||||
except Exception as e:
|
||||
error_result = {
|
||||
"group_index": group_index,
|
||||
"group_size": len(grouped_trajectories[group_index]),
|
||||
"error": str(e),
|
||||
}
|
||||
results.append(error_result)
|
||||
print(f"❌ Group {group_index} failed: {e}")
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function to convert JSONL to memories using ReMe service."""
|
||||
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:8002", help="ReMe service URL")
|
||||
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"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,
|
||||
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_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 = {
|
||||
"jsonl_file": args.jsonl_file,
|
||||
"total_groups": len(grouped_trajectories),
|
||||
"success_count": success_count,
|
||||
"error_count": error_count,
|
||||
"total_task_memories": total_memories,
|
||||
"results": results,
|
||||
}
|
||||
|
||||
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:
|
||||
print(f"Error saving results: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,73 +0,0 @@
|
|||
# pylint: disable=W0621
|
||||
"""Preprocess multi-turn test cases"""
|
||||
|
||||
import json
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from bfcl_eval.model_handler.model_style import ModelStyle
|
||||
from bfcl_eval.eval_checker.eval_runner_helper import load_file
|
||||
from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI
|
||||
from bfcl_eval.constants.eval_config import MULTI_TURN_FUNC_DOC_PATH
|
||||
from bfcl_eval.constants.category_mapping import MULTI_TURN_FUNC_DOC_FILE_MAPPING
|
||||
from bfcl_eval.model_handler.utils import (
|
||||
convert_to_tool,
|
||||
func_doc_language_specific_pre_processing,
|
||||
)
|
||||
|
||||
|
||||
def process_multi_turn_test_case(file_path, output_path):
|
||||
"""
|
||||
Multi-turn test cases don't have the function doc in the prompt. We need to add them here.
|
||||
"""
|
||||
test_cases = []
|
||||
with open(output_path, "w", encoding="utf-8") as outf:
|
||||
with open(file_path, encoding="utf-8") as f:
|
||||
file = f.readlines()
|
||||
for line in file:
|
||||
entry = json.loads(line)
|
||||
if "multi_turn" not in entry["id"]:
|
||||
continue
|
||||
test_category: str = entry["id"].rsplit("_", 1)[0]
|
||||
involved_classes = entry["involved_classes"]
|
||||
entry["function"] = []
|
||||
for func_collection in involved_classes:
|
||||
# func_doc is a list of dict
|
||||
func_doc = load_file(
|
||||
MULTI_TURN_FUNC_DOC_PATH / MULTI_TURN_FUNC_DOC_FILE_MAPPING[func_collection],
|
||||
)
|
||||
entry["function"].extend(func_doc)
|
||||
|
||||
# Handle Miss Func category; we need to remove the holdout function doc
|
||||
if "missed_function" in entry:
|
||||
for turn_index, missed_func_names in entry["missed_function"].items():
|
||||
entry["missed_function"][turn_index] = []
|
||||
for missed_func_name in missed_func_names:
|
||||
for i, func_doc in enumerate(entry["function"]):
|
||||
if func_doc["name"] == missed_func_name:
|
||||
# Add the missed function doc to the missed_function list
|
||||
entry["missed_function"][turn_index].append(func_doc)
|
||||
# Remove it from the function list
|
||||
entry["function"].pop(i)
|
||||
break
|
||||
|
||||
functions = func_doc_language_specific_pre_processing(entry["function"], test_category)
|
||||
tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions)
|
||||
|
||||
test_cases.append(
|
||||
{
|
||||
"id": entry["id"],
|
||||
"messages": entry["question"][0],
|
||||
"tools": tools,
|
||||
"extra": entry,
|
||||
},
|
||||
)
|
||||
outf.write(json.dumps(test_cases[-1], ensure_ascii=False) + "\n")
|
||||
|
||||
return test_cases
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
file_path = Path("./gorilla/berkeley-function-call-leaderboard/bfcl_eval/data/BFCL_v3_multi_turn_base.json")
|
||||
output_path = "data/multiturn_data_base.jsonl"
|
||||
preprocessed_test_cases = process_multi_turn_test_case(file_path, output_path)
|
||||
|
|
@ -1,151 +0,0 @@
|
|||
"""Run evaluation on BFCL-V3-Multi-Turn-Base dataset."""
|
||||
|
||||
import time
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import ray
|
||||
import requests
|
||||
from loguru import logger
|
||||
from dotenv import load_dotenv
|
||||
from bfcl_agent import BFCLAgent
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
|
||||
def run_agent(
|
||||
max_workers: int,
|
||||
dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
model_name: str = "qwen3-8b",
|
||||
enable_thinking: bool = False,
|
||||
data_path: str = "data/multiturn_data_base_val.jsonl",
|
||||
answer_path: Path = Path("data/possible_answer"),
|
||||
num_trials: int = 1,
|
||||
use_memory: bool = False,
|
||||
memory_base_url: str = "http://0.0.0.0:8002/",
|
||||
use_memory_addition: bool = True,
|
||||
use_memory_deletion: bool = False,
|
||||
delete_freq: int = 10,
|
||||
freq_threshold: int = 5,
|
||||
utility_threshold: float = 0.5,
|
||||
):
|
||||
"""Run the agent"""
|
||||
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.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
task_ids = [json.loads(line)["id"] for line in f]
|
||||
|
||||
result: list = []
|
||||
|
||||
def dump_file():
|
||||
with open(path / f"{experiment_name}.jsonl", "a", encoding="utf-8") as f:
|
||||
for x in result:
|
||||
f.write(json.dumps(x) + "\n")
|
||||
|
||||
future_list: list = []
|
||||
for i in range(max_workers):
|
||||
actor = BFCLAgent.remote(
|
||||
index=i,
|
||||
model_name=model_name,
|
||||
task_ids=task_ids[i::max_workers],
|
||||
experiment_name=experiment_name,
|
||||
data_path=data_path,
|
||||
answer_path=answer_path,
|
||||
num_trials=num_trials,
|
||||
use_memory=use_memory,
|
||||
memory_base_url=memory_base_url,
|
||||
use_memory_addition=use_memory_addition,
|
||||
use_memory_deletion=use_memory_deletion,
|
||||
delete_freq=delete_freq,
|
||||
freq_threshold=freq_threshold,
|
||||
utility_threshold=utility_threshold,
|
||||
enable_thinking=enable_thinking,
|
||||
)
|
||||
future = actor.execute.remote()
|
||||
future_list.append(future)
|
||||
time.sleep(1)
|
||||
logger.info("submit complete")
|
||||
|
||||
for i, future in enumerate(future_list):
|
||||
t_result = ray.get(future)
|
||||
if t_result:
|
||||
if isinstance(t_result, list):
|
||||
result.extend(t_result)
|
||||
else:
|
||||
result.append(t_result)
|
||||
|
||||
logger.info(f"{i + 1}/{len(task_ids)} complete")
|
||||
dump_file()
|
||||
|
||||
|
||||
def handle_api_response(response: requests.Response):
|
||||
"""Handle API response with proper error checking"""
|
||||
if response.status_code != 200:
|
||||
print(f"Error: {response.status_code}")
|
||||
print(response.text)
|
||||
return None
|
||||
|
||||
return response.json()
|
||||
|
||||
|
||||
def load_memory(path: str = "docs/library", api_url: str = "http://0.0.0.0:8002/"):
|
||||
"""Load memories from disk into the vector store"""
|
||||
response = requests.post(
|
||||
url=f"{api_url}load_memory",
|
||||
json={
|
||||
"load_file_path": path,
|
||||
"clear_existing": True,
|
||||
},
|
||||
)
|
||||
|
||||
result = handle_api_response(response)
|
||||
if result:
|
||||
print(f"Memory loaded from {path}")
|
||||
|
||||
|
||||
def main():
|
||||
"""Main function"""
|
||||
max_workers = 4
|
||||
if max_workers > 1:
|
||||
ray.init(num_cpus=max_workers)
|
||||
|
||||
num_runs = 4
|
||||
num_trials = 1
|
||||
model_name = "qwen3-8b"
|
||||
enable_thinking = True
|
||||
use_memory = True
|
||||
use_memory_addition = False
|
||||
use_memory_deletion = False
|
||||
memory_base_url = "http://0.0.0.0:8003/"
|
||||
|
||||
if use_memory:
|
||||
load_file_path = "docs/library/paper_data/task/bfcl_qwen3_8b.jsonl"
|
||||
load_memory(load_file_path, memory_base_url)
|
||||
|
||||
for _ in range(num_runs):
|
||||
run_agent(
|
||||
max_workers=max_workers,
|
||||
model_name=model_name,
|
||||
dataset_name="bfcl-multi-turn-base-val",
|
||||
experiment_suffix="w-fixed-memory",
|
||||
data_path="data/multiturn_data_base_val.jsonl",
|
||||
answer_path=Path("data/possible_answer"),
|
||||
enable_thinking=enable_thinking,
|
||||
num_trials=num_trials,
|
||||
use_memory=use_memory,
|
||||
memory_base_url=memory_base_url,
|
||||
use_memory_addition=use_memory_addition,
|
||||
use_memory_deletion=use_memory_deletion,
|
||||
delete_freq=5,
|
||||
freq_threshold=5,
|
||||
utility_threshold=0.5,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,163 +0,0 @@
|
|||
"""Run the experiment statistic."""
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
|
||||
|
||||
def calculate_best_at_k(scores: list, k: int) -> float:
|
||||
"""
|
||||
Calculate best@k
|
||||
Divide scores into groups of size k, take the maximum value in each group,
|
||||
then average these maximum values
|
||||
|
||||
Args:
|
||||
scores: List of after_score values for all runs of a task
|
||||
k: Group size
|
||||
|
||||
Returns:
|
||||
best@k value
|
||||
"""
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i : i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def calculate_pass_at_k(scores: list, k: int) -> float:
|
||||
"""Calculate pass@k."""
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
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_maxs.append(is_pass)
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def get_possible_k_values(total_runs: int) -> list:
|
||||
"""
|
||||
Get all possible k values (factors of total_runs)
|
||||
|
||||
Args:
|
||||
total_runs: Total number of runs
|
||||
|
||||
Returns:
|
||||
List of k values in descending order
|
||||
"""
|
||||
k_values = []
|
||||
for k in range(1, total_runs + 1):
|
||||
if total_runs % k == 0:
|
||||
k_values.append(k)
|
||||
return sorted(k_values, reverse=True) # Sort from large to small
|
||||
|
||||
|
||||
def run_exp_statistic():
|
||||
"""Run the experiment statistic."""
|
||||
path: Path = Path("./exp_result/qwen3-8b/with_think")
|
||||
|
||||
# Store results for all experiments
|
||||
all_results = {}
|
||||
for file in path.glob("*.jsonl"):
|
||||
# Group results by task_id
|
||||
task_results = defaultdict(list)
|
||||
print(file)
|
||||
with open(file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
data = json.loads(line)
|
||||
|
||||
if isinstance(data, list):
|
||||
for part_data in data:
|
||||
task_id = part_data["task_id"]
|
||||
after_score = part_data["reward"]
|
||||
task_results[task_id].append(after_score)
|
||||
else:
|
||||
task_id = data["task_id"]
|
||||
after_score = data["reward"]
|
||||
task_results[task_id].append(after_score)
|
||||
|
||||
if not task_results:
|
||||
logger.warning(f"No valid data found in file {file}")
|
||||
continue
|
||||
|
||||
# Check if each task has consistent number of runs
|
||||
run_counts = [len(scores) for scores in task_results.values()]
|
||||
if len(set(run_counts)) > 1:
|
||||
logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}")
|
||||
continue
|
||||
|
||||
num_runs = run_counts[0]
|
||||
logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task")
|
||||
|
||||
# Get all possible k values
|
||||
k_values = get_possible_k_values(num_runs)
|
||||
logger.info(f"Calculable best@k values: {k_values}")
|
||||
|
||||
# Calculate various best@k values
|
||||
file_results = {"file": file.name}
|
||||
|
||||
for k in k_values:
|
||||
best_at_k_scores = []
|
||||
pass_at_k_scores = []
|
||||
for task_id, scores in task_results.items():
|
||||
try:
|
||||
best_k_score = calculate_best_at_k(scores, k)
|
||||
pass_at_k_score = calculate_pass_at_k(scores, k)
|
||||
pass_at_k_scores.append(pass_at_k_score)
|
||||
best_at_k_scores.append(best_k_score)
|
||||
except ValueError as e:
|
||||
logger.error(f"Error calculating best@{k} for task {task_id}: {e}")
|
||||
continue
|
||||
|
||||
if best_at_k_scores:
|
||||
avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores)
|
||||
file_results[f"best@{k}"] = avg_best_at_k
|
||||
logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}")
|
||||
|
||||
if pass_at_k_scores:
|
||||
avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores)
|
||||
file_results[f"pass@{k}"] = avg_pass_at_k
|
||||
logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}")
|
||||
|
||||
all_results[file.name] = file_results
|
||||
|
||||
# Create and display table
|
||||
if all_results:
|
||||
df = pd.DataFrame(list(all_results.values()))
|
||||
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@')]
|
||||
best_columns = list(df.columns)
|
||||
best_columns.sort(key=lambda x: x, reverse=False)
|
||||
df = df[best_columns]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Experiment Results Summary Table")
|
||||
print("=" * 80)
|
||||
print(df.round(4))
|
||||
print("=" * 80)
|
||||
|
||||
# Save table to CSV
|
||||
output_path = path / "experiment_summary.csv"
|
||||
df.to_csv(output_path)
|
||||
logger.info(f"Results table saved to: {output_path}")
|
||||
else:
|
||||
logger.warning("No valid experiment results found")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_exp_statistic()
|
||||
|
|
@ -1,69 +0,0 @@
|
|||
"""Split the JSONL file into train and validation sets."""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import random
|
||||
|
||||
from default_ids import DEFAULT_TRAIN_IDS, DEFAULT_VAL_IDS
|
||||
|
||||
|
||||
def split_jsonl(
|
||||
input_file: str,
|
||||
train_file: str,
|
||||
val_file: str,
|
||||
ratio: float = 0.75,
|
||||
random_split: bool = False,
|
||||
) -> None:
|
||||
"""Split the JSONL file into train and validation sets."""
|
||||
with open(input_file, "r", encoding="utf-8") as f:
|
||||
data = [json.loads(line) for line in f]
|
||||
|
||||
if random_split:
|
||||
random.shuffle(data)
|
||||
split_idx = int(len(data) * ratio)
|
||||
train_data = data[:split_idx]
|
||||
val_data = data[split_idx:]
|
||||
else:
|
||||
train_data = []
|
||||
val_data = []
|
||||
unknown_ids: list[str] = []
|
||||
for obj in data:
|
||||
if "id" not in obj:
|
||||
raise ValueError(f"Missing 'id' field in input file: {input_file}")
|
||||
obj_id = str(obj["id"])
|
||||
if obj_id in DEFAULT_TRAIN_IDS:
|
||||
train_data.append(obj)
|
||||
elif obj_id in DEFAULT_VAL_IDS:
|
||||
val_data.append(obj)
|
||||
else:
|
||||
unknown_ids.append(obj_id)
|
||||
|
||||
if len(train_data) + len(val_data) != len(data):
|
||||
missing = len(data) - (len(train_data) + len(val_data))
|
||||
examples = ", ".join(unknown_ids) if unknown_ids else "(none)"
|
||||
raise ValueError(
|
||||
f"{missing} samples in {input_file} not found in train_ref/val_ref id sets. Examples: {examples}",
|
||||
)
|
||||
|
||||
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:
|
||||
for item in val_data:
|
||||
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.add_argument(
|
||||
"--random",
|
||||
action="store_true",
|
||||
help="Whether to randomly split input into train/val. "
|
||||
"If false, split strictly by default train/val id sets (see default_ids.py).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
split_jsonl(args.input, args.train, args.val, args.ratio, args.random)
|
||||
File diff suppressed because it is too large
Load diff
File diff suppressed because it is too large
Load diff
|
|
@ -1,346 +0,0 @@
|
|||
"""
|
||||
LongMemEval Evaluation Statistics Analyzer
|
||||
|
||||
Computes detailed statistics from evaluation results including:
|
||||
- Overall accuracy
|
||||
- Accuracy by question type
|
||||
- Timing statistics (summary, retrieval)
|
||||
- Memory extraction statistics
|
||||
|
||||
Usage:
|
||||
python bench/longmemeval/compute_stats.py \
|
||||
--results_dir bench/longmemeval/bench_results/longmemeval_reme
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def load_results(results_dir: str) -> list[dict]:
|
||||
"""Load all question result files from the directory.
|
||||
|
||||
Args:
|
||||
results_dir: Path to the results directory
|
||||
|
||||
Returns:
|
||||
List of result dictionaries
|
||||
"""
|
||||
results_path = Path(results_dir)
|
||||
results = []
|
||||
|
||||
# Load individual question files
|
||||
question_files = sorted(results_path.glob("question_*.json"))
|
||||
|
||||
for file_path in question_files:
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
result = json.load(f)
|
||||
results.append(result)
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error loading {file_path}: {e}")
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def compute_accuracy_stats(results: list[dict]) -> dict[str, Any]:
|
||||
"""Compute overall and per-type accuracy statistics.
|
||||
|
||||
Args:
|
||||
results: List of result dictionaries
|
||||
|
||||
Returns:
|
||||
Dictionary with accuracy statistics
|
||||
"""
|
||||
total = len(results)
|
||||
correct = 0
|
||||
incorrect = 0
|
||||
error = 0
|
||||
|
||||
# Per question type statistics
|
||||
type_stats = defaultdict(lambda: {"total": 0, "correct": 0, "incorrect": 0, "error": 0})
|
||||
|
||||
for r in results:
|
||||
qtype = r.get("question_type", "unknown")
|
||||
judgment = r.get("judgment", {})
|
||||
is_correct = judgment.get("is_correct")
|
||||
|
||||
type_stats[qtype]["total"] += 1
|
||||
|
||||
if is_correct is True:
|
||||
correct += 1
|
||||
type_stats[qtype]["correct"] += 1
|
||||
elif is_correct is False:
|
||||
incorrect += 1
|
||||
type_stats[qtype]["incorrect"] += 1
|
||||
else:
|
||||
error += 1
|
||||
type_stats[qtype]["error"] += 1
|
||||
|
||||
# Compute accuracies
|
||||
overall = {
|
||||
"total": total,
|
||||
"correct": correct,
|
||||
"incorrect": incorrect,
|
||||
"error": error,
|
||||
"accuracy": correct / total if total > 0 else 0,
|
||||
"accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0,
|
||||
}
|
||||
|
||||
by_type = {}
|
||||
for qtype, stats in type_stats.items():
|
||||
valid = stats["correct"] + stats["incorrect"]
|
||||
by_type[qtype] = {
|
||||
**stats,
|
||||
"accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0,
|
||||
"accuracy_valid": stats["correct"] / valid if valid > 0 else 0,
|
||||
}
|
||||
|
||||
return {
|
||||
"overall": overall,
|
||||
"by_question_type": by_type,
|
||||
}
|
||||
|
||||
|
||||
def compute_timing_stats(results: list[dict]) -> dict[str, Any]:
|
||||
"""Compute timing statistics.
|
||||
|
||||
Args:
|
||||
results: List of result dictionaries
|
||||
|
||||
Returns:
|
||||
Dictionary with timing statistics
|
||||
"""
|
||||
summary_times = []
|
||||
retrieve_times = []
|
||||
|
||||
for r in results:
|
||||
summary_ms = r.get("summary_duration_ms", 0)
|
||||
retrieve_ms = r.get("retrieve_duration_ms", 0)
|
||||
|
||||
if summary_ms > 0:
|
||||
summary_times.append(summary_ms)
|
||||
if retrieve_ms > 0:
|
||||
retrieve_times.append(retrieve_ms)
|
||||
|
||||
def compute_stats(times: list[float]) -> dict:
|
||||
if not times:
|
||||
return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0}
|
||||
|
||||
return {
|
||||
"count": len(times),
|
||||
"total_ms": sum(times),
|
||||
"total_min": sum(times) / 1000 / 60,
|
||||
"avg_ms": sum(times) / len(times),
|
||||
"min_ms": min(times),
|
||||
"max_ms": max(times),
|
||||
}
|
||||
|
||||
return {
|
||||
"summary": compute_stats(summary_times),
|
||||
"retrieve": compute_stats(retrieve_times),
|
||||
"total_time_min": (sum(summary_times) + sum(retrieve_times)) / 1000 / 60,
|
||||
}
|
||||
|
||||
|
||||
def compute_memory_stats(results: list[dict]) -> dict[str, Any]:
|
||||
"""Compute memory extraction statistics.
|
||||
|
||||
Args:
|
||||
results: List of result dictionaries
|
||||
|
||||
Returns:
|
||||
Dictionary with memory statistics
|
||||
"""
|
||||
memory_counts = []
|
||||
session_counts = []
|
||||
|
||||
for r in results:
|
||||
memories = r.get("extracted_memories", [])
|
||||
num_sessions = r.get("num_sessions", 0)
|
||||
|
||||
memory_counts.append(len(memories))
|
||||
session_counts.append(num_sessions)
|
||||
|
||||
def compute_stats(counts: list[int]) -> dict:
|
||||
if not counts:
|
||||
return {"count": 0, "total": 0, "avg": 0, "min": 0, "max": 0}
|
||||
|
||||
return {
|
||||
"count": len(counts),
|
||||
"total": sum(counts),
|
||||
"avg": sum(counts) / len(counts),
|
||||
"min": min(counts),
|
||||
"max": max(counts),
|
||||
}
|
||||
|
||||
return {
|
||||
"memories_per_question": compute_stats(memory_counts),
|
||||
"sessions_per_question": compute_stats(session_counts),
|
||||
}
|
||||
|
||||
|
||||
def print_report(
|
||||
accuracy_stats: dict,
|
||||
timing_stats: dict,
|
||||
memory_stats: dict,
|
||||
results_dir: str,
|
||||
):
|
||||
"""Print formatted statistics report.
|
||||
|
||||
Args:
|
||||
accuracy_stats: Accuracy statistics
|
||||
timing_stats: Timing statistics
|
||||
memory_stats: Memory statistics
|
||||
results_dir: Path to results directory
|
||||
"""
|
||||
print("\n" + "=" * 80)
|
||||
print("LONGMEMEVAL EVALUATION STATISTICS")
|
||||
print(f"Results Directory: {results_dir}")
|
||||
print("=" * 80)
|
||||
|
||||
# Overall accuracy
|
||||
overall = accuracy_stats["overall"]
|
||||
print("\n📊 Overall Accuracy:")
|
||||
print(f" Total Questions: {overall['total']}")
|
||||
print(f" ✅ Correct: {overall['correct']} ({100 * overall['accuracy']:.2f}%)")
|
||||
print(
|
||||
f" ❌ Incorrect: {overall['incorrect']} "
|
||||
f"({100 * overall['incorrect'] / overall['total'] if overall['total'] > 0 else 0:.2f}%)",
|
||||
)
|
||||
if overall["error"] > 0:
|
||||
print(f" ⚠️ Error: {overall['error']} ({100 * overall['error'] / overall['total']:.2f}%)")
|
||||
print(f" Accuracy (valid): {100 * overall['accuracy_valid']:.2f}%")
|
||||
|
||||
# Accuracy by question type
|
||||
print("\n📊 Accuracy by Question Type:")
|
||||
print("-" * 60)
|
||||
print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}")
|
||||
print("-" * 60)
|
||||
|
||||
by_type = accuracy_stats["by_question_type"]
|
||||
for qtype in sorted(by_type.keys()):
|
||||
stats = by_type[qtype]
|
||||
print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100 * stats['accuracy']:.2f}%")
|
||||
|
||||
print("-" * 60)
|
||||
|
||||
# Timing statistics
|
||||
print("\n⏱️ Timing Statistics:")
|
||||
summary = timing_stats["summary"]
|
||||
retrieve = timing_stats["retrieve"]
|
||||
|
||||
print(" Memory Summarization:")
|
||||
print(f" Total Time: {summary['total_min']:.2f} min")
|
||||
print(f" Avg per Q: {summary['avg_ms']:.0f} ms")
|
||||
print(f" Min/Max: {summary['min_ms']:.0f} / {summary['max_ms']:.0f} ms")
|
||||
|
||||
print(" Memory Retrieval:")
|
||||
print(f" Total Time: {retrieve['total_min']:.2f} min")
|
||||
print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms")
|
||||
print(f" Min/Max: {retrieve['min_ms']:.0f} / {retrieve['max_ms']:.0f} ms")
|
||||
|
||||
print(f" Total Time: {timing_stats['total_time_min']:.2f} min")
|
||||
|
||||
# Memory statistics
|
||||
print("\n📝 Memory Statistics:")
|
||||
mem = memory_stats["memories_per_question"]
|
||||
sess = memory_stats["sessions_per_question"]
|
||||
|
||||
print(" Extracted Memories per Question:")
|
||||
print(f" Total: {mem['total']}")
|
||||
print(f" Average: {mem['avg']:.1f}")
|
||||
print(f" Min/Max: {mem['min']} / {mem['max']}")
|
||||
|
||||
print(" Sessions per Question:")
|
||||
print(f" Average: {sess['avg']:.1f}")
|
||||
print(f" Min/Max: {sess['min']} / {sess['max']}")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
def save_statistics(
|
||||
accuracy_stats: dict,
|
||||
timing_stats: dict,
|
||||
memory_stats: dict,
|
||||
output_file: str,
|
||||
):
|
||||
"""Save statistics to JSON file.
|
||||
|
||||
Args:
|
||||
accuracy_stats: Accuracy statistics
|
||||
timing_stats: Timing statistics
|
||||
memory_stats: Memory statistics
|
||||
output_file: Path to output file
|
||||
"""
|
||||
stats = {
|
||||
"accuracy": accuracy_stats,
|
||||
"timing": timing_stats,
|
||||
"memory": memory_stats,
|
||||
}
|
||||
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
json.dump(stats, f, indent=4, ensure_ascii=False)
|
||||
|
||||
print(f"\n📁 Statistics saved to: {output_file}")
|
||||
|
||||
|
||||
def main(results_dir: str, output_file: str = None):
|
||||
"""Main function to compute and display statistics.
|
||||
|
||||
Args:
|
||||
results_dir: Path to results directory
|
||||
output_file: Optional path to save statistics JSON
|
||||
"""
|
||||
print(f"\nLoading results from: {results_dir}")
|
||||
|
||||
results = load_results(results_dir)
|
||||
|
||||
if not results:
|
||||
print("❌ No results found!")
|
||||
return
|
||||
|
||||
print(f"Loaded {len(results)} question results")
|
||||
|
||||
# Compute statistics
|
||||
accuracy_stats = compute_accuracy_stats(results)
|
||||
timing_stats = compute_timing_stats(results)
|
||||
memory_stats = compute_memory_stats(results)
|
||||
|
||||
# Print report
|
||||
print_report(accuracy_stats, timing_stats, memory_stats, results_dir)
|
||||
|
||||
# Save to file if specified
|
||||
if output_file:
|
||||
save_statistics(accuracy_stats, timing_stats, memory_stats, output_file)
|
||||
else:
|
||||
# Default output file in results directory
|
||||
default_output = Path(results_dir) / "statistics.json"
|
||||
save_statistics(accuracy_stats, timing_stats, memory_stats, str(default_output))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Compute statistics from LongMemEval evaluation results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--results_dir",
|
||||
type=str,
|
||||
default="bench_results/longmemeval_reme",
|
||||
help="Path to results directory containing question_*.json files",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_file",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to save statistics JSON (default: <results_dir>/statistics.json)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
results_dir=args.results_dir,
|
||||
output_file=args.output_file,
|
||||
)
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,921 +0,0 @@
|
|||
"""
|
||||
LongMemEval Benchmark Evaluator for ReMe - Retrieve Only
|
||||
|
||||
A simplified evaluation pipeline that only runs the retrieve and judge phases:
|
||||
1. Loads LongMemEval benchmark data
|
||||
2. Skips memory summarization (assumes memories are already in vector store)
|
||||
3. Uses questions to query memory and generate answers
|
||||
4. Uses LLM to judge answer correctness
|
||||
5. Generates comprehensive metrics
|
||||
|
||||
This is useful for debugging/tuning the retrieve phase without re-running summary.
|
||||
|
||||
Usage:
|
||||
python benchmark/longmemeval/eval_longmemeval_reme_retrieve.py \
|
||||
--data_path dataset/longmemeval/longmemeval_s_cleaned.json \
|
||||
--top_k 20 --start_index 0 --end_index 10
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from reme.reme import ReMe
|
||||
|
||||
|
||||
# ==================== Configuration ====================
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetrieveEvalConfig:
|
||||
"""Evaluation configuration parameters for retrieve-only mode."""
|
||||
|
||||
data_path: str
|
||||
top_k: int = 10
|
||||
start_index: int = 0
|
||||
end_index: Optional[int] = None
|
||||
max_concurrency: int = 1
|
||||
output_dir: str = "cache/bench_results/longmemeval_reme_retrieve"
|
||||
reme_model_name: str = "qwen-flash"
|
||||
eval_model_name: str = "qwen3-max"
|
||||
algo_version: str = "v1"
|
||||
samples_per_type: int = -1 # Number of samples per question type, -1 for all
|
||||
enable_thinking_params: bool = False
|
||||
# Optional: path to previous summary results to reload memories
|
||||
summary_results_dir: Optional[str] = None
|
||||
|
||||
|
||||
# ==================== Answer Judge Prompts ====================
|
||||
|
||||
|
||||
def get_anscheck_prompt(task: str, question: str, answer: str, response: str, abstention: bool = False) -> str:
|
||||
"""Generate the answer checking prompt based on question type.
|
||||
|
||||
Args:
|
||||
task: Question type, e.g. 'single-session-user', 'multi-session', 'temporal-reasoning'
|
||||
question: The question content
|
||||
answer: The reference answer
|
||||
response: The model's response
|
||||
abstention: Whether this is an unanswerable question
|
||||
|
||||
Returns:
|
||||
Prompt for judging answer correctness
|
||||
"""
|
||||
if not abstention:
|
||||
if task in ["single-session-user", "single-session-assistant", "multi-session"]:
|
||||
template = (
|
||||
"I will give you a question, a correct answer, and a response from a model. Please answer yes i"
|
||||
"f the response contains the correct answer. Otherwise, answer no. If the response is equival"
|
||||
"ent to the correct answer or contains all the intermediate steps to get the correct answer, "
|
||||
"you should also answer yes. If the response only contains a subset of the information required"
|
||||
" by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs"
|
||||
" the model response correct? Answer yes or no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
elif task == "temporal-reasoning":
|
||||
template = (
|
||||
"I will give you a question, a correct answer, and a response from a model. Please answer yes"
|
||||
" if the response contains the correct answer. Otherwise, answer no. If the response is equiv"
|
||||
"alent to the correct answer or contains all the intermediate steps to get the correct answer"
|
||||
", you should also answer yes. If the response only contains a subset of the information requ"
|
||||
"ired by the answer, answer no. In addition, do not penalize off-by-one errors for the numbe"
|
||||
"r of days. If the question asks for the number of days/weeks/months, etc., and the model ma"
|
||||
"kes off-by-one errors (e.g., predicting 19 days when the answer is 18), the model's respon"
|
||||
"se is still correct. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel Response: {}\n\nIs th"
|
||||
"e model response correct? Answer yes or no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
elif task == "knowledge-update":
|
||||
template = (
|
||||
"I will give you a question, a correct answer, and a response from a model. Please answer yes "
|
||||
"if the response contains the correct answer. Otherwise, answer no. If the response contains "
|
||||
"some previous information along with an updated answer, the response should be considered "
|
||||
"as correct as long as the updated answer is the required answer.\n\nQuestion: {}\n\nCorrec"
|
||||
"t Answer: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes or no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
elif task == "single-session-preference":
|
||||
template = (
|
||||
"I will give you a question, a rubric for desired personalized response, and a response from a"
|
||||
" model. Please answer yes if the response satisfies the desired response. Otherwise, answer"
|
||||
" no. The model does not need to reflect all the points in the rubric. The response is corr"
|
||||
"ect as long as it recalls and utilizes the user's personal information correctly.\n\nQuest"
|
||||
"ion: {}\n\nRubric: {}\n\nModel Response: {}\n\nIs the model response correct? Answer yes o"
|
||||
"r no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
else:
|
||||
# Default template
|
||||
template = (
|
||||
"I will give you a question, a correct answer, and a response from a model. Please answer y"
|
||||
"es if the response contains the correct answer. Otherwise, answer no. If the response is "
|
||||
"equivalent to the correct answer or contains all the intermediate steps to get the correc"
|
||||
"t answer, you should also answer yes. If the response only contains a subset of the infor"
|
||||
"mation required by the answer, answer no. \n\nQuestion: {}\n\nCorrect Answer: {}\n\nModel"
|
||||
" Response: {}\n\nIs the model response correct? Answer yes or no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
|
||||
else:
|
||||
template = (
|
||||
"I will give you an unanswerable question, an explanation, and a response from a model. Please "
|
||||
"answer yes if the model correctly identifies the question as unanswerable. The model could say "
|
||||
"that the information is incomplete, or some other information is given but the asked informati"
|
||||
"on is not.\n\nQuestion: {}\n\nExplanation: {}\n\nModel Response: {}\n\nDoes the model correct"
|
||||
"ly identify the question as unanswerable? Answer yes or no only."
|
||||
)
|
||||
prompt = template.format(question, answer, response)
|
||||
return prompt
|
||||
|
||||
|
||||
# ==================== Utilities ====================
|
||||
|
||||
|
||||
class DataLoader:
|
||||
"""Handles loading and parsing of LongMemEval data."""
|
||||
|
||||
@staticmethod
|
||||
def load_json(file_path: str) -> list[dict]:
|
||||
"""Load all entries from a JSON file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
@staticmethod
|
||||
def filter_by_type(data: list[dict], samples_per_type: int = -1) -> list[tuple[int, dict]]:
|
||||
"""Filter data by question type with specified number of samples per type.
|
||||
|
||||
Args:
|
||||
data: List of question entries
|
||||
samples_per_type: Number of samples per type, -1 for all
|
||||
|
||||
Returns:
|
||||
List of tuples (original_index, entry) for selected samples
|
||||
"""
|
||||
if samples_per_type == -1:
|
||||
# Return all with original indices
|
||||
return list(enumerate(data))
|
||||
|
||||
# Group by question type
|
||||
type_groups: dict[str, list[tuple[int, dict]]] = {}
|
||||
for i, entry in enumerate(data):
|
||||
qtype = entry.get("question_type", "unknown")
|
||||
if qtype not in type_groups:
|
||||
type_groups[qtype] = []
|
||||
type_groups[qtype].append((i, entry))
|
||||
|
||||
# Select samples from each type
|
||||
selected = []
|
||||
for qtype, entries in type_groups.items():
|
||||
count = min(samples_per_type, len(entries))
|
||||
selected.extend(entries[:count])
|
||||
logger.info(f" {qtype}: selected {count}/{len(entries)} samples")
|
||||
|
||||
# Sort by original index to maintain order
|
||||
selected.sort(key=lambda x: x[0])
|
||||
return selected
|
||||
|
||||
|
||||
class FileManager:
|
||||
"""Manages file I/O operations."""
|
||||
|
||||
def __init__(self, base_dir: str):
|
||||
self.base_dir = Path(base_dir)
|
||||
self.base_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def save_question_result(self, idx: int, question_id: str, data: dict):
|
||||
"""Save result for a single question."""
|
||||
file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json"
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=4, ensure_ascii=False)
|
||||
logger.info(f"✅ Saved question result to {file_path}")
|
||||
|
||||
def load_question_result(self, idx: int, question_id: str) -> Optional[dict]:
|
||||
"""Load result for a single question if exists."""
|
||||
file_path = self.base_dir / f"question_{idx:04d}_{question_id}.json"
|
||||
if not file_path.exists():
|
||||
return None
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
def save_summary(self, results: list[dict]):
|
||||
"""Save summary of all results."""
|
||||
file_path = self.base_dir / "summary.json"
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, indent=4, ensure_ascii=False)
|
||||
logger.info(f"✅ Saved summary to {file_path}")
|
||||
|
||||
|
||||
# ==================== Evaluation Functions ====================
|
||||
|
||||
|
||||
async def answer_question_with_memories(
|
||||
reme: ReMe,
|
||||
question: str,
|
||||
memories: str,
|
||||
user_id: str = None,
|
||||
model_name: str = "qwen3-max",
|
||||
):
|
||||
"""
|
||||
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
|
||||
|
||||
Args:
|
||||
reme: ReMe instance with default_llm and prompt_handler
|
||||
question: The question to answer
|
||||
memories: The retrieved memories (formatted as context)
|
||||
user_id: Optional user ID for context formatting
|
||||
model_name: Model name to use for LLM request
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'answer' fields
|
||||
"""
|
||||
# Format context with memories
|
||||
if user_id:
|
||||
context = reme.prompt_handler.prompt_format(
|
||||
"TEMPLATE_MEMOS",
|
||||
user_id=user_id,
|
||||
memories=memories,
|
||||
)
|
||||
else:
|
||||
context = f"Memories:\n{memories}"
|
||||
|
||||
# Use PROMPT_MEMZERO_JSON template for structured JSON response
|
||||
prompt = reme.prompt_handler.prompt_format(
|
||||
"PROMPT_MEMZERO_JSON",
|
||||
context=context,
|
||||
question=question,
|
||||
)
|
||||
|
||||
result = await reme.get_llm(name=model_name).simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=None,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
# ==================== Memory Operations ====================
|
||||
|
||||
|
||||
class RetrieveProcessor:
|
||||
"""Handles ReMe memory retrieve operations only."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reme: ReMe,
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
algo_version: str = "v1",
|
||||
enable_thinking_params: bool = False,
|
||||
):
|
||||
self.reme = reme
|
||||
self.reme_model_name = reme_model_name
|
||||
self.eval_model_name = eval_model_name
|
||||
self.algo_version = algo_version
|
||||
self.enable_thinking_params = enable_thinking_params
|
||||
|
||||
async def search_memory(
|
||||
self,
|
||||
query: str,
|
||||
user_id: str,
|
||||
top_k: int = 20,
|
||||
) -> tuple[dict, list, float]:
|
||||
"""
|
||||
Search memory using ReMe and return structured answer with reasoning.
|
||||
|
||||
Returns:
|
||||
tuple: (answer_dict, agent_messages, duration_ms)
|
||||
answer_dict contains: {"reasoning": str, "answer": str, "memories": str}
|
||||
"""
|
||||
start = time.time()
|
||||
|
||||
# Retrieve memories from ReMe using new API
|
||||
result = await self.reme.retrieve_memory(
|
||||
llm_config_name="qwen3-max",
|
||||
query=query,
|
||||
retrieve_top_k=top_k,
|
||||
user_name=user_id,
|
||||
version=self.algo_version,
|
||||
return_dict=True,
|
||||
enable_time_filter=True,
|
||||
enable_thinking_params=True,
|
||||
)
|
||||
|
||||
# Extract memories from response
|
||||
memories = result["answer"]
|
||||
agent_messages = [x.simple_dump(enable_argument_dict=True) for x in result["messages"]]
|
||||
retrieved_nodes = [x.model_dump(exclude_none=True) for x in result["retrieved_nodes"]]
|
||||
|
||||
# Use LLM to generate structured answer from memories
|
||||
answer_result = await answer_question_with_memories(
|
||||
reme=self.reme,
|
||||
question=query,
|
||||
memories=memories,
|
||||
user_id=user_id,
|
||||
model_name=self.eval_model_name,
|
||||
)
|
||||
|
||||
# Add original memories to the result
|
||||
answer_result["memories"] = memories
|
||||
answer_result["retrieved_nodes"] = retrieved_nodes
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return answer_result, agent_messages, duration_ms
|
||||
|
||||
|
||||
# ==================== Answer Judge ====================
|
||||
|
||||
|
||||
class LongMemEvalJudge:
|
||||
"""LongMemEval answer judge using LLM."""
|
||||
|
||||
def __init__(self, reme: ReMe, model: str = "qwen3-max"):
|
||||
self.reme = reme
|
||||
self.model = model
|
||||
|
||||
async def judge_answer(
|
||||
self,
|
||||
question_type: str,
|
||||
question: str,
|
||||
answer: str,
|
||||
response: str,
|
||||
abstention: bool = False,
|
||||
) -> dict:
|
||||
"""
|
||||
Judge if the model's response is correct.
|
||||
|
||||
Returns:
|
||||
dict with is_correct, llm_response, and judge_prompt
|
||||
"""
|
||||
prompt = get_anscheck_prompt(question_type, question, answer, response, abstention)
|
||||
|
||||
try:
|
||||
llm_response = await self.reme.get_llm("default").simple_request(
|
||||
prompt=prompt,
|
||||
model_name=self.model,
|
||||
)
|
||||
llm_response_lower = llm_response.strip().lower()
|
||||
is_correct = llm_response_lower.startswith("yes")
|
||||
|
||||
return {
|
||||
"is_correct": is_correct,
|
||||
"llm_response": llm_response,
|
||||
"judge_prompt": prompt,
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"is_correct": None,
|
||||
"error": str(e),
|
||||
"judge_prompt": prompt,
|
||||
}
|
||||
|
||||
|
||||
# ==================== Metrics ====================
|
||||
|
||||
|
||||
class MetricsAggregator:
|
||||
"""Aggregates evaluation metrics for LongMemEval."""
|
||||
|
||||
@staticmethod
|
||||
def compute_metrics(results: list[dict]) -> dict[str, Any]:
|
||||
"""Compute overall and per-type metrics."""
|
||||
total = len(results)
|
||||
correct = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is True)
|
||||
incorrect = sum(1 for r in results if r.get("judgment", {}).get("is_correct") is False)
|
||||
error = total - correct - incorrect
|
||||
|
||||
metrics = {
|
||||
"total": total,
|
||||
"correct": correct,
|
||||
"incorrect": incorrect,
|
||||
"error": error,
|
||||
"accuracy": correct / total if total > 0 else 0,
|
||||
"accuracy_valid": correct / (correct + incorrect) if (correct + incorrect) > 0 else 0,
|
||||
}
|
||||
|
||||
# Per question type statistics
|
||||
type_stats = {}
|
||||
for r in results:
|
||||
qtype = r.get("question_type", "unknown")
|
||||
if qtype not in type_stats:
|
||||
type_stats[qtype] = {"total": 0, "correct": 0, "incorrect": 0}
|
||||
type_stats[qtype]["total"] += 1
|
||||
if r.get("judgment", {}).get("is_correct") is True:
|
||||
type_stats[qtype]["correct"] += 1
|
||||
elif r.get("judgment", {}).get("is_correct") is False:
|
||||
type_stats[qtype]["incorrect"] += 1
|
||||
|
||||
metrics["by_question_type"] = {
|
||||
qtype: {
|
||||
**stats,
|
||||
"accuracy": stats["correct"] / stats["total"] if stats["total"] > 0 else 0,
|
||||
"accuracy_valid": (
|
||||
stats["correct"] / (stats["correct"] + stats["incorrect"])
|
||||
if (stats["correct"] + stats["incorrect"]) > 0
|
||||
else 0
|
||||
),
|
||||
}
|
||||
for qtype, stats in type_stats.items()
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_timing_stats(results: list[dict]) -> dict[str, Any]:
|
||||
"""Compute timing statistics."""
|
||||
retrieve_times = []
|
||||
|
||||
for r in results:
|
||||
retrieve_ms = r.get("retrieve_duration_ms", 0)
|
||||
if retrieve_ms > 0:
|
||||
retrieve_times.append(retrieve_ms)
|
||||
|
||||
def compute_stats(times: list[float]) -> dict:
|
||||
if not times:
|
||||
return {"count": 0, "total_ms": 0, "avg_ms": 0, "min_ms": 0, "max_ms": 0}
|
||||
|
||||
return {
|
||||
"count": len(times),
|
||||
"total_ms": sum(times),
|
||||
"total_min": sum(times) / 1000 / 60,
|
||||
"avg_ms": sum(times) / len(times),
|
||||
"min_ms": min(times),
|
||||
"max_ms": max(times),
|
||||
}
|
||||
|
||||
return {
|
||||
"retrieve": compute_stats(retrieve_times),
|
||||
"total_time_min": sum(retrieve_times) / 1000 / 60,
|
||||
}
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
|
||||
class LongMemEvalRetrieveEvaluator:
|
||||
"""Retrieve-only evaluator for LongMemEval benchmark using ReMe."""
|
||||
|
||||
def __init__(self, config: RetrieveEvalConfig):
|
||||
self.config = config
|
||||
self.reme = ReMe(
|
||||
default_llm_config={
|
||||
"model_name": self.config.reme_model_name,
|
||||
},
|
||||
llms={
|
||||
"qwen-plus-think": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen-plus",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"qwen3-max-think": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"qwen3-max": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
"extra_body": {
|
||||
"enable_thinking": False,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
# Load evaluation prompts into ReMe's prompt handler
|
||||
prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml"
|
||||
self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path)
|
||||
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.retrieve_processor = RetrieveProcessor(
|
||||
self.reme,
|
||||
config.reme_model_name,
|
||||
config.eval_model_name,
|
||||
config.algo_version,
|
||||
config.enable_thinking_params,
|
||||
)
|
||||
self.judge = LongMemEvalJudge(self.reme, config.eval_model_name)
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry."""
|
||||
await self.reme.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Async context manager exit with cleanup."""
|
||||
await self.reme.close()
|
||||
return False
|
||||
|
||||
async def process_question_entry(self, entry: dict, idx: int) -> dict:
|
||||
"""Process a single question entry (retrieve + judge only).
|
||||
|
||||
Args:
|
||||
entry: A question entry from LongMemEval dataset
|
||||
idx: Index of the question
|
||||
|
||||
Returns:
|
||||
Result dictionary
|
||||
"""
|
||||
question_id = entry["question_id"]
|
||||
question = entry["question"]
|
||||
answer = entry["answer"]
|
||||
question_type = entry["question_type"]
|
||||
question_date = entry.get("question_date", "")
|
||||
haystack_dates = entry["haystack_dates"]
|
||||
haystack_session_ids = entry["haystack_session_ids"]
|
||||
haystack_sessions = entry["haystack_sessions"]
|
||||
|
||||
# Use question_id as user_id for isolation (same as full eval)
|
||||
user_id = f"longmemeval_{question_id}"
|
||||
|
||||
logger.info(f"\n{'='*60}")
|
||||
logger.info(f"Question ID: {question_id}")
|
||||
logger.info(f"Question Type: {question_type}")
|
||||
logger.info(f"Question: {question}")
|
||||
logger.info(f"Question_date: {question_date}")
|
||||
logger.info(f"Answer: {answer}")
|
||||
logger.info(f"Number of sessions: {len(haystack_sessions)}")
|
||||
logger.info(f"{'='*60}")
|
||||
|
||||
# Skip summary phase - directly search memory and answer question
|
||||
logger.info(" Retrieving and answering question using ReMe...")
|
||||
answer_dict, retrieve_messages, retrieve_duration_ms = await self.retrieve_processor.search_memory(
|
||||
query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question,
|
||||
user_id=user_id,
|
||||
top_k=self.config.top_k,
|
||||
)
|
||||
|
||||
# Extract answer and reasoning from the structured response
|
||||
model_response = answer_dict.get("answer", "")
|
||||
model_reasoning = answer_dict.get("reasoning", "")
|
||||
retrieved_memories = answer_dict.get("memories", "")
|
||||
retrieved_nodes = answer_dict.get("retrieved_nodes", [])
|
||||
|
||||
# Judge answer correctness
|
||||
logger.info(" Judging answer correctness...")
|
||||
judgment = await self.judge.judge_answer(
|
||||
question_type=question_type,
|
||||
question=question,
|
||||
answer=answer,
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
is_correct = judgment.get("is_correct")
|
||||
logger.info(
|
||||
f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}",
|
||||
)
|
||||
|
||||
result = {
|
||||
"question_id": question_id,
|
||||
"question_type": question_type,
|
||||
"question": question,
|
||||
"answer": answer,
|
||||
"question_date": question_date,
|
||||
"haystack_dates": haystack_dates,
|
||||
"haystack_session_ids": haystack_session_ids,
|
||||
"num_sessions": len(haystack_sessions),
|
||||
"model_response": model_response,
|
||||
"model_reasoning": model_reasoning,
|
||||
"retrieved_memories": retrieved_memories,
|
||||
"retrieved_nodes": retrieved_nodes,
|
||||
"judgment": judgment,
|
||||
"retrieve_duration_ms": retrieve_duration_ms,
|
||||
"retrieve_messages": retrieve_messages,
|
||||
}
|
||||
|
||||
# Save individual result
|
||||
self.file_manager.save_question_result(idx, question_id, result)
|
||||
|
||||
logger.info(f" Question {question_id} - Completed")
|
||||
|
||||
return result
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the retrieve-only evaluation pipeline with parallel processing."""
|
||||
start_time = time.time()
|
||||
|
||||
# NOTE: Do NOT clear vector store - we assume memories are already there from previous summary run
|
||||
|
||||
# Load dataset
|
||||
logger.info(f"Loading dataset from: {self.config.data_path}")
|
||||
all_data = self.data_loader.load_json(self.config.data_path)
|
||||
logger.info(f"Total questions in dataset: {len(all_data)}")
|
||||
|
||||
# Filter by question type
|
||||
logger.info(f"Filtering by type (samples_per_type={self.config.samples_per_type}):")
|
||||
filtered_data = self.data_loader.filter_by_type(all_data, self.config.samples_per_type)
|
||||
logger.info(f"Selected {len(filtered_data)} questions after filtering")
|
||||
|
||||
# Apply start_index and end_index on filtered data
|
||||
end_index = self.config.end_index or len(filtered_data)
|
||||
start_index = self.config.start_index
|
||||
end_index = min(end_index, len(filtered_data))
|
||||
|
||||
# Get the slice we want to process
|
||||
data_to_process = filtered_data[start_index:end_index]
|
||||
total_questions = len(data_to_process)
|
||||
|
||||
logger.info(f"Processing {total_questions} questions (index {start_index} to {end_index - 1})")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("LONGMEMEVAL EVALUATION - REME (RETRIEVE ONLY)")
|
||||
print(f"Samples per type: {self.config.samples_per_type} (-1 = all)")
|
||||
print(f"Questions to process: {total_questions} | Top-K: {self.config.top_k}")
|
||||
print(f"Max Concurrency: {self.config.max_concurrency}")
|
||||
print(f"ReMe Model: {self.config.reme_model_name} | Eval Model: {self.config.eval_model_name}")
|
||||
print(f"Algo Version: {self.config.algo_version}")
|
||||
print("⚠️ NOTE: Assumes memories are already in vector store from previous summary run")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Use semaphore to control concurrency
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
|
||||
async def process_with_semaphore(idx: int, original_idx: int, entry: dict) -> Optional[dict]:
|
||||
"""Process a question with semaphore for concurrency control."""
|
||||
async with semaphore:
|
||||
question_id = entry["question_id"]
|
||||
|
||||
# Check cache first (use original index for cache file naming)
|
||||
cached_result = self.file_manager.load_question_result(original_idx, question_id)
|
||||
if cached_result:
|
||||
print(f"⚡ [{idx}/{total_questions}] Skipping question {original_idx} (cached)")
|
||||
return cached_result
|
||||
|
||||
print(f"\n{'#'*60}")
|
||||
print(f"### [{idx}/{total_questions}] Processing Question {original_idx} ###")
|
||||
print(f"{'#'*60}")
|
||||
|
||||
try:
|
||||
result = await self.process_question_entry(entry, original_idx)
|
||||
print(f"✅ [{idx}/{total_questions}] Completed question {original_idx}")
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error(f"❌ Error processing question {original_idx}: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
return {
|
||||
"question_id": question_id,
|
||||
"error": str(e),
|
||||
"question_type": entry.get("question_type", "unknown"),
|
||||
"question": entry.get("question", ""),
|
||||
"answer": entry.get("answer", ""),
|
||||
"judgment": {"is_correct": None, "error": str(e)},
|
||||
}
|
||||
|
||||
# Create all tasks from filtered data (each item is a tuple of (original_idx, entry))
|
||||
tasks = [
|
||||
process_with_semaphore(idx + 1, original_idx, entry)
|
||||
for idx, (original_idx, entry) in enumerate(data_to_process)
|
||||
]
|
||||
|
||||
# Execute in parallel with controlled concurrency
|
||||
all_results = await asyncio.gather(*tasks, return_exceptions=False)
|
||||
|
||||
# Filter out None results if any
|
||||
all_results = [r for r in all_results if r is not None]
|
||||
|
||||
# Save summary
|
||||
self.file_manager.save_summary(all_results)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
if total_questions > 0:
|
||||
print(f" Average time per question: {elapsed / total_questions:.2f}s")
|
||||
|
||||
# Compute and report metrics
|
||||
self._report_metrics(all_results)
|
||||
|
||||
return all_results
|
||||
|
||||
def _report_metrics(self, results: list[dict]):
|
||||
"""Report evaluation metrics."""
|
||||
metrics = MetricsAggregator.compute_metrics(results)
|
||||
timing_stats = MetricsAggregator.compute_timing_stats(results)
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("EVALUATION SUMMARY - LONGMEMEVAL - REME (RETRIEVE ONLY)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Overall Results:")
|
||||
print(f" ✅ Correct: {metrics['correct']}/{metrics['total']} ({100*metrics['accuracy']:.2f}%)")
|
||||
print(
|
||||
f" ❌ Incorrect: {metrics['incorrect']}/{metrics['total']}"
|
||||
f" ({100*metrics['incorrect']/metrics['total'] if metrics['total'] > 0 else 0:.2f}%)",
|
||||
)
|
||||
if metrics["error"] > 0:
|
||||
print(f" ⚠️ Error: {metrics['error']}/{metrics['total']} ({100*metrics['error']/metrics['total']:.2f}%)")
|
||||
print(f" Accuracy (valid): {100*metrics['accuracy_valid']:.2f}%")
|
||||
|
||||
print("\n📊 Accuracy by Question Type:")
|
||||
print("-" * 60)
|
||||
print(f"{'Question Type':<30} {'Correct':<10} {'Total':<10} {'Accuracy':<10}")
|
||||
print("-" * 60)
|
||||
for qtype in sorted(metrics["by_question_type"].keys()):
|
||||
stats = metrics["by_question_type"][qtype]
|
||||
print(f"{qtype:<30} {stats['correct']:<10} {stats['total']:<10} {100*stats['accuracy']:.2f}%")
|
||||
print("-" * 60)
|
||||
|
||||
print("\n⏱️ Timing Statistics (Retrieve Only):")
|
||||
retrieve = timing_stats["retrieve"]
|
||||
print(" Memory Retrieval:")
|
||||
print(f" Total Time: {retrieve['total_min']:.2f} min")
|
||||
print(f" Avg per Q: {retrieve['avg_ms']:.0f} ms")
|
||||
print(f" Total Time: {timing_stats['total_time_min']:.2f} min")
|
||||
|
||||
# Save metrics
|
||||
final_results = {
|
||||
"accuracy": metrics,
|
||||
"timing": timing_stats,
|
||||
}
|
||||
metrics_file = self.file_manager.base_dir / "eval_statistics.json"
|
||||
with open(metrics_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, indent=4, ensure_ascii=False)
|
||||
print(f"\n📁 Statistics saved to: {metrics_file}")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
# ==================== Entry Point ====================
|
||||
|
||||
|
||||
async def main_async(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
start_index: int = 0,
|
||||
end_index: Optional[int] = None,
|
||||
max_concurrency: int = 1,
|
||||
output_dir: str = "bench_results/longmemeval_reme_retrieve",
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
algo_version: str = "v1",
|
||||
samples_per_type: int = -1,
|
||||
enable_thinking_params: bool = False,
|
||||
summary_results_dir: Optional[str] = None,
|
||||
):
|
||||
"""Main async entry point for LongMemEval retrieve-only evaluation with proper resource cleanup."""
|
||||
config = RetrieveEvalConfig(
|
||||
data_path=data_path,
|
||||
top_k=top_k,
|
||||
start_index=start_index,
|
||||
end_index=end_index,
|
||||
max_concurrency=max_concurrency,
|
||||
output_dir=output_dir,
|
||||
reme_model_name=reme_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version,
|
||||
samples_per_type=samples_per_type,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
summary_results_dir=summary_results_dir,
|
||||
)
|
||||
|
||||
# Use async context manager for automatic cleanup
|
||||
async with LongMemEvalRetrieveEvaluator(config) as evaluator:
|
||||
await evaluator.run_evaluation()
|
||||
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
start_index: int = 0,
|
||||
end_index: Optional[int] = None,
|
||||
max_concurrency: int = 1,
|
||||
output_dir: str = "bench_results/longmemeval_reme_retrieve",
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
algo_version: str = "v1",
|
||||
samples_per_type: int = -1,
|
||||
enable_thinking_params: bool = False,
|
||||
summary_results_dir: Optional[str] = None,
|
||||
):
|
||||
"""Main entry point for LongMemEval retrieve-only evaluation."""
|
||||
asyncio.run(
|
||||
main_async(
|
||||
data_path=data_path,
|
||||
top_k=top_k,
|
||||
start_index=start_index,
|
||||
end_index=end_index,
|
||||
max_concurrency=max_concurrency,
|
||||
output_dir=output_dir,
|
||||
reme_model_name=reme_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version,
|
||||
samples_per_type=samples_per_type,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
summary_results_dir=summary_results_dir,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate ReMe on LongMemEval benchmark (Retrieve Phase Only)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
# default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_s_cleaned.json",
|
||||
default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_oracle.json",
|
||||
help="Path to LongMemEval JSON file",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=10,
|
||||
help="Number of memories to retrieve (default: 10)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_index",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Start index for processing questions (default: 0)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--end_index",
|
||||
type=int,
|
||||
default=None,
|
||||
help="End index for processing questions (default: None, process all)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Maximum concurrent question processing (default: 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default="bench_results/longmemeval_reme_retrieve_gpt4",
|
||||
help="Output directory for results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reme_model_name",
|
||||
type=str,
|
||||
default="gpt-4o-mini-2024-07-18",
|
||||
help="Model name for ReMe operations (default: gpt-4o-mini-2024-07-18)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eval_model_name",
|
||||
type=str,
|
||||
default="gpt-4o-mini-2024-07-18",
|
||||
help="Model name for evaluation/judgment (default: gpt-4o-mini-2024-07-18)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--algo_version",
|
||||
type=str,
|
||||
default="longmemeval",
|
||||
help="Algorithm version for retrieval (default: longmemeval)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--samples_per_type",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of samples per question type, -1 for all (default: 4)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable_thinking_params",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Enable thinking parameters for retrieval (default: False)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--summary_results_dir",
|
||||
type=str,
|
||||
default="/Users/zhouwk/PycharmProjects/ReMe/benchmark/longmemeval/bench_results",
|
||||
help="Optional: path to previous summary results directory (for reference)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_cache",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Ignore cached results and re-run all questions (default: False)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
print(f"args={args}!")
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
top_k=args.top_k,
|
||||
start_index=args.start_index,
|
||||
end_index=args.end_index,
|
||||
max_concurrency=args.max_concurrency,
|
||||
output_dir=args.output_dir,
|
||||
reme_model_name=args.reme_model_name,
|
||||
eval_model_name=args.eval_model_name,
|
||||
algo_version=args.algo_version,
|
||||
samples_per_type=args.samples_per_type,
|
||||
enable_thinking_params=args.enable_thinking_params,
|
||||
summary_results_dir=args.summary_results_dir,
|
||||
)
|
||||
|
|
@ -1,232 +0,0 @@
|
|||
"""Evaluation tools for ReMe LongMemEval benchmark."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
|
||||
from reme.reme import ReMe
|
||||
|
||||
|
||||
# Load prompts from YAML file
|
||||
_YAML_PATH = Path(__file__).parent / "eval_reme.yaml"
|
||||
with open(_YAML_PATH, "r", encoding="utf-8") as f:
|
||||
_PROMPTS = yaml.safe_load(f)
|
||||
|
||||
|
||||
async def evaluation_for_memory_integrity(
|
||||
reme: ReMe,
|
||||
extract_memories: str,
|
||||
target_memory: str,
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Memory Integrity Evaluation
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
extract_memories: A formatted string concatenating all memory points extracted by the memory system.
|
||||
target_memory: The target key memory point.
|
||||
model_name: Model name for evaluation
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'score' fields
|
||||
"""
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY"].format(
|
||||
memories=extract_memories,
|
||||
expected_memory_point=target_memory,
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_memory_accuracy(
|
||||
reme: ReMe,
|
||||
dialogue: str,
|
||||
golden_memories: str,
|
||||
candidate_memory: str,
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Memory Accuracy Evaluation
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
dialogue: The complete human-machine dialogue record.
|
||||
golden_memories: The core memory points for this dialogue segment in the evaluation set .
|
||||
candidate_memory: A specific memory point extracted by the memory system being evaluated.
|
||||
model_name: Model name for evaluation
|
||||
|
||||
Returns:
|
||||
dict with 'accuracy_score', 'is_included_in_golden_memories', and 'reason' fields
|
||||
"""
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_ACCURACY"].format(
|
||||
dialogue=dialogue,
|
||||
golden_memories=golden_memories,
|
||||
candidate_memory=candidate_memory,
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_update_memory(
|
||||
reme: ReMe,
|
||||
extract_memories: str,
|
||||
target_update_memory: str,
|
||||
original_memory: str,
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Memory Update Evaluation
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
extract_memories: A formatted string concatenating all memory points extracted by the memory system .
|
||||
target_update_memory: The target updated memory point.
|
||||
original_memory: A formatted string concatenating all original memory points corresponding.
|
||||
model_name: Model name for evaluation
|
||||
|
||||
Returns:
|
||||
dict with 'reason' and 'evaluation_result' fields
|
||||
"""
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_UPDATE_MEMORY"].format(
|
||||
memories=extract_memories,
|
||||
updated_memory=target_update_memory,
|
||||
original_memory=original_memory,
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_question(
|
||||
reme: ReMe,
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Question-Answering Evaluation
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
question: The question string to be evaluated.
|
||||
reference_answer: The reference (gold-standard) answer.
|
||||
key_memory_points: The memory points used to derive the reference answer.
|
||||
response: The answer produced by the memory system.
|
||||
model_name: Model name for evaluation
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'evaluation_result' fields
|
||||
"""
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_question2(
|
||||
reme: ReMe,
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
dialogue: str = "",
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Question-Answering Evaluation with Dialogue Context (Version 2)
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
question: The question string to be evaluated.
|
||||
reference_answer: The reference (gold-standard) answer.
|
||||
key_memory_points: The memory points used to derive the reference answer.
|
||||
response: The answer produced by the memory system.
|
||||
dialogue: The formatted dialogue history (role, content, time_created).
|
||||
model_name: Model name for evaluation
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'evaluation_result' fields
|
||||
"""
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
dialogue=dialogue if dialogue else "",
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def answer_question_with_memories(
|
||||
reme: ReMe,
|
||||
question: str,
|
||||
memories: str,
|
||||
user_id: str = None,
|
||||
model_name: str = "qwen3-max",
|
||||
) -> dict:
|
||||
"""
|
||||
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
question: The question to answer
|
||||
memories: The retrieved memories (formatted as context)
|
||||
user_id: Optional user ID for context formatting
|
||||
model_name: Model name for LLM request
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'answer' fields
|
||||
"""
|
||||
# Format context with memories
|
||||
if user_id:
|
||||
context = _PROMPTS["TEMPLATE_MEMOS"].format(
|
||||
user_id=user_id,
|
||||
memories=memories,
|
||||
)
|
||||
else:
|
||||
context = f"Memories:\n{memories}"
|
||||
|
||||
# Use PROMPT_MEMZERO_JSON template for structured JSON response
|
||||
prompt = _PROMPTS["PROMPT_MEMZERO_JSON"].format(
|
||||
context=context,
|
||||
question=question,
|
||||
)
|
||||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
|
@ -1,97 +0,0 @@
|
|||
"""LLM utilities for LongMemEval benchmark evaluation."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
|
||||
from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log
|
||||
|
||||
from reme.core.schema import Message
|
||||
from reme.core.utils import load_env
|
||||
from reme.reme import ReMe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
load_env()
|
||||
|
||||
WAIT_TIME_LOWER = 1
|
||||
WAIT_TIME_UPPER = 60
|
||||
RETRY_TIMES = 5
|
||||
|
||||
|
||||
@retry(
|
||||
wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER),
|
||||
stop=stop_after_attempt(3),
|
||||
reraise=True,
|
||||
before_sleep=before_sleep_log(logger, logging.WARNING),
|
||||
)
|
||||
async def llm_request(reme: ReMe, prompt: str, model_name: str = "qwen3-max", **kwargs) -> str:
|
||||
"""Make an LLM request using ReMe's LLM with optional model override.
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
prompt: The prompt to send to the LLM
|
||||
model_name: Optional model name to override the default model (default: "qwen3-max")
|
||||
**kwargs: Additional arguments to pass to the chat method
|
||||
|
||||
Returns:
|
||||
The assistant's response content
|
||||
"""
|
||||
assistant_message = await reme.llm.chat(
|
||||
messages=[
|
||||
Message(role="user", content=prompt),
|
||||
],
|
||||
model_name=model_name,
|
||||
**kwargs,
|
||||
)
|
||||
return assistant_message.content
|
||||
|
||||
|
||||
@retry(
|
||||
wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER),
|
||||
stop=stop_after_attempt(RETRY_TIMES),
|
||||
reraise=True,
|
||||
before_sleep=before_sleep_log(logger, logging.WARNING),
|
||||
)
|
||||
async def llm_request_for_json(reme: ReMe, prompt: str, model_name: str = "qwen-flash", **kwargs) -> dict:
|
||||
"""Make an LLM request expecting JSON response using ReMe's LLM.
|
||||
|
||||
Args:
|
||||
reme: ReMe instance
|
||||
prompt: The prompt to send to the LLM
|
||||
model_name: Optional model name to override the default model (default: "qwen-flash")
|
||||
**kwargs: Additional arguments to pass to the chat method
|
||||
|
||||
Returns:
|
||||
Parsed JSON object from the LLM response
|
||||
|
||||
Raises:
|
||||
ValueError: If no JSON block is found in the model output
|
||||
"""
|
||||
content = await llm_request(reme, prompt, model_name=model_name, **kwargs)
|
||||
|
||||
match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL)
|
||||
if not match:
|
||||
raise ValueError(f"No JSON block found in model output: {content}")
|
||||
|
||||
json_str = match.group(1).strip()
|
||||
return json.loads(json_str)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
async def test():
|
||||
"""Simple manual test for JSON LLM request."""
|
||||
reme = ReMe()
|
||||
await reme.start()
|
||||
try:
|
||||
r = await llm_request_for_json(
|
||||
reme,
|
||||
'hello? answer in ```json\n{"answer": "..."}```',
|
||||
)
|
||||
print(r)
|
||||
finally:
|
||||
await reme.close()
|
||||
|
||||
asyncio.run(test())
|
||||
|
|
@ -1,23 +1,26 @@
|
|||
"""ReMe"""
|
||||
"""ReMe CLI package."""
|
||||
|
||||
__version__ = "0.4.0.0"
|
||||
|
||||
from . import config
|
||||
from . import core
|
||||
from . import extension
|
||||
from . import memory
|
||||
from . import constants
|
||||
from . import enumeration
|
||||
from . import schema
|
||||
from . import steps
|
||||
from . import utils
|
||||
from .application import Application
|
||||
from .components import BaseComponent
|
||||
from .reme import ReMe
|
||||
|
||||
__version__ = "0.3.1.10"
|
||||
|
||||
__all__ = [
|
||||
"config",
|
||||
"core",
|
||||
"extension",
|
||||
"memory",
|
||||
"Application",
|
||||
"BaseComponent",
|
||||
"ReMe",
|
||||
# submodules
|
||||
"config",
|
||||
"constants",
|
||||
"enumeration",
|
||||
"schema",
|
||||
"steps",
|
||||
"utils",
|
||||
]
|
||||
|
||||
"""
|
||||
conda create -n fl_test2 python=3.10
|
||||
conda activate fl_test2
|
||||
conda env remove -n fl_test2
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
"""config"""
|
||||
"""Config"""
|
||||
|
||||
from .reme_config_parser import ReMeConfigParser
|
||||
from .config_parser import parse_args, resolve_app_config
|
||||
|
||||
__all__ = ["ReMeConfigParser"]
|
||||
__all__ = [
|
||||
"parse_args",
|
||||
"resolve_app_config",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,7 +0,0 @@
|
|||
"""Configuration parser for ReMe framework."""
|
||||
|
||||
from ..core.utils import PydanticConfigParser
|
||||
|
||||
|
||||
class ReMeConfigParser(PydanticConfigParser):
|
||||
"""Configuration parser for ReMe framework."""
|
||||
|
|
@ -1,51 +0,0 @@
|
|||
"""Core"""
|
||||
|
||||
from . import as_llm
|
||||
from . import as_llm_formatter
|
||||
from . import as_token_counter
|
||||
from . import embedding
|
||||
from . import enumeration
|
||||
from . import file_store
|
||||
from . import file_watcher
|
||||
from . import flow
|
||||
from . import llm
|
||||
from . import op
|
||||
from . import schema
|
||||
from . import service
|
||||
from . import token_counter
|
||||
from . import utils
|
||||
from . import vector_store
|
||||
from .application import Application
|
||||
from .base_dict import BaseDict
|
||||
from .prompt_handler import PromptHandler
|
||||
from .registry_factory import R, Registry, RegistryFactory
|
||||
from .runtime_context import RuntimeContext
|
||||
from .service_context import ServiceContext
|
||||
|
||||
__all__ = [
|
||||
# Submodules
|
||||
"as_llm",
|
||||
"as_llm_formatter",
|
||||
"as_token_counter",
|
||||
"embedding",
|
||||
"enumeration",
|
||||
"file_watcher",
|
||||
"flow",
|
||||
"llm",
|
||||
"file_store",
|
||||
"op",
|
||||
"schema",
|
||||
"service",
|
||||
"token_counter",
|
||||
"utils",
|
||||
"vector_store",
|
||||
# Classes
|
||||
"Application",
|
||||
"BaseDict",
|
||||
"PromptHandler",
|
||||
"R",
|
||||
"Registry",
|
||||
"RegistryFactory",
|
||||
"RuntimeContext",
|
||||
"ServiceContext",
|
||||
]
|
||||
|
|
@ -1,657 +0,0 @@
|
|||
"""High-level entry point for configuring and running ReMe services and flows."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
from .embedding import BaseEmbeddingModel
|
||||
from .file_store import BaseFileStore
|
||||
from .file_watcher import BaseFileWatcher
|
||||
from .flow import BaseFlow
|
||||
from .llm import BaseLLM
|
||||
from .prompt_handler import PromptHandler
|
||||
from .registry_factory import R
|
||||
from .schema import (
|
||||
EmbeddingModelConfig,
|
||||
Response,
|
||||
ServiceConfig,
|
||||
LLMConfig,
|
||||
VectorStoreConfig,
|
||||
FileStoreConfig,
|
||||
FileWatcherConfig,
|
||||
TokenCounterConfig,
|
||||
)
|
||||
from .service_context import ServiceContext
|
||||
from .token_counter import BaseTokenCounter
|
||||
from .utils import execute_stream_task, PydanticConfigParser, init_logger, MCPClient, print_logo, get_logger, load_env
|
||||
from .vector_store import BaseVectorStore
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class Application:
|
||||
"""Application wrapper that wires together service context, flows, and runtimes."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
working_dir: str | None = None,
|
||||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
log_to_file: bool = True,
|
||||
enable_load_env: bool = True,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
default_as_llm_config: dict | None = None,
|
||||
default_as_llm_formatter_config: dict | None = None,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
default_file_store_config: dict | None = None,
|
||||
default_token_counter_config: dict | None = None,
|
||||
default_file_watcher_config: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
if enable_load_env:
|
||||
load_env()
|
||||
|
||||
self.llm_api_key = llm_api_key or os.getenv("LLM_API_KEY", "")
|
||||
self.llm_base_url = llm_base_url or os.getenv("LLM_BASE_URL", "")
|
||||
self.embedding_api_key = embedding_api_key or os.getenv("EMBEDDING_API_KEY", "")
|
||||
self.embedding_base_url = embedding_base_url or os.getenv("EMBEDDING_BASE_URL", "")
|
||||
|
||||
self.service_context = ServiceContext(
|
||||
*args,
|
||||
service_config=None,
|
||||
parser=parser,
|
||||
working_dir=working_dir,
|
||||
config_path=config_path,
|
||||
enable_logo=enable_logo,
|
||||
log_to_console=log_to_console,
|
||||
log_to_file=log_to_file,
|
||||
default_as_llm_config=default_as_llm_config,
|
||||
default_as_llm_formatter_config=default_as_llm_formatter_config,
|
||||
default_llm_config=default_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
default_vector_store_config=default_vector_store_config,
|
||||
default_file_store_config=default_file_store_config,
|
||||
default_token_counter_config=default_token_counter_config,
|
||||
default_file_watcher_config=default_file_watcher_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.prompt_handler = PromptHandler(language=self.service_config.language)
|
||||
|
||||
# NOTE: flows are initialized here to start service!
|
||||
self.init_flows()
|
||||
|
||||
self._started: bool = False
|
||||
|
||||
@classmethod
|
||||
async def create(cls, *args, **kwargs) -> "Application":
|
||||
"""Create and start an Application instance asynchronously."""
|
||||
instance = cls(*args, **kwargs)
|
||||
await instance.start()
|
||||
return instance
|
||||
|
||||
def init_flows(self):
|
||||
"""Initialize flows."""
|
||||
expression_flow_cls = None
|
||||
for name, flow_cls in R.flows.items():
|
||||
if not self._filter_flows(name):
|
||||
continue
|
||||
|
||||
if name == "ExpressionFlow":
|
||||
expression_flow_cls = flow_cls
|
||||
else:
|
||||
flow: "BaseFlow" = flow_cls(name=name, service_context=self.service_context)
|
||||
self.service_context.flows[flow.name] = flow
|
||||
|
||||
if expression_flow_cls is not None:
|
||||
for name, flow_config in self.service_config.flows.items():
|
||||
if not self._filter_flows(name):
|
||||
continue
|
||||
flow_config.name = name
|
||||
flow: BaseFlow = expression_flow_cls( # noqa
|
||||
flow_config=flow_config,
|
||||
service_context=self.service_context,
|
||||
)
|
||||
self.service_context.flows[flow.name] = flow
|
||||
else:
|
||||
logger.info("No expression flow found, please check your configuration.")
|
||||
|
||||
def _filter_flows(self, name: str) -> bool:
|
||||
"""Filter flows based on enabled_flows and disabled_flows configuration."""
|
||||
if self.service_config.enabled_flows:
|
||||
return name in self.service_config.enabled_flows
|
||||
elif self.service_config.disabled_flows:
|
||||
return name not in self.service_config.disabled_flows
|
||||
else:
|
||||
return True
|
||||
|
||||
@property
|
||||
def service_config(self) -> ServiceConfig:
|
||||
"""Get the service configuration."""
|
||||
return self.service_context.service_config
|
||||
|
||||
async def start(self):
|
||||
"""Start the service context by initializing all configured components."""
|
||||
if self._started:
|
||||
logger.warning("Application has already started.")
|
||||
return self
|
||||
|
||||
init_logger(
|
||||
log_to_console=self.service_config.log_to_console,
|
||||
log_to_file=self.service_config.log_to_file,
|
||||
)
|
||||
logger.info(f"Init ReMe with config: {self.service_config.model_dump_json()}")
|
||||
|
||||
working_path = Path(self.service_config.working_dir)
|
||||
working_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if self.service_config.ray_max_workers > 1:
|
||||
import ray
|
||||
|
||||
if not ray.is_initialized():
|
||||
ray.init(num_cpus=self.service_config.ray_max_workers)
|
||||
|
||||
if self.service_config.thread_pool_max_workers > 0 and (
|
||||
self.service_context.thread_pool is None
|
||||
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
|
||||
):
|
||||
self.service_context.thread_pool = ThreadPoolExecutor(
|
||||
max_workers=self.service_config.thread_pool_max_workers,
|
||||
)
|
||||
elif self.service_config.thread_pool_max_workers <= 0:
|
||||
logger.info("Thread pool is disabled (thread_pool_max_workers <= 0)")
|
||||
|
||||
if self.service_context.service_config.enable_logo:
|
||||
print_logo(service_config=self.service_config)
|
||||
|
||||
for name, config in self.service_config.as_llms.items():
|
||||
if config.backend not in R.as_llms:
|
||||
logger.warning(f"AS LLM backend {config.backend} is not supported.")
|
||||
else:
|
||||
try:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
if not config_dict.get("api_key", ""):
|
||||
config_dict["api_key"] = self.llm_api_key
|
||||
if "client_kwargs" not in config_dict:
|
||||
config_dict["client_kwargs"] = {}
|
||||
if not config_dict["client_kwargs"].get("base_url", ""):
|
||||
config_dict["client_kwargs"]["base_url"] = self.llm_base_url
|
||||
self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize AS LLM '{name}': {e}")
|
||||
|
||||
for name, config in self.service_config.as_llm_formatters.items():
|
||||
if config.backend not in R.as_llm_formatters:
|
||||
logger.warning(f"AS LLM formatter backend {config.backend} is not supported.")
|
||||
else:
|
||||
try:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config.backend](**config_dict)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize AS LLM formatter '{name}': {e}")
|
||||
|
||||
for name, config in self.service_config.as_token_counters.items():
|
||||
if config.backend not in R.as_token_counters:
|
||||
logger.warning(f"Token counter backend {config.backend} is not supported.")
|
||||
else:
|
||||
try:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
self.service_context.as_token_counters[name] = R.as_token_counters[config.backend](**config_dict)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize AS token counter '{name}': {e}")
|
||||
|
||||
for name, config in self.service_config.llms.items():
|
||||
if config.backend not in R.llms:
|
||||
logger.warning(f"LLM backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict.setdefault("api_key", self.llm_api_key)
|
||||
config_dict.setdefault("base_url", self.llm_base_url)
|
||||
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
|
||||
await self.service_context.llms[name].start()
|
||||
|
||||
for name, config in self.service_config.embedding_models.items():
|
||||
if config.backend not in R.embedding_models:
|
||||
logger.warning(f"Embedding model backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict.setdefault("api_key", self.embedding_api_key)
|
||||
config_dict.setdefault("base_url", self.embedding_base_url)
|
||||
config_dict.setdefault("cache_dir", working_path / "embedding_cache")
|
||||
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
|
||||
await self.service_context.embedding_models[name].start()
|
||||
|
||||
for name, config in self.service_config.token_counters.items():
|
||||
if config.backend not in R.token_counters:
|
||||
logger.warning(f"Token counter backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict)
|
||||
|
||||
for name, config in self.service_config.vector_stores.items():
|
||||
if config.backend not in R.vector_stores:
|
||||
logger.warning(f"Vector store backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict.update(
|
||||
{
|
||||
"embedding_model": self.service_context.embedding_models[config.embedding_model],
|
||||
"db_path": working_path / "vector_store",
|
||||
},
|
||||
)
|
||||
self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
|
||||
await self.service_context.vector_stores[name].start()
|
||||
|
||||
for name, config in self.service_config.file_stores.items():
|
||||
if config.backend not in R.file_stores:
|
||||
logger.warning(f"File store backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict.update(
|
||||
{
|
||||
"embedding_model": self.service_context.embedding_models[config.embedding_model],
|
||||
"db_path": working_path / "file_store",
|
||||
},
|
||||
)
|
||||
self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict)
|
||||
await self.service_context.file_stores[name].start()
|
||||
|
||||
for name, config in self.service_config.file_watchers.items():
|
||||
if config.backend not in R.file_watchers:
|
||||
logger.warning(f"File watcher backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend", "file_store"})
|
||||
config_dict["file_store"] = self.service_context.file_stores[config.file_store]
|
||||
self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
|
||||
await self.service_context.file_watchers[name].start()
|
||||
|
||||
if self.service_config.mcp_servers:
|
||||
await self.prepare_mcp_servers()
|
||||
|
||||
self._started = True
|
||||
logger.info("ReMe Application started")
|
||||
return self
|
||||
|
||||
# pylint: disable=too-many-statements
|
||||
async def restart(self, restart_config: dict):
|
||||
"""Restart the application with new config."""
|
||||
|
||||
working_path = Path(self.service_config.working_dir)
|
||||
working_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# as_llms
|
||||
if "as_llms" in restart_config:
|
||||
as_llms_config = restart_config["as_llms"]
|
||||
assert isinstance(as_llms_config, dict)
|
||||
for name, config in as_llms_config.items():
|
||||
if name in self.service_context.as_llms:
|
||||
del self.service_context.as_llms[name]
|
||||
|
||||
if config.get("backend") not in R.as_llms:
|
||||
logger.warning(f"AS LLM backend {config.get('backend')} is not supported.")
|
||||
continue
|
||||
|
||||
try:
|
||||
config_dict = {k: v for k, v in config.items() if k != "backend"}
|
||||
if not config_dict.get("api_key", ""):
|
||||
config_dict["api_key"] = self.llm_api_key
|
||||
if "client_kwargs" not in config_dict:
|
||||
config_dict["client_kwargs"] = {}
|
||||
if not config_dict["client_kwargs"].get("base_url", ""):
|
||||
config_dict["client_kwargs"]["base_url"] = self.llm_base_url
|
||||
self.service_context.as_llms[name] = R.as_llms[config["backend"]](**config_dict)
|
||||
logger.info(f"Restarted AS LLM: {name}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to restart AS LLM '{name}': {e}")
|
||||
|
||||
# as_llm_formatters
|
||||
if "as_llm_formatters" in restart_config:
|
||||
as_llm_formatters_config = restart_config["as_llm_formatters"]
|
||||
assert isinstance(as_llm_formatters_config, dict)
|
||||
for name, config in as_llm_formatters_config.items():
|
||||
if name in self.service_context.as_llm_formatters:
|
||||
del self.service_context.as_llm_formatters[name]
|
||||
|
||||
if config.get("backend") not in R.as_llm_formatters:
|
||||
logger.warning(f"AS LLM formatter backend {config.get('backend')} is not supported.")
|
||||
continue
|
||||
try:
|
||||
config_dict = {k: v for k, v in config.items() if k != "backend"}
|
||||
self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config["backend"]](**config_dict)
|
||||
logger.info(f"Restarted AS LLM formatter: {name}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to restart AS LLM formatter '{name}': {e}")
|
||||
|
||||
# as_token_counters
|
||||
if "as_token_counters" in restart_config:
|
||||
as_token_counters_config = restart_config["as_token_counters"]
|
||||
assert isinstance(as_token_counters_config, dict)
|
||||
for name, config in as_token_counters_config.items():
|
||||
if name in self.service_context.as_token_counters:
|
||||
del self.service_context.as_token_counters[name]
|
||||
|
||||
if config.get("backend") not in R.as_token_counters:
|
||||
logger.warning(f"Token counter backend {config.get('backend')} is not supported.")
|
||||
continue
|
||||
try:
|
||||
config_dict = {k: v for k, v in config.items() if k != "backend"}
|
||||
self.service_context.as_token_counters[name] = R.as_token_counters[config["backend"]](**config_dict)
|
||||
logger.info(f"Restarted AS token counter: {name}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to restart AS token counter '{name}': {e}")
|
||||
|
||||
# llms
|
||||
if "llms" in restart_config:
|
||||
llms_config = restart_config["llms"]
|
||||
assert isinstance(llms_config, dict)
|
||||
for name, config in llms_config.items():
|
||||
if name in self.service_context.llms:
|
||||
llm = self.service_context.llms.pop(name)
|
||||
await llm.close()
|
||||
|
||||
if isinstance(config, dict):
|
||||
config = LLMConfig(**config)
|
||||
if config.backend not in R.llms:
|
||||
logger.warning(f"LLM backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict.setdefault("api_key", self.llm_api_key)
|
||||
config_dict.setdefault("base_url", self.llm_base_url)
|
||||
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
|
||||
await self.service_context.llms[name].start()
|
||||
logger.info(f"Restarted LLM: {name}")
|
||||
|
||||
# embedding_models
|
||||
if "embedding_models" in restart_config:
|
||||
embedding_models_config = restart_config["embedding_models"]
|
||||
assert isinstance(embedding_models_config, dict)
|
||||
updated_names = set()
|
||||
for name, config in embedding_models_config.items():
|
||||
if name in self.service_context.embedding_models:
|
||||
embedding_model = self.service_context.embedding_models.pop(name)
|
||||
await embedding_model.close()
|
||||
|
||||
if isinstance(config, dict):
|
||||
config = EmbeddingModelConfig(**config)
|
||||
if config.backend not in R.embedding_models:
|
||||
logger.warning(f"Embedding model backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict.setdefault("api_key", self.embedding_api_key)
|
||||
config_dict.setdefault("base_url", self.embedding_base_url)
|
||||
config_dict.setdefault("cache_dir", working_path / "embedding_cache")
|
||||
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
|
||||
await self.service_context.embedding_models[name].start()
|
||||
logger.info(f"Restarted embedding model: {name}")
|
||||
updated_names.add(name)
|
||||
|
||||
# update embedding_model attribute for existing vector_stores and file_stores
|
||||
for name in updated_names:
|
||||
for vs_name, vs_config in self.service_config.vector_stores.items():
|
||||
if vs_config.embedding_model == name and vs_name in self.service_context.vector_stores:
|
||||
self.service_context.vector_stores[vs_name].embedding_model = (
|
||||
self.service_context.embedding_models[name]
|
||||
)
|
||||
logger.info(f"Updated embedding model for vector store: {vs_name}")
|
||||
for fs_name, fs_config in self.service_config.file_stores.items():
|
||||
if fs_config.embedding_model == name and fs_name in self.service_context.file_stores:
|
||||
self.service_context.file_stores[fs_name].embedding_model = (
|
||||
self.service_context.embedding_models[name]
|
||||
)
|
||||
logger.info(f"Updated embedding model for file store: {fs_name}")
|
||||
|
||||
# token_counters
|
||||
if "token_counters" in restart_config:
|
||||
token_counters_config = restart_config["token_counters"]
|
||||
assert isinstance(token_counters_config, dict)
|
||||
for name, config in token_counters_config.items():
|
||||
if name in self.service_context.token_counters:
|
||||
del self.service_context.token_counters[name]
|
||||
|
||||
if isinstance(config, dict):
|
||||
config = TokenCounterConfig(**config)
|
||||
if config.backend not in R.token_counters:
|
||||
logger.warning(f"Token counter backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
self.service_context.token_counters[name] = R.token_counters[config.backend](**config_dict)
|
||||
logger.info(f"Restarted token counter: {name}")
|
||||
|
||||
# vector_stores
|
||||
if "vector_stores" in restart_config:
|
||||
vector_stores_config = restart_config["vector_stores"]
|
||||
assert isinstance(vector_stores_config, dict)
|
||||
for name, config in vector_stores_config.items():
|
||||
if name in self.service_context.vector_stores:
|
||||
vector_store = self.service_context.vector_stores.pop(name)
|
||||
await vector_store.close()
|
||||
if isinstance(config, dict):
|
||||
config = VectorStoreConfig(**config)
|
||||
if config.backend not in R.vector_stores:
|
||||
logger.warning(f"Vector store backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict.update(
|
||||
{
|
||||
"embedding_model": self.service_context.embedding_models[config.embedding_model],
|
||||
"db_path": working_path / "vector_store",
|
||||
},
|
||||
)
|
||||
self.service_context.vector_stores[name] = R.vector_stores[config.backend](**config_dict)
|
||||
await self.service_context.vector_stores[name].start()
|
||||
logger.info(f"Restarted vector store: {name}")
|
||||
|
||||
# file_stores
|
||||
if "file_stores" in restart_config:
|
||||
file_stores_config = restart_config["file_stores"]
|
||||
assert isinstance(file_stores_config, dict)
|
||||
for name, config in file_stores_config.items():
|
||||
if name in self.service_context.file_stores:
|
||||
file_store = self.service_context.file_stores.pop(name)
|
||||
await file_store.close()
|
||||
if isinstance(config, dict):
|
||||
config = FileStoreConfig(**config)
|
||||
if config.backend not in R.file_stores:
|
||||
logger.warning(f"File store backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend", "embedding_model"})
|
||||
config_dict.update(
|
||||
{
|
||||
"embedding_model": self.service_context.embedding_models[config.embedding_model],
|
||||
"db_path": working_path / "file_store",
|
||||
},
|
||||
)
|
||||
self.service_context.file_stores[name] = R.file_stores[config.backend](**config_dict)
|
||||
await self.service_context.file_stores[name].start()
|
||||
logger.info(f"Restarted file store: {name}")
|
||||
|
||||
# file_watchers
|
||||
if "file_watchers" in restart_config:
|
||||
file_watchers_config = restart_config["file_watchers"]
|
||||
assert isinstance(file_watchers_config, dict)
|
||||
for name, config in file_watchers_config.items():
|
||||
if name in self.service_context.file_watchers:
|
||||
file_watcher = self.service_context.file_watchers.pop(name)
|
||||
await file_watcher.close()
|
||||
if isinstance(config, dict):
|
||||
config = FileWatcherConfig(**config)
|
||||
if config.backend not in R.file_watchers:
|
||||
logger.warning(f"File watcher backend {config.backend} is not supported.")
|
||||
continue
|
||||
config_dict = config.model_dump(exclude={"backend", "file_store"})
|
||||
config_dict["file_store"] = self.service_context.file_stores[config.file_store]
|
||||
self.service_context.file_watchers[name] = R.file_watchers[config.backend](**config_dict)
|
||||
await self.service_context.file_watchers[name].start()
|
||||
logger.info(f"Restarted file watcher: {name}")
|
||||
|
||||
async def prepare_mcp_servers(self):
|
||||
"""Prepare and initialize MCP server connections."""
|
||||
mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers})
|
||||
for server_name in self.service_config.mcp_servers.keys():
|
||||
try:
|
||||
tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False)
|
||||
self.service_context.mcp_server_mapping[server_name] = {
|
||||
tool_call.name: tool_call for tool_call in tool_calls
|
||||
}
|
||||
for tool_call in tool_calls:
|
||||
logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}")
|
||||
except Exception as e:
|
||||
logger.exception(f"list_tool_calls: {server_name} error: {e}")
|
||||
|
||||
async def close(self) -> bool:
|
||||
"""Close all service components asynchronously."""
|
||||
if not self._started:
|
||||
logger.warning("Application is not started")
|
||||
return True
|
||||
|
||||
for name, file_watcher in self.service_context.file_watchers.items():
|
||||
logger.info(f"Closing file watcher: {name}")
|
||||
await file_watcher.close()
|
||||
|
||||
for name, file_store in self.service_context.file_stores.items():
|
||||
logger.info(f"Closing file store: {name}")
|
||||
await file_store.close()
|
||||
|
||||
for name, vector_store in self.service_context.vector_stores.items():
|
||||
logger.info(f"Closing vector store: {name}")
|
||||
await vector_store.close()
|
||||
|
||||
for name, llm in self.service_context.llms.items():
|
||||
logger.info(f"Closing LLM: {name}")
|
||||
await llm.close()
|
||||
|
||||
for name, embedding_model in self.service_context.embedding_models.items():
|
||||
logger.info(f"Closing embedding model: {name}")
|
||||
await embedding_model.close()
|
||||
|
||||
self.shutdown_thread_pool()
|
||||
self.shutdown_ray()
|
||||
|
||||
self._started = False
|
||||
logger.info("ReMe Application closed")
|
||||
return False
|
||||
|
||||
def shutdown_thread_pool(self, wait: bool = True):
|
||||
"""Shutdown the thread pool executor."""
|
||||
if self.service_context.thread_pool is not None:
|
||||
self.service_context.thread_pool.shutdown(wait=wait)
|
||||
|
||||
def shutdown_ray(self, wait: bool = True):
|
||||
"""Shutdown Ray cluster if it was initialized."""
|
||||
if self.service_config and self.service_config.ray_max_workers > 1:
|
||||
import ray
|
||||
|
||||
ray.shutdown(_exiting_interpreter=not wait)
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry."""
|
||||
return await self.start()
|
||||
|
||||
async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None):
|
||||
"""Async context manager exit."""
|
||||
return await self.close()
|
||||
|
||||
async def execute_flow(self, name: str, **kwargs) -> Response:
|
||||
"""Execute a flow with the given name and parameters."""
|
||||
assert name in self.service_context.flows, f"Flow {name} not found"
|
||||
flow: BaseFlow = self.service_context.flows[name]
|
||||
return await flow.call(**kwargs)
|
||||
|
||||
async def execute_stream_flow(self, name: str, **kwargs):
|
||||
"""Execute a stream flow with the given name and parameters."""
|
||||
assert name in self.service_context.flows, f"Flow {name} not found"
|
||||
flow: BaseFlow = self.service_context.flows[name]
|
||||
assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!"
|
||||
stream_queue = asyncio.Queue()
|
||||
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
|
||||
async for chunk in execute_stream_task(
|
||||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name=name,
|
||||
output_format="str",
|
||||
):
|
||||
yield chunk
|
||||
|
||||
@property
|
||||
def default_llm(self) -> BaseLLM:
|
||||
"""Get the default LLM instance."""
|
||||
return self.service_context.llms.get("default")
|
||||
|
||||
def get_llm(self, name: str):
|
||||
"""Get an LLM instance by name."""
|
||||
return self.service_context.llms.get(name)
|
||||
|
||||
def update_default_llm_name(self, name: str):
|
||||
"""Update the default LLM name."""
|
||||
self.default_llm.model_name = name
|
||||
|
||||
@property
|
||||
def default_embedding_model(self) -> BaseEmbeddingModel:
|
||||
"""Get the default embedding model instance."""
|
||||
return self.service_context.embedding_models.get("default")
|
||||
|
||||
def get_embedding_model(self, name: str):
|
||||
"""Get an embedding model instance by name."""
|
||||
return self.service_context.embedding_models.get(name)
|
||||
|
||||
def update_default_embedding_name(self, name: str):
|
||||
"""Update the default embedding model name."""
|
||||
self.default_embedding_model.model_name = name
|
||||
|
||||
@property
|
||||
def default_vector_store(self) -> BaseVectorStore:
|
||||
"""Get the default vector store instance."""
|
||||
return self.service_context.vector_stores.get("default")
|
||||
|
||||
def get_vector_store(self, name: str):
|
||||
"""Get a vector store instance by name."""
|
||||
return self.service_context.vector_stores.get(name)
|
||||
|
||||
@property
|
||||
def default_file_store(self) -> BaseFileStore:
|
||||
"""Get the default file store instance."""
|
||||
return self.service_context.file_stores.get("default")
|
||||
|
||||
def get_file_store(self, name: str):
|
||||
"""Get a file store instance by name."""
|
||||
return self.service_context.file_stores.get(name)
|
||||
|
||||
@property
|
||||
def default_file_watcher(self) -> BaseFileWatcher:
|
||||
"""Get the default file watcher instance."""
|
||||
return self.service_context.file_watchers.get("default")
|
||||
|
||||
def get_file_watcher(self, name: str):
|
||||
"""Get a file watcher instance by name."""
|
||||
return self.service_context.file_watchers.get(name)
|
||||
|
||||
@property
|
||||
def default_token_counter(self) -> BaseTokenCounter:
|
||||
"""Get the default token counter instance."""
|
||||
return self.service_context.token_counters.get("default")
|
||||
|
||||
def get_token_counter(self, name: str):
|
||||
"""Get a token counter instance by name."""
|
||||
return self.service_context.token_counters.get(name)
|
||||
|
||||
def run_service(self):
|
||||
"""Run the configured service (HTTP, MCP, or CMD)."""
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning)
|
||||
service = R.services[self.service_config.backend](app=self)
|
||||
service.run()
|
||||
|
||||
async def reset_default_collection(self, collection_name: str):
|
||||
"""Reset the default vector store."""
|
||||
await self.service_context.vector_stores["default"].reset_collection(collection_name)
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
"""Module for registering AgentScope LLM models."""
|
||||
|
||||
from agentscope.model import DashScopeChatModel
|
||||
from agentscope.model import OpenAIChatModel
|
||||
|
||||
from ..registry_factory import R
|
||||
|
||||
R.as_llms.register("openai")(OpenAIChatModel)
|
||||
R.as_llms.register("dashscope")(DashScopeChatModel)
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
"""Module for registering AgentScope LLM formatters."""
|
||||
|
||||
from agentscope.formatter import DashScopeChatFormatter
|
||||
|
||||
from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter
|
||||
from ..registry_factory import R
|
||||
|
||||
R.as_llm_formatters.register("openai")(ReMeOpenAIChatFormatter)
|
||||
R.as_llm_formatters.register("dashscope")(DashScopeChatFormatter)
|
||||
|
|
@ -1,215 +0,0 @@
|
|||
"""ReMeOpenAIChatFormatter"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from agentscope.formatter import OpenAIChatFormatter
|
||||
from agentscope.formatter._openai_formatter import (
|
||||
_format_openai_image_block,
|
||||
_to_openai_audio_data,
|
||||
)
|
||||
from agentscope.message import Msg, TextBlock, ImageBlock, URLSource
|
||||
from loguru import logger
|
||||
|
||||
|
||||
def _format_openai_video_block(video_block: dict) -> dict[str, Any]:
|
||||
"""Format a video block for OpenAI API.
|
||||
|
||||
Args:
|
||||
video_block: The video block to format.
|
||||
|
||||
Returns:
|
||||
A dictionary with video content in OpenAI format.
|
||||
"""
|
||||
source = video_block["source"]
|
||||
if source["type"] == "url":
|
||||
url = source["url"]
|
||||
elif source["type"] == "base64":
|
||||
data = source["data"]
|
||||
media_type = source["media_type"]
|
||||
url = f"data:{media_type};base64,{data}"
|
||||
else:
|
||||
raise ValueError(f"Unsupported video source type: {source['type']}")
|
||||
|
||||
return {
|
||||
"type": "video_url",
|
||||
"video_url": {
|
||||
"url": url,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
|
||||
"""ReMeOpenAIChatFormatter"""
|
||||
|
||||
async def _format(
|
||||
self,
|
||||
msgs: list[Msg],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Format message objects into OpenAI API required format.
|
||||
|
||||
Args:
|
||||
msgs (`list[Msg]`):
|
||||
The list of Msg objects to format.
|
||||
|
||||
Returns:
|
||||
`list[dict[str, Any]]`:
|
||||
A list of dictionaries, where each dictionary has "name",
|
||||
"role", and "content" keys.
|
||||
"""
|
||||
self.assert_list_of_msgs(msgs)
|
||||
|
||||
messages: list[dict] = []
|
||||
i = 0
|
||||
while i < len(msgs):
|
||||
msg = msgs[i]
|
||||
content_blocks = []
|
||||
tool_calls = []
|
||||
reasoning_content_blocks = []
|
||||
|
||||
for block in msg.get_content_blocks():
|
||||
typ = block.get("type")
|
||||
if typ == "text":
|
||||
content_blocks.append({**block})
|
||||
|
||||
elif typ == "thinking":
|
||||
# Collect thinking blocks for reasoning_content field
|
||||
# This is compatible with models like DeepSeek that support
|
||||
# extended thinking via reasoning_content field
|
||||
reasoning_content_blocks.append({**block})
|
||||
|
||||
elif typ == "tool_use":
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": block.get("id"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.get("name"),
|
||||
"arguments": json.dumps(
|
||||
block.get("input", {}),
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
elif typ == "tool_result":
|
||||
(
|
||||
textual_output,
|
||||
multimodal_data,
|
||||
) = self.convert_tool_result_to_string(block["output"])
|
||||
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": block.get("id"),
|
||||
"content": (textual_output), # type: ignore[arg-type]
|
||||
"name": block.get("name"),
|
||||
},
|
||||
)
|
||||
|
||||
# Then, handle the multimodal data if any
|
||||
promoted_blocks: list = []
|
||||
for url, multimodal_block in multimodal_data:
|
||||
if multimodal_block["type"] == "image" and self.promote_tool_result_images:
|
||||
promoted_blocks.extend(
|
||||
[
|
||||
TextBlock(
|
||||
type="text",
|
||||
text=f"\n- The image from '{url}': ",
|
||||
),
|
||||
ImageBlock(
|
||||
type="image",
|
||||
source=URLSource(
|
||||
type="url",
|
||||
url=url,
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
if promoted_blocks:
|
||||
# Insert promoted blocks as new user message(s)
|
||||
promoted_blocks = [
|
||||
TextBlock(
|
||||
type="text",
|
||||
text="<system-info>The following are "
|
||||
"the image contents from the tool "
|
||||
f"result of '{block['name']}':",
|
||||
),
|
||||
*promoted_blocks,
|
||||
TextBlock(
|
||||
type="text",
|
||||
text="</system-info>",
|
||||
),
|
||||
]
|
||||
|
||||
msgs.insert(
|
||||
i + 1,
|
||||
Msg(
|
||||
name="user",
|
||||
content=promoted_blocks,
|
||||
role="user",
|
||||
),
|
||||
)
|
||||
|
||||
elif typ == "image":
|
||||
content_blocks.append(
|
||||
_format_openai_image_block(
|
||||
block, # type: ignore[arg-type]
|
||||
),
|
||||
)
|
||||
|
||||
elif typ == "audio":
|
||||
# Filter out audio content when the multimodal model
|
||||
# outputs both text and audio, to prevent errors in
|
||||
# subsequent model calls
|
||||
if msg.role == "assistant":
|
||||
continue
|
||||
input_audio = _to_openai_audio_data(block["source"])
|
||||
content_blocks.append(
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": input_audio,
|
||||
},
|
||||
)
|
||||
|
||||
elif typ == "video":
|
||||
# Filter out video content when the multimodal model
|
||||
# outputs both text and video, to prevent errors in
|
||||
# subsequent model calls
|
||||
if msg.role == "assistant":
|
||||
continue
|
||||
content_blocks.append(
|
||||
_format_openai_video_block(block),
|
||||
)
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
"Unsupported block type %s in the message, skipped.",
|
||||
typ,
|
||||
)
|
||||
|
||||
msg_openai = {
|
||||
"role": msg.role,
|
||||
"name": msg.name,
|
||||
"content": content_blocks or None,
|
||||
}
|
||||
|
||||
if tool_calls:
|
||||
msg_openai["tool_calls"] = tool_calls
|
||||
|
||||
# Add reasoning_content for thinking blocks (compatible with DeepSeek, etc.)
|
||||
if reasoning_content_blocks:
|
||||
reasoning_msg = "\n".join(reasoning.get("thinking", "") for reasoning in reasoning_content_blocks)
|
||||
if reasoning_msg:
|
||||
msg_openai["reasoning_content"] = reasoning_msg
|
||||
|
||||
# When both content and tool_calls are None, skipped
|
||||
if msg_openai["content"] or msg_openai.get("tool_calls"):
|
||||
messages.append(msg_openai)
|
||||
|
||||
# Move to next message
|
||||
i += 1
|
||||
|
||||
return messages
|
||||
|
|
@ -1,8 +0,0 @@
|
|||
"""Module for registering AgentScope token counters."""
|
||||
|
||||
from .reme_token_counter import ReMeTokenCounter
|
||||
from .rule_token_counter import RuleTokenCounter
|
||||
from ..registry_factory import R
|
||||
|
||||
R.as_token_counters.register("hf")(ReMeTokenCounter)
|
||||
R.as_token_counters.register("rule")(RuleTokenCounter)
|
||||
|
|
@ -1,123 +0,0 @@
|
|||
"""Token counter for ReMe."""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from ..utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class ReMeTokenCounter(HuggingFaceTokenCounter):
|
||||
"""Token counter for CoPaw with configurable tokenizer support.
|
||||
|
||||
This class extends HuggingFaceTokenCounter to provide token counting
|
||||
functionality with support for both local and remote tokenizers,
|
||||
as well as HuggingFace mirror for users in China.
|
||||
|
||||
Attributes:
|
||||
pretrained_model_name_or_path: The tokenizer model path or "default" for local tokenizer.
|
||||
use_mirror: Whether to use HuggingFace mirror.
|
||||
token_count_estimate_divisor: Divisor for token estimation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pretrained_model_name_or_path: str,
|
||||
use_mirror: bool = True,
|
||||
token_count_estimate_divisor: float = 3.75,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize the token counter with the specified configuration.
|
||||
|
||||
Args:
|
||||
pretrained_model_name_or_path: The tokenizer model path.
|
||||
use_mirror: Whether to use the HuggingFace mirror
|
||||
(https://hf-mirror.com) for downloading tokenizers. Useful for
|
||||
users in China.
|
||||
token_count_estimate_divisor: Divisor for estimating tokens when
|
||||
tokenizer is unavailable. Defaults to 3.75.
|
||||
**kwargs: Additional keyword arguments passed to HuggingFaceTokenCounter.
|
||||
"""
|
||||
self.pretrained_model_name_or_path = pretrained_model_name_or_path
|
||||
self.use_mirror = use_mirror
|
||||
self.token_count_estimate_divisor = token_count_estimate_divisor
|
||||
|
||||
# Set HuggingFace endpoint for mirror support
|
||||
if use_mirror:
|
||||
mirror = "https://hf-mirror.com"
|
||||
else:
|
||||
mirror = "https://huggingface.co"
|
||||
|
||||
os.environ["HF_ENDPOINT"] = mirror
|
||||
|
||||
# if the huggingface is already imported in other dependencies,
|
||||
# we need to set the endpoint manually
|
||||
import huggingface_hub.constants
|
||||
|
||||
huggingface_hub.constants.ENDPOINT = mirror
|
||||
huggingface_hub.constants.HUGGINGFACE_CO_URL_TEMPLATE = mirror + "/{repo_id}/resolve/{revision}/{filename}"
|
||||
|
||||
try:
|
||||
super().__init__(
|
||||
pretrained_model_name_or_path=self.pretrained_model_name_or_path,
|
||||
use_mirror=use_mirror,
|
||||
use_fast=True,
|
||||
trust_remote_code=True,
|
||||
**kwargs,
|
||||
)
|
||||
self._tokenizer_available = True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize tokenizer {e}")
|
||||
self._tokenizer_available = False
|
||||
|
||||
async def count(
|
||||
self,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None = None,
|
||||
text: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> int:
|
||||
"""Count tokens in messages or text.
|
||||
|
||||
If text is provided, counts tokens directly in the text string.
|
||||
Otherwise, counts tokens in the messages using the parent class method.
|
||||
|
||||
Args:
|
||||
messages: List of message dictionaries in chat format.
|
||||
tools: Optional list of tool definitions for token counting.
|
||||
text: Optional text string to count tokens directly.
|
||||
**kwargs: Additional keyword arguments passed to parent count method.
|
||||
|
||||
Returns:
|
||||
The number of tokens, guaranteed to be at least the estimated minimum.
|
||||
"""
|
||||
if text:
|
||||
if self._tokenizer_available:
|
||||
try:
|
||||
token_ids = self.tokenizer.encode(text)
|
||||
return max(len(token_ids), self.estimate_tokens(text))
|
||||
except Exception as e:
|
||||
logger.exception("Failed to encode text with tokenizer: %s", e)
|
||||
return self.estimate_tokens(text)
|
||||
else:
|
||||
return self.estimate_tokens(text)
|
||||
else:
|
||||
return await super().count(messages, tools, **kwargs)
|
||||
|
||||
def estimate_tokens(self, text: str) -> int:
|
||||
"""Estimate the number of tokens in a text string.
|
||||
|
||||
Provides a fast character-based estimation as a fallback or lower bound.
|
||||
Uses the configured divisor from instance settings.
|
||||
|
||||
Args:
|
||||
text: The text string to estimate tokens for.
|
||||
|
||||
Returns:
|
||||
The estimated number of tokens in the text string.
|
||||
"""
|
||||
return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5)
|
||||
|
|
@ -1,78 +0,0 @@
|
|||
"""Rule-based token counter for fast estimation without loading tokenizer."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
|
||||
class RuleTokenCounter(HuggingFaceTokenCounter):
|
||||
"""Lightweight token counter using rule-based estimation only.
|
||||
|
||||
This class provides fast token estimation without loading any tokenizer,
|
||||
useful when exact token counts are not critical or for quick approximations.
|
||||
|
||||
Attributes:
|
||||
token_count_estimate_divisor: Divisor for token estimation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token_count_estimate_divisor: float = 3.75,
|
||||
**_kwargs,
|
||||
):
|
||||
"""Initialize the rule-based token counter.
|
||||
|
||||
Args:
|
||||
token_count_estimate_divisor: Divisor for estimating tokens.
|
||||
Defaults to 3.75 (approximately 4 characters per token).
|
||||
**kwargs: Additional keyword arguments (ignored).
|
||||
"""
|
||||
self.token_count_estimate_divisor = token_count_estimate_divisor
|
||||
# Skip tokenizer initialization from parent
|
||||
self._tokenizer_available = False
|
||||
|
||||
async def count(
|
||||
self,
|
||||
messages: list[dict],
|
||||
_tools: list[dict] | None = None,
|
||||
text: str | None = None,
|
||||
**_kwargs: Any,
|
||||
) -> int:
|
||||
"""Count tokens using rule-based estimation.
|
||||
|
||||
Args:
|
||||
messages: List of message dictionaries in chat format.
|
||||
_tools: Optional list of tool definitions (ignored).
|
||||
text: Optional text string to count tokens directly.
|
||||
**_kwargs: Additional keyword arguments (ignored).
|
||||
|
||||
Returns:
|
||||
The estimated number of tokens.
|
||||
"""
|
||||
if text:
|
||||
return self.estimate_tokens(text)
|
||||
|
||||
# Estimate from messages
|
||||
total_text = ""
|
||||
for msg in messages:
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
total_text += content
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and "text" in part:
|
||||
total_text += part["text"]
|
||||
return self.estimate_tokens(total_text)
|
||||
|
||||
def estimate_tokens(self, text: str) -> int:
|
||||
"""Estimate the number of tokens in a text string.
|
||||
|
||||
Uses character-based estimation with the configured divisor.
|
||||
|
||||
Args:
|
||||
text: The text string to estimate tokens for.
|
||||
|
||||
Returns:
|
||||
The estimated number of tokens in the text string.
|
||||
"""
|
||||
return int(len(text.encode("utf-8")) / self.token_count_estimate_divisor + 0.5)
|
||||
|
|
@ -1,41 +0,0 @@
|
|||
"""Module providing a dictionary subclass with attribute-style access and pickling support."""
|
||||
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
_KT = TypeVar("_KT")
|
||||
_VT = TypeVar("_VT")
|
||||
|
||||
|
||||
class BaseDict(dict, Generic[_KT, _VT]):
|
||||
"""A dictionary subclass that enables accessing and modifying keys as attributes."""
|
||||
|
||||
def __getattr__(self, name: str) -> _VT:
|
||||
"""Retrieve a dictionary item as an attribute."""
|
||||
try:
|
||||
return self[name]
|
||||
except KeyError as e:
|
||||
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
|
||||
|
||||
def __setattr__(self, name: str, value: _VT) -> None:
|
||||
"""Assign a value to a dictionary item using attribute syntax."""
|
||||
self[name] = value
|
||||
|
||||
def __delattr__(self, name: str) -> None:
|
||||
"""Remove a dictionary item using attribute syntax."""
|
||||
try:
|
||||
# Delete item from dict via key
|
||||
del self[name]
|
||||
except KeyError as e:
|
||||
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
|
||||
|
||||
def __getstate__(self) -> dict:
|
||||
"""Return the dictionary representation for pickling."""
|
||||
return dict(self)
|
||||
|
||||
def __setstate__(self, state: dict) -> None:
|
||||
"""Restore the dictionary state from a pickled object."""
|
||||
self.update(state)
|
||||
|
||||
def __reduce__(self):
|
||||
"""Define the reconstruction logic for pickling processes."""
|
||||
return self.__class__, (), self.__getstate__()
|
||||
|
|
@ -1,15 +0,0 @@
|
|||
"""embedding"""
|
||||
|
||||
from .base_embedding_model import BaseEmbeddingModel
|
||||
from .openai_embedding_model import OpenAIEmbeddingModel
|
||||
from .openai_embedding_model_sync import OpenAIEmbeddingModelSync
|
||||
from ..registry_factory import R
|
||||
|
||||
__all__ = [
|
||||
"BaseEmbeddingModel",
|
||||
"OpenAIEmbeddingModel",
|
||||
"OpenAIEmbeddingModelSync",
|
||||
]
|
||||
|
||||
R.embedding_models.register("openai")(OpenAIEmbeddingModel)
|
||||
R.embedding_models.register("openai_sync")(OpenAIEmbeddingModelSync)
|
||||
|
|
@ -1,599 +0,0 @@
|
|||
"""Base embedding model interface for ReMe.
|
||||
|
||||
Defines the abstract base class and standard API for all embedding model implementations.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..schema import VectorNode, MemoryChunk
|
||||
|
||||
|
||||
class BaseEmbeddingModel(ABC):
|
||||
"""Abstract base class for embedding model implementations.
|
||||
|
||||
Provides a standard interface for text-to-vector generation with
|
||||
built-in batching, retry logic, and error handling.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
model_name: str = "",
|
||||
dimensions: int = 1024,
|
||||
use_dimensions: bool = False,
|
||||
max_batch_size: int = 10,
|
||||
max_retries: int = 3,
|
||||
raise_exception: bool = True,
|
||||
max_input_length: int = 8192,
|
||||
cache_dir: str | Path = ".reme",
|
||||
max_cache_size: int = 2000,
|
||||
enable_cache: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize model configuration and parameters.
|
||||
|
||||
Args:
|
||||
api_key: API key for the embedding service
|
||||
base_url: Base URL for the embedding service
|
||||
model_name: Name of the embedding model
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
use_dimensions: Whether to pass dimensions parameter to API (some APIs don't support it)
|
||||
max_batch_size: Maximum batch size for embedding requests
|
||||
max_retries: Maximum number of retry attempts on failure
|
||||
raise_exception: Whether to raise exceptions on failure
|
||||
max_input_length: Maximum input text length
|
||||
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
|
||||
enable_cache: Whether to enable embedding cache
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self.api_key: str | None = api_key
|
||||
self.base_url: str | None = base_url
|
||||
self.model_name = model_name
|
||||
self.dimensions = dimensions
|
||||
self.use_dimensions = use_dimensions
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_retries = max_retries
|
||||
self.raise_exception = raise_exception
|
||||
self.max_input_length = max_input_length
|
||||
self.cache_dir = cache_dir
|
||||
self.max_cache_size = max_cache_size
|
||||
self.enable_cache = enable_cache
|
||||
self.kwargs = kwargs
|
||||
|
||||
# Initialize LRU cache for embeddings
|
||||
self._embedding_cache: OrderedDict[str, list[float]] = OrderedDict()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
self.cache_path: Path = Path(self.cache_dir)
|
||||
self.cache_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _truncate_text(self, text: str) -> str:
|
||||
"""Truncate text to max_input_length if it exceeds the limit."""
|
||||
if len(text) > self.max_input_length:
|
||||
logger.warning(f"Text length {len(text)} exceeds {self.max_input_length}, truncating")
|
||||
return text[: self.max_input_length]
|
||||
return text
|
||||
|
||||
def _truncate_texts(self, texts: list[str]) -> list[str]:
|
||||
"""Truncate a list of texts to max_input_length."""
|
||||
return [self._truncate_text(text) for text in texts]
|
||||
|
||||
def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]:
|
||||
"""Validate and adjust embedding dimensions to match expected dimensions.
|
||||
|
||||
Args:
|
||||
embedding: The embedding vector to validate
|
||||
|
||||
Returns:
|
||||
Embedding vector adjusted to match self.dimensions
|
||||
"""
|
||||
actual_len = len(embedding)
|
||||
if actual_len == self.dimensions:
|
||||
return embedding
|
||||
|
||||
elif actual_len < self.dimensions:
|
||||
logger.warning(
|
||||
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, "
|
||||
f"padding with zeros",
|
||||
)
|
||||
return embedding + [0.0] * (self.dimensions - actual_len)
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, "
|
||||
f"truncating to {self.dimensions}",
|
||||
)
|
||||
return embedding[: self.dimensions]
|
||||
|
||||
def _get_cache_key(self, text: str, dimensions: int) -> str:
|
||||
"""Generate a cache key by hashing text + model_name + dimensions.
|
||||
|
||||
This ensures that the same text produces different cache keys when
|
||||
using different models or dimensions.
|
||||
|
||||
Args:
|
||||
text: Input text to hash
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
|
||||
Returns:
|
||||
SHA256 hash combining text, model name, and dimensions
|
||||
"""
|
||||
# Combine text, model_name, and dimensions to create unique cache key
|
||||
cache_string = f"{text}|{self.model_name}|{dimensions}"
|
||||
return hashlib.sha256(cache_string.encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_cache_file_path(self) -> Path:
|
||||
"""Get the path to the cache file.
|
||||
|
||||
Returns:
|
||||
Path to the embedding cache JSONL file
|
||||
"""
|
||||
return self.cache_path / "embedding_cache.jsonl"
|
||||
|
||||
def _load_cache(self) -> None:
|
||||
"""Load embedding cache from disk (JSONL format).
|
||||
|
||||
Each line in the JSONL file contains a JSON object with:
|
||||
- key: the cache key (SHA256 hash)
|
||||
- embedding: the embedding vector (list of floats)
|
||||
|
||||
Loads in reverse order (newest first) to prioritize recent embeddings
|
||||
when max_cache_size is smaller than the file content.
|
||||
"""
|
||||
if not self.enable_cache:
|
||||
return
|
||||
|
||||
cache_file = self._get_cache_file_path()
|
||||
if not cache_file.exists():
|
||||
logger.info(f"No cache file found at {cache_file}, starting with empty cache")
|
||||
return
|
||||
|
||||
try:
|
||||
load_start = time.time()
|
||||
# Read all lines first (to load in reverse order)
|
||||
with open(cache_file, "r", encoding="utf-8") as f:
|
||||
lines = f.readlines()
|
||||
|
||||
loaded_count = 0
|
||||
# Load in reverse order (newest entries first)
|
||||
for _, line in enumerate(reversed(lines), 1):
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line)
|
||||
if not data:
|
||||
continue
|
||||
# Each line is {cache_key: embedding}
|
||||
cache_key, embedding = next(iter(data.items()))
|
||||
|
||||
if cache_key and embedding and isinstance(embedding, list):
|
||||
# Skip if already loaded (keep the newest)
|
||||
if cache_key in self._embedding_cache:
|
||||
continue
|
||||
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got {len(embedding)}",
|
||||
)
|
||||
continue
|
||||
|
||||
# Respect max_cache_size during loading
|
||||
if len(self._embedding_cache) >= self.max_cache_size:
|
||||
logger.info(
|
||||
f"Cache size limit reached ({self.max_cache_size}), "
|
||||
f"loaded {loaded_count} newest entries",
|
||||
)
|
||||
break
|
||||
self._embedding_cache[cache_key] = embedding
|
||||
loaded_count += 1
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"Failed to parse line in cache file: {e}")
|
||||
continue
|
||||
|
||||
logger.info(
|
||||
f"Loaded {loaded_count} embeddings from cache file: {cache_file} in {time.time() - load_start:.2f}s",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load cache from {cache_file}: {e}, deleting cache file")
|
||||
try:
|
||||
cache_file.unlink()
|
||||
logger.info(f"Deleted corrupted cache file: {cache_file}")
|
||||
except Exception as del_e:
|
||||
logger.error(f"Failed to delete cache file {cache_file}: {del_e}")
|
||||
|
||||
def _save_cache(self) -> None:
|
||||
"""Save embedding cache to disk (JSONL format).
|
||||
|
||||
Each line contains a JSON object with the cache key and embedding vector.
|
||||
Only saves if cache is non-empty.
|
||||
"""
|
||||
if not self.enable_cache:
|
||||
return
|
||||
|
||||
logger.info(f"Attempting to save cache, current size: {len(self._embedding_cache)}")
|
||||
if not self._embedding_cache:
|
||||
logger.info("Cache is empty, skipping save")
|
||||
return
|
||||
|
||||
cache_file = self._get_cache_file_path()
|
||||
try:
|
||||
with open(cache_file, "w", encoding="utf-8") as f:
|
||||
for cache_key, embedding in self._embedding_cache.items():
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got {len(embedding)}",
|
||||
)
|
||||
continue
|
||||
cache_entry = {cache_key: embedding}
|
||||
f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n")
|
||||
|
||||
logger.info(f"Saved {len(self._embedding_cache)} embeddings to cache file: {cache_file}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save cache to {cache_file}: {e}")
|
||||
|
||||
def _get_from_cache(self, text: str) -> list[float] | None:
|
||||
"""Retrieve embedding from cache if it exists.
|
||||
|
||||
Args:
|
||||
text: Input text to look up
|
||||
|
||||
Returns:
|
||||
Cached embedding vector or None if not found
|
||||
"""
|
||||
if not self.enable_cache:
|
||||
return None
|
||||
|
||||
cache_key = self._get_cache_key(text, self.dimensions)
|
||||
if cache_key in self._embedding_cache:
|
||||
embeddings: list[float] = self._embedding_cache[cache_key]
|
||||
|
||||
# Validate embedding dimensions match expected dimensions
|
||||
if len(embeddings) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Cached embedding dimensions mismatch: expected {self.dimensions}, "
|
||||
f"got {len(embeddings)}. Removing invalid cache entry.",
|
||||
)
|
||||
del self._embedding_cache[cache_key]
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
# Move to end (most recently used)
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
self._cache_hits += 1
|
||||
text_preview = text[:50] + "..." if len(text) > 50 else text
|
||||
logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})")
|
||||
return embeddings
|
||||
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
def _put_to_cache(self, text: str, embedding: list[float]) -> None:
|
||||
"""Store embedding in cache with LRU eviction.
|
||||
|
||||
Args:
|
||||
text: Input text used as cache key
|
||||
embedding: Embedding vector to cache
|
||||
"""
|
||||
if not self.enable_cache:
|
||||
return
|
||||
|
||||
if self.max_cache_size <= 0:
|
||||
return
|
||||
|
||||
cache_key = self._get_cache_key(text, self.dimensions)
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got real length {len(embedding)}",
|
||||
)
|
||||
return
|
||||
|
||||
# Remove the oldest entry if cache is full
|
||||
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
|
||||
self._embedding_cache.popitem(last=False)
|
||||
|
||||
self._embedding_cache[cache_key] = embedding
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
|
||||
def get_cache_stats(self) -> dict[str, int]:
|
||||
"""Get cache statistics.
|
||||
|
||||
Returns:
|
||||
Dictionary with cache size, hits, misses, and hit rate
|
||||
"""
|
||||
total_requests = self._cache_hits + self._cache_misses
|
||||
hit_rate = self._cache_hits / total_requests if total_requests > 0 else 0.0
|
||||
return {
|
||||
"cache_size": len(self._embedding_cache),
|
||||
"max_cache_size": self.max_cache_size,
|
||||
"cache_hits": self._cache_hits,
|
||||
"cache_misses": self._cache_misses,
|
||||
"hit_rate": hit_rate,
|
||||
}
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""Clear the embedding cache and reset statistics."""
|
||||
self._embedding_cache.clear()
|
||||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Internal async implementation for calling the embedding API with batch input."""
|
||||
|
||||
def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Internal synchronous implementation for calling the embedding API with batch input."""
|
||||
|
||||
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Async get embedding for a single text with exponential backoff retries."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = await self._get_embeddings([truncated_text], **kwargs)
|
||||
embedding = self._validate_and_adjust_embedding(result[0])
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} failed: {e}")
|
||||
if i == self.max_retries - 1:
|
||||
if self.raise_exception:
|
||||
raise
|
||||
return []
|
||||
await asyncio.sleep(i + 1)
|
||||
return []
|
||||
|
||||
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Async get embeddings with automatic batching and exponential backoff retries."""
|
||||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_embeddings = await self._get_embeddings(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
|
||||
results[orig_idx] = adjusted_embedding
|
||||
self._put_to_cache(text, adjusted_embedding)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} batch failed: {e}")
|
||||
if retry == self.max_retries - 1:
|
||||
if self.raise_exception:
|
||||
raise
|
||||
else:
|
||||
await asyncio.sleep(retry + 1)
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
|
||||
"""Synchronous get embedding for a single text with retry logic."""
|
||||
truncated_text = self._truncate_text(input_text)
|
||||
|
||||
# Check cache first
|
||||
cached_embedding = self._get_from_cache(truncated_text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
|
||||
# Cache miss - compute embedding
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = self._get_embeddings_sync([truncated_text], **kwargs)
|
||||
embedding = self._validate_and_adjust_embedding(result[0])
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} failed: {exc}")
|
||||
if i == self.max_retries - 1:
|
||||
if self.raise_exception:
|
||||
raise
|
||||
return []
|
||||
time.sleep(i + 1)
|
||||
return []
|
||||
|
||||
def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Synchronous get embeddings with automatic batching and retry logic."""
|
||||
# Truncate all input texts first
|
||||
truncated_texts = self._truncate_texts(input_text)
|
||||
|
||||
# Check cache for each text and separate cached vs uncached
|
||||
results: list[list[float] | None] = [None] * len(truncated_texts)
|
||||
texts_to_compute: list[tuple[int, str]] = [] # (original_index, text)
|
||||
|
||||
for idx, text in enumerate(truncated_texts):
|
||||
cached = self._get_from_cache(text)
|
||||
if cached is not None:
|
||||
results[idx] = cached
|
||||
else:
|
||||
texts_to_compute.append((idx, text))
|
||||
|
||||
# If all texts were cached, return early
|
||||
if not texts_to_compute:
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
# Compute embeddings for uncached texts in batches
|
||||
uncached_texts = [text for _, text in texts_to_compute]
|
||||
for i in range(0, len(uncached_texts), self.max_batch_size):
|
||||
batch_texts = uncached_texts[i : i + self.max_batch_size]
|
||||
batch_indices = [idx for idx, _ in texts_to_compute[i : i + self.max_batch_size]]
|
||||
|
||||
# Process each batch with retry logic
|
||||
for retry in range(self.max_retries):
|
||||
try:
|
||||
batch_embeddings = self._get_embeddings_sync(batch_texts, **kwargs)
|
||||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
|
||||
results[orig_idx] = adjusted_embedding
|
||||
self._put_to_cache(text, adjusted_embedding)
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} batch failed: {exc}")
|
||||
if retry == self.max_retries - 1:
|
||||
if self.raise_exception:
|
||||
raise
|
||||
else:
|
||||
time.sleep(retry + 1)
|
||||
|
||||
return [r for r in results if r is not None]
|
||||
|
||||
async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
|
||||
"""Async generate and populate vector field for a single VectorNode object."""
|
||||
node.vector = await self.get_embedding(node.content, **kwargs)
|
||||
return node
|
||||
|
||||
async def get_node_embeddings(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
|
||||
"""Async generate and populate vector fields for a batch of VectorNode objects."""
|
||||
contents = [node.content for node in nodes]
|
||||
embeddings: list[list[float]] = await self.get_embeddings(contents, **kwargs)
|
||||
|
||||
if len(embeddings) == len(nodes):
|
||||
for node, vec in zip(nodes, embeddings):
|
||||
node.vector = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
|
||||
return nodes
|
||||
|
||||
def get_node_embedding_sync(self, node: VectorNode, **kwargs) -> VectorNode:
|
||||
"""Synchronously generate and populate vector field for a single VectorNode object."""
|
||||
node.vector = self.get_embedding_sync(node.content, **kwargs)
|
||||
return node
|
||||
|
||||
def get_node_embeddings_sync(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
|
||||
"""Synchronously generate and populate vector fields for a batch of VectorNode objects."""
|
||||
contents = [node.content for node in nodes]
|
||||
embeddings: list[list[float]] = self.get_embeddings_sync(contents, **kwargs)
|
||||
|
||||
if len(embeddings) == len(nodes):
|
||||
for node, vec in zip(nodes, embeddings):
|
||||
node.vector = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
|
||||
return nodes
|
||||
|
||||
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Async generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
chunk.embedding = await self.get_embedding(chunk.text, **kwargs)
|
||||
return chunk
|
||||
|
||||
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Async generate and populate embedding fields for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
texts = [chunk.text for chunk in chunks]
|
||||
embeddings: list[list[float]] = await self.get_embeddings(texts, **kwargs)
|
||||
|
||||
if len(embeddings) == len(chunks):
|
||||
for chunk, vec in zip(chunks, embeddings):
|
||||
chunk.embedding = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
|
||||
return chunks
|
||||
|
||||
def get_chunk_embedding_sync(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Synchronously generate and populate embedding field for a single MemoryChunk object.
|
||||
|
||||
Args:
|
||||
chunk: MemoryChunk object containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same MemoryChunk object with populated embedding field
|
||||
"""
|
||||
chunk.embedding = self.get_embedding_sync(chunk.text, **kwargs)
|
||||
return chunk
|
||||
|
||||
def get_chunk_embeddings_sync(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Synchronously generate embeddings for a batch of MemoryChunk objects.
|
||||
|
||||
Args:
|
||||
chunks: List of MemoryChunk objects containing text to embed
|
||||
**kwargs: Additional arguments passed to the embedding model
|
||||
|
||||
Returns:
|
||||
The same list of MemoryChunk objects with populated embedding fields
|
||||
"""
|
||||
texts = [chunk.text for chunk in chunks]
|
||||
embeddings: list[list[float]] = self.get_embeddings_sync(texts, **kwargs)
|
||||
|
||||
if len(embeddings) == len(chunks):
|
||||
for chunk, vec in zip(chunks, embeddings):
|
||||
chunk.embedding = vec
|
||||
else:
|
||||
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(chunks)} chunks")
|
||||
return chunks
|
||||
|
||||
def start_sync(self):
|
||||
"""Synchronously initialize resources and load cache."""
|
||||
self._load_cache()
|
||||
|
||||
async def start(self):
|
||||
"""Asynchronously initialize resources and load cache."""
|
||||
self._load_cache()
|
||||
|
||||
def close_sync(self):
|
||||
"""Synchronously release resources and close connections."""
|
||||
self._save_cache()
|
||||
|
||||
async def close(self):
|
||||
"""Asynchronously release resources and close connections."""
|
||||
self._save_cache()
|
||||
|
|
@ -1,62 +0,0 @@
|
|||
"""Asynchronous OpenAI-compatible embedding model implementation for ReMe."""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from .base_embedding_model import BaseEmbeddingModel
|
||||
|
||||
|
||||
class OpenAIEmbeddingModel(BaseEmbeddingModel):
|
||||
"""Asynchronous embedding model implementation compatible with OpenAI-style APIs."""
|
||||
|
||||
def __init__(self, encoding_format: Literal["float", "base64"] = "float", **kwargs):
|
||||
"""Initialize the OpenAI async embedding model with API credentials and configuration."""
|
||||
super().__init__(**kwargs)
|
||||
self.encoding_format: Literal["float", "base64"] = encoding_format
|
||||
|
||||
# Lazy client initialization
|
||||
self._client = None
|
||||
|
||||
def _create_client(self):
|
||||
"""Create and return an internal AsyncOpenAI client instance."""
|
||||
return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||
|
||||
@property
|
||||
def client(self):
|
||||
"""Lazily create and return the AsyncOpenAI client."""
|
||||
if self._client is None:
|
||||
self._client = self._create_client()
|
||||
return self._client
|
||||
|
||||
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Fetch embeddings from the API for a batch of strings."""
|
||||
create_kwargs: dict = {
|
||||
"model": self.model_name,
|
||||
"input": input_text,
|
||||
"encoding_format": self.encoding_format,
|
||||
**self.kwargs,
|
||||
**kwargs,
|
||||
}
|
||||
if self.use_dimensions:
|
||||
create_kwargs["dimensions"] = self.dimensions
|
||||
|
||||
completion = await self.client.embeddings.create(**create_kwargs)
|
||||
|
||||
result_emb = [[] for _ in range(len(input_text))]
|
||||
for emb in completion.data:
|
||||
# BGE-M3 returns dense_embedding instead of embedding; use as fallback
|
||||
vec = getattr(emb, "embedding", None) or getattr(emb, "dense_embedding", None)
|
||||
result_emb[emb.index] = list(vec) if vec is not None else []
|
||||
return result_emb
|
||||
|
||||
async def start(self):
|
||||
"""Initialize the asynchronous OpenAI embedding model and load cache."""
|
||||
await super().start()
|
||||
|
||||
async def close(self):
|
||||
"""Close the asynchronous OpenAI client and release network resources."""
|
||||
if self._client is not None:
|
||||
await self._client.close()
|
||||
self._client = None
|
||||
await super().close()
|
||||
|
|
@ -1,39 +0,0 @@
|
|||
"""Synchronous OpenAI-compatible embedding model implementation for ReMe."""
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
from .openai_embedding_model import OpenAIEmbeddingModel
|
||||
|
||||
|
||||
class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel):
|
||||
"""Synchronous embedding model implementation that extends the asynchronous OpenAI model."""
|
||||
|
||||
def _create_client(self):
|
||||
"""Create and return an internal synchronous OpenAI client instance."""
|
||||
return OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||
|
||||
def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Fetch embeddings synchronously from the API for a batch of strings."""
|
||||
create_kwargs: dict = {
|
||||
"model": self.model_name,
|
||||
"input": input_text,
|
||||
"encoding_format": self.encoding_format,
|
||||
**self.kwargs,
|
||||
**kwargs,
|
||||
}
|
||||
if self.use_dimensions:
|
||||
create_kwargs["dimensions"] = self.dimensions
|
||||
|
||||
completion = self.client.embeddings.create(**create_kwargs)
|
||||
|
||||
result_emb = [[] for _ in range(len(input_text))]
|
||||
for emb in completion.data:
|
||||
result_emb[emb.index] = emb.embedding
|
||||
return result_emb
|
||||
|
||||
def close_sync(self):
|
||||
"""Close the synchronous OpenAI client and release network resources."""
|
||||
if self._client is not None:
|
||||
self._client.close()
|
||||
self._client = None
|
||||
super().close_sync()
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
"""enumeration"""
|
||||
|
||||
from .chunk_enum import ChunkEnum
|
||||
from .http_enum import HttpEnum
|
||||
from .json_schema_enum import JsonSchemaEnum
|
||||
from .memory_source import MemorySource
|
||||
from .memory_type import MemoryType
|
||||
from .registry_enum import RegistryEnum
|
||||
from .role import Role
|
||||
|
||||
__all__ = [
|
||||
"ChunkEnum",
|
||||
"HttpEnum",
|
||||
"JsonSchemaEnum",
|
||||
"MemorySource",
|
||||
"MemoryType",
|
||||
"RegistryEnum",
|
||||
"Role",
|
||||
]
|
||||
|
|
@ -1,31 +0,0 @@
|
|||
"""Defines the types of data chunks used in streaming responses."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ChunkEnum(str, Enum):
|
||||
"""Enumeration of possible chunk categories for stream processing."""
|
||||
|
||||
# Internal reasoning or chain-of-thought process
|
||||
THINK = "think"
|
||||
|
||||
# The final generated response content
|
||||
ANSWER = "answer"
|
||||
|
||||
# Metadata or calls related to external tools
|
||||
TOOL = "tool"
|
||||
|
||||
# Resource consumption and token usage statistics
|
||||
USAGE = "usage"
|
||||
|
||||
# Error messages or exception details
|
||||
ERROR = "error"
|
||||
|
||||
# Signal indicating the start of a new ReAct step
|
||||
STEP_START = "step_start"
|
||||
|
||||
# Tool execution result
|
||||
TOOL_RESULT = "tool_result"
|
||||
|
||||
# Final signal indicating the completion of the stream
|
||||
DONE = "done"
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
"""Provides a collection of standard HTTP request methods."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class HttpEnum(str, Enum):
|
||||
"""Enumeration of supported HTTP methods for network requests."""
|
||||
|
||||
# Retrieves data from a specified resource
|
||||
GET = "get"
|
||||
|
||||
# Submits data to be processed to a specified resource
|
||||
POST = "post"
|
||||
|
||||
# Identical to GET but only retrieves the response headers
|
||||
HEAD = "head"
|
||||
|
||||
# Uploads or replaces the representation of a target resource
|
||||
PUT = "put"
|
||||
|
||||
# Deletes the specified resource from the server
|
||||
DELETE = "delete"
|
||||
|
|
@ -1,38 +0,0 @@
|
|||
"""Defines the standard data types supported by JSON Schema.
|
||||
|
||||
This enum maps common JSON Schema primitive types to their corresponding
|
||||
Python runtime types, and provides a convenient string representation
|
||||
compatible with JSON Schema (`"string"`, `"number"`, etc.).
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class JsonSchemaEnum(Enum):
|
||||
"""Enumeration of valid JSON Schema data types.
|
||||
|
||||
The enum value is the corresponding Python type, while the string
|
||||
representation (`str(...)`) is the canonical JSON Schema type name.
|
||||
"""
|
||||
|
||||
# Textual data
|
||||
STRING = str
|
||||
|
||||
# Numeric values, including integers and floats
|
||||
NUMBER = float
|
||||
|
||||
# Integer-only numeric values
|
||||
INTEGER = int
|
||||
|
||||
# JSON objects (key-value mappings)
|
||||
OBJECT = dict
|
||||
|
||||
# Ordered JSON lists/arrays
|
||||
ARRAY = list
|
||||
|
||||
# Boolean values: true / false
|
||||
BOOLEAN = bool
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""Return the lowercase JSON Schema type name for this enum member."""
|
||||
return self.name.lower()
|
||||
|
|
@ -1,11 +0,0 @@
|
|||
"""Memory source types."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class MemorySource(str, Enum):
|
||||
"""Source of memory data."""
|
||||
|
||||
MEMORY = "memory"
|
||||
|
||||
SESSIONS = "sessions"
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
"""Defines the high-level categories of memory managed by ReMe.
|
||||
|
||||
This enumeration is used across the system to tag, route, and store different
|
||||
kinds of memories (identity, personal context, procedures, tools, etc.).
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class MemoryType(str, Enum):
|
||||
"""Enumeration of memory categories used by the memory subsystem.
|
||||
|
||||
These types describe *what* a piece of memory is about, which guides
|
||||
storage, retrieval, and summarization strategies.
|
||||
"""
|
||||
|
||||
# Long‑term, relatively stable attributes about the user (name, roles, etc.)
|
||||
IDENTITY = "identity"
|
||||
|
||||
# User-specific preferences, habits, and evolving personal context
|
||||
PERSONAL = "personal"
|
||||
|
||||
# How‑to knowledge, workflows, and step‑by‑step instructions
|
||||
PROCEDURAL = "procedural"
|
||||
|
||||
# Information learned about tools, APIs, and their usage patterns
|
||||
TOOL = "tool"
|
||||
|
||||
# Condensed representation of larger memory collections
|
||||
SUMMARY = "summary"
|
||||
|
||||
# Raw chronological interaction history, typically before summarization
|
||||
HISTORY = "history"
|
||||
|
|
@ -1,31 +0,0 @@
|
|||
"""Defines the registry categories for core components of the system."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class RegistryEnum(str, Enum):
|
||||
"""Enumeration of component types registered within the application lifecycle."""
|
||||
|
||||
# Large Language Model interfaces
|
||||
LLM = "llm"
|
||||
|
||||
# Models used for generating vector embeddings
|
||||
EMBEDDING_MODEL = "embedding_model"
|
||||
|
||||
# Databases or storage systems for vector search
|
||||
VECTOR_STORE = "vector_store"
|
||||
|
||||
# Databases or storage systems for long-term file storage
|
||||
FILE_STORE = "file_store"
|
||||
|
||||
# Atomic operations or functional units
|
||||
OP = "op"
|
||||
|
||||
# Orchestrated sequences of operations or workflows
|
||||
FLOW = "flow"
|
||||
|
||||
# External APIs or shared internal services
|
||||
SERVICE = "service"
|
||||
|
||||
# Utilities for tracking and limiting token consumption
|
||||
TOKEN_COUNTER = "token_counter"
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
"""Defines the participant roles in a chat completion sequence."""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class Role(str, Enum):
|
||||
"""Enumeration of standard personas involved in a conversation flow."""
|
||||
|
||||
# High-level instructions to guide the model's behavior
|
||||
SYSTEM = "system"
|
||||
|
||||
# Input or queries provided by the human user
|
||||
USER = "user"
|
||||
|
||||
# Responses or messages generated by the AI model
|
||||
ASSISTANT = "assistant"
|
||||
|
||||
# Output or results returned from external tool executions
|
||||
TOOL = "tool"
|
||||
|
|
@ -1,34 +0,0 @@
|
|||
"""File store module for persistent memory management.
|
||||
|
||||
This module provides storage backends for memory chunks and file metadata,
|
||||
including SQLite-based, ChromaDB-based, seekdb-based, and pure-Python local
|
||||
implementations with vector and full-text search.
|
||||
"""
|
||||
|
||||
from .base_file_store import BaseFileStore
|
||||
from .chroma_file_store import ChromaFileStore
|
||||
from .local_file_store import LocalFileStore
|
||||
from .sqlite_file_store import SqliteFileStore
|
||||
from .zvec_file_store import ZvecFileStore
|
||||
from ..registry_factory import R
|
||||
|
||||
__all__ = [
|
||||
"BaseFileStore",
|
||||
"ChromaFileStore",
|
||||
"LocalFileStore",
|
||||
"SqliteFileStore",
|
||||
"ZvecFileStore",
|
||||
]
|
||||
|
||||
R.file_stores.register("sqlite")(SqliteFileStore)
|
||||
R.file_stores.register("chroma")(ChromaFileStore)
|
||||
R.file_stores.register("local")(LocalFileStore)
|
||||
R.file_stores.register("zvec")(ZvecFileStore)
|
||||
|
||||
try:
|
||||
from .seekdb_file_store import SeekdbFileStore
|
||||
|
||||
R.file_stores.register("seekdb")(SeekdbFileStore)
|
||||
__all__.append("SeekdbFileStore")
|
||||
except ImportError:
|
||||
pass
|
||||
|
|
@ -1,227 +0,0 @@
|
|||
"""Base storage interface for file store."""
|
||||
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
|
||||
from ..embedding import BaseEmbeddingModel
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
from ..utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class BaseFileStore(ABC):
|
||||
"""Abstract base class for file storage backends."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store_name: str,
|
||||
db_path: str | Path,
|
||||
embedding_model: BaseEmbeddingModel | None = None,
|
||||
vector_enabled: bool = False,
|
||||
fts_enabled: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize"""
|
||||
# Validate store_name to prevent SQL injection
|
||||
# Only allow alphanumeric characters and underscores
|
||||
if not re.match(r"^[a-zA-Z0-9_]+$", store_name):
|
||||
raise ValueError(f"Invalid '{store_name}'. Only alphanumeric characters and underscores are allowed.")
|
||||
|
||||
# Ensure at least one search method is enabled
|
||||
if not vector_enabled and not fts_enabled:
|
||||
raise ValueError("At least one of vector_enabled or fts_enabled must be True.")
|
||||
|
||||
# Ensure embedding_model is provided when vector search is enabled
|
||||
if vector_enabled and embedding_model is None:
|
||||
raise ValueError("embedding_model is required when vector_enabled is True.")
|
||||
|
||||
self.store_name: str = store_name
|
||||
self.db_path: Path = Path(db_path)
|
||||
self.db_path.mkdir(parents=True, exist_ok=True)
|
||||
self.embedding_model: BaseEmbeddingModel | None = embedding_model
|
||||
self.vector_enabled: bool = vector_enabled
|
||||
self.fts_enabled: bool = fts_enabled
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
@property
|
||||
def embedding_dim(self) -> int:
|
||||
"""Get the embedding model's dimensionality."""
|
||||
if self.embedding_model is None:
|
||||
return 1024
|
||||
return self.embedding_model.dimensions
|
||||
|
||||
def _get_mock_embedding(self) -> list[float]:
|
||||
"""Generate a zero vector based on embedding model dimensions."""
|
||||
return [0.0] * self.embedding_dim
|
||||
|
||||
def _disable_vector_search(self, reason: str = "embedding API error") -> None:
|
||||
"""Disable vector search and log a warning."""
|
||||
if self.vector_enabled:
|
||||
logger.warning(
|
||||
f"[{self.store_name}] Disabling vector search due to {reason}. "
|
||||
"Falling back to full-text search only.",
|
||||
)
|
||||
self.vector_enabled = False
|
||||
|
||||
async def get_embedding(self, query: str, **kwargs) -> list[float]:
|
||||
"""Get embedding for a single query string."""
|
||||
if not self.vector_enabled:
|
||||
return self._get_mock_embedding()
|
||||
try:
|
||||
return await self.embedding_model.get_embedding(query, **kwargs)
|
||||
except Exception as e:
|
||||
self._disable_vector_search(str(e))
|
||||
return self._get_mock_embedding()
|
||||
|
||||
async def get_embeddings(self, queries: list[str], **kwargs) -> list[list[float]]:
|
||||
"""Get embeddings for a batch of query strings."""
|
||||
if not self.vector_enabled:
|
||||
return [self._get_mock_embedding() for _ in queries]
|
||||
try:
|
||||
return await self.embedding_model.get_embeddings(queries, **kwargs)
|
||||
except Exception as e:
|
||||
self._disable_vector_search(str(e))
|
||||
return [self._get_mock_embedding() for _ in queries]
|
||||
|
||||
async def get_chunk_embedding(self, chunk: MemoryChunk, **kwargs) -> MemoryChunk:
|
||||
"""Generate and populate embedding field for a single MemoryChunk object."""
|
||||
if not self.vector_enabled:
|
||||
chunk.embedding = self._get_mock_embedding()
|
||||
return chunk
|
||||
try:
|
||||
return await self.embedding_model.get_chunk_embedding(chunk, **kwargs)
|
||||
except Exception as e:
|
||||
self._disable_vector_search(str(e))
|
||||
chunk.embedding = self._get_mock_embedding()
|
||||
return chunk
|
||||
|
||||
async def get_chunk_embeddings(self, chunks: list[MemoryChunk], **kwargs) -> list[MemoryChunk]:
|
||||
"""Generate and populate embedding fields for a batch of MemoryChunk objects."""
|
||||
if not self.vector_enabled:
|
||||
mock_embedding = self._get_mock_embedding()
|
||||
for chunk in chunks:
|
||||
chunk.embedding = mock_embedding.copy()
|
||||
return chunks
|
||||
try:
|
||||
return await self.embedding_model.get_chunk_embeddings(chunks, **kwargs)
|
||||
except Exception as e:
|
||||
self._disable_vector_search(str(e))
|
||||
mock_embedding = self._get_mock_embedding()
|
||||
for chunk in chunks:
|
||||
chunk.embedding = mock_embedding.copy()
|
||||
return chunks
|
||||
|
||||
@abstractmethod
|
||||
async def start(self):
|
||||
"""Initialize the storage backend."""
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_file(self, file_meta: FileMetadata, source: MemorySource, chunks: list[MemoryChunk]):
|
||||
"""Insert or update a file and its chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_file(self, path: str, source: MemorySource):
|
||||
"""Delete a file and all its chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]):
|
||||
"""Delete chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
async def upsert_chunks(self, chunks: list[MemoryChunk], source: MemorySource):
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
|
||||
@abstractmethod
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed file paths for a source."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_metadata(self, path: str, source: MemorySource) -> FileMetadata | None:
|
||||
"""Get full file metadata with statistics."""
|
||||
|
||||
@abstractmethod
|
||||
async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None:
|
||||
"""Update file metadata without affecting chunks.
|
||||
|
||||
This is useful for incremental updates where only metadata needs to be updated
|
||||
(e.g., after adding/removing chunks in delta file watcher).
|
||||
|
||||
Args:
|
||||
file_meta: Updated file metadata (hash, mtime_ms, size, chunk_count)
|
||||
source: Memory source
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_file_chunks(self, path: str, source: MemorySource) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
|
||||
@abstractmethod
|
||||
async def vector_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search.
|
||||
|
||||
Args:
|
||||
query: Query embedding vector
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
|
||||
Returns:
|
||||
List of search results sorted by similarity
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def keyword_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform keyword/full-text search.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
|
||||
Returns:
|
||||
List of search results sorted by relevance
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def hybrid_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
vector_weight: float = 0.7,
|
||||
candidate_multiplier: float = 3.0,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform hybrid search combining vector and keyword search.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
vector_weight: Weight for vector search results (0.0-1.0).
|
||||
Keyword weight = 1.0 - vector_weight.
|
||||
candidate_multiplier: Multiplier for candidate pool size.
|
||||
candidates = limit * candidate_multiplier
|
||||
|
||||
Returns:
|
||||
List of search results sorted by combined relevance score
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def clear_all(self):
|
||||
"""Clear all indexed data."""
|
||||
|
||||
@abstractmethod
|
||||
async def close(self):
|
||||
"""Close storage and release resources."""
|
||||
|
|
@ -1,633 +0,0 @@
|
|||
"""ChromaDB storage backend for file store."""
|
||||
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from .base_file_store import BaseFileStore
|
||||
from ..enumeration import MemorySource
|
||||
from ..schema import FileMetadata, MemoryChunk, MemorySearchResult
|
||||
from ..utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
try:
|
||||
import chromadb
|
||||
from chromadb.config import Settings
|
||||
|
||||
_CHROMADB_IMPORT_ERROR: Exception | None = None
|
||||
except Exception as e:
|
||||
_CHROMADB_IMPORT_ERROR = e
|
||||
chromadb = None
|
||||
Settings = None
|
||||
|
||||
|
||||
class ChromaFileStore(BaseFileStore):
|
||||
"""ChromaDB file storage with vector and full-text search.
|
||||
|
||||
Inherits embedding methods from BaseFileStore:
|
||||
- get_chunk_embedding / get_chunk_embeddings (async)
|
||||
- get_embedding / get_embeddings (async)
|
||||
|
||||
Provides ChromaDB-backed persistent storage with:
|
||||
- Vector similarity search (native ChromaDB)
|
||||
- Full-text search (via ChromaDB where_document filter)
|
||||
- Efficient chunk and file metadata management
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
**kwargs,
|
||||
):
|
||||
if _CHROMADB_IMPORT_ERROR is not None:
|
||||
raise _CHROMADB_IMPORT_ERROR
|
||||
|
||||
super().__init__(**kwargs)
|
||||
self.client: "chromadb.ClientAPI | None" = None
|
||||
self.chunks_collection: "chromadb.Collection | None" = None
|
||||
# Initialize metadata file path (db_path and store_name are set by base class)
|
||||
self._metadata_file: Path = self.db_path.parent / f"{self.store_name}_file_metadata.json"
|
||||
self._metadata_cache: dict[str, dict[str, FileMetadata]] = {}
|
||||
|
||||
@property
|
||||
def collection_name(self) -> str:
|
||||
"""Get the name of the ChromaDB collection for this store."""
|
||||
return f"chunks_{self.store_name}"
|
||||
|
||||
async def _load_metadata(self) -> dict[str, dict[str, FileMetadata]]:
|
||||
"""Load file metadata from disk.
|
||||
|
||||
Returns:
|
||||
Dictionary mapping source -> path -> FileMetadata
|
||||
"""
|
||||
if not self._metadata_file.exists():
|
||||
return {}
|
||||
|
||||
try:
|
||||
data = self._metadata_file.read_text(encoding="utf-8")
|
||||
metadata_dict = json.loads(data)
|
||||
|
||||
# Convert dict to FileMetadata objects
|
||||
result = {}
|
||||
for source, files in metadata_dict.items():
|
||||
result[source] = {}
|
||||
for path, meta in files.items():
|
||||
result[source][path] = FileMetadata(**meta)
|
||||
|
||||
logger.debug(f"Loaded file metadata from {self._metadata_file}")
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load file metadata from {self._metadata_file}: {e}")
|
||||
return {}
|
||||
|
||||
async def _save_metadata(self, metadata: dict[str, dict[str, FileMetadata]]) -> None:
|
||||
"""Save file metadata to disk.
|
||||
|
||||
Args:
|
||||
metadata: Dictionary mapping source -> path -> FileMetadata
|
||||
"""
|
||||
try:
|
||||
# Convert FileMetadata objects to dict for JSON serialization
|
||||
metadata_dict = {}
|
||||
for source, files in metadata.items():
|
||||
metadata_dict[source] = {}
|
||||
for path, meta in files.items():
|
||||
metadata_dict[source][path] = {
|
||||
"path": meta.path,
|
||||
"hash": meta.hash,
|
||||
"mtime_ms": meta.mtime_ms,
|
||||
"size": meta.size,
|
||||
"chunk_count": meta.chunk_count,
|
||||
}
|
||||
|
||||
data = json.dumps(metadata_dict, indent=2, ensure_ascii=False)
|
||||
self._metadata_file.write_text(data, encoding="utf-8")
|
||||
logger.debug(f"Saved file metadata to {self._metadata_file}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save file metadata to {self._metadata_file}: {e}")
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Initialize ChromaDB client and collection."""
|
||||
if self.client is not None:
|
||||
return
|
||||
|
||||
# Initialize persistent ChromaDB client
|
||||
self.client = chromadb.PersistentClient(
|
||||
path=str(self.db_path),
|
||||
settings=Settings(
|
||||
anonymized_telemetry=False,
|
||||
allow_reset=True,
|
||||
),
|
||||
)
|
||||
|
||||
# Get or create the chunks collection
|
||||
# ChromaDB uses cosine distance by default for similarity
|
||||
self.chunks_collection = self.client.get_or_create_collection(
|
||||
name=self.collection_name,
|
||||
metadata={"hnsw:space": "cosine"},
|
||||
)
|
||||
|
||||
# Load metadata into cache
|
||||
self._metadata_cache = await self._load_metadata()
|
||||
|
||||
logger.info(f"ChromaDB initialized with collection: {self.collection_name}")
|
||||
logger.info(f"File metadata will be persisted to: {self._metadata_file}")
|
||||
|
||||
async def upsert_file(
|
||||
self,
|
||||
file_meta: FileMetadata,
|
||||
source: MemorySource,
|
||||
chunks: list[MemoryChunk],
|
||||
) -> None:
|
||||
"""Insert or update file and its chunks."""
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
# Delete existing chunks for this file first
|
||||
await self.delete_file(file_meta.path, source)
|
||||
|
||||
# Batch generate embeddings for all chunks
|
||||
# (base class returns mock embeddings when vector_enabled=False)
|
||||
chunks = await self.get_chunk_embeddings(chunks)
|
||||
|
||||
# Prepare data for ChromaDB batch upsert
|
||||
ids = []
|
||||
documents = []
|
||||
embeddings = []
|
||||
metadatas = []
|
||||
|
||||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
ids.append(chunk.id)
|
||||
documents.append(chunk.text)
|
||||
embeddings.append(chunk.embedding)
|
||||
metadatas.append(
|
||||
{
|
||||
"path": file_meta.path,
|
||||
"source": source.value,
|
||||
"start_line": chunk.start_line,
|
||||
"end_line": chunk.end_line,
|
||||
"hash": chunk.hash,
|
||||
"updated_at": now,
|
||||
},
|
||||
)
|
||||
|
||||
# Batch upsert to ChromaDB (always pass embeddings to prevent default embedding function)
|
||||
self.chunks_collection.upsert(
|
||||
ids=ids,
|
||||
documents=documents,
|
||||
embeddings=embeddings,
|
||||
metadatas=metadatas,
|
||||
)
|
||||
|
||||
# Update file metadata in cache
|
||||
if source.value not in self._metadata_cache:
|
||||
self._metadata_cache[source.value] = {}
|
||||
self._metadata_cache[source.value][file_meta.path] = FileMetadata(
|
||||
hash=file_meta.hash,
|
||||
mtime_ms=file_meta.mtime_ms,
|
||||
size=file_meta.size,
|
||||
path=file_meta.path,
|
||||
chunk_count=len(chunks),
|
||||
)
|
||||
|
||||
async def delete_file(self, path: str, source: MemorySource) -> None:
|
||||
"""Delete file and all its chunks."""
|
||||
# Query for all chunks with this path and source
|
||||
results = self.chunks_collection.get(
|
||||
where={"$and": [{"path": path}, {"source": source.value}]},
|
||||
include=[],
|
||||
)
|
||||
|
||||
if results["ids"]:
|
||||
self.chunks_collection.delete(
|
||||
ids=results["ids"],
|
||||
)
|
||||
|
||||
# Remove from file metadata cache
|
||||
if source.value in self._metadata_cache:
|
||||
self._metadata_cache[source.value].pop(path, None)
|
||||
|
||||
async def delete_file_chunks(self, path: str, chunk_ids: list[str]) -> None:
|
||||
"""Delete specific chunks for a file."""
|
||||
if not chunk_ids:
|
||||
return
|
||||
|
||||
self.chunks_collection.delete(
|
||||
ids=chunk_ids,
|
||||
)
|
||||
|
||||
# Update chunk count in file metadata cache
|
||||
for source_meta in self._metadata_cache.values():
|
||||
if path in source_meta:
|
||||
# Recalculate chunk count
|
||||
results = self.chunks_collection.get(
|
||||
where={"path": path},
|
||||
include=[],
|
||||
)
|
||||
source_meta[path].chunk_count = len(results["ids"])
|
||||
break
|
||||
|
||||
async def upsert_chunks(
|
||||
self,
|
||||
chunks: list[MemoryChunk],
|
||||
source: MemorySource,
|
||||
) -> None:
|
||||
"""Insert or update specific chunks without affecting other chunks."""
|
||||
if not chunks:
|
||||
return
|
||||
|
||||
# Batch generate embeddings for all chunks
|
||||
# (base class returns mock embeddings when vector_enabled=False)
|
||||
chunks = await self.get_chunk_embeddings(chunks)
|
||||
|
||||
ids = []
|
||||
documents = []
|
||||
embeddings = []
|
||||
metadatas = []
|
||||
|
||||
now = int(time.time() * 1000)
|
||||
for chunk in chunks:
|
||||
ids.append(chunk.id)
|
||||
documents.append(chunk.text)
|
||||
embeddings.append(chunk.embedding)
|
||||
metadatas.append(
|
||||
{
|
||||
"path": chunk.path,
|
||||
"source": source.value,
|
||||
"start_line": chunk.start_line,
|
||||
"end_line": chunk.end_line,
|
||||
"hash": chunk.hash,
|
||||
"updated_at": now,
|
||||
},
|
||||
)
|
||||
|
||||
# Always pass embeddings to prevent default embedding function
|
||||
self.chunks_collection.upsert(
|
||||
ids=ids,
|
||||
documents=documents,
|
||||
embeddings=embeddings,
|
||||
metadatas=metadatas,
|
||||
)
|
||||
|
||||
async def list_files(self, source: MemorySource) -> list[str]:
|
||||
"""List all indexed files for a source."""
|
||||
if source.value not in self._metadata_cache:
|
||||
return []
|
||||
return list(self._metadata_cache[source.value].keys())
|
||||
|
||||
async def get_file_metadata(
|
||||
self,
|
||||
path: str,
|
||||
source: MemorySource,
|
||||
) -> FileMetadata | None:
|
||||
"""Get file metadata with chunk count."""
|
||||
if source.value not in self._metadata_cache:
|
||||
return None
|
||||
return self._metadata_cache[source.value].get(path)
|
||||
|
||||
async def update_file_metadata(self, file_meta: FileMetadata, source: MemorySource) -> None:
|
||||
"""Update file metadata without affecting chunks."""
|
||||
if source.value not in self._metadata_cache:
|
||||
self._metadata_cache[source.value] = {}
|
||||
|
||||
self._metadata_cache[source.value][file_meta.path] = FileMetadata(
|
||||
hash=file_meta.hash,
|
||||
mtime_ms=file_meta.mtime_ms,
|
||||
size=file_meta.size,
|
||||
path=file_meta.path,
|
||||
chunk_count=file_meta.chunk_count,
|
||||
)
|
||||
|
||||
async def get_file_chunks(
|
||||
self,
|
||||
path: str,
|
||||
source: MemorySource,
|
||||
) -> list[MemoryChunk]:
|
||||
"""Get all chunks for a file."""
|
||||
results = self.chunks_collection.get(
|
||||
where={"$and": [{"path": path}, {"source": source.value}]},
|
||||
include=["documents", "embeddings", "metadatas"],
|
||||
)
|
||||
|
||||
chunks = []
|
||||
for i, chunk_id in enumerate(results["ids"]):
|
||||
metadata = results["metadatas"][i]
|
||||
chunks.append(
|
||||
MemoryChunk(
|
||||
id=chunk_id,
|
||||
path=metadata["path"],
|
||||
source=MemorySource(metadata["source"]),
|
||||
start_line=metadata["start_line"],
|
||||
end_line=metadata["end_line"],
|
||||
text=results["documents"][i],
|
||||
hash=metadata["hash"],
|
||||
embedding=results["embeddings"][i] if results["embeddings"] is not None else None,
|
||||
),
|
||||
)
|
||||
|
||||
# Sort by start_line
|
||||
chunks.sort(key=lambda c: c.start_line)
|
||||
return chunks
|
||||
|
||||
async def vector_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform vector similarity search."""
|
||||
if not self.vector_enabled or not query:
|
||||
return []
|
||||
|
||||
# Get query embedding
|
||||
query_embedding = await self.get_embedding(query)
|
||||
if not query_embedding:
|
||||
return []
|
||||
|
||||
# Build where filter for sources
|
||||
where_filter = None
|
||||
if sources:
|
||||
if len(sources) == 1:
|
||||
where_filter = {"source": sources[0].value}
|
||||
else:
|
||||
where_filter = {"source": {"$in": [s.value for s in sources]}}
|
||||
|
||||
# Perform vector search
|
||||
try:
|
||||
results = self.chunks_collection.query(
|
||||
query_embeddings=[query_embedding],
|
||||
n_results=limit,
|
||||
where=where_filter,
|
||||
include=["documents", "metadatas", "distances"],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Vector search failed: {e}, falling back to random results")
|
||||
# Fallback: get some documents without vector search and assign random scores
|
||||
try:
|
||||
fallback_results = self.chunks_collection.get(
|
||||
where=where_filter,
|
||||
limit=limit,
|
||||
include=["documents", "metadatas"],
|
||||
)
|
||||
search_results = []
|
||||
if fallback_results["ids"]:
|
||||
for i, _ in enumerate(fallback_results["ids"]):
|
||||
metadata = fallback_results["metadatas"][i]
|
||||
search_results.append(
|
||||
MemorySearchResult(
|
||||
path=metadata["path"],
|
||||
start_line=metadata["start_line"],
|
||||
end_line=metadata["end_line"],
|
||||
score=random.uniform(0.3, 0.7), # Random score in middle range
|
||||
snippet=fallback_results["documents"][i],
|
||||
source=MemorySource(metadata["source"]),
|
||||
raw_metric=None,
|
||||
),
|
||||
)
|
||||
return search_results
|
||||
except Exception as fallback_e:
|
||||
logger.error(f"Fallback search also failed: {fallback_e}")
|
||||
return []
|
||||
|
||||
search_results = []
|
||||
if results["ids"] and results["ids"][0]:
|
||||
for i, _ in enumerate(results["ids"][0]):
|
||||
metadata = results["metadatas"][0][i]
|
||||
distance = results["distances"][0][i]
|
||||
|
||||
# Convert cosine distance to similarity score
|
||||
# Cosine distance range is [0, 2], convert to [1, 0] score
|
||||
score = max(0.0, 1.0 - distance / 2.0)
|
||||
|
||||
search_results.append(
|
||||
MemorySearchResult(
|
||||
path=metadata["path"],
|
||||
start_line=metadata["start_line"],
|
||||
end_line=metadata["end_line"],
|
||||
score=score,
|
||||
snippet=results["documents"][0][i],
|
||||
source=MemorySource(metadata["source"]),
|
||||
raw_metric=distance,
|
||||
),
|
||||
)
|
||||
|
||||
# Sort by score descending
|
||||
search_results.sort(key=lambda r: r.score, reverse=True)
|
||||
return search_results
|
||||
|
||||
async def keyword_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform keyword/full-text search.
|
||||
|
||||
ChromaDB supports where_document filter for text matching.
|
||||
Note: ChromaDB's $contains is case-sensitive, so we generate multiple
|
||||
case variants (original, lowercase, capitalized) for each word to
|
||||
improve recall while maintaining case-insensitive scoring.
|
||||
"""
|
||||
if not self.fts_enabled or not query:
|
||||
return []
|
||||
|
||||
# Normalize whitespace and split into words
|
||||
words = query.split()
|
||||
if not words:
|
||||
return []
|
||||
|
||||
# Generate case variants for each word to handle case-sensitive $contains
|
||||
# Include: original, lowercase, and capitalized forms
|
||||
word_variants = set()
|
||||
for word in words:
|
||||
word_variants.add(word) # original
|
||||
word_variants.add(word.lower()) # lowercase
|
||||
word_variants.add(word.capitalize()) # Capitalized
|
||||
word_variants.add(word.upper()) # UPPERCASE
|
||||
word_variants_list = list(word_variants)
|
||||
|
||||
# Build where filter for sources
|
||||
where_filter = None
|
||||
if sources:
|
||||
if len(sources) == 1:
|
||||
where_filter = {"source": sources[0].value}
|
||||
else:
|
||||
where_filter = {"source": {"$in": [s.value for s in sources]}}
|
||||
|
||||
# ChromaDB where_document uses $contains for substring matching (case-sensitive)
|
||||
# Use multiple case variants to improve recall
|
||||
if len(word_variants_list) == 1:
|
||||
where_document: dict = {"$contains": word_variants_list[0]}
|
||||
else:
|
||||
where_document = {"$or": [{"$contains": w} for w in word_variants_list]}
|
||||
|
||||
# Get all matching documents
|
||||
results = self.chunks_collection.get(
|
||||
where=where_filter,
|
||||
where_document=where_document,
|
||||
include=["documents", "metadatas"],
|
||||
)
|
||||
|
||||
search_results = []
|
||||
query_lower = query.lower()
|
||||
words_lower = [w.lower() for w in words] # lowercase words for scoring
|
||||
n_words = len(words)
|
||||
|
||||
for i, _ in enumerate(results["ids"]):
|
||||
metadata = results["metadatas"][i]
|
||||
text = results["documents"][i]
|
||||
text_lower = text.lower()
|
||||
|
||||
# Calculate relevance score based on word matches
|
||||
match_count = sum(1 for w in words_lower if w in text_lower)
|
||||
base_score = match_count / n_words
|
||||
|
||||
# Bonus for full phrase match (only applies to multi-word queries)
|
||||
phrase_bonus = 0.2 if n_words > 1 and query_lower in text_lower else 0.0
|
||||
# Scale base_score and add phrase bonus, max score is 1.0
|
||||
score = min(1.0, base_score + phrase_bonus)
|
||||
|
||||
search_results.append(
|
||||
MemorySearchResult(
|
||||
path=metadata["path"],
|
||||
start_line=metadata["start_line"],
|
||||
end_line=metadata["end_line"],
|
||||
score=score,
|
||||
snippet=text,
|
||||
source=MemorySource(metadata["source"]),
|
||||
),
|
||||
)
|
||||
|
||||
# Sort by score descending and limit results
|
||||
search_results.sort(key=lambda r: r.score, reverse=True)
|
||||
return search_results[:limit]
|
||||
|
||||
async def hybrid_search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int,
|
||||
sources: list[MemorySource] | None = None,
|
||||
vector_weight: float = 0.7,
|
||||
candidate_multiplier: float = 3.0,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Perform hybrid search combining vector and keyword search.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
limit: Maximum number of results
|
||||
sources: Optional list of sources to filter
|
||||
vector_weight: Weight for vector search results (0.0-1.0).
|
||||
Keyword weight = 1.0 - vector_weight.
|
||||
candidate_multiplier: Multiplier for candidate pool size.
|
||||
|
||||
Returns:
|
||||
List of search results sorted by combined relevance score
|
||||
"""
|
||||
assert 0.0 <= vector_weight <= 1.0, f"vector_weight must be between 0 and 1, got {vector_weight}"
|
||||
|
||||
candidates = min(200, max(1, int(limit * candidate_multiplier)))
|
||||
text_weight = 1.0 - vector_weight
|
||||
|
||||
# Perform search based on enabled backends
|
||||
if self.vector_enabled and self.fts_enabled:
|
||||
keyword_results = await self.keyword_search(query, candidates, sources)
|
||||
vector_results = await self.vector_search(query, candidates, sources)
|
||||
|
||||
# Log original vector results
|
||||
logger.info("\n=== Vector Search Results ===")
|
||||
for i, r in enumerate(vector_results[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
# Log original keyword results
|
||||
logger.info("\n=== Keyword Search Results ===")
|
||||
for i, r in enumerate(keyword_results[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
if not keyword_results:
|
||||
return vector_results[:limit]
|
||||
elif not vector_results:
|
||||
return keyword_results[:limit]
|
||||
else:
|
||||
merged = self._merge_hybrid_results(
|
||||
vector=vector_results,
|
||||
keyword=keyword_results,
|
||||
vector_weight=vector_weight,
|
||||
text_weight=text_weight,
|
||||
)
|
||||
|
||||
# Log merged results
|
||||
logger.info("\n=== Merged Hybrid Results ===")
|
||||
for i, r in enumerate(merged[:10], 1):
|
||||
snippet_preview = (r.snippet[:100] + "...") if len(r.snippet) > 100 else r.snippet
|
||||
logger.info(f"{i}. Score: {r.score:.4f} | Snippet: {snippet_preview}")
|
||||
|
||||
return merged[:limit]
|
||||
elif self.vector_enabled:
|
||||
vector_results = await self.vector_search(query, limit, sources)
|
||||
return vector_results
|
||||
elif self.fts_enabled:
|
||||
keyword_results = await self.keyword_search(query, limit, sources)
|
||||
return keyword_results
|
||||
else:
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _merge_hybrid_results(
|
||||
vector: list[MemorySearchResult],
|
||||
keyword: list[MemorySearchResult],
|
||||
vector_weight: float,
|
||||
text_weight: float,
|
||||
) -> list[MemorySearchResult]:
|
||||
"""Merge vector and keyword search results with weighted scoring."""
|
||||
merged: dict[str, MemorySearchResult] = {}
|
||||
|
||||
# Process vector results
|
||||
for result in vector:
|
||||
result.score = result.score * vector_weight
|
||||
merged[result.merge_key] = result
|
||||
|
||||
# Process keyword results
|
||||
for result in keyword:
|
||||
key = result.merge_key
|
||||
if key in merged:
|
||||
merged[key].score += result.score * text_weight
|
||||
else:
|
||||
result.score = result.score * text_weight
|
||||
merged[key] = result
|
||||
|
||||
# Sort by score and return
|
||||
results = list(merged.values())
|
||||
results.sort(key=lambda r: r.score, reverse=True)
|
||||
return results
|
||||
|
||||
async def clear_all(self) -> None:
|
||||
"""Clear all indexed data."""
|
||||
# Delete and recreate the collection
|
||||
self.client.delete_collection(
|
||||
name=self.collection_name,
|
||||
)
|
||||
self.chunks_collection = self.client.get_or_create_collection(
|
||||
name=self.collection_name,
|
||||
metadata={"hnsw:space": "cosine"},
|
||||
)
|
||||
|
||||
# Clear file metadata cache and disk
|
||||
self._metadata_cache = {}
|
||||
await self._save_metadata({})
|
||||
|
||||
logger.info(f"Cleared all data from ChromaDB collection: {self.collection_name}")
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close ChromaDB client and release resources."""
|
||||
# Persist metadata cache to disk before closing
|
||||
if self._metadata_cache:
|
||||
await self._save_metadata(self._metadata_cache)
|
||||
|
||||
# ChromaDB PersistentClient handles persistence automatically
|
||||
self.client = None
|
||||
self.chunks_collection = None
|
||||
await super().close()
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue