mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
commit
777e08ecc9
36 changed files with 787 additions and 7548 deletions
|
|
@ -1,335 +0,0 @@
|
|||
"""ReMe evaluation script for HaluMem-like benchmarks."""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from reme_ai.core.enumeration import Role
|
||||
from reme_ai.core.schema import Message, MemoryNode
|
||||
from reme_ai.reme import ReMe
|
||||
|
||||
TEMPLATE_REME = """Memories for user {user_id}:
|
||||
|
||||
{memories}
|
||||
"""
|
||||
|
||||
RETRY_TIMES = 3
|
||||
WAIT_TIME = 2
|
||||
|
||||
# Default prompt for answering questions with memory context
|
||||
PROMPT_REME = """You are a helpful AI assistant with access to the user's memories.
|
||||
Use the following context to answer the user's question accurately.
|
||||
|
||||
Context:
|
||||
{context}
|
||||
|
||||
Question: {question}
|
||||
|
||||
Please provide a detailed and accurate answer based on the available context.
|
||||
If the context doesn't contain enough information to answer the question, say so clearly."""
|
||||
|
||||
|
||||
async def add_memory_async(
|
||||
reme: ReMe,
|
||||
user_id: str,
|
||||
messages: list[dict],
|
||||
description: str = "",
|
||||
):
|
||||
"""Add memory to ReMe system asynchronously."""
|
||||
start = time.time()
|
||||
|
||||
result = await reme.summary(
|
||||
messages=messages,
|
||||
user_id=user_id,
|
||||
description=description,
|
||||
memory_mode="personal",
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return result, duration_ms
|
||||
|
||||
|
||||
async def search_memory_async(
|
||||
reme: ReMe,
|
||||
query: str,
|
||||
user_id: str,
|
||||
top_k: int = 20,
|
||||
):
|
||||
"""Search memory from ReMe system asynchronously."""
|
||||
start = time.time()
|
||||
|
||||
result = await reme.retrieve(
|
||||
query=query,
|
||||
user_id=user_id,
|
||||
memory_mode="personal",
|
||||
top_k=top_k,
|
||||
)
|
||||
|
||||
# Format the context
|
||||
context = TEMPLATE_REME.format(
|
||||
user_id=user_id,
|
||||
memories=result if isinstance(result, str) else json.dumps(result, indent=4, ensure_ascii=False),
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
|
||||
return context, result, duration_ms
|
||||
|
||||
|
||||
async def llm_request_async(reme: ReMe, prompt: str):
|
||||
"""Make LLM request using ReMe's llm."""
|
||||
messages = [
|
||||
Message(role=Role.SYSTEM, content="You are a helpful assistant."),
|
||||
Message(role=Role.USER, content=prompt),
|
||||
]
|
||||
|
||||
response = await reme.llm.chat(messages=messages)
|
||||
return response.content
|
||||
|
||||
|
||||
def extract_user_name(persona_info: str):
|
||||
"""Extract user name from persona info."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
|
||||
if match:
|
||||
username = match.group(1).strip()
|
||||
return username
|
||||
else:
|
||||
raise ValueError("No name found.")
|
||||
|
||||
|
||||
async def _process_session_questions(
|
||||
session: dict,
|
||||
new_session: dict,
|
||||
reme: ReMe,
|
||||
user_name: str,
|
||||
top_k_value: int,
|
||||
) -> None:
|
||||
"""Process questions for a session."""
|
||||
if "questions" not in session:
|
||||
return
|
||||
|
||||
new_session["questions"] = []
|
||||
|
||||
for qa in session["questions"]:
|
||||
context, _, duration_ms = await search_memory_async(
|
||||
reme=reme,
|
||||
query=qa["question"],
|
||||
user_id=user_name,
|
||||
top_k=top_k_value,
|
||||
)
|
||||
|
||||
new_qa = copy.deepcopy(qa)
|
||||
new_qa["context"] = context
|
||||
new_qa["search_duration_ms"] = duration_ms
|
||||
|
||||
prompt = PROMPT_REME.format(
|
||||
context=context,
|
||||
question=qa["question"],
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
response = await llm_request_async(reme, prompt)
|
||||
new_qa["system_response"] = response
|
||||
new_qa["response_duration_ms"] = (time.time() - start_time) * 1000
|
||||
|
||||
new_session["questions"].append(new_qa)
|
||||
|
||||
|
||||
async def process_user_async(
|
||||
user_data: dict,
|
||||
top_k_value: int,
|
||||
save_path: str,
|
||||
reme: ReMe,
|
||||
):
|
||||
"""Process a single user's data asynchronously."""
|
||||
user_name = extract_user_name(user_data["persona_info"])
|
||||
sessions = user_data["sessions"]
|
||||
|
||||
tmp_dir = os.path.join(save_path, "tmp")
|
||||
os.makedirs(tmp_dir, exist_ok=True)
|
||||
|
||||
tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json")
|
||||
|
||||
# Clear existing memories for this user
|
||||
collection_name = f"reme_eval_{user_name}".replace(" ", "_").lower()
|
||||
await reme.vector_store.delete_collection(collection_name)
|
||||
|
||||
# Update collection name for this user
|
||||
reme.vector_store.set_collection_name(collection_name)
|
||||
|
||||
new_user_data = {
|
||||
"uuid": user_data["uuid"],
|
||||
"user_name": user_name,
|
||||
"sessions": [],
|
||||
}
|
||||
|
||||
for session in tqdm(sessions, total=len(sessions), desc=f"Processing user {user_name}"):
|
||||
new_session = {
|
||||
"memory_points": session["memory_points"],
|
||||
"dialogue": session["dialogue"],
|
||||
}
|
||||
|
||||
# Add messages to ReMe
|
||||
dialogue = session["dialogue"]
|
||||
formatted_dialogue = [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(turn["timestamp"], "%b %d, %Y, %H:%M:%S")
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
]
|
||||
|
||||
# Add memory - process every 2 messages
|
||||
result = []
|
||||
total_duration_ms = 0
|
||||
batch_size = 4
|
||||
|
||||
for i in range(0, len(formatted_dialogue), batch_size):
|
||||
batch = formatted_dialogue[i : i + batch_size]
|
||||
batch_result, duration_ms = await add_memory_async(
|
||||
reme=reme,
|
||||
user_id=user_name,
|
||||
messages=batch,
|
||||
)
|
||||
result.extend(batch_result)
|
||||
total_duration_ms += duration_ms
|
||||
|
||||
duration_ms = total_duration_ms
|
||||
|
||||
memories = []
|
||||
for memory_mode in result:
|
||||
if not isinstance(memory_mode, MemoryNode):
|
||||
continue
|
||||
|
||||
memories.append(memory_mode.content)
|
||||
|
||||
print(memories)
|
||||
|
||||
if session.get("is_generated_qa_session", False):
|
||||
new_session["add_dialogue_duration_ms"] = duration_ms
|
||||
new_session["is_generated_qa_session"] = True
|
||||
del new_session["dialogue"]
|
||||
del new_session["memory_points"]
|
||||
new_user_data["sessions"].append(new_session)
|
||||
continue
|
||||
|
||||
# Store the result from summary
|
||||
new_session["extracted_memories"] = memories
|
||||
new_session["add_dialogue_duration_ms"] = duration_ms
|
||||
|
||||
# Search updated memories for memory points
|
||||
# for memory in new_session["memory_points"]:
|
||||
# if memory["is_update"] == "False" or not memory["original_memories"]:
|
||||
# continue
|
||||
#
|
||||
# _, memories_from_system, duration_ms = await search_memory_async(
|
||||
# reme=reme,
|
||||
# query=memory["memory_content"],
|
||||
# user_id=user_name,
|
||||
# top_k=10,
|
||||
# )
|
||||
#
|
||||
# memory["memories_from_system"] = str(memories_from_system)
|
||||
|
||||
# Process questions
|
||||
await _process_session_questions(session, new_session, reme, user_name, top_k_value)
|
||||
|
||||
new_user_data["sessions"].append(new_session)
|
||||
with open(tmp_file, "w", encoding="utf-8") as f:
|
||||
json.dump(new_user_data, f, ensure_ascii=False, indent=2)
|
||||
# raise NotImplementedError
|
||||
|
||||
# Save results
|
||||
with open(tmp_file, "w", encoding="utf-8") as f:
|
||||
json.dump(new_user_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print(f"✅ Saved user {user_name} to {tmp_file}")
|
||||
return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file}
|
||||
|
||||
|
||||
def iter_jsonl(file_path: str):
|
||||
"""Iterate over lines in a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
yield json.loads(line)
|
||||
|
||||
|
||||
async def main_async(
|
||||
data_path_arg: str,
|
||||
version_arg: str = "default",
|
||||
top_k_arg: int = 20,
|
||||
):
|
||||
"""Main evaluation function."""
|
||||
frame = "reme"
|
||||
save_path = f"bench_results/{frame}-{version_arg}/"
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
|
||||
output_file = os.path.join(save_path, f"{frame}_eval_results.jsonl")
|
||||
tmp_dir = os.path.join(save_path, "tmp")
|
||||
os.makedirs(tmp_dir, exist_ok=True)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Initialize ReMe instance (will reuse for all users)
|
||||
reme = ReMe()
|
||||
|
||||
# Load all user data
|
||||
user_data_list = list(iter_jsonl(data_path_arg))
|
||||
total_users = len(user_data_list)
|
||||
|
||||
print(f"Processing {total_users} users sequentially...")
|
||||
|
||||
# Sequential processing
|
||||
for idx, user_data in enumerate(user_data_list, 1):
|
||||
result = await process_user_async(user_data, top_k_arg, save_path, reme)
|
||||
print(f"[{idx}/{total_users}] ✅ Finished {user_data['uuid']} ({result['status']})")
|
||||
|
||||
# Combine all results into final output
|
||||
with open(output_file, "w", encoding="utf-8") as f_out:
|
||||
for file in os.listdir(tmp_dir):
|
||||
if file.endswith(".json"):
|
||||
file_path = os.path.join(tmp_dir, file)
|
||||
with open(file_path, "r", encoding="utf-8") as f_in:
|
||||
data = json.load(f_in)
|
||||
f_out.write(json.dumps(data, ensure_ascii=False) + "\n")
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"✅ All done in {elapsed:.2f}s")
|
||||
print(f"✅ Final results saved to: {output_file}")
|
||||
|
||||
|
||||
def main(
|
||||
data_path_arg: str,
|
||||
version_arg: str = "default",
|
||||
top_k_arg: int = 20,
|
||||
):
|
||||
"""Synchronous entry point for main evaluation."""
|
||||
asyncio.run(main_async(data_path_arg, version_arg, top_k_arg))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Example usage - update these paths as needed
|
||||
# Note: Don't use HaluMem-long.jsonl directly as each line is too large
|
||||
# Instead, create a smaller test dataset or use a different data file
|
||||
|
||||
DEFAULT_DATA_PATH = "/Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl"
|
||||
DEFAULT_VERSION = "test"
|
||||
DEFAULT_TOP_K = 20
|
||||
|
||||
main(
|
||||
data_path_arg=DEFAULT_DATA_PATH,
|
||||
version_arg=DEFAULT_VERSION,
|
||||
top_k_arg=DEFAULT_TOP_K,
|
||||
)
|
||||
|
|
@ -1,588 +0,0 @@
|
|||
"""
|
||||
HaluMem Dataset Statistics Analyzer
|
||||
|
||||
统计 HaluMem 数据集的各项指标:
|
||||
- 每个用户的 session 数量
|
||||
- 每个 session 的对话数量
|
||||
- 每个 session 的对话总长度
|
||||
|
||||
Usage:
|
||||
python bench/halumem/analyze_dataset_stats.py --data_path /path/to/HaluMem-Medium.jsonl
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserStats:
|
||||
"""单个用户的统计数据"""
|
||||
user_name: str
|
||||
uuid: str
|
||||
num_sessions: int
|
||||
dialogues_per_session: list[int] # 每个 session 的对话数量
|
||||
dialogue_lengths_per_session: list[int] # 每个 session 的对话总长度(字符数)
|
||||
num_chunks_after_split: int # 按 5000 字符分割后的 chunk 数量
|
||||
session_time_ranges: list[tuple[Any, Any]] # 每个 session 的 (开始时间, 结束时间)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DatasetStats:
|
||||
"""整体数据集统计"""
|
||||
total_users: int
|
||||
total_sessions: int
|
||||
total_dialogues: int
|
||||
|
||||
avg_sessions_per_user: float
|
||||
avg_dialogues_per_session: float
|
||||
avg_dialogue_length_per_session: float
|
||||
|
||||
# 详细分布
|
||||
sessions_per_user_list: list[int]
|
||||
dialogues_per_session_list: list[int]
|
||||
dialogue_lengths_per_session_list: list[int]
|
||||
|
||||
# Content 统计
|
||||
total_contents: int # 所有对话回合的 content 总数
|
||||
content_sizes: list[int] # 每个 content 的大小(字符数)
|
||||
min_content_size: int
|
||||
max_content_size: int
|
||||
percentiles: dict[str, float] # 分位点统计(全部)
|
||||
|
||||
# 按 role 分类的 Content 统计
|
||||
total_user_contents: int
|
||||
total_assistant_contents: int
|
||||
user_percentiles: dict[str, float] # user 角色的分位点
|
||||
assistant_percentiles: dict[str, float] # assistant 角色的分位点
|
||||
|
||||
# Session 分割统计
|
||||
total_chunks_after_split: int # 按 5000 字符分割后的总 chunk 数
|
||||
chunks_per_user_list: list[int] # 每个用户分割后的 chunk 数量
|
||||
avg_chunks_per_user: float # 平均每个用户的 chunk 数量
|
||||
|
||||
|
||||
class DatasetAnalyzer:
|
||||
"""数据集分析器"""
|
||||
|
||||
def __init__(self, data_path: str):
|
||||
self.data_path = data_path
|
||||
self.user_stats_list: list[UserStats] = []
|
||||
self.all_content_sizes: list[int] = [] # 收集所有 content 的大小
|
||||
self.user_content_sizes: list[int] = [] # user 角色的 content 大小
|
||||
self.assistant_content_sizes: list[int] = [] # assistant 角色的 content 大小
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""从 persona_info 中提取用户名"""
|
||||
match = re.search(r"Name:\s*(.*?);", persona_info)
|
||||
if not match:
|
||||
return "Unknown"
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def calculate_dialogue_length(dialogue: list[dict]) -> int:
|
||||
"""计算对话的总长度(字符数)"""
|
||||
total_length = 0
|
||||
for turn in dialogue:
|
||||
content = turn.get("content", "")
|
||||
total_length += len(content)
|
||||
return total_length
|
||||
|
||||
@staticmethod
|
||||
def split_session_into_chunks(dialogue: list[dict], max_length: int = 5000) -> int:
|
||||
"""
|
||||
将一个 session 按照 max_length 分割成多个 chunks。
|
||||
规则:
|
||||
1. 每次添加 2 个对话回合(user-assistant 对)
|
||||
2. 如果添加后超过 max_length,就开始新的 chunk
|
||||
3. 但是每个 chunk 至少包含 2 个对话回合
|
||||
|
||||
返回分割后的 chunk 数量
|
||||
"""
|
||||
if not dialogue:
|
||||
return 0
|
||||
|
||||
chunks = []
|
||||
current_chunk = []
|
||||
current_length = 0
|
||||
|
||||
# 每次处理 2 个对话回合
|
||||
i = 0
|
||||
while i < len(dialogue):
|
||||
# 取 2 个对话回合(如果不足 2 个,取剩余的)
|
||||
pair = dialogue[i:i+2]
|
||||
pair_length = sum(len(turn.get("content", "")) for turn in pair)
|
||||
|
||||
# 如果当前 chunk 为空,直接添加(保证至少 2 个)
|
||||
if not current_chunk:
|
||||
current_chunk.extend(pair)
|
||||
current_length += pair_length
|
||||
i += len(pair)
|
||||
else:
|
||||
# 如果添加这一对后会超过限制
|
||||
if current_length + pair_length > max_length:
|
||||
# 保存当前 chunk,开始新的 chunk
|
||||
chunks.append(current_chunk)
|
||||
current_chunk = pair
|
||||
current_length = pair_length
|
||||
i += len(pair)
|
||||
else:
|
||||
# 否则添加到当前 chunk
|
||||
current_chunk.extend(pair)
|
||||
current_length += pair_length
|
||||
i += len(pair)
|
||||
|
||||
# 添加最后一个 chunk
|
||||
if current_chunk:
|
||||
chunks.append(current_chunk)
|
||||
|
||||
return len(chunks)
|
||||
|
||||
def load_and_analyze(self):
|
||||
"""加载并分析数据集"""
|
||||
logger.info(f"Loading data from: {self.data_path}")
|
||||
|
||||
with open(self.data_path, "r", encoding="utf-8") as f:
|
||||
for line_num, line in enumerate(f, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
|
||||
try:
|
||||
user_data = json.loads(line)
|
||||
self._analyze_user(user_data)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.error(f"Error parsing line {line_num}: {e}")
|
||||
continue
|
||||
|
||||
logger.info(f"Analyzed {len(self.user_stats_list)} users")
|
||||
|
||||
def _analyze_user(self, user_data: dict):
|
||||
"""分析单个用户的数据"""
|
||||
user_name = self.extract_user_name(user_data.get("persona_info", ""))
|
||||
uuid = user_data.get("uuid", "")
|
||||
sessions = user_data.get("sessions", [])
|
||||
|
||||
dialogues_per_session = []
|
||||
dialogue_lengths_per_session = []
|
||||
session_time_ranges = []
|
||||
total_chunks = 0
|
||||
|
||||
for session in sessions:
|
||||
dialogue = session.get("dialogue", [])
|
||||
num_dialogues = len(dialogue)
|
||||
dialogue_length = self.calculate_dialogue_length(dialogue)
|
||||
|
||||
dialogues_per_session.append(num_dialogues)
|
||||
dialogue_lengths_per_session.append(dialogue_length)
|
||||
|
||||
# 收集 session 的时间范围
|
||||
start_time = session.get("start_time", None)
|
||||
end_time = session.get("end_time", None)
|
||||
session_time_ranges.append((start_time, end_time))
|
||||
|
||||
# 计算这个 session 分割后的 chunk 数量
|
||||
num_chunks = self.split_session_into_chunks(dialogue, max_length=5000)
|
||||
total_chunks += num_chunks
|
||||
|
||||
# 收集每个 content 的大小,并按 role 分类
|
||||
for turn in dialogue:
|
||||
content = turn.get("content", "")
|
||||
content_size = len(content)
|
||||
role = turn.get("role", "")
|
||||
|
||||
self.all_content_sizes.append(content_size)
|
||||
|
||||
if role == "user":
|
||||
self.user_content_sizes.append(content_size)
|
||||
elif role == "assistant":
|
||||
self.assistant_content_sizes.append(content_size)
|
||||
|
||||
user_stats = UserStats(
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
num_sessions=len(sessions),
|
||||
dialogues_per_session=dialogues_per_session,
|
||||
dialogue_lengths_per_session=dialogue_lengths_per_session,
|
||||
num_chunks_after_split=total_chunks,
|
||||
session_time_ranges=session_time_ranges
|
||||
)
|
||||
|
||||
self.user_stats_list.append(user_stats)
|
||||
|
||||
def compute_dataset_stats(self) -> DatasetStats:
|
||||
"""计算整体数据集统计"""
|
||||
total_users = len(self.user_stats_list)
|
||||
|
||||
sessions_per_user_list = [u.num_sessions for u in self.user_stats_list]
|
||||
total_sessions = sum(sessions_per_user_list)
|
||||
|
||||
dialogues_per_session_list = []
|
||||
dialogue_lengths_per_session_list = []
|
||||
|
||||
for user in self.user_stats_list:
|
||||
dialogues_per_session_list.extend(user.dialogues_per_session)
|
||||
dialogue_lengths_per_session_list.extend(user.dialogue_lengths_per_session)
|
||||
|
||||
total_dialogues = sum(dialogues_per_session_list)
|
||||
|
||||
# 计算平均值
|
||||
avg_sessions_per_user = total_sessions / total_users if total_users > 0 else 0
|
||||
avg_dialogues_per_session = (
|
||||
total_dialogues / total_sessions if total_sessions > 0 else 0
|
||||
)
|
||||
avg_dialogue_length_per_session = (
|
||||
sum(dialogue_lengths_per_session_list) / len(dialogue_lengths_per_session_list)
|
||||
if dialogue_lengths_per_session_list else 0
|
||||
)
|
||||
|
||||
# Content 统计
|
||||
total_contents = len(self.all_content_sizes)
|
||||
min_content_size = min(self.all_content_sizes) if self.all_content_sizes else 0
|
||||
max_content_size = max(self.all_content_sizes) if self.all_content_sizes else 0
|
||||
|
||||
# 计算分位点 (10%, 15%, 20%, ..., 95%)
|
||||
percentile_points = list(range(10, 100, 5)) # 10, 15, 20, ..., 95
|
||||
|
||||
# 全部 content 的分位点
|
||||
percentiles = {}
|
||||
if self.all_content_sizes:
|
||||
content_array = np.array(self.all_content_sizes)
|
||||
for p in percentile_points:
|
||||
percentiles[f"p{p}"] = float(np.percentile(content_array, p))
|
||||
|
||||
# user 角色的分位点
|
||||
user_percentiles = {}
|
||||
if self.user_content_sizes:
|
||||
user_array = np.array(self.user_content_sizes)
|
||||
for p in percentile_points:
|
||||
user_percentiles[f"p{p}"] = float(np.percentile(user_array, p))
|
||||
|
||||
# assistant 角色的分位点
|
||||
assistant_percentiles = {}
|
||||
if self.assistant_content_sizes:
|
||||
assistant_array = np.array(self.assistant_content_sizes)
|
||||
for p in percentile_points:
|
||||
assistant_percentiles[f"p{p}"] = float(np.percentile(assistant_array, p))
|
||||
|
||||
# Session 分割统计
|
||||
chunks_per_user_list = [u.num_chunks_after_split for u in self.user_stats_list]
|
||||
total_chunks_after_split = sum(chunks_per_user_list)
|
||||
avg_chunks_per_user = (
|
||||
total_chunks_after_split / total_users if total_users > 0 else 0
|
||||
)
|
||||
|
||||
return DatasetStats(
|
||||
total_users=total_users,
|
||||
total_sessions=total_sessions,
|
||||
total_dialogues=total_dialogues,
|
||||
avg_sessions_per_user=avg_sessions_per_user,
|
||||
avg_dialogues_per_session=avg_dialogues_per_session,
|
||||
avg_dialogue_length_per_session=avg_dialogue_length_per_session,
|
||||
sessions_per_user_list=sessions_per_user_list,
|
||||
dialogues_per_session_list=dialogues_per_session_list,
|
||||
dialogue_lengths_per_session_list=dialogue_lengths_per_session_list,
|
||||
total_contents=total_contents,
|
||||
content_sizes=self.all_content_sizes,
|
||||
min_content_size=min_content_size,
|
||||
max_content_size=max_content_size,
|
||||
percentiles=percentiles,
|
||||
total_user_contents=len(self.user_content_sizes),
|
||||
total_assistant_contents=len(self.assistant_content_sizes),
|
||||
user_percentiles=user_percentiles,
|
||||
assistant_percentiles=assistant_percentiles,
|
||||
total_chunks_after_split=total_chunks_after_split,
|
||||
chunks_per_user_list=chunks_per_user_list,
|
||||
avg_chunks_per_user=avg_chunks_per_user
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _print_percentiles(percentiles: dict[str, float]):
|
||||
"""打印分位点统计(辅助函数)"""
|
||||
if not percentiles:
|
||||
print(" (无数据)")
|
||||
return
|
||||
|
||||
sorted_percentiles = sorted(percentiles.keys(), key=lambda x: int(x[1:]))
|
||||
|
||||
# 每行显示 5 个分位点,让输出更紧凑
|
||||
for i in range(0, len(sorted_percentiles), 5):
|
||||
line_items = []
|
||||
for percentile_key in sorted_percentiles[i:i+5]:
|
||||
percentile_value = percentiles[percentile_key]
|
||||
p_num = percentile_key[1:] # 去掉 'p' 前缀
|
||||
line_items.append(f"{p_num}%: {percentile_value:.0f}")
|
||||
print(f" {' | '.join(line_items)}")
|
||||
|
||||
def print_summary(self, stats: DatasetStats):
|
||||
"""打印统计摘要"""
|
||||
print("\n" + "=" * 80)
|
||||
print("HALUMEM DATASET STATISTICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 总体统计:")
|
||||
print(f" 总用户数: {stats.total_users}")
|
||||
print(f" 总 Session 数: {stats.total_sessions}")
|
||||
print(f" 总对话数: {stats.total_dialogues}")
|
||||
|
||||
print(f"\n📈 平均值:")
|
||||
print(f" 每个用户的平均 Session 数: {stats.avg_sessions_per_user:.2f}")
|
||||
print(f" 每个 Session 的平均对话数: {stats.avg_dialogues_per_session:.2f}")
|
||||
print(f" 每个 Session 的平均对话长度(字符): {stats.avg_dialogue_length_per_session:.2f}")
|
||||
|
||||
print(f"\n📊 分布统计:")
|
||||
if stats.sessions_per_user_list:
|
||||
print(f" 每用户 Session 数 - 最小: {min(stats.sessions_per_user_list)}, "
|
||||
f"最大: {max(stats.sessions_per_user_list)}")
|
||||
|
||||
if stats.dialogues_per_session_list:
|
||||
print(f" 每 Session 对话数 - 最小: {min(stats.dialogues_per_session_list)}, "
|
||||
f"最大: {max(stats.dialogues_per_session_list)}")
|
||||
|
||||
if stats.dialogue_lengths_per_session_list:
|
||||
print(f" 每 Session 对话长度 - 最小: {min(stats.dialogue_lengths_per_session_list)}, "
|
||||
f"最大: {max(stats.dialogue_lengths_per_session_list)}")
|
||||
|
||||
print(f"\n💬 Content 详细统计:")
|
||||
print(f" 总 Content 数量: {stats.total_contents}")
|
||||
print(f" User 消息数: {stats.total_user_contents}")
|
||||
print(f" Assistant 消息数: {stats.total_assistant_contents}")
|
||||
print(f" Content 大小(字符数):")
|
||||
print(f" 最小值: {stats.min_content_size}")
|
||||
print(f" 最大值: {stats.max_content_size}")
|
||||
|
||||
if stats.content_sizes:
|
||||
avg_content_size = sum(stats.content_sizes) / len(stats.content_sizes)
|
||||
print(f" 平均值: {avg_content_size:.2f}")
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (全部):")
|
||||
self._print_percentiles(stats.percentiles)
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (User 角色):")
|
||||
self._print_percentiles(stats.user_percentiles)
|
||||
|
||||
print(f"\n📈 Content 大小分位点 (Assistant 角色):")
|
||||
self._print_percentiles(stats.assistant_percentiles)
|
||||
|
||||
print(f"\n✂️ Session 分割统计 (按 5000 字符分割):")
|
||||
print(f" 原始 Session 总数: {stats.total_sessions}")
|
||||
print(f" 分割后 Chunk 总数: {stats.total_chunks_after_split}")
|
||||
print(f" 每个用户平均 Chunk 数: {stats.avg_chunks_per_user:.2f}")
|
||||
print(f" Chunk/Session 比例: {stats.total_chunks_after_split / stats.total_sessions:.2f}x")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
def print_per_user_stats(self):
|
||||
"""打印每个用户的详细统计"""
|
||||
print("\n" + "=" * 80)
|
||||
print("PER-USER STATISTICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for idx, user_stats in enumerate(self.user_stats_list, 1):
|
||||
avg_dialogues = (
|
||||
sum(user_stats.dialogues_per_session) / len(user_stats.dialogues_per_session)
|
||||
if user_stats.dialogues_per_session else 0
|
||||
)
|
||||
avg_length = (
|
||||
sum(user_stats.dialogue_lengths_per_session) / len(user_stats.dialogue_lengths_per_session)
|
||||
if user_stats.dialogue_lengths_per_session else 0
|
||||
)
|
||||
|
||||
print(f"[{idx}] {user_stats.user_name} (UUID: {user_stats.uuid[:8]}...)")
|
||||
print(f" Session 数: {user_stats.num_sessions}")
|
||||
print(f" 分割后 Chunk 数: {user_stats.num_chunks_after_split}")
|
||||
print(f" 平均每 Session 对话数: {avg_dialogues:.2f}")
|
||||
print(f" 平均每 Session 对话长度: {avg_length:.2f} 字符")
|
||||
print()
|
||||
|
||||
def print_first_user_session_times(self):
|
||||
"""打印第一个用户的每个 session 的时间范围"""
|
||||
if not self.user_stats_list:
|
||||
print("\n没有用户数据")
|
||||
return
|
||||
|
||||
first_user = self.user_stats_list[0]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print(f"第一个用户的 Session 时间统计")
|
||||
print("=" * 80 + "\n")
|
||||
print(f"用户名: {first_user.user_name}")
|
||||
print(f"UUID: {first_user.uuid}")
|
||||
print(f"总 Session 数: {first_user.num_sessions}\n")
|
||||
|
||||
print("-" * 80)
|
||||
print(f"{'Session #':<12} {'开始时间':<30} {'结束时间':<30}")
|
||||
print("-" * 80)
|
||||
|
||||
for idx, (start_time, end_time) in enumerate(first_user.session_time_ranges, 1):
|
||||
start_str = str(start_time) if start_time is not None else "无"
|
||||
end_str = str(end_time) if end_time is not None else "无"
|
||||
print(f"{idx:<12} {start_str:<30} {end_str:<30}")
|
||||
|
||||
print("=" * 80)
|
||||
|
||||
def print_user_split_summary(self):
|
||||
"""打印每个用户的分割统计摘要(表格形式)"""
|
||||
print("\n" + "=" * 80)
|
||||
print("PER-USER SESSION SPLIT SUMMARY (按 5000 字符分割)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# 表头
|
||||
print(f"{'序号':<6} {'用户名':<25} {'原始Sessions':<15} {'分割后Chunks':<15} {'比例':<10}")
|
||||
print("-" * 80)
|
||||
|
||||
# 每个用户的数据
|
||||
for idx, user_stats in enumerate(self.user_stats_list, 1):
|
||||
ratio = (
|
||||
user_stats.num_chunks_after_split / user_stats.num_sessions
|
||||
if user_stats.num_sessions > 0 else 0
|
||||
)
|
||||
print(f"{idx:<6} {user_stats.user_name[:24]:<25} {user_stats.num_sessions:<15} "
|
||||
f"{user_stats.num_chunks_after_split:<15} {ratio:.2f}x")
|
||||
|
||||
print("-" * 80)
|
||||
|
||||
# 总计
|
||||
total_sessions = sum(u.num_sessions for u in self.user_stats_list)
|
||||
total_chunks = sum(u.num_chunks_after_split for u in self.user_stats_list)
|
||||
overall_ratio = total_chunks / total_sessions if total_sessions > 0 else 0
|
||||
|
||||
print(f"{'总计':<6} {'':<25} {total_sessions:<15} {total_chunks:<15} {overall_ratio:.2f}x")
|
||||
print("=" * 80)
|
||||
|
||||
def save_results(self, output_path: str, stats: DatasetStats):
|
||||
"""保存统计结果到 JSON 文件"""
|
||||
results = {
|
||||
"summary": {
|
||||
"total_users": stats.total_users,
|
||||
"total_sessions": stats.total_sessions,
|
||||
"total_dialogues": stats.total_dialogues,
|
||||
"avg_sessions_per_user": stats.avg_sessions_per_user,
|
||||
"avg_dialogues_per_session": stats.avg_dialogues_per_session,
|
||||
"avg_dialogue_length_per_session": stats.avg_dialogue_length_per_session,
|
||||
"session_split_stats": {
|
||||
"total_chunks_after_split": stats.total_chunks_after_split,
|
||||
"avg_chunks_per_user": stats.avg_chunks_per_user,
|
||||
"chunk_to_session_ratio": (
|
||||
stats.total_chunks_after_split / stats.total_sessions
|
||||
if stats.total_sessions > 0 else 0
|
||||
)
|
||||
},
|
||||
"content_stats": {
|
||||
"total_contents": stats.total_contents,
|
||||
"total_user_contents": stats.total_user_contents,
|
||||
"total_assistant_contents": stats.total_assistant_contents,
|
||||
"min_content_size": stats.min_content_size,
|
||||
"max_content_size": stats.max_content_size,
|
||||
"avg_content_size": (
|
||||
sum(stats.content_sizes) / len(stats.content_sizes)
|
||||
if stats.content_sizes else 0
|
||||
),
|
||||
"percentiles_all": stats.percentiles,
|
||||
"percentiles_user": stats.user_percentiles,
|
||||
"percentiles_assistant": stats.assistant_percentiles
|
||||
}
|
||||
},
|
||||
"per_user_stats": [
|
||||
{
|
||||
"user_name": u.user_name,
|
||||
"uuid": u.uuid,
|
||||
"num_sessions": u.num_sessions,
|
||||
"num_chunks_after_split": u.num_chunks_after_split,
|
||||
"avg_dialogues_per_session": (
|
||||
sum(u.dialogues_per_session) / len(u.dialogues_per_session)
|
||||
if u.dialogues_per_session else 0
|
||||
),
|
||||
"avg_dialogue_length_per_session": (
|
||||
sum(u.dialogue_lengths_per_session) / len(u.dialogue_lengths_per_session)
|
||||
if u.dialogue_lengths_per_session else 0
|
||||
),
|
||||
"dialogues_per_session": u.dialogues_per_session,
|
||||
"dialogue_lengths_per_session": u.dialogue_lengths_per_session,
|
||||
"session_time_ranges": [
|
||||
{"start_time": start, "end_time": end}
|
||||
for start, end in u.session_time_ranges
|
||||
]
|
||||
}
|
||||
for u in self.user_stats_list
|
||||
]
|
||||
}
|
||||
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, ensure_ascii=False, indent=2)
|
||||
|
||||
logger.info(f"Results saved to: {output_path}")
|
||||
|
||||
|
||||
def main(data_path: str, output_path: str = None, show_per_user: bool = False):
|
||||
"""主函数"""
|
||||
# 检查文件是否存在
|
||||
if not Path(data_path).exists():
|
||||
logger.error(f"File not found: {data_path}")
|
||||
return
|
||||
|
||||
# 创建分析器并执行分析
|
||||
analyzer = DatasetAnalyzer(data_path)
|
||||
analyzer.load_and_analyze()
|
||||
|
||||
# 计算统计数据
|
||||
stats = analyzer.compute_dataset_stats()
|
||||
|
||||
# 打印摘要
|
||||
analyzer.print_summary(stats)
|
||||
|
||||
# 打印第一个用户的 session 时间统计
|
||||
analyzer.print_first_user_session_times()
|
||||
|
||||
# 打印每个用户的分割统计摘要(始终显示)
|
||||
analyzer.print_user_split_summary()
|
||||
|
||||
# 打印每个用户的详细统计(可选)
|
||||
if show_per_user:
|
||||
analyzer.print_per_user_stats()
|
||||
|
||||
# 保存结果到文件
|
||||
if output_path:
|
||||
analyzer.save_results(output_path, stats)
|
||||
else:
|
||||
# 默认保存到与数据文件相同目录
|
||||
default_output = str(Path(data_path).parent / "dataset_statistics.json")
|
||||
analyzer.save_results(default_output, stats)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Analyze HaluMem dataset statistics"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to save statistics JSON (default: dataset_statistics.json in same dir)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--show_per_user",
|
||||
action="store_true",
|
||||
help="Show detailed statistics for each user"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
output_path=args.output_path,
|
||||
show_per_user=args.show_per_user
|
||||
)
|
||||
|
|
@ -1,180 +0,0 @@
|
|||
"""
|
||||
分析 bench_results/reme_simple/tmp 目录下的评估结果
|
||||
|
||||
统计所有用户session中的result_type分布,并输出非Correct结果的详细位置信息。
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
|
||||
def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
|
||||
"""
|
||||
分析评估结果目录。
|
||||
|
||||
Args:
|
||||
tmp_dir: 临时结果目录路径
|
||||
"""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
if not tmp_path.exists():
|
||||
print(f"❌ 目录不存在: {tmp_dir}")
|
||||
return
|
||||
|
||||
# 统计数据
|
||||
result_counter = Counter()
|
||||
non_correct_results = [] # 存储非Correct结果的详细信息
|
||||
|
||||
# 遍历所有用户目录
|
||||
user_dirs = sorted([d for d in tmp_path.iterdir() if d.is_dir()])
|
||||
|
||||
if not user_dirs:
|
||||
print(f"❌ {tmp_dir} 下没有用户目录")
|
||||
return
|
||||
|
||||
print(f"📁 找到 {len(user_dirs)} 个用户目录\n")
|
||||
print("=" * 80)
|
||||
print("开始分析...")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
total_sessions = 0
|
||||
total_questions = 0
|
||||
|
||||
# 遍历每个用户目录
|
||||
for user_dir in user_dirs:
|
||||
user_name = user_dir.name
|
||||
|
||||
# 获取该用户的所有session文件
|
||||
session_files = sorted([
|
||||
f for f in user_dir.iterdir()
|
||||
if f.name.startswith("session_") and f.suffix == ".json"
|
||||
])
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
# 遍历每个session
|
||||
for session_file in session_files:
|
||||
try:
|
||||
with open(session_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
|
||||
session_id = session_data.get("session_id", -1)
|
||||
total_sessions += 1
|
||||
|
||||
# 跳过生成的QA session
|
||||
if session_data.get("is_generated_qa_session", False):
|
||||
continue
|
||||
|
||||
# 获取评估结果
|
||||
eval_results = session_data.get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
|
||||
# 分析每个问题的结果
|
||||
for qa_idx, qa_record in enumerate(qa_records):
|
||||
result_type = qa_record.get("result_type", "Unknown")
|
||||
|
||||
# 统计result_type
|
||||
result_counter[result_type] += 1
|
||||
total_questions += 1
|
||||
|
||||
# 如果不是Correct,记录详细信息
|
||||
if result_type != "Correct":
|
||||
non_correct_results.append({
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"question_id": qa_idx,
|
||||
"result_type": result_type,
|
||||
"question": qa_record.get("question", ""),
|
||||
"answer": qa_record.get("answer", ""),
|
||||
"system_response": qa_record.get("system_response", "")
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
print(f"⚠️ 读取文件失败: {session_file}, 错误: {e}")
|
||||
continue
|
||||
|
||||
# 输出统计结果
|
||||
print("\n" + "=" * 80)
|
||||
print("统计结果")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print(f"📊 总用户数: {len(user_dirs)}")
|
||||
print(f"📊 总Session数: {total_sessions}")
|
||||
print(f"📊 总问题数: {total_questions}\n")
|
||||
|
||||
if total_questions == 0:
|
||||
print("❌ 没有找到任何问题数据")
|
||||
return
|
||||
|
||||
# 输出result_type分布
|
||||
print("=" * 80)
|
||||
print("Result Type 分布")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# 按数量降序排列
|
||||
sorted_results = sorted(result_counter.items(), key=lambda x: x[1], reverse=True)
|
||||
|
||||
for result_type, count in sorted_results:
|
||||
ratio = count / total_questions * 100
|
||||
print(f" {result_type:20s}: {count:5d} ({ratio:6.2f}%)")
|
||||
|
||||
# 输出非Correct结果的详细信息
|
||||
if non_correct_results:
|
||||
print("\n" + "=" * 80)
|
||||
print(f"非 Correct 结果详情 (共 {len(non_correct_results)} 条)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for idx, result in enumerate(non_correct_results, 1):
|
||||
print(f"[{idx}] {result['result_type']}")
|
||||
print(f" 用户: {result['user_name']}")
|
||||
print(f" 位置: Session {result['session_id']}, Question {result['question_id']}")
|
||||
print(f" 问题: {result['question']}")
|
||||
print(f" 正确答案: {result['answer']}")
|
||||
print(f" 系统回答: {result['system_response'][:200]}{'...' if len(result['system_response']) > 200 else ''}")
|
||||
print()
|
||||
|
||||
else:
|
||||
print("\n🎉 所有问题都是 Correct!")
|
||||
|
||||
# 保存详细报告到文件
|
||||
report_file = Path(tmp_dir).parent / "analysis_report.json"
|
||||
report_data = {
|
||||
"summary": {
|
||||
"total_users": len(user_dirs),
|
||||
"total_sessions": total_sessions,
|
||||
"total_questions": total_questions,
|
||||
"result_type_distribution": dict(result_counter),
|
||||
"result_type_ratio": {
|
||||
result_type: count / total_questions
|
||||
for result_type, count in result_counter.items()
|
||||
}
|
||||
},
|
||||
"non_correct_results": non_correct_results
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(report_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print("=" * 80)
|
||||
print(f"📄 详细报告已保存到: {report_file}")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="分析 ReMe 评估结果中的 result_type 分布"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tmp_dir",
|
||||
type=str,
|
||||
default="bench_results/reme_simple/tmp",
|
||||
help="临时结果目录路径 (默认: bench_results/reme_simple/tmp)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
analyze_results(args.tmp_dir)
|
||||
|
|
@ -1,285 +0,0 @@
|
|||
"""
|
||||
Compute Question Answering statistics from eval_reme_simple_v4.py results.
|
||||
|
||||
Usage:
|
||||
python bench/halumem/compute_qa_stats_v4.py --results_file bench_results/reme_simple_v4/eval_results.jsonl
|
||||
python bench/halumem/compute_qa_stats_v4.py --tmp_dir bench_results/reme_simple_v4/tmp
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = hallucination = omission = valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
valid += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
valid += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
valid += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"correct_qa_ratio(valid)": correct / valid if valid > 0 else 0,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0,
|
||||
"omission_qa_ratio(valid)": omission / valid if valid > 0 else 0,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
def compute_time_metrics(results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
add_duration = search_duration = 0
|
||||
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data.get("sessions", []):
|
||||
add_duration += session.get("add_dialogue_duration_ms", 0)
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
search_duration += qa.get("search_duration_ms", 0)
|
||||
|
||||
return {
|
||||
"add_dialogue_duration_time": add_duration / 1000 / 60,
|
||||
"search_memory_duration_time": search_duration / 1000 / 60,
|
||||
"total_duration_time": (add_duration + search_duration) / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
def load_from_tmp_dir(tmp_dir: str) -> str:
|
||||
"""Load data from tmp directory and generate eval_results.jsonl file."""
|
||||
tmp_path = Path(tmp_dir)
|
||||
eval_results_file = tmp_path.parent / "eval_results.jsonl"
|
||||
|
||||
print(f"\n📁 Loading from: {tmp_dir}")
|
||||
print(f"📝 Generating: {eval_results_file}")
|
||||
|
||||
user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()]
|
||||
print(f" Found {len(user_dirs)} users")
|
||||
|
||||
users_data = []
|
||||
for user_dir in user_dirs:
|
||||
session_files = sorted(
|
||||
[f for f in user_dir.iterdir() if f.name.startswith("session_") and f.suffix == ".json"],
|
||||
key=lambda f: int(f.stem.split("_")[1])
|
||||
)
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
with open(session_files[0], "r", encoding="utf-8") as f:
|
||||
first_session = json.load(f)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
users_data.append(user_data)
|
||||
print(f" ✓ {user_dir.name}: {len(session_files)} sessions")
|
||||
|
||||
with open(eval_results_file, "w", encoding="utf-8") as f:
|
||||
for user_data in users_data:
|
||||
f.write(json.dumps(user_data, ensure_ascii=False) + "\n")
|
||||
|
||||
print(f" ✅ Generated: {eval_results_file}")
|
||||
return str(eval_results_file)
|
||||
|
||||
|
||||
def main(input_path: str):
|
||||
"""Main function to compute statistics from eval results."""
|
||||
|
||||
if not os.path.exists(input_path):
|
||||
print(f"❌ Error: Path not found: {input_path}")
|
||||
return
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("REME V4 - QUESTION ANSWERING STATISTICS")
|
||||
print("=" * 80)
|
||||
|
||||
# Load or generate eval_results.jsonl
|
||||
if os.path.isdir(input_path):
|
||||
results_file = load_from_tmp_dir(input_path)
|
||||
else:
|
||||
results_file = input_path
|
||||
print(f"\n📁 Using: {results_file}")
|
||||
|
||||
# Collect QA records with metadata
|
||||
qa_records = []
|
||||
qa_with_metadata = []
|
||||
user_count = session_count = 0
|
||||
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
user_count += 1
|
||||
user_name = user_data.get("user_name", "Unknown")
|
||||
|
||||
valid_session_idx = 0
|
||||
for original_idx, session in enumerate(user_data.get("sessions", [])):
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
session_count += 1
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
|
||||
for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])):
|
||||
qa_records.append(qa)
|
||||
qa_with_metadata.append({
|
||||
"user_name": user_name,
|
||||
"session_idx": valid_session_idx,
|
||||
"question_idx": qa_idx,
|
||||
"qa_record": qa
|
||||
})
|
||||
|
||||
valid_session_idx += 1
|
||||
|
||||
print(f"\n📊 Data Summary:")
|
||||
print(f" Users: {user_count}")
|
||||
print(f" Sessions: {session_count}")
|
||||
print(f" QA Records: {len(qa_records)}")
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = compute_qa_metrics(qa_records)
|
||||
time_metrics = compute_time_metrics(results_file)
|
||||
|
||||
# Save results
|
||||
output_dir = Path(results_file).parent
|
||||
report_file = output_dir / "reme_eval_stat_result.json"
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
print(f"\n✅ Results saved to: {report_file}")
|
||||
|
||||
# Print metrics
|
||||
print("\n" + "=" * 80)
|
||||
print("📊 QUESTION ANSWERING METRICS")
|
||||
print("=" * 80)
|
||||
print(f"\n Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ TIME METRICS")
|
||||
print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min")
|
||||
print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min")
|
||||
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
|
||||
|
||||
# Print error records
|
||||
print("\n" + "=" * 80)
|
||||
print("❌ ERROR RECORDS (Non-Correct)")
|
||||
print("=" * 80)
|
||||
|
||||
error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]]
|
||||
|
||||
if not error_records:
|
||||
print("\n✅ All QA records are correct!")
|
||||
else:
|
||||
print(f"\nFound {len(error_records)} error records:\n")
|
||||
|
||||
for idx, record in enumerate(error_records, 1):
|
||||
qa = record["qa_record"]
|
||||
|
||||
print(f"\n{'━' * 80}")
|
||||
print(f"❌ ERROR #{idx}")
|
||||
print(f"{'━' * 80}")
|
||||
print(f"👤 User: {record['user_name']}")
|
||||
print(f"📅 Session: {record['session_idx']} | Question: {record['question_idx']}")
|
||||
print(f"🏷️ Result Type: {qa.get('result_type', 'Unknown')}")
|
||||
print(f"\n❓ Question:")
|
||||
print(f" {qa.get('question', 'N/A')}")
|
||||
print(f"\n✅ Expected Answer:")
|
||||
print(f" {qa.get('answer', 'N/A')}")
|
||||
print(f"\n🤖 System Response:")
|
||||
print(f" {qa.get('system_response', 'N/A')}")
|
||||
print(f"\n💭 Reasoning:")
|
||||
reason = qa.get('question_answering_reasoning', 'N/A')
|
||||
# Wrap long reasoning text
|
||||
if len(reason) > 80:
|
||||
words = reason.split()
|
||||
lines = []
|
||||
current_line = " "
|
||||
for word in words:
|
||||
if len(current_line) + len(word) + 1 <= 80:
|
||||
current_line += word + " "
|
||||
else:
|
||||
lines.append(current_line.rstrip())
|
||||
current_line = " " + word + " "
|
||||
if current_line.strip():
|
||||
lines.append(current_line.rstrip())
|
||||
print("\n".join(lines))
|
||||
else:
|
||||
print(f" {reason}")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Compute QA statistics from eval_reme_simple_v4.py results")
|
||||
parser.add_argument("--results_file", type=str, help="Path to eval_results.jsonl file")
|
||||
parser.add_argument("--tmp_dir", type=str, help="Path to tmp directory")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.tmp_dir:
|
||||
main(input_path=args.tmp_dir)
|
||||
elif args.results_file:
|
||||
main(input_path=args.results_file)
|
||||
else:
|
||||
parser.error("Either --results_file or --tmp_dir must be provided")
|
||||
|
|
@ -1,556 +0,0 @@
|
|||
"""
|
||||
Compute statistics from existing tmp JSON files (Stage 1 results).
|
||||
|
||||
This script assumes that process_user_stage1 has already been run and JSON files
|
||||
are available in the tmp directory. It will:
|
||||
1. Load all JSON files from tmp directory
|
||||
2. Generate the combined JSONL file
|
||||
3. Run Stage 2 evaluation (extraction only, no new API calls)
|
||||
4. Aggregate results and compute metrics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme_simple_v4/tmp
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import (
|
||||
evaluation_for_memory_accuracy,
|
||||
evaluation_for_memory_integrity,
|
||||
evaluation_for_question,
|
||||
evaluation_for_update_memory,
|
||||
)
|
||||
|
||||
|
||||
def compute_f1(precision: float, recall: float) -> float:
|
||||
"""Compute F1-score from precision and recall."""
|
||||
if precision + recall == 0:
|
||||
return 0.0
|
||||
return 2 * (precision * recall) / (precision + recall)
|
||||
|
||||
|
||||
def iter_jsonl(file_path: str):
|
||||
"""Iterate over lines in a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
yield json.loads(line)
|
||||
|
||||
|
||||
async def process_user_stage2(idx: int, user_data: dict):
|
||||
"""Stage 2: Extract evaluation results from user's sessions (already computed in Stage 1)."""
|
||||
user_name = user_data["user_name"]
|
||||
|
||||
eval_results = {
|
||||
"memory_integrity_records": [],
|
||||
"memory_accuracy_records": [],
|
||||
"memory_update_records": [],
|
||||
"question_answering_records": [],
|
||||
}
|
||||
|
||||
logger.info(f"[{idx}]{user_name}: Extracting evaluation results from sessions...")
|
||||
|
||||
# Extract evaluation results from each session
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session", False):
|
||||
continue
|
||||
|
||||
if "evaluation_results" not in session:
|
||||
logger.warning(f"[{idx}]{user_name}: Session missing evaluation_results, skipping...")
|
||||
continue
|
||||
|
||||
session_eval = session["evaluation_results"]
|
||||
eval_results["memory_integrity_records"].extend(session_eval.get("memory_integrity_records", []))
|
||||
eval_results["memory_accuracy_records"].extend(session_eval.get("memory_accuracy_records", []))
|
||||
eval_results["memory_update_records"].extend(session_eval.get("memory_update_records", []))
|
||||
eval_results["question_answering_records"].extend(session_eval.get("question_answering_records", []))
|
||||
|
||||
logger.info(
|
||||
f"[{idx}]{user_name}: Extracted {len(eval_results['memory_integrity_records'])} integrity, "
|
||||
f"{len(eval_results['memory_accuracy_records'])} accuracy, "
|
||||
f"{len(eval_results['memory_update_records'])} update, "
|
||||
f"{len(eval_results['question_answering_records'])} QA records",
|
||||
)
|
||||
|
||||
return eval_results
|
||||
|
||||
|
||||
def aggregate_eval_results(eval_results):
|
||||
"""Aggregate evaluation results and compute metrics."""
|
||||
|
||||
# Memory Integrity Evaluation
|
||||
memory_integrity_scores = 0
|
||||
memory_integrity_weighted_scores = 0
|
||||
memory_integrity_valid_num = 0
|
||||
memory_integrity_num = 0
|
||||
memory_integrity_weighted_valid_num = 0
|
||||
memory_integrity_weighted_num = 0
|
||||
interference_memory_scores = 0
|
||||
interference_memory_valid_num = 0
|
||||
interference_memory_num = 0
|
||||
|
||||
for item in eval_results["memory_integrity_records"]:
|
||||
item["is_valid"] = True
|
||||
|
||||
if item["memory_source"] != "interference":
|
||||
memory_integrity_num += 1
|
||||
memory_integrity_weighted_num += item["importance"]
|
||||
else:
|
||||
interference_memory_num += 1
|
||||
|
||||
if item["memory_integrity_score"] is None:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["memory_source"] != "interference":
|
||||
if item["memory_integrity_score"] == 2:
|
||||
memory_integrity_scores += 1
|
||||
memory_integrity_weighted_scores += 0.5 * item["memory_integrity_score"] * item["importance"]
|
||||
memory_integrity_valid_num += 1
|
||||
memory_integrity_weighted_valid_num += item["importance"]
|
||||
else:
|
||||
if item["memory_integrity_score"] == 0:
|
||||
interference_memory_scores += 1
|
||||
interference_memory_valid_num += 1
|
||||
|
||||
eval_results["overall_score"]["memory_integrity"]["recall(all)"] = (
|
||||
memory_integrity_scores / memory_integrity_num if memory_integrity_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["recall(valid)"] = (
|
||||
memory_integrity_scores / memory_integrity_valid_num if memory_integrity_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["weighted_recall(all)"] = (
|
||||
memory_integrity_weighted_scores / memory_integrity_weighted_num if memory_integrity_weighted_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["weighted_recall(valid)"] = (
|
||||
memory_integrity_weighted_scores / memory_integrity_weighted_valid_num
|
||||
if memory_integrity_weighted_valid_num > 0
|
||||
else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"][
|
||||
"memory_valid_importance_sum"
|
||||
] = memory_integrity_weighted_valid_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_importance_sum"] = memory_integrity_weighted_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_valid_num"] = memory_integrity_valid_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_num"] = memory_integrity_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(all)"] = (
|
||||
interference_memory_scores / interference_memory_num if interference_memory_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(valid)"] = (
|
||||
interference_memory_scores / interference_memory_valid_num if interference_memory_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_memory_valid_num"] = interference_memory_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_memory_num"] = interference_memory_num
|
||||
|
||||
# Memory Accuracy Evaluation
|
||||
target_memory_accuracy_scores = 0
|
||||
memory_accuracy_weighted_scores = 0
|
||||
target_memory_accuracy_valid_num = 0
|
||||
target_memory_accuracy_num = 0
|
||||
memory_accuracy_valid_num = 0
|
||||
memory_accuracy_num = 0
|
||||
|
||||
for item in eval_results["memory_accuracy_records"]:
|
||||
item["is_valid"] = True
|
||||
memory_accuracy_num += 1
|
||||
|
||||
if item["is_included_in_golden_memories"] in ["true", "True"]:
|
||||
target_memory_accuracy_num += 1
|
||||
|
||||
if item["memory_accuracy_score"] is None:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["is_included_in_golden_memories"] in ["true", "True"]:
|
||||
target_memory_accuracy_scores += 0.5 * item["memory_accuracy_score"]
|
||||
target_memory_accuracy_valid_num += 1
|
||||
|
||||
memory_accuracy_weighted_scores += 0.5 * item["memory_accuracy_score"]
|
||||
memory_accuracy_valid_num += 1
|
||||
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"] = (
|
||||
target_memory_accuracy_scores / target_memory_accuracy_num if target_memory_accuracy_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_accuracy(valid)"] = (
|
||||
target_memory_accuracy_scores / target_memory_accuracy_valid_num if target_memory_accuracy_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_memory_valid_num"] = target_memory_accuracy_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_memory_num"] = target_memory_accuracy_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(all)"] = (
|
||||
memory_accuracy_weighted_scores / memory_accuracy_num if memory_accuracy_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(valid)"] = (
|
||||
memory_accuracy_weighted_scores / memory_accuracy_valid_num if memory_accuracy_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["memory_valid_num"] = memory_accuracy_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["memory_num"] = memory_accuracy_num
|
||||
|
||||
# Memory Extraction F1-score
|
||||
eval_results["overall_score"]["memory_extraction_f1"] = compute_f1(
|
||||
precision=eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"],
|
||||
recall=eval_results["overall_score"]["memory_integrity"]["recall(all)"],
|
||||
)
|
||||
|
||||
# Memory Update Evaluation
|
||||
correct_update_memory_num = 0
|
||||
hallucination_update_memory_num = 0
|
||||
omission_update_memory_num = 0
|
||||
other_update_memory_num = 0
|
||||
update_memory_num = 0
|
||||
update_memory_valid_num = 0
|
||||
|
||||
for item in eval_results["memory_update_records"]:
|
||||
item["is_valid"] = True
|
||||
update_memory_num += 1
|
||||
|
||||
if item["memory_update_type"] not in ["Correct", "Hallucination", "Omission", "Other"]:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["memory_update_type"] == "Correct":
|
||||
correct_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Hallucination":
|
||||
hallucination_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Omission":
|
||||
omission_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Other":
|
||||
other_update_memory_num += 1
|
||||
|
||||
update_memory_valid_num += 1
|
||||
|
||||
if update_memory_num > 0:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = (
|
||||
correct_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = (
|
||||
hallucination_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = (
|
||||
omission_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = (
|
||||
other_update_memory_num / update_memory_num
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = 0
|
||||
|
||||
if update_memory_valid_num > 0:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = (
|
||||
correct_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = (
|
||||
hallucination_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = (
|
||||
omission_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = (
|
||||
other_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = 0
|
||||
|
||||
eval_results["overall_score"]["memory_update"]["update_memory_valid_num"] = update_memory_valid_num
|
||||
eval_results["overall_score"]["memory_update"]["update_memory_num"] = update_memory_num
|
||||
|
||||
# Question-Answering Evaluation
|
||||
correct_qa_num = 0
|
||||
hallucination_qa_num = 0
|
||||
omission_qa_num = 0
|
||||
qa_num = 0
|
||||
qa_valid_num = 0
|
||||
|
||||
for item in eval_results["question_answering_records"]:
|
||||
item["is_valid"] = True
|
||||
qa_num += 1
|
||||
|
||||
if item["result_type"] not in ["Correct", "Hallucination", "Omission"]:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["result_type"] == "Correct":
|
||||
correct_qa_num += 1
|
||||
elif item["result_type"] == "Hallucination":
|
||||
hallucination_qa_num += 1
|
||||
elif item["result_type"] == "Omission":
|
||||
omission_qa_num += 1
|
||||
|
||||
qa_valid_num += 1
|
||||
|
||||
if qa_num > 0:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = correct_qa_num / qa_num
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = (
|
||||
hallucination_qa_num / qa_num
|
||||
)
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = omission_qa_num / qa_num
|
||||
else:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = 0
|
||||
|
||||
if qa_valid_num > 0:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = correct_qa_num / qa_valid_num
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = (
|
||||
hallucination_qa_num / qa_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = omission_qa_num / qa_valid_num
|
||||
else:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = 0
|
||||
|
||||
eval_results["overall_score"]["question_answering"]["qa_valid_num"] = qa_valid_num
|
||||
eval_results["overall_score"]["question_answering"]["qa_num"] = qa_num
|
||||
|
||||
# Memory Type Accuracy
|
||||
for item in eval_results["memory_integrity_records"]:
|
||||
if "memory_integrity_score" not in item or "importance" not in item:
|
||||
continue
|
||||
score = 1 if item["memory_integrity_score"] == 2 else 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_integrity_acc"] += score
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1
|
||||
|
||||
for item in eval_results["memory_update_records"]:
|
||||
if "memory_update_type" not in item or "importance" not in item:
|
||||
continue
|
||||
score = 1 if item["memory_update_type"] == "Correct" else 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_update_acc"] += score
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1
|
||||
|
||||
for key in eval_results["overall_score"]["memory_type_accuracy"]:
|
||||
if eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] > 0:
|
||||
total = eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"]
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] / total
|
||||
)
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] / total
|
||||
)
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"]
|
||||
+ eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"]
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = 0
|
||||
|
||||
return eval_results
|
||||
|
||||
|
||||
async def main_async(tmp_dir: str):
|
||||
"""Main function to compute statistics from tmp directory."""
|
||||
start_time = time.time()
|
||||
|
||||
# Determine paths
|
||||
parent_dir = os.path.dirname(tmp_dir)
|
||||
frame = "reme"
|
||||
|
||||
output_file_stage1 = os.path.join(parent_dir, f"{frame}_eval_results.jsonl")
|
||||
output_file_stage2 = os.path.join(parent_dir, f"{frame}_eval_stat_result.json")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("LOADING STAGE 1 RESULTS FROM TMP DIRECTORY")
|
||||
print(f"Tmp Directory: {tmp_dir}")
|
||||
print("=" * 80)
|
||||
|
||||
# Step 1: Combine all tmp JSON files into the stage1 output JSONL
|
||||
json_files = [f for f in os.listdir(tmp_dir) if f.endswith(".json")]
|
||||
print(f"\n📁 Found {len(json_files)} JSON files in tmp directory")
|
||||
|
||||
with open(output_file_stage1, "w", encoding="utf-8") as f_out:
|
||||
for file_name in json_files:
|
||||
file_path = os.path.join(tmp_dir, file_name)
|
||||
with open(file_path, "r", encoding="utf-8") as f_in:
|
||||
data = json.load(f_in)
|
||||
f_out.write(json.dumps(data, ensure_ascii=False) + "\n")
|
||||
|
||||
print(f"✅ Combined results saved to: {output_file_stage1}")
|
||||
|
||||
# Step 2: Run Stage 2 evaluation (extraction only)
|
||||
print("\n" + "=" * 80)
|
||||
print("STAGE 2: EXTRACTING AND AGGREGATING EVALUATION RESULTS")
|
||||
print("=" * 80)
|
||||
|
||||
tmp_dir2 = os.path.join(parent_dir, "tmp2")
|
||||
os.makedirs(tmp_dir2, exist_ok=True)
|
||||
|
||||
start_stage2 = time.time()
|
||||
|
||||
# Load all users and process
|
||||
user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1))
|
||||
|
||||
for idx, user_data in user_data_list:
|
||||
uuid = user_data["uuid"]
|
||||
tmp_file = os.path.join(tmp_dir2, f"{uuid}.json")
|
||||
|
||||
if os.path.exists(tmp_file):
|
||||
print(f"⚡ Skipping user {uuid} ({idx}/{len(user_data_list)}) — cached result found.")
|
||||
continue
|
||||
|
||||
print(f"[{idx}/{len(user_data_list)}] Processing user {uuid}...")
|
||||
t_user_result = await process_user_stage2(idx, user_data)
|
||||
|
||||
with open(tmp_file, "w", encoding="utf-8") as f:
|
||||
json.dump(t_user_result, f, ensure_ascii=False, indent=4)
|
||||
|
||||
elapsed = time.time() - start_stage2
|
||||
print(f"[{idx}/{len(user_data_list)}] ✅ Finished user {uuid}, elapsed {elapsed:.2f}s.")
|
||||
|
||||
# Calculate time consuming
|
||||
add_dialogue_duration_time = 0
|
||||
search_memory_duration_time = 0
|
||||
|
||||
for user_data in iter_jsonl(output_file_stage1):
|
||||
sessions = user_data["sessions"]
|
||||
|
||||
for session in sessions:
|
||||
if "add_dialogue_duration_ms" in session:
|
||||
add_dialogue_duration_time += session["add_dialogue_duration_ms"]
|
||||
|
||||
if "questions" in session:
|
||||
for question in session["questions"]:
|
||||
if "search_duration_ms" in question:
|
||||
search_memory_duration_time += question["search_duration_ms"]
|
||||
|
||||
add_dialogue_duration_time = add_dialogue_duration_time / 1000 / 60
|
||||
search_memory_duration_time = search_memory_duration_time / 1000 / 60
|
||||
|
||||
print("\n🔄 Aggregating all user results...")
|
||||
|
||||
eval_results = {
|
||||
"overall_score": {
|
||||
"memory_integrity": {},
|
||||
"memory_accuracy": {},
|
||||
"memory_extraction_f1": 0,
|
||||
"memory_update": {},
|
||||
"question_answering": {},
|
||||
"memory_type_accuracy": {
|
||||
"Event Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
"Persona Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
"Relationship Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
},
|
||||
"time_consuming": {
|
||||
"add_dialogue_duration_time": add_dialogue_duration_time,
|
||||
"search_memory_duration_time": search_memory_duration_time,
|
||||
"total_duration_time": add_dialogue_duration_time + search_memory_duration_time,
|
||||
},
|
||||
},
|
||||
"memory_integrity_records": [],
|
||||
"memory_accuracy_records": [],
|
||||
"memory_update_records": [],
|
||||
"question_answering_records": [],
|
||||
}
|
||||
|
||||
for file_name in os.listdir(tmp_dir2):
|
||||
if not file_name.endswith(".json"):
|
||||
continue
|
||||
user_file = os.path.join(tmp_dir2, file_name)
|
||||
with open(user_file, "r", encoding="utf-8") as f:
|
||||
user_result = json.load(f)
|
||||
|
||||
eval_results["memory_accuracy_records"].extend(user_result.get("memory_accuracy_records", []))
|
||||
eval_results["memory_integrity_records"].extend(user_result.get("memory_integrity_records", []))
|
||||
eval_results["memory_update_records"].extend(user_result.get("memory_update_records", []))
|
||||
eval_results["question_answering_records"].extend(user_result.get("question_answering_records", []))
|
||||
|
||||
eval_results = aggregate_eval_results(eval_results)
|
||||
|
||||
with open(output_file_stage2, "w", encoding="utf-8") as f:
|
||||
json.dump(eval_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
elapsed_total = time.time() - start_time
|
||||
print(f"\n✅ All done in {elapsed_total:.2f}s. Results saved to {output_file_stage2}")
|
||||
|
||||
# Print summary
|
||||
print("\n" + "=" * 80)
|
||||
print("EVALUATION SUMMARY")
|
||||
print("=" * 80)
|
||||
print("\n📊 Memory Integrity:")
|
||||
print(f" - Recall (all): {eval_results['overall_score']['memory_integrity'].get('recall(all)', 0):.4f}")
|
||||
print(f" - Recall (valid): {eval_results['overall_score']['memory_integrity'].get('recall(valid)', 0):.4f}")
|
||||
print(f" - Weighted Recall (all): "
|
||||
f"{eval_results['overall_score']['memory_integrity'].get('weighted_recall(all)', 0):.4f}")
|
||||
|
||||
print(f"\n📊 Memory Accuracy:")
|
||||
print(f" - Target Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(all)', 0):.4f}")
|
||||
print(f" - Target Accuracy (valid): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(valid)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Weighted Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('weighted_accuracy(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n📊 Memory Extraction F1: {eval_results['overall_score']['memory_extraction_f1']:.4f}")
|
||||
|
||||
print(f"\n📊 Memory Update:")
|
||||
print(
|
||||
f" - Correct (all): {eval_results['overall_score']['memory_update'].get('correct_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Hallucination (all): {eval_results['overall_score']['memory_update'].get('hallucination_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Omission (all): {eval_results['overall_score']['memory_update'].get('omission_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n📊 Question Answering:")
|
||||
print(
|
||||
f" - Correct (all): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Hallucination (all): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Omission (all): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n⏱️ Time Consuming:")
|
||||
print(f" - Add Dialogue: {add_dialogue_duration_time:.2f} min")
|
||||
print(f" - Search Memory: {search_memory_duration_time:.2f} min")
|
||||
print(f" - Total: {add_dialogue_duration_time + search_memory_duration_time:.2f} min")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
def main(tmp_dir: str):
|
||||
"""Synchronous entry point."""
|
||||
asyncio.run(main_async(tmp_dir))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Compute statistics from existing Stage 1 tmp results")
|
||||
parser.add_argument(
|
||||
"--tmp_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to tmp directory containing Stage 1 JSON results (e.g., bench_results/reme/tmp)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(tmp_dir=args.tmp_dir)
|
||||
|
|
@ -1,613 +0,0 @@
|
|||
"""
|
||||
HaluMem Benchmark Evaluator - Baseline (Direct QA without Memory System)
|
||||
|
||||
A simple baseline evaluation pipeline that:
|
||||
1. Loads HaluMem benchmark data
|
||||
2. Directly uses dialogue history to answer questions (no memory system)
|
||||
3. Evaluates question answering performance
|
||||
4. Generates comprehensive metrics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/eval_baseline_simple.py \
|
||||
--data_path /path/to/HaluMem-Medium.jsonl \
|
||||
--user_num 100 --max_concurrency 20
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import evaluation_for_question2
|
||||
from llms import llm_request_for_json
|
||||
|
||||
|
||||
# ==================== Configuration ====================
|
||||
|
||||
@dataclass
|
||||
class EvalConfig:
|
||||
"""Evaluation configuration parameters."""
|
||||
data_path: str
|
||||
user_num: int = 1
|
||||
max_concurrency: int = 2
|
||||
output_dir: str = "bench_results/baseline_simple"
|
||||
|
||||
|
||||
# ==================== Utilities ====================
|
||||
|
||||
class DataLoader:
|
||||
"""Handles loading and parsing of HaluMem data."""
|
||||
|
||||
@staticmethod
|
||||
def load_jsonl(file_path: str) -> list[dict]:
|
||||
"""Load all entries from a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return [json.loads(line.strip()) for line in f if line.strip()]
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""Extract user name from persona info string."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
if not match:
|
||||
raise ValueError(f"No name found in persona_info: {persona_info}")
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
|
||||
"""Format dialogue into ReMe message format."""
|
||||
return [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
)
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
|
||||
"""Format dialogue into string for evaluation (only user messages)."""
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
# Skip assistant messages - only include user messages
|
||||
if turn['role'] != 'user':
|
||||
continue
|
||||
|
||||
timestamp = datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
# Use user_name if provided
|
||||
role = user_name if user_name else 'user'
|
||||
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\n"
|
||||
f"Content: {turn['content']}\n"
|
||||
f"Time: {timestamp}"
|
||||
)
|
||||
return "\n\n".join(formatted_turns)
|
||||
|
||||
|
||||
class FileManager:
|
||||
"""Manages file I/O operations."""
|
||||
|
||||
def __init__(self, base_dir: str):
|
||||
self.base_dir = Path(base_dir)
|
||||
self.tmp_dir = self.base_dir / "tmp"
|
||||
self.tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def get_user_dir(self, user_name: str) -> Path:
|
||||
"""Get the directory path for a user."""
|
||||
user_dir = self.tmp_dir / user_name
|
||||
user_dir.mkdir(parents=True, exist_ok=True)
|
||||
return user_dir
|
||||
|
||||
def get_session_file(self, user_name: str, session_id: int) -> Path:
|
||||
"""Get the file path for a specific session."""
|
||||
return self.get_user_dir(user_name) / f"session_{session_id}.json"
|
||||
|
||||
def save_session(self, user_name: str, session_id: int, data: dict):
|
||||
"""Save session data to file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
logger.info(f"✅ Saved session {session_id} to {file_path}")
|
||||
|
||||
def load_session(self, user_name: str, session_id: int) -> dict | None:
|
||||
"""Load session data from file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
if not file_path.exists():
|
||||
return None
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
def user_has_cache(self, user_name: str) -> bool:
|
||||
"""Check if user has cached results."""
|
||||
user_dir = self.get_user_dir(user_name)
|
||||
return any(f.name.startswith("session_") and f.suffix == ".json"
|
||||
for f in user_dir.iterdir())
|
||||
|
||||
def combine_results(self, output_file: str):
|
||||
"""Combine all user session files into a single JSONL file."""
|
||||
with open(output_file, "w", encoding="utf-8") as f_out:
|
||||
for user_dir in self.tmp_dir.iterdir():
|
||||
if not user_dir.is_dir():
|
||||
continue
|
||||
|
||||
session_files = sorted([
|
||||
f for f in user_dir.iterdir()
|
||||
if f.name.startswith("session_") and f.suffix == ".json"
|
||||
])
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
# Load first session to get user metadata
|
||||
with open(session_files[0], "r", encoding="utf-8") as f_in:
|
||||
first_session = json.load(f_in)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
# Load all sessions
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f_in:
|
||||
session_data = json.load(f_in)
|
||||
# Remove redundant user metadata
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n")
|
||||
|
||||
|
||||
# ==================== Question Answering Prompt ====================
|
||||
|
||||
BASELINE_QA_PROMPT = """You are a helpful AI assistant. Based on the dialogue history provided below, please answer the question.
|
||||
|
||||
**Dialogue History:**
|
||||
{dialogue}
|
||||
|
||||
**Question:**
|
||||
{question}
|
||||
|
||||
**Instructions:**
|
||||
- Carefully read through the dialogue history
|
||||
- Answer the question based ONLY on information present in the dialogue
|
||||
- If the information needed to answer the question is NOT in the dialogue, respond with "I don't know" or "The information is not available in the dialogue"
|
||||
- Do NOT make up or hallucinate information that is not explicitly mentioned in the dialogue
|
||||
- Provide your reasoning process before giving the final answer
|
||||
|
||||
**Response Format:**
|
||||
Please respond in JSON format with the following structure:
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Your step-by-step reasoning process",
|
||||
"answer": "Your final answer (or 'I don't know' if information is not available)"
|
||||
}}
|
||||
```"""
|
||||
|
||||
|
||||
# ==================== Evaluation ====================
|
||||
|
||||
class BaselineQuestionAnsweringEvaluator:
|
||||
"""Evaluates question answering performance using direct LLM inference (no memory system)."""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
async def answer_question(
|
||||
self,
|
||||
question: str,
|
||||
formatted_dialogue: str
|
||||
) -> tuple[str, str, float]:
|
||||
"""
|
||||
Answer a question using the dialogue history directly.
|
||||
|
||||
Returns:
|
||||
tuple: (answer, reasoning, duration_ms)
|
||||
"""
|
||||
start = time.time()
|
||||
|
||||
# Format prompt
|
||||
prompt = BASELINE_QA_PROMPT.format(
|
||||
dialogue=formatted_dialogue,
|
||||
question=question
|
||||
)
|
||||
|
||||
# Get answer from LLM
|
||||
try:
|
||||
# model_name = "qwen3-max"
|
||||
model_name = "qwen3-30b-a3b-instruct-2507"
|
||||
result = await llm_request_for_json(prompt, model_name=model_name)
|
||||
answer = result.get("answer", "I don't know")
|
||||
reasoning = result.get("reasoning", "")
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting answer from LLM: {e}")
|
||||
answer = "Error: Failed to get answer"
|
||||
reasoning = str(e)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return answer, reasoning, duration_ms
|
||||
|
||||
async def evaluate_questions(
|
||||
self,
|
||||
questions: list[dict],
|
||||
user_name: str,
|
||||
uuid: str,
|
||||
session_id: int,
|
||||
formatted_dialogue: str
|
||||
) -> list[dict]:
|
||||
"""Evaluate all questions for a session."""
|
||||
results = []
|
||||
|
||||
for qa in questions:
|
||||
# Get answer directly from LLM
|
||||
answer, reasoning, duration_ms = await self.answer_question(
|
||||
question=qa["question"],
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
# Evaluate response
|
||||
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
|
||||
eval_result = await evaluation_for_question2(
|
||||
qa["question"],
|
||||
qa["answer"],
|
||||
evidence_text,
|
||||
answer,
|
||||
formatted_dialogue
|
||||
)
|
||||
|
||||
# Build result record
|
||||
qa_result = {
|
||||
**qa,
|
||||
"uuid": uuid,
|
||||
"session_id": session_id,
|
||||
"system_response": answer,
|
||||
"reasoning": reasoning,
|
||||
"answer_duration_ms": duration_ms,
|
||||
"result_type": eval_result.get("evaluation_result"),
|
||||
"question_answering_reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
results.append(qa_result)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class MetricsAggregator:
|
||||
"""Aggregates evaluation metrics."""
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = 0
|
||||
hallucination = 0
|
||||
omission = 0
|
||||
valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
|
||||
if result_type in ["Correct", "Hallucination", "Omission"]:
|
||||
valid += 1
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
if valid > 0:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": correct / valid,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid,
|
||||
"omission_qa_ratio(valid)": omission / valid
|
||||
})
|
||||
else:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0
|
||||
})
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
answer_duration = 0
|
||||
|
||||
with open(eval_results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
answer_duration += qa.get("answer_duration_ms", 0)
|
||||
|
||||
# Convert to minutes
|
||||
return {
|
||||
"answer_duration_time": answer_duration / 1000 / 60,
|
||||
"total_duration_time": answer_duration / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
class HaluMemBaselineEvaluator:
|
||||
"""Main evaluator orchestrating the baseline evaluation pipeline."""
|
||||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.qa_evaluator = BaselineQuestionAnsweringEvaluator()
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
async def process_session(
|
||||
self,
|
||||
session: dict,
|
||||
session_id: int,
|
||||
user_name: str,
|
||||
uuid: str
|
||||
) -> dict:
|
||||
"""Process a single session."""
|
||||
session_data = {
|
||||
"uuid": uuid,
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"memory_points": session["memory_points"]
|
||||
}
|
||||
|
||||
# Skip generated QA sessions
|
||||
if session.get("is_generated_qa_session", False):
|
||||
session_data["is_generated_qa_session"] = True
|
||||
return session_data
|
||||
|
||||
# Store dialogue
|
||||
dialogue = session["dialogue"]
|
||||
session_data["dialogue"] = dialogue
|
||||
|
||||
# Evaluate questions if present
|
||||
if "questions" in session:
|
||||
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
|
||||
qa_results = await self.qa_evaluator.evaluate_questions(
|
||||
questions=session["questions"],
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
session_id=session_id,
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
session_data["evaluation_results"] = {
|
||||
"question_answering_records": qa_results
|
||||
}
|
||||
|
||||
return session_data
|
||||
|
||||
async def process_user(self, user_data: dict) -> dict:
|
||||
"""Process all sessions for a user."""
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
uuid = user_data["uuid"]
|
||||
|
||||
total_sessions = len(user_data["sessions"])
|
||||
logger.info(f"Processing user: {user_name} ({total_sessions} sessions)")
|
||||
|
||||
# Semaphore for concurrency control within user sessions
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
completed_count = [0] # Use list to allow modification in nested async function
|
||||
|
||||
async def process_session_with_log(idx: int, session: dict):
|
||||
async with semaphore:
|
||||
session_data = await self.process_session(
|
||||
session=session,
|
||||
session_id=idx,
|
||||
user_name=user_name,
|
||||
uuid=uuid
|
||||
)
|
||||
|
||||
self.file_manager.save_session(user_name, idx, session_data)
|
||||
|
||||
# Update and log completion
|
||||
completed_count[0] += 1
|
||||
print(f"✅ {user_name} complete {completed_count[0]}/{total_sessions}")
|
||||
|
||||
# Process all sessions in parallel
|
||||
tasks = [
|
||||
process_session_with_log(idx, session)
|
||||
for idx, session in enumerate(user_data["sessions"])
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the complete evaluation pipeline."""
|
||||
start_time = time.time()
|
||||
|
||||
# Load user data
|
||||
all_users = self.data_loader.load_jsonl(self.config.data_path)
|
||||
users_to_process = all_users[:self.config.user_num]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("HALUMEM BASELINE EVALUATION - DIRECT QA WITHOUT MEMORY SYSTEM")
|
||||
print(f"Users: {len(users_to_process)} | Session Concurrency: {self.config.max_concurrency}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Process users sequentially (for loop)
|
||||
for idx, user_data in enumerate(users_to_process, 1):
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
|
||||
# Check cache
|
||||
if self.file_manager.user_has_cache(user_name):
|
||||
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
|
||||
continue
|
||||
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] User {user_name} completed\n")
|
||||
|
||||
# Combine results
|
||||
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
self.file_manager.combine_results(output_file)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
print(f"📁 Results: {output_file}\n")
|
||||
|
||||
# Aggregate metrics
|
||||
await self.aggregate_and_report(output_file)
|
||||
|
||||
async def aggregate_and_report(self, results_file: str):
|
||||
"""Aggregate results and generate final report."""
|
||||
print("=" * 80)
|
||||
print("AGGREGATING METRICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Collect all QA records
|
||||
qa_records = []
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
qa_records.extend(
|
||||
eval_results.get("question_answering_records", [])
|
||||
)
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
|
||||
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
# Save final report
|
||||
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
print(f"📊 Statistics saved to: {report_file}\n")
|
||||
|
||||
# Print summary
|
||||
self._print_summary(qa_metrics, time_metrics)
|
||||
|
||||
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
|
||||
"""Print evaluation summary."""
|
||||
print("=" * 80)
|
||||
print("EVALUATION SUMMARY")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Question Answering:")
|
||||
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ Time Metrics:")
|
||||
print(f" Answer Duration: {time_metrics['answer_duration_time']:.2f} min")
|
||||
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
# ==================== Entry Point ====================
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2
|
||||
):
|
||||
"""Main entry point."""
|
||||
config = EvalConfig(
|
||||
data_path=data_path,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency
|
||||
)
|
||||
|
||||
evaluator = HaluMemBaselineEvaluator(config)
|
||||
asyncio.run(evaluator.run_evaluation())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate Baseline (Direct QA) on HaluMem benchmark"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--user_num",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of users to evaluate (default: 1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Maximum concurrent user processing (default: 2)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency
|
||||
)
|
||||
|
|
@ -1,925 +0,0 @@
|
|||
"""
|
||||
Complete evaluation script for ReMe on HaluMem benchmark.
|
||||
|
||||
This script performs the full evaluation pipeline:
|
||||
1. Load HaluMem data
|
||||
2. Process each user's sessions with ReMe (summary + retrieve) - Stage 1 (Parallel)
|
||||
3. Evaluate memory integrity, accuracy, updates, and question answering - Stage 2 (Sequential)
|
||||
4. Generate metrics and statistics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/eval_reme.py --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl \
|
||||
--top_k 20 --user_num 100 --max_concurrency 20
|
||||
python bench/halumem/eval_reme.py --data_path ./HaluMem-Long.jsonl \
|
||||
--top_k 20 --user_num 100 --max_concurrency 20
|
||||
|
||||
python bench/halumem/eval_reme.py --data_path /Users/yuli/workspace/HaluMem/data/tmp_14.jsonl \
|
||||
--top_k 20 --user_num 1 --max_concurrency 1
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import (
|
||||
_PROMPTS,
|
||||
evaluation_for_memory_accuracy,
|
||||
evaluation_for_memory_integrity,
|
||||
evaluation_for_question,
|
||||
evaluation_for_update_memory,
|
||||
)
|
||||
from llms import llm_request
|
||||
from reme_ai.core.enumeration import MemoryType
|
||||
from reme_ai.core.schema import MemoryNode
|
||||
from reme_ai.reme import ReMe
|
||||
|
||||
# Template for formatting memories (from shared YAML config)
|
||||
TEMPLATE_MEMOS = _PROMPTS["TEMPLATE_MEMOS"]
|
||||
|
||||
# Prompt for question answering (using optimized PROMPT_MEMOS)
|
||||
PROMPT_MEMOS = _PROMPTS["PROMPT_MEMOS"]
|
||||
|
||||
# Initialize ReMe with rate limiting configuration
|
||||
# The default LLM can be overridden at call time using model_name parameter
|
||||
reme: ReMe = ReMe()
|
||||
|
||||
|
||||
def extract_user_name(persona_info: str):
|
||||
"""Extract user name from persona info."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
if match:
|
||||
username = match.group(1).strip()
|
||||
return username
|
||||
else:
|
||||
raise ValueError("No name found.")
|
||||
|
||||
|
||||
def iter_jsonl(file_path: str):
|
||||
"""Iterate over lines in a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
yield json.loads(line)
|
||||
|
||||
|
||||
def compute_f1(precision: float, recall: float) -> float:
|
||||
"""Compute F1-score from precision and recall."""
|
||||
if precision + recall == 0:
|
||||
return 0.0
|
||||
return 2 * (precision * recall) / (precision + recall)
|
||||
|
||||
|
||||
# ==================== Stage 1: Data Processing ====================
|
||||
|
||||
|
||||
async def add_memory_async(user_id: str, messages: list[dict]) -> tuple[list[MemoryNode], float]:
|
||||
"""Add memory to ReMe system asynchronously."""
|
||||
start = time.time()
|
||||
result = await reme.summary_v2(messages=messages, user_id=user_id)
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return result, duration_ms
|
||||
|
||||
|
||||
async def search_memory_async(query: str, user_id: str, top_k: int = 20):
|
||||
"""Search memory from ReMe system asynchronously."""
|
||||
start = time.time()
|
||||
memories = await reme.retrieve_v2(query=query, user_id=user_id, top_k=top_k)
|
||||
|
||||
# Format the context
|
||||
context = TEMPLATE_MEMOS.format(user_id=user_id, memories=memories)
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return context, memories, duration_ms
|
||||
|
||||
|
||||
async def process_user_stage1(
|
||||
user_data: dict,
|
||||
top_k_value: int,
|
||||
save_path: str,
|
||||
):
|
||||
"""Stage 1: Process user data through ReMe (summary + retrieve)."""
|
||||
user_name = extract_user_name(user_data["persona_info"])
|
||||
sessions = user_data["sessions"]
|
||||
|
||||
tmp_dir = os.path.join(save_path, "tmp")
|
||||
os.makedirs(tmp_dir, exist_ok=True)
|
||||
tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json")
|
||||
|
||||
new_user_data = {
|
||||
"uuid": user_data["uuid"],
|
||||
"user_name": user_name,
|
||||
"sessions": [],
|
||||
}
|
||||
|
||||
for idx, session in enumerate(sessions):
|
||||
logger.info(f"Processing user {user_name}: session {idx}/{len(sessions)}")
|
||||
new_session = {
|
||||
"memory_points": session["memory_points"],
|
||||
"dialogue": session["dialogue"],
|
||||
}
|
||||
|
||||
# Format dialogue
|
||||
dialogue = session["dialogue"]
|
||||
formatted_dialogue = [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(turn["timestamp"], "%b %d, %Y, %H:%M:%S")
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
]
|
||||
|
||||
# Process in batches
|
||||
result = []
|
||||
total_duration_ms = 0
|
||||
batch_size = 20
|
||||
|
||||
for i in range(0, len(formatted_dialogue), batch_size):
|
||||
batch = formatted_dialogue[i : i + batch_size]
|
||||
batch_result, duration_ms = await add_memory_async(
|
||||
user_id=user_name,
|
||||
messages=batch,
|
||||
)
|
||||
if batch_result:
|
||||
result.extend(batch_result)
|
||||
total_duration_ms += duration_ms
|
||||
|
||||
duration_ms = total_duration_ms
|
||||
|
||||
# Extract memory content
|
||||
memories = []
|
||||
for memory_node in result:
|
||||
if isinstance(memory_node, MemoryNode) and memory_node.memory_type is not MemoryType.HISTORY:
|
||||
memories.append(memory_node.content)
|
||||
|
||||
if session.get("is_generated_qa_session", False):
|
||||
new_session["add_dialogue_duration_ms"] = duration_ms
|
||||
new_session["is_generated_qa_session"] = True
|
||||
del new_session["dialogue"]
|
||||
del new_session["memory_points"]
|
||||
new_user_data["sessions"].append(new_session)
|
||||
continue
|
||||
|
||||
# Store extracted memories
|
||||
new_session["extracted_memories"] = memories
|
||||
new_session["add_dialogue_duration_ms"] = duration_ms
|
||||
|
||||
# Search updated memories for memory points
|
||||
for memory in new_session["memory_points"]:
|
||||
if memory["is_update"] == "False" or not memory.get("original_memories"):
|
||||
continue
|
||||
|
||||
_, memories_from_system, duration_ms = await search_memory_async(
|
||||
query=memory["memory_content"],
|
||||
user_id=user_name,
|
||||
top_k=10,
|
||||
)
|
||||
|
||||
memory["memories_from_system"] = memories_from_system
|
||||
|
||||
# Process questions
|
||||
if "questions" not in session:
|
||||
new_user_data["sessions"].append(new_session)
|
||||
continue
|
||||
|
||||
new_session["questions"] = []
|
||||
|
||||
for qa in session["questions"]:
|
||||
context, _, duration_ms = await search_memory_async(
|
||||
query=qa["question"],
|
||||
user_id=user_name,
|
||||
top_k=top_k_value,
|
||||
)
|
||||
|
||||
new_qa = copy.deepcopy(qa)
|
||||
new_qa["context"] = context
|
||||
new_qa["search_duration_ms"] = duration_ms
|
||||
|
||||
prompt = PROMPT_MEMOS.format(
|
||||
context=context,
|
||||
question=qa["question"],
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
response = await llm_request(prompt)
|
||||
new_qa["system_response"] = response
|
||||
new_qa["response_duration_ms"] = (time.time() - start_time) * 1000
|
||||
|
||||
new_session["questions"].append(new_qa)
|
||||
|
||||
# ==================== Evaluation for this session ====================
|
||||
session_eval_results = {
|
||||
"memory_integrity_records": [],
|
||||
"memory_accuracy_records": [],
|
||||
"memory_update_records": [],
|
||||
"question_answering_records": [],
|
||||
}
|
||||
|
||||
uuid = user_data["uuid"]
|
||||
golden_memories = session["memory_points"]
|
||||
extract_memories = new_session["extracted_memories"]
|
||||
extract_memories_str = "\n".join(extract_memories)
|
||||
|
||||
# Evaluate Memory Integrity
|
||||
logger.info(f"Evaluating Memory Integrity for session {idx}...")
|
||||
for memory in golden_memories:
|
||||
if memory["is_update"] == "True" and memory.get("memories_from_system", []):
|
||||
# Skip update memories for integrity check
|
||||
continue
|
||||
|
||||
new_memory = copy.deepcopy(memory)
|
||||
new_memory["uuid"] = uuid
|
||||
new_memory["session_id"] = idx
|
||||
|
||||
if extract_memories_str.strip() == "":
|
||||
new_memory["memory_integrity_score"] = 0
|
||||
new_memory["memory_integrity_reasoning"] = "No memories extracted"
|
||||
session_eval_results["memory_integrity_records"].append(new_memory)
|
||||
continue
|
||||
|
||||
result = await evaluation_for_memory_integrity(extract_memories_str, memory["memory_content"])
|
||||
score = int(result.get("score"))
|
||||
reasoning = result.get("reasoning", "")
|
||||
new_memory["memory_integrity_score"] = score
|
||||
new_memory["memory_integrity_reasoning"] = reasoning
|
||||
session_eval_results["memory_integrity_records"].append(new_memory)
|
||||
|
||||
# Evaluate Memory Accuracy
|
||||
logger.info(f"Evaluating Memory Accuracy for session {idx}...")
|
||||
dialogue = session["dialogue"]
|
||||
dialogue_str = []
|
||||
for turn in dialogue:
|
||||
dialogue_str.append(f'[{turn["timestamp"]}]{turn["role"]}: {turn["content"]}')
|
||||
if turn["role"] == "assistant":
|
||||
dialogue_str.append("")
|
||||
dialogue_str = "\n".join(dialogue_str)
|
||||
|
||||
golden_memories_str = "\n".join(
|
||||
[m["memory_content"] for m in golden_memories if m["memory_source"] != "interference"],
|
||||
)
|
||||
|
||||
for memory in extract_memories:
|
||||
new_memory = {
|
||||
"uuid": uuid,
|
||||
"session_id": idx,
|
||||
"memory_content": memory,
|
||||
}
|
||||
result = await evaluation_for_memory_accuracy(dialogue_str, golden_memories_str, memory)
|
||||
score = int(result.get("accuracy_score"))
|
||||
is_included_in_golden_memories = result.get("is_included_in_golden_memories", "false")
|
||||
reason = result.get("reason", "")
|
||||
new_memory["memory_accuracy_score"] = score
|
||||
new_memory["is_included_in_golden_memories"] = is_included_in_golden_memories
|
||||
new_memory["memory_accuracy_reason"] = reason
|
||||
session_eval_results["memory_accuracy_records"].append(new_memory)
|
||||
|
||||
# Evaluate Memory Update
|
||||
logger.info(f"Evaluating Memory Update for session {idx}...")
|
||||
for memory in golden_memories:
|
||||
if memory["is_update"] == "False" or not memory.get("original_memories"):
|
||||
continue
|
||||
|
||||
if not memory.get("memories_from_system", []):
|
||||
continue
|
||||
|
||||
update_memory = copy.deepcopy(memory)
|
||||
update_memory["uuid"] = uuid
|
||||
update_memory["session_id"] = idx
|
||||
|
||||
result = await evaluation_for_update_memory(
|
||||
"\n".join(update_memory["memories_from_system"]),
|
||||
update_memory["memory_content"],
|
||||
"\n".join(update_memory["original_memories"]),
|
||||
)
|
||||
update_type = result.get("evaluation_result")
|
||||
reason = result.get("reason", "")
|
||||
update_memory["memory_update_type"] = update_type
|
||||
update_memory["memory_update_reason"] = reason
|
||||
session_eval_results["memory_update_records"].append(update_memory)
|
||||
|
||||
# Evaluate Question Answering
|
||||
if "questions" in new_session:
|
||||
logger.info(f"Evaluating Question Answering for session {idx}...")
|
||||
for qa in new_session["questions"]:
|
||||
new_qa = copy.deepcopy(qa)
|
||||
new_qa["uuid"] = uuid
|
||||
new_qa["session_id"] = idx
|
||||
|
||||
result = await evaluation_for_question(
|
||||
qa["question"],
|
||||
qa["answer"],
|
||||
"\n".join([i["memory_content"] for i in qa["evidence"]]),
|
||||
qa["system_response"],
|
||||
)
|
||||
result_type = result.get("evaluation_result")
|
||||
reasoning = result.get("reasoning", "")
|
||||
new_qa["result_type"] = result_type
|
||||
new_qa["question_answering_reasoning"] = reasoning
|
||||
session_eval_results["question_answering_records"].append(new_qa)
|
||||
|
||||
# Store evaluation results in session
|
||||
new_session["evaluation_results"] = session_eval_results
|
||||
|
||||
new_user_data["sessions"].append(new_session)
|
||||
|
||||
# Save results
|
||||
with open(tmp_file, "w", encoding="utf-8") as f:
|
||||
json.dump(new_user_data, f, ensure_ascii=False, indent=2)
|
||||
session_size = len(new_user_data["sessions"])
|
||||
logger.info(f"✅ Saved user {user_name} to {tmp_file} session_size={session_size}")
|
||||
|
||||
logger.info(f"✅ Saved user {user_name} to {tmp_file} all!")
|
||||
return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file}
|
||||
|
||||
|
||||
# ==================== Stage 2: Evaluation ====================
|
||||
|
||||
|
||||
async def process_user_stage2(idx: int, user_data: dict):
|
||||
"""Stage 2: Extract evaluation results from user's sessions (already computed in Stage 1)."""
|
||||
user_name = user_data["user_name"]
|
||||
|
||||
eval_results = {
|
||||
"memory_integrity_records": [],
|
||||
"memory_accuracy_records": [],
|
||||
"memory_update_records": [],
|
||||
"question_answering_records": [],
|
||||
}
|
||||
|
||||
logger.info(f"[{idx}]{user_name}: Extracting evaluation results from sessions...")
|
||||
|
||||
# Extract evaluation results from each session
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session", False):
|
||||
continue
|
||||
|
||||
if "evaluation_results" not in session:
|
||||
logger.warning(f"[{idx}]{user_name}: Session missing evaluation_results, skipping...")
|
||||
continue
|
||||
|
||||
session_eval = session["evaluation_results"]
|
||||
eval_results["memory_integrity_records"].extend(session_eval.get("memory_integrity_records", []))
|
||||
eval_results["memory_accuracy_records"].extend(session_eval.get("memory_accuracy_records", []))
|
||||
eval_results["memory_update_records"].extend(session_eval.get("memory_update_records", []))
|
||||
eval_results["question_answering_records"].extend(session_eval.get("question_answering_records", []))
|
||||
|
||||
logger.info(
|
||||
f"[{idx}]{user_name}: Extracted {len(eval_results['memory_integrity_records'])} integrity, "
|
||||
f"{len(eval_results['memory_accuracy_records'])} accuracy, "
|
||||
f"{len(eval_results['memory_update_records'])} update, "
|
||||
f"{len(eval_results['question_answering_records'])} QA records",
|
||||
)
|
||||
|
||||
return eval_results
|
||||
|
||||
|
||||
def aggregate_eval_results(eval_results):
|
||||
"""Aggregate evaluation results and compute metrics."""
|
||||
|
||||
# Memory Integrity Evaluation
|
||||
memory_integrity_scores = 0
|
||||
memory_integrity_weighted_scores = 0
|
||||
memory_integrity_valid_num = 0
|
||||
memory_integrity_num = 0
|
||||
memory_integrity_weighted_valid_num = 0
|
||||
memory_integrity_weighted_num = 0
|
||||
interference_memory_scores = 0
|
||||
interference_memory_valid_num = 0
|
||||
interference_memory_num = 0
|
||||
|
||||
for item in eval_results["memory_integrity_records"]:
|
||||
item["is_valid"] = True
|
||||
|
||||
if item["memory_source"] != "interference":
|
||||
memory_integrity_num += 1
|
||||
memory_integrity_weighted_num += item["importance"]
|
||||
else:
|
||||
interference_memory_num += 1
|
||||
|
||||
if item["memory_integrity_score"] is None:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["memory_source"] != "interference":
|
||||
if item["memory_integrity_score"] == 2:
|
||||
memory_integrity_scores += 1
|
||||
memory_integrity_weighted_scores += 0.5 * item["memory_integrity_score"] * item["importance"]
|
||||
memory_integrity_valid_num += 1
|
||||
memory_integrity_weighted_valid_num += item["importance"]
|
||||
else:
|
||||
if item["memory_integrity_score"] == 0:
|
||||
interference_memory_scores += 1
|
||||
interference_memory_valid_num += 1
|
||||
|
||||
eval_results["overall_score"]["memory_integrity"]["recall(all)"] = (
|
||||
memory_integrity_scores / memory_integrity_num if memory_integrity_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["recall(valid)"] = (
|
||||
memory_integrity_scores / memory_integrity_valid_num if memory_integrity_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["weighted_recall(all)"] = (
|
||||
memory_integrity_weighted_scores / memory_integrity_weighted_num if memory_integrity_weighted_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"]["weighted_recall(valid)"] = (
|
||||
memory_integrity_weighted_scores / memory_integrity_weighted_valid_num
|
||||
if memory_integrity_weighted_valid_num > 0
|
||||
else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_integrity"][
|
||||
"memory_valid_importance_sum"
|
||||
] = memory_integrity_weighted_valid_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_importance_sum"] = memory_integrity_weighted_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_valid_num"] = memory_integrity_valid_num
|
||||
eval_results["overall_score"]["memory_integrity"]["memory_num"] = memory_integrity_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(all)"] = (
|
||||
interference_memory_scores / interference_memory_num if interference_memory_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(valid)"] = (
|
||||
interference_memory_scores / interference_memory_valid_num if interference_memory_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_memory_valid_num"] = interference_memory_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["interference_memory_num"] = interference_memory_num
|
||||
|
||||
# Memory Accuracy Evaluation
|
||||
target_memory_accuracy_scores = 0
|
||||
memory_accuracy_weighted_scores = 0
|
||||
target_memory_accuracy_valid_num = 0
|
||||
target_memory_accuracy_num = 0
|
||||
memory_accuracy_valid_num = 0
|
||||
memory_accuracy_num = 0
|
||||
|
||||
for item in eval_results["memory_accuracy_records"]:
|
||||
item["is_valid"] = True
|
||||
memory_accuracy_num += 1
|
||||
|
||||
if item["is_included_in_golden_memories"] in ["true", "True"]:
|
||||
target_memory_accuracy_num += 1
|
||||
|
||||
if item["memory_accuracy_score"] is None:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["is_included_in_golden_memories"] in ["true", "True"]:
|
||||
target_memory_accuracy_scores += 0.5 * item["memory_accuracy_score"]
|
||||
target_memory_accuracy_valid_num += 1
|
||||
|
||||
memory_accuracy_weighted_scores += 0.5 * item["memory_accuracy_score"]
|
||||
memory_accuracy_valid_num += 1
|
||||
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"] = (
|
||||
target_memory_accuracy_scores / target_memory_accuracy_num if target_memory_accuracy_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_accuracy(valid)"] = (
|
||||
target_memory_accuracy_scores / target_memory_accuracy_valid_num if target_memory_accuracy_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_memory_valid_num"] = target_memory_accuracy_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["target_memory_num"] = target_memory_accuracy_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(all)"] = (
|
||||
memory_accuracy_weighted_scores / memory_accuracy_num if memory_accuracy_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(valid)"] = (
|
||||
memory_accuracy_weighted_scores / memory_accuracy_valid_num if memory_accuracy_valid_num > 0 else 0
|
||||
)
|
||||
eval_results["overall_score"]["memory_accuracy"]["memory_valid_num"] = memory_accuracy_valid_num
|
||||
eval_results["overall_score"]["memory_accuracy"]["memory_num"] = memory_accuracy_num
|
||||
|
||||
# Memory Extraction F1-score
|
||||
eval_results["overall_score"]["memory_extraction_f1"] = compute_f1(
|
||||
precision=eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"],
|
||||
recall=eval_results["overall_score"]["memory_integrity"]["recall(all)"],
|
||||
)
|
||||
|
||||
# Memory Update Evaluation
|
||||
correct_update_memory_num = 0
|
||||
hallucination_update_memory_num = 0
|
||||
omission_update_memory_num = 0
|
||||
other_update_memory_num = 0
|
||||
update_memory_num = 0
|
||||
update_memory_valid_num = 0
|
||||
|
||||
for item in eval_results["memory_update_records"]:
|
||||
item["is_valid"] = True
|
||||
update_memory_num += 1
|
||||
|
||||
if item["memory_update_type"] not in ["Correct", "Hallucination", "Omission", "Other"]:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["memory_update_type"] == "Correct":
|
||||
correct_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Hallucination":
|
||||
hallucination_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Omission":
|
||||
omission_update_memory_num += 1
|
||||
elif item["memory_update_type"] == "Other":
|
||||
other_update_memory_num += 1
|
||||
|
||||
update_memory_valid_num += 1
|
||||
|
||||
if update_memory_num > 0:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = (
|
||||
correct_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = (
|
||||
hallucination_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = (
|
||||
omission_update_memory_num / update_memory_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = (
|
||||
other_update_memory_num / update_memory_num
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = 0
|
||||
|
||||
if update_memory_valid_num > 0:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = (
|
||||
correct_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = (
|
||||
hallucination_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = (
|
||||
omission_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = (
|
||||
other_update_memory_num / update_memory_valid_num
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = 0
|
||||
|
||||
eval_results["overall_score"]["memory_update"]["update_memory_valid_num"] = update_memory_valid_num
|
||||
eval_results["overall_score"]["memory_update"]["update_memory_num"] = update_memory_num
|
||||
|
||||
# Question-Answering Evaluation
|
||||
correct_qa_num = 0
|
||||
hallucination_qa_num = 0
|
||||
omission_qa_num = 0
|
||||
qa_num = 0
|
||||
qa_valid_num = 0
|
||||
|
||||
for item in eval_results["question_answering_records"]:
|
||||
item["is_valid"] = True
|
||||
qa_num += 1
|
||||
|
||||
if item["result_type"] not in ["Correct", "Hallucination", "Omission"]:
|
||||
item["is_valid"] = False
|
||||
continue
|
||||
|
||||
if item["result_type"] == "Correct":
|
||||
correct_qa_num += 1
|
||||
elif item["result_type"] == "Hallucination":
|
||||
hallucination_qa_num += 1
|
||||
elif item["result_type"] == "Omission":
|
||||
omission_qa_num += 1
|
||||
|
||||
qa_valid_num += 1
|
||||
|
||||
if qa_num > 0:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = correct_qa_num / qa_num
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = (
|
||||
hallucination_qa_num / qa_num
|
||||
)
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = omission_qa_num / qa_num
|
||||
else:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = 0
|
||||
|
||||
if qa_valid_num > 0:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = correct_qa_num / qa_valid_num
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = (
|
||||
hallucination_qa_num / qa_valid_num
|
||||
)
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = omission_qa_num / qa_valid_num
|
||||
else:
|
||||
eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = 0
|
||||
eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = 0
|
||||
|
||||
eval_results["overall_score"]["question_answering"]["qa_valid_num"] = qa_valid_num
|
||||
eval_results["overall_score"]["question_answering"]["qa_num"] = qa_num
|
||||
|
||||
# Memory Type Accuracy
|
||||
for item in eval_results["memory_integrity_records"]:
|
||||
if "memory_integrity_score" not in item or "importance" not in item:
|
||||
continue
|
||||
score = 1 if item["memory_integrity_score"] == 2 else 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_integrity_acc"] += score
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1
|
||||
|
||||
for item in eval_results["memory_update_records"]:
|
||||
if "memory_update_type" not in item or "importance" not in item:
|
||||
continue
|
||||
score = 1 if item["memory_update_type"] == "Correct" else 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_update_acc"] += score
|
||||
eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1
|
||||
|
||||
for key in eval_results["overall_score"]["memory_type_accuracy"]:
|
||||
if eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] > 0:
|
||||
total = eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"]
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] / total
|
||||
)
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] / total
|
||||
)
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = (
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"]
|
||||
+ eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"]
|
||||
)
|
||||
else:
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = 0
|
||||
eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = 0
|
||||
|
||||
return eval_results
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
|
||||
async def main_async(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2,
|
||||
):
|
||||
"""Main evaluation pipeline."""
|
||||
frame = "reme"
|
||||
save_path = f"bench_results/{frame}/"
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
|
||||
output_file_stage1 = os.path.join(save_path, f"{frame}_eval_results.jsonl")
|
||||
output_file_stage2 = os.path.join(save_path, f"{frame}_eval_stat_result.json")
|
||||
|
||||
start_time = time.time()
|
||||
await reme.vector_store.delete_all()
|
||||
|
||||
# ==================== Stage 1: Data Processing ====================
|
||||
print("\n" + "=" * 80)
|
||||
print("STAGE 1: PROCESSING DATA WITH ReMe")
|
||||
print(f"Max Concurrency: {max_concurrency}")
|
||||
print("=" * 80)
|
||||
|
||||
tmp_dir = os.path.join(save_path, "tmp")
|
||||
os.makedirs(tmp_dir, exist_ok=True)
|
||||
|
||||
# Load all user data
|
||||
user_data_list = list(iter_jsonl(data_path))
|
||||
total_users = min(len(user_data_list), user_num)
|
||||
user_data_list = user_data_list[:total_users]
|
||||
|
||||
print(f"Processing {total_users} users with max concurrency {max_concurrency}...")
|
||||
|
||||
# Create semaphore to limit concurrency for Stage 1
|
||||
semaphore_stage1 = asyncio.Semaphore(max_concurrency)
|
||||
|
||||
async def process_single_user_stage1(idx: int, user_data: dict):
|
||||
"""Process a single user in Stage 1 with semaphore control."""
|
||||
async with semaphore_stage1:
|
||||
uuid = user_data['uuid']
|
||||
tmp_file = os.path.join(tmp_dir, f"{uuid}.json")
|
||||
|
||||
if os.path.exists(tmp_file):
|
||||
print(f"⚡ Skipping user {uuid} ({idx}/{total_users}) — cached result found.")
|
||||
return {"uuid": uuid, "status": "cached", "path": tmp_file}
|
||||
|
||||
print(f"[{idx}/{total_users}] Processing user {uuid}...")
|
||||
result = await process_user_stage1(user_data, top_k, save_path)
|
||||
print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})")
|
||||
return result
|
||||
|
||||
# Process users in parallel with controlled concurrency
|
||||
tasks = [process_single_user_stage1(idx, user_data) for idx, user_data in enumerate(user_data_list, 1)]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Combine all results into final output
|
||||
with open(output_file_stage1, "w", encoding="utf-8") as f_out:
|
||||
for file in os.listdir(tmp_dir):
|
||||
if file.endswith(".json"):
|
||||
file_path = os.path.join(tmp_dir, file)
|
||||
with open(file_path, "r", encoding="utf-8") as f_in:
|
||||
data = json.load(f_in)
|
||||
f_out.write(json.dumps(data, ensure_ascii=False) + "\n")
|
||||
|
||||
elapsed_stage1 = time.time() - start_time
|
||||
print(f"\n✅ Stage 1 completed in {elapsed_stage1:.2f}s")
|
||||
print(f"✅ Results saved to: {output_file_stage1}")
|
||||
|
||||
# ==================== Stage 2: Evaluation ====================
|
||||
print("\n" + "=" * 80)
|
||||
print("STAGE 2: EVALUATING MEMORY PERFORMANCE (Sequential)")
|
||||
print("=" * 80)
|
||||
|
||||
tmp_dir2 = os.path.join(save_path, "tmp2")
|
||||
os.makedirs(tmp_dir2, exist_ok=True)
|
||||
|
||||
start_stage2 = time.time()
|
||||
|
||||
# Load all users and process sequentially
|
||||
user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1))
|
||||
|
||||
for idx, user_data in user_data_list:
|
||||
uuid = user_data["uuid"]
|
||||
tmp_file = os.path.join(tmp_dir2, f"{uuid}.json")
|
||||
|
||||
if os.path.exists(tmp_file):
|
||||
print(f"⚡ Skipping user {uuid} ({idx}/{len(user_data_list)}) — cached result found.")
|
||||
continue
|
||||
|
||||
print(f"[{idx}/{len(user_data_list)}] Processing user {uuid}...")
|
||||
t_user_result = await process_user_stage2(idx, user_data)
|
||||
|
||||
with open(tmp_file, "w", encoding="utf-8") as f:
|
||||
json.dump(t_user_result, f, ensure_ascii=False, indent=4)
|
||||
|
||||
elapsed = time.time() - start_stage2
|
||||
print(f"[{idx}/{len(user_data_list)}] ✅ Finished user {uuid}, elapsed {elapsed:.2f}s.")
|
||||
|
||||
# Calculate time consuming
|
||||
add_dialogue_duration_time = 0
|
||||
search_memory_duration_time = 0
|
||||
|
||||
for user_data in iter_jsonl(output_file_stage1):
|
||||
sessions = user_data["sessions"]
|
||||
|
||||
for session in sessions:
|
||||
if "add_dialogue_duration_ms" in session:
|
||||
add_dialogue_duration_time += session["add_dialogue_duration_ms"]
|
||||
|
||||
if "questions" in session:
|
||||
for question in session["questions"]:
|
||||
if "search_duration_ms" in question:
|
||||
search_memory_duration_time += question["search_duration_ms"]
|
||||
|
||||
add_dialogue_duration_time = add_dialogue_duration_time / 1000 / 60
|
||||
search_memory_duration_time = search_memory_duration_time / 1000 / 60
|
||||
|
||||
print("\n🔄 Aggregating all user results...")
|
||||
|
||||
eval_results = {
|
||||
"overall_score": {
|
||||
"memory_integrity": {},
|
||||
"memory_accuracy": {},
|
||||
"memory_extraction_f1": 0,
|
||||
"memory_update": {},
|
||||
"question_answering": {},
|
||||
"memory_type_accuracy": {
|
||||
"Event Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
"Persona Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
"Relationship Memory": {
|
||||
"memory_integrity_acc": 0,
|
||||
"memory_update_acc": 0,
|
||||
"total_num": 0,
|
||||
},
|
||||
},
|
||||
"time_consuming": {
|
||||
"add_dialogue_duration_time": add_dialogue_duration_time,
|
||||
"search_memory_duration_time": search_memory_duration_time,
|
||||
"total_duration_time": add_dialogue_duration_time + search_memory_duration_time,
|
||||
},
|
||||
},
|
||||
"memory_integrity_records": [],
|
||||
"memory_accuracy_records": [],
|
||||
"memory_update_records": [],
|
||||
"question_answering_records": [],
|
||||
}
|
||||
|
||||
for file_name in os.listdir(tmp_dir2):
|
||||
if not file_name.endswith(".json"):
|
||||
continue
|
||||
user_file = os.path.join(tmp_dir2, file_name)
|
||||
with open(user_file, "r", encoding="utf-8") as f:
|
||||
user_result = json.load(f)
|
||||
|
||||
eval_results["memory_accuracy_records"].extend(user_result.get("memory_accuracy_records", []))
|
||||
eval_results["memory_integrity_records"].extend(user_result.get("memory_integrity_records", []))
|
||||
eval_results["memory_update_records"].extend(user_result.get("memory_update_records", []))
|
||||
eval_results["question_answering_records"].extend(user_result.get("question_answering_records", []))
|
||||
|
||||
eval_results = aggregate_eval_results(eval_results)
|
||||
|
||||
with open(output_file_stage2, "w", encoding="utf-8") as f:
|
||||
json.dump(eval_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
elapsed_total = time.time() - start_time
|
||||
print(f"\n✅ All done in {elapsed_total:.2f}s. Results saved to {output_file_stage2}")
|
||||
|
||||
# Print summary
|
||||
print("\n" + "=" * 80)
|
||||
print("EVALUATION SUMMARY")
|
||||
print("=" * 80)
|
||||
print("\n📊 Memory Integrity:")
|
||||
print(f" - Recall (all): {eval_results['overall_score']['memory_integrity'].get('recall(all)', 0):.4f}")
|
||||
print(f" - Recall (valid): {eval_results['overall_score']['memory_integrity'].get('recall(valid)', 0):.4f}")
|
||||
print(f" - Weighted Recall (all): "
|
||||
f"{eval_results['overall_score']['memory_integrity'].get('weighted_recall(all)', 0):.4f}")
|
||||
|
||||
print(f"\n📊 Memory Accuracy:")
|
||||
print(f" - Target Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(all)', 0):.4f}")
|
||||
print(f" - Target Accuracy (valid): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(valid)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Weighted Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('weighted_accuracy(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n📊 Memory Extraction F1: {eval_results['overall_score']['memory_extraction_f1']:.4f}")
|
||||
|
||||
print(f"\n📊 Memory Update:")
|
||||
print(
|
||||
f" - Correct (all): {eval_results['overall_score']['memory_update'].get('correct_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Hallucination (all): {eval_results['overall_score']['memory_update'].get('hallucination_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Omission (all): {eval_results['overall_score']['memory_update'].get('omission_update_memory_ratio(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n📊 Question Answering:")
|
||||
print(
|
||||
f" - Correct (all): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Hallucination (all): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
print(
|
||||
f" - Omission (all): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(all)', 0):.4f}",
|
||||
)
|
||||
|
||||
print(f"\n⏱️ Time Consuming:")
|
||||
print(f" - Add Dialogue: {add_dialogue_duration_time:.2f} min")
|
||||
print(f" - Search Memory: {search_memory_duration_time:.2f} min")
|
||||
print(f" - Total: {add_dialogue_duration_time + search_memory_duration_time:.2f} min")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2,
|
||||
):
|
||||
"""Synchronous entry point."""
|
||||
asyncio.run(main_async(data_path, top_k, user_num, max_concurrency))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Complete evaluation for ReMe on HaluMem benchmark")
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem data file (e.g., HaluMem-medium.jsonl)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=20,
|
||||
help="Number of top memories to retrieve (default: 20)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--user_num",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of users to evaluate (default: 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Maximum concurrency for stage 1 processing (default: 2)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
top_k=args.top_k,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency,
|
||||
)
|
||||
|
|
@ -1,662 +0,0 @@
|
|||
"""
|
||||
HaluMem Benchmark Evaluator for ReMe - Question Answering
|
||||
|
||||
A modular evaluation pipeline that:
|
||||
1. Loads HaluMem benchmark data
|
||||
2. Processes user sessions through ReMe (summarization + retrieval)
|
||||
3. Evaluates question answering performance
|
||||
4. Generates comprehensive metrics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/eval_reme_simple.py \
|
||||
--data_path /path/to/HaluMem-Medium.jsonl \
|
||||
--top_k 20 --user_num 100 --max_concurrency 20
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import evaluation_for_question2
|
||||
from reme_ai.core.enumeration import MemoryType
|
||||
from reme_ai.core.schema import MemoryNode
|
||||
from reme_ai.reme import ReMe
|
||||
|
||||
|
||||
# ==================== Configuration ====================
|
||||
|
||||
@dataclass
|
||||
class EvalConfig:
|
||||
"""Evaluation configuration parameters."""
|
||||
data_path: str
|
||||
top_k: int = 20
|
||||
user_num: int = 1
|
||||
max_concurrency: int = 2
|
||||
batch_size: int = 20
|
||||
output_dir: str = "bench_results/reme_simple"
|
||||
|
||||
|
||||
# ==================== Utilities ====================
|
||||
|
||||
class DataLoader:
|
||||
"""Handles loading and parsing of HaluMem data."""
|
||||
|
||||
@staticmethod
|
||||
def load_jsonl(file_path: str) -> list[dict]:
|
||||
"""Load all entries from a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return [json.loads(line.strip()) for line in f if line.strip()]
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""Extract user name from persona info string."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
if not match:
|
||||
raise ValueError(f"No name found in persona_info: {persona_info}")
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
|
||||
"""Format dialogue into ReMe message format."""
|
||||
return [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
)
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
|
||||
"""Format dialogue into string for evaluation."""
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
timestamp = datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
# Use user_name if role is 'user' and user_name is provided
|
||||
role = user_name if turn['role'] == 'user' and user_name else turn['role']
|
||||
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\n"
|
||||
f"Content: {turn['content']}\n"
|
||||
f"Time: {timestamp}"
|
||||
)
|
||||
return "\n\n".join(formatted_turns)
|
||||
|
||||
|
||||
class FileManager:
|
||||
"""Manages file I/O operations."""
|
||||
|
||||
def __init__(self, base_dir: str):
|
||||
self.base_dir = Path(base_dir)
|
||||
self.tmp_dir = self.base_dir / "tmp"
|
||||
self.tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def get_user_dir(self, user_name: str) -> Path:
|
||||
"""Get the directory path for a user."""
|
||||
user_dir = self.tmp_dir / user_name
|
||||
user_dir.mkdir(parents=True, exist_ok=True)
|
||||
return user_dir
|
||||
|
||||
def get_session_file(self, user_name: str, session_id: int) -> Path:
|
||||
"""Get the file path for a specific session."""
|
||||
return self.get_user_dir(user_name) / f"session_{session_id}.json"
|
||||
|
||||
def save_session(self, user_name: str, session_id: int, data: dict):
|
||||
"""Save session data to file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
logger.info(f"✅ Saved session {session_id} to {file_path}")
|
||||
|
||||
def load_session(self, user_name: str, session_id: int) -> dict | None:
|
||||
"""Load session data from file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
if not file_path.exists():
|
||||
return None
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
def user_has_cache(self, user_name: str) -> bool:
|
||||
"""Check if user has cached results."""
|
||||
user_dir = self.get_user_dir(user_name)
|
||||
return any(f.name.startswith("session_") and f.suffix == ".json"
|
||||
for f in user_dir.iterdir())
|
||||
|
||||
def combine_results(self, output_file: str):
|
||||
"""Combine all user session files into a single JSONL file."""
|
||||
with open(output_file, "w", encoding="utf-8") as f_out:
|
||||
for user_dir in self.tmp_dir.iterdir():
|
||||
if not user_dir.is_dir():
|
||||
continue
|
||||
|
||||
session_files = sorted([
|
||||
f for f in user_dir.iterdir()
|
||||
if f.name.startswith("session_") and f.suffix == ".json"
|
||||
])
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
# Load first session to get user metadata
|
||||
with open(session_files[0], "r", encoding="utf-8") as f_in:
|
||||
first_session = json.load(f_in)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
# Load all sessions
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f_in:
|
||||
session_data = json.load(f_in)
|
||||
# Remove redundant user metadata
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n")
|
||||
|
||||
|
||||
# ==================== Memory Operations ====================
|
||||
|
||||
class MemoryProcessor:
|
||||
"""Handles ReMe memory operations."""
|
||||
|
||||
def __init__(self, reme: ReMe):
|
||||
self.reme = reme
|
||||
|
||||
async def add_memories(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list[dict],
|
||||
batch_size: int = 20
|
||||
) -> tuple[list[str], list[list[dict]], float]:
|
||||
"""
|
||||
Add memories in batches and return extracted memory contents.
|
||||
|
||||
Returns:
|
||||
tuple: (extracted_memories, agent_messages, total_duration_ms)
|
||||
"""
|
||||
added_memories: list[MemoryNode] = []
|
||||
deleted_memories: list[str] = []
|
||||
all_agent_messages: list = []
|
||||
total_duration_ms = 0
|
||||
for i in range(0, len(messages), batch_size):
|
||||
batch = messages[i:i + batch_size]
|
||||
start = time.time()
|
||||
|
||||
memory_nodes, agent_messages, success = await self.reme.summary_v2(
|
||||
messages=batch,
|
||||
user_id=user_id
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
total_duration_ms += duration_ms
|
||||
|
||||
# Save agent messages for this batch
|
||||
if agent_messages:
|
||||
all_agent_messages.extend(agent_messages)
|
||||
|
||||
if memory_nodes:
|
||||
for node in memory_nodes:
|
||||
if isinstance(node, MemoryNode) and node.memory_type == MemoryType.HISTORY:
|
||||
continue
|
||||
|
||||
if isinstance(node, MemoryNode):
|
||||
added_memories.append(node)
|
||||
|
||||
if isinstance(node, str):
|
||||
deleted_memories.append(node)
|
||||
|
||||
extracted_memories = deleted_memories
|
||||
extracted_memories += ["[delete]" + n.format_memory() for n in added_memories if n.memory_id in deleted_memories]
|
||||
extracted_memories += ["[add]" + n.format_memory() for n in added_memories if n.memory_id not in deleted_memories]
|
||||
return extracted_memories, all_agent_messages, total_duration_ms
|
||||
|
||||
async def search_memory(
|
||||
self,
|
||||
query: str,
|
||||
user_id: str,
|
||||
top_k: int = 20
|
||||
) -> tuple[str, list, float]:
|
||||
"""
|
||||
Search memory and return response.
|
||||
|
||||
Returns:
|
||||
tuple: (response, agent_messages, duration_ms)
|
||||
"""
|
||||
start = time.time()
|
||||
response, agent_messages, success = await self.reme.retrieve_v2(
|
||||
query=query,
|
||||
user_id=user_id,
|
||||
top_k=top_k
|
||||
)
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return response, agent_messages, duration_ms
|
||||
|
||||
|
||||
# ==================== Evaluation ====================
|
||||
|
||||
class QuestionAnsweringEvaluator:
|
||||
"""Evaluates question answering performance."""
|
||||
|
||||
def __init__(self, memory_processor: MemoryProcessor, top_k: int):
|
||||
self.memory_processor = memory_processor
|
||||
self.top_k = top_k
|
||||
|
||||
async def evaluate_questions(
|
||||
self,
|
||||
questions: list[dict],
|
||||
user_name: str,
|
||||
uuid: str,
|
||||
session_id: int,
|
||||
formatted_dialogue: str
|
||||
) -> list[dict]:
|
||||
"""Evaluate all questions for a session."""
|
||||
results = []
|
||||
|
||||
for qa in questions:
|
||||
# Search memory for answer
|
||||
response, agent_messages, duration_ms = await self.memory_processor.search_memory(
|
||||
query=qa["question"],
|
||||
user_id=user_name,
|
||||
top_k=self.top_k
|
||||
)
|
||||
|
||||
# Evaluate response
|
||||
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
|
||||
eval_result = await evaluation_for_question2(
|
||||
qa["question"],
|
||||
qa["answer"],
|
||||
evidence_text,
|
||||
response,
|
||||
formatted_dialogue
|
||||
)
|
||||
|
||||
# Build result record
|
||||
qa_result = {
|
||||
**qa,
|
||||
"uuid": uuid,
|
||||
"session_id": session_id,
|
||||
"system_response": response,
|
||||
"retrieve_messages": [m.model_dump() for m in agent_messages],
|
||||
"search_duration_ms": duration_ms,
|
||||
"result_type": eval_result.get("evaluation_result"),
|
||||
"question_answering_reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
results.append(qa_result)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class MetricsAggregator:
|
||||
"""Aggregates evaluation metrics."""
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = 0
|
||||
hallucination = 0
|
||||
omission = 0
|
||||
valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
|
||||
if result_type in ["Correct", "Hallucination", "Omission"]:
|
||||
valid += 1
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
if valid > 0:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": correct / valid,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid,
|
||||
"omission_qa_ratio(valid)": omission / valid
|
||||
})
|
||||
else:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0
|
||||
})
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
add_duration = 0
|
||||
search_duration = 0
|
||||
|
||||
with open(eval_results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
add_duration += session.get("add_dialogue_duration_ms", 0)
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
search_duration += qa.get("search_duration_ms", 0)
|
||||
|
||||
# Convert to minutes
|
||||
return {
|
||||
"add_dialogue_duration_time": add_duration / 1000 / 60,
|
||||
"search_memory_duration_time": search_duration / 1000 / 60,
|
||||
"total_duration_time": (add_duration + search_duration) / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
class HaluMemEvaluator:
|
||||
"""Main evaluator orchestrating the entire pipeline."""
|
||||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.reme = ReMe()
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.memory_processor = MemoryProcessor(self.reme)
|
||||
self.qa_evaluator = QuestionAnsweringEvaluator(
|
||||
self.memory_processor,
|
||||
config.top_k
|
||||
)
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
async def process_session(
|
||||
self,
|
||||
session: dict,
|
||||
session_id: int,
|
||||
user_name: str,
|
||||
uuid: str
|
||||
) -> dict:
|
||||
"""Process a single session."""
|
||||
session_data = {
|
||||
"uuid": uuid,
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"memory_points": session["memory_points"]
|
||||
}
|
||||
|
||||
# Skip generated QA sessions
|
||||
if session.get("is_generated_qa_session", False):
|
||||
session_data["is_generated_qa_session"] = True
|
||||
return session_data
|
||||
|
||||
# Format and add dialogue to memory
|
||||
dialogue = session["dialogue"]
|
||||
formatted_messages = self.data_loader.format_dialogue_messages(dialogue)
|
||||
|
||||
extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories(
|
||||
user_id=user_name,
|
||||
messages=formatted_messages,
|
||||
batch_size=self.config.batch_size
|
||||
)
|
||||
|
||||
session_data.update({
|
||||
"dialogue": dialogue,
|
||||
"extracted_memories": extracted_memories,
|
||||
"summary_messages": [m.model_dump() for m in agent_messages],
|
||||
"add_dialogue_duration_ms": duration_ms
|
||||
})
|
||||
|
||||
# Evaluate questions if present
|
||||
if "questions" in session:
|
||||
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
|
||||
qa_results = await self.qa_evaluator.evaluate_questions(
|
||||
questions=session["questions"],
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
session_id=session_id,
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
session_data["evaluation_results"] = {
|
||||
"question_answering_records": qa_results
|
||||
}
|
||||
|
||||
return session_data
|
||||
|
||||
async def process_user(self, user_data: dict) -> dict:
|
||||
"""Process all sessions for a user."""
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
uuid = user_data["uuid"]
|
||||
|
||||
logger.info(f"Processing user: {user_name}")
|
||||
|
||||
for idx, session in enumerate(user_data["sessions"]):
|
||||
logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}")
|
||||
|
||||
session_data = await self.process_session(
|
||||
session=session,
|
||||
session_id=idx,
|
||||
user_name=user_name,
|
||||
uuid=uuid
|
||||
)
|
||||
|
||||
self.file_manager.save_session(user_name, idx, session_data)
|
||||
|
||||
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the complete evaluation pipeline."""
|
||||
start_time = time.time()
|
||||
|
||||
# Clear existing data
|
||||
await self.reme.vector_store.delete_all()
|
||||
|
||||
# Load user data
|
||||
all_users = self.data_loader.load_jsonl(self.config.data_path)
|
||||
users_to_process = all_users[:self.config.user_num]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("HALUMEM EVALUATION - QUESTION ANSWERING")
|
||||
print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Process users with concurrency control
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
|
||||
async def process_with_cache_check(idx: int, user_data: dict):
|
||||
async with semaphore:
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
|
||||
# Check cache
|
||||
if self.file_manager.user_has_cache(user_name):
|
||||
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
|
||||
return {"user_name": user_name, "status": "cached"}
|
||||
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
result = await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
|
||||
return result
|
||||
|
||||
tasks = [
|
||||
process_with_cache_check(idx, user)
|
||||
for idx, user in enumerate(users_to_process, 1)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Combine results
|
||||
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
self.file_manager.combine_results(output_file)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
print(f"📁 Results: {output_file}\n")
|
||||
|
||||
# Aggregate metrics
|
||||
await self.aggregate_and_report(output_file)
|
||||
|
||||
async def aggregate_and_report(self, results_file: str):
|
||||
"""Aggregate results and generate final report."""
|
||||
print("=" * 80)
|
||||
print("AGGREGATING METRICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Collect all QA records
|
||||
qa_records = []
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
qa_records.extend(
|
||||
eval_results.get("question_answering_records", [])
|
||||
)
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
|
||||
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
# Save final report
|
||||
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
print(f"📊 Statistics saved to: {report_file}\n")
|
||||
|
||||
# Print summary
|
||||
self._print_summary(qa_metrics, time_metrics)
|
||||
|
||||
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
|
||||
"""Print evaluation summary."""
|
||||
print("=" * 80)
|
||||
print("EVALUATION SUMMARY")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Question Answering:")
|
||||
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ Time Metrics:")
|
||||
print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min")
|
||||
print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min")
|
||||
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
# ==================== Entry Point ====================
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2
|
||||
):
|
||||
"""Main entry point."""
|
||||
config = EvalConfig(
|
||||
data_path=data_path,
|
||||
top_k=top_k,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency
|
||||
)
|
||||
|
||||
evaluator = HaluMemEvaluator(config)
|
||||
asyncio.run(evaluator.run_evaluation())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate ReMe on HaluMem benchmark (Question Answering)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=20,
|
||||
help="Number of memories to retrieve (default: 20)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--user_num",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of users to evaluate (default: 1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Maximum concurrent user processing (default: 2)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
top_k=args.top_k,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency
|
||||
)
|
||||
|
|
@ -1,676 +0,0 @@
|
|||
"""
|
||||
HaluMem Benchmark Evaluator for ReMe V3 - Question Answering
|
||||
|
||||
A modular evaluation pipeline that:
|
||||
1. Loads HaluMem benchmark data
|
||||
2. Processes user sessions through ReMe V3 (summarization + retrieval)
|
||||
3. Evaluates question answering performance
|
||||
4. Generates comprehensive metrics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/eval_reme_simple_v3.py \
|
||||
--data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \
|
||||
--top_k 20 --user_num 100 --max_concurrency 20
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import evaluation_for_question2
|
||||
from reme_ai.core.enumeration import MemoryType
|
||||
from reme_ai.core.schema import MemoryNode
|
||||
from reme_ai.reme import ReMe
|
||||
|
||||
|
||||
# ==================== Configuration ====================
|
||||
|
||||
@dataclass
|
||||
class EvalConfig:
|
||||
"""Evaluation configuration parameters."""
|
||||
data_path: str
|
||||
top_k: int = 20
|
||||
user_num: int = 1
|
||||
max_concurrency: int = 2
|
||||
batch_size: int = 20
|
||||
output_dir: str = "bench_results/reme_simple_v3"
|
||||
|
||||
|
||||
# ==================== Utilities ====================
|
||||
|
||||
class DataLoader:
|
||||
"""Handles loading and parsing of HaluMem data."""
|
||||
|
||||
@staticmethod
|
||||
def load_jsonl(file_path: str) -> list[dict]:
|
||||
"""Load all entries from a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return [json.loads(line.strip()) for line in f if line.strip()]
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""Extract user name from persona info string."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
if not match:
|
||||
raise ValueError(f"No name found in persona_info: {persona_info}")
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
|
||||
"""Format dialogue into ReMe message format with conversation_time (user messages only)."""
|
||||
return [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
)
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
if turn["role"] == "user" # Only include user messages
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
|
||||
"""Format dialogue into string for evaluation."""
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
timestamp = datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
# Use user_name if role is 'user' and user_name is provided
|
||||
role = user_name if turn['role'] == 'user' and user_name else turn['role']
|
||||
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\n"
|
||||
f"Content: {turn['content']}\n"
|
||||
f"Time: {timestamp}"
|
||||
)
|
||||
return "\n\n".join(formatted_turns)
|
||||
|
||||
|
||||
class FileManager:
|
||||
"""Manages file I/O operations."""
|
||||
|
||||
def __init__(self, base_dir: str):
|
||||
self.base_dir = Path(base_dir)
|
||||
self.tmp_dir = self.base_dir / "tmp"
|
||||
self.tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def get_user_dir(self, user_name: str) -> Path:
|
||||
"""Get the directory path for a user."""
|
||||
user_dir = self.tmp_dir / user_name
|
||||
user_dir.mkdir(parents=True, exist_ok=True)
|
||||
return user_dir
|
||||
|
||||
def get_session_file(self, user_name: str, session_id: int) -> Path:
|
||||
"""Get the file path for a specific session."""
|
||||
return self.get_user_dir(user_name) / f"session_{session_id}.json"
|
||||
|
||||
def save_session(self, user_name: str, session_id: int, data: dict):
|
||||
"""Save session data to file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
with open(file_path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
logger.info(f"✅ Saved session {session_id} to {file_path}")
|
||||
|
||||
def load_session(self, user_name: str, session_id: int) -> dict | None:
|
||||
"""Load session data from file."""
|
||||
file_path = self.get_session_file(user_name, session_id)
|
||||
if not file_path.exists():
|
||||
return None
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
def user_has_cache(self, user_name: str) -> bool:
|
||||
"""Check if user has cached results."""
|
||||
user_dir = self.get_user_dir(user_name)
|
||||
return any(f.name.startswith("session_") and f.suffix == ".json"
|
||||
for f in user_dir.iterdir())
|
||||
|
||||
def combine_results(self, output_file: str):
|
||||
"""Combine all user session files into a single JSONL file."""
|
||||
with open(output_file, "w", encoding="utf-8") as f_out:
|
||||
for user_dir in self.tmp_dir.iterdir():
|
||||
if not user_dir.is_dir():
|
||||
continue
|
||||
|
||||
session_files = sorted([
|
||||
f for f in user_dir.iterdir()
|
||||
if f.name.startswith("session_") and f.suffix == ".json"
|
||||
])
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
# Load first session to get user metadata
|
||||
with open(session_files[0], "r", encoding="utf-8") as f_in:
|
||||
first_session = json.load(f_in)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
# Load all sessions
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f_in:
|
||||
session_data = json.load(f_in)
|
||||
# Remove redundant user metadata
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n")
|
||||
|
||||
|
||||
# ==================== Memory Operations ====================
|
||||
|
||||
class MemoryProcessor:
|
||||
"""Handles ReMe V3 memory operations."""
|
||||
|
||||
def __init__(self, reme: ReMe):
|
||||
self.reme = reme
|
||||
|
||||
async def add_memories(
|
||||
self,
|
||||
user_id: str,
|
||||
messages: list[dict],
|
||||
batch_size: int = 10000
|
||||
) -> tuple[list[str], list[list[dict]], float]:
|
||||
"""
|
||||
Add memories in batches using ReMe V3 and return extracted memory contents.
|
||||
|
||||
Returns:
|
||||
tuple: (extracted_memories, agent_messages, total_duration_ms)
|
||||
"""
|
||||
added_memories: list[MemoryNode] = []
|
||||
deleted_memories: list[str] = []
|
||||
all_agent_messages: list = []
|
||||
total_duration_ms = 0
|
||||
|
||||
for i in range(0, len(messages), batch_size):
|
||||
batch = messages[i:i + batch_size]
|
||||
start = time.time()
|
||||
|
||||
# Use summary_v3 instead of summary_v2
|
||||
memory_nodes, agent_messages, success = await self.reme.summary_v3(
|
||||
messages=batch,
|
||||
user_id=user_id
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
total_duration_ms += duration_ms
|
||||
|
||||
# Save agent messages for this batch
|
||||
if agent_messages:
|
||||
all_agent_messages.extend(agent_messages)
|
||||
|
||||
if memory_nodes:
|
||||
for node in memory_nodes:
|
||||
if isinstance(node, MemoryNode) and node.memory_type == MemoryType.HISTORY:
|
||||
continue
|
||||
|
||||
if isinstance(node, MemoryNode):
|
||||
added_memories.append(node)
|
||||
|
||||
if isinstance(node, str):
|
||||
deleted_memories.append(node)
|
||||
|
||||
extracted_memories = deleted_memories
|
||||
extracted_memories += ["[delete]" + n.format_memory() for n in added_memories if n.memory_id in deleted_memories]
|
||||
extracted_memories += ["[add]" + n.format_memory() for n in added_memories if n.memory_id not in deleted_memories]
|
||||
return extracted_memories, all_agent_messages, total_duration_ms
|
||||
|
||||
async def search_memory(
|
||||
self,
|
||||
query: str,
|
||||
user_id: str,
|
||||
top_k: int = 20
|
||||
) -> tuple[str, list, float]:
|
||||
"""
|
||||
Search memory using ReMe V3 and return response.
|
||||
|
||||
Returns:
|
||||
tuple: (response, agent_messages, duration_ms)
|
||||
"""
|
||||
start = time.time()
|
||||
|
||||
# Use retrieve_v3 instead of retrieve_v2
|
||||
response, agent_messages, success = await self.reme.retrieve_v3(
|
||||
query=query,
|
||||
user_id=user_id,
|
||||
top_k=top_k
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return response, agent_messages, duration_ms
|
||||
|
||||
|
||||
# ==================== Evaluation ====================
|
||||
|
||||
class QuestionAnsweringEvaluator:
|
||||
"""Evaluates question answering performance."""
|
||||
|
||||
def __init__(self, memory_processor: MemoryProcessor, top_k: int):
|
||||
self.memory_processor = memory_processor
|
||||
self.top_k = top_k
|
||||
|
||||
async def evaluate_questions(
|
||||
self,
|
||||
questions: list[dict],
|
||||
user_name: str,
|
||||
uuid: str,
|
||||
session_id: int,
|
||||
formatted_dialogue: str
|
||||
) -> list[dict]:
|
||||
"""Evaluate all questions for a session."""
|
||||
results = []
|
||||
|
||||
for qa in questions:
|
||||
# Search memory for answer using V3
|
||||
response, agent_messages, duration_ms = await self.memory_processor.search_memory(
|
||||
query=qa["question"],
|
||||
user_id=user_name,
|
||||
top_k=self.top_k
|
||||
)
|
||||
|
||||
# Evaluate response
|
||||
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
|
||||
eval_result = await evaluation_for_question2(
|
||||
qa["question"],
|
||||
qa["answer"],
|
||||
evidence_text,
|
||||
response,
|
||||
formatted_dialogue
|
||||
)
|
||||
|
||||
# Build result record
|
||||
qa_result = {
|
||||
**qa,
|
||||
"uuid": uuid,
|
||||
"session_id": session_id,
|
||||
"system_response": response,
|
||||
"retrieve_messages": [m.model_dump() for m in agent_messages],
|
||||
"search_duration_ms": duration_ms,
|
||||
"result_type": eval_result.get("evaluation_result"),
|
||||
"question_answering_reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
results.append(qa_result)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class MetricsAggregator:
|
||||
"""Aggregates evaluation metrics."""
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = 0
|
||||
hallucination = 0
|
||||
omission = 0
|
||||
valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
|
||||
if result_type in ["Correct", "Hallucination", "Omission"]:
|
||||
valid += 1
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
if valid > 0:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": correct / valid,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid,
|
||||
"omission_qa_ratio(valid)": omission / valid
|
||||
})
|
||||
else:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0
|
||||
})
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
add_duration = 0
|
||||
search_duration = 0
|
||||
|
||||
with open(eval_results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
add_duration += session.get("add_dialogue_duration_ms", 0)
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
search_duration += qa.get("search_duration_ms", 0)
|
||||
|
||||
# Convert to minutes
|
||||
return {
|
||||
"add_dialogue_duration_time": add_duration / 1000 / 60,
|
||||
"search_memory_duration_time": search_duration / 1000 / 60,
|
||||
"total_duration_time": (add_duration + search_duration) / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
class HaluMemEvaluatorV3:
|
||||
"""Main evaluator orchestrating the entire ReMe V3 pipeline."""
|
||||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.reme = ReMe()
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.memory_processor = MemoryProcessor(self.reme)
|
||||
self.qa_evaluator = QuestionAnsweringEvaluator(
|
||||
self.memory_processor,
|
||||
config.top_k
|
||||
)
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
async def process_session(
|
||||
self,
|
||||
session: dict,
|
||||
session_id: int,
|
||||
user_name: str,
|
||||
uuid: str
|
||||
) -> dict:
|
||||
"""Process a single session using ReMe V3."""
|
||||
session_data = {
|
||||
"uuid": uuid,
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"memory_points": session["memory_points"]
|
||||
}
|
||||
|
||||
# Skip generated QA sessions
|
||||
if session.get("is_generated_qa_session", False):
|
||||
session_data["is_generated_qa_session"] = True
|
||||
return session_data
|
||||
|
||||
# Format and add dialogue to memory using V3
|
||||
dialogue = session["dialogue"]
|
||||
formatted_messages = self.data_loader.format_dialogue_messages(dialogue)
|
||||
|
||||
extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories(
|
||||
user_id=user_name,
|
||||
messages=formatted_messages,
|
||||
batch_size=self.config.batch_size
|
||||
)
|
||||
|
||||
session_data.update({
|
||||
"dialogue": dialogue,
|
||||
"extracted_memories": extracted_memories,
|
||||
"summary_messages": [m.model_dump() for m in agent_messages],
|
||||
"add_dialogue_duration_ms": duration_ms
|
||||
})
|
||||
|
||||
# Evaluate questions if present
|
||||
if "questions" in session:
|
||||
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
|
||||
qa_results = await self.qa_evaluator.evaluate_questions(
|
||||
questions=session["questions"],
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
session_id=session_id,
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
session_data["evaluation_results"] = {
|
||||
"question_answering_records": qa_results
|
||||
}
|
||||
|
||||
return session_data
|
||||
|
||||
async def process_user(self, user_data: dict) -> dict:
|
||||
"""Process all sessions for a user."""
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
uuid = user_data["uuid"]
|
||||
|
||||
logger.info(f"Processing user: {user_name}")
|
||||
|
||||
for idx, session in enumerate(user_data["sessions"]):
|
||||
logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}")
|
||||
|
||||
session_data = await self.process_session(
|
||||
session=session,
|
||||
session_id=idx,
|
||||
user_name=user_name,
|
||||
uuid=uuid
|
||||
)
|
||||
|
||||
self.file_manager.save_session(user_name, idx, session_data)
|
||||
|
||||
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the complete evaluation pipeline using ReMe V3."""
|
||||
start_time = time.time()
|
||||
|
||||
# Clear existing data
|
||||
await self.reme.vector_store.delete_all()
|
||||
|
||||
# Clear meta_memory directory
|
||||
meta_memory_path = Path(f"meta_memory/{self.reme.vector_store.collection_name}")
|
||||
if meta_memory_path.exists():
|
||||
shutil.rmtree(meta_memory_path)
|
||||
logger.info(f"Cleared meta_memory directory: {meta_memory_path}")
|
||||
meta_memory_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Load user data
|
||||
all_users = self.data_loader.load_jsonl(self.config.data_path)
|
||||
users_to_process = all_users[:self.config.user_num]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("HALUMEM EVALUATION - REME V3 - QUESTION ANSWERING")
|
||||
print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Process users with concurrency control
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
|
||||
async def process_with_cache_check(idx: int, user_data: dict):
|
||||
async with semaphore:
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
|
||||
# Check cache
|
||||
if self.file_manager.user_has_cache(user_name):
|
||||
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
|
||||
return {"user_name": user_name, "status": "cached"}
|
||||
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
result = await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
|
||||
return result
|
||||
|
||||
tasks = [
|
||||
process_with_cache_check(idx, user)
|
||||
for idx, user in enumerate(users_to_process, 1)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Combine results
|
||||
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
self.file_manager.combine_results(output_file)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
print(f"📁 Results: {output_file}\n")
|
||||
|
||||
# Aggregate metrics
|
||||
await self.aggregate_and_report(output_file)
|
||||
|
||||
async def aggregate_and_report(self, results_file: str):
|
||||
"""Aggregate results and generate final report."""
|
||||
print("=" * 80)
|
||||
print("AGGREGATING METRICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Collect all QA records
|
||||
qa_records = []
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
qa_records.extend(
|
||||
eval_results.get("question_answering_records", [])
|
||||
)
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
|
||||
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
# Save final report
|
||||
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
print(f"📊 Statistics saved to: {report_file}\n")
|
||||
|
||||
# Print summary
|
||||
self._print_summary(qa_metrics, time_metrics)
|
||||
|
||||
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
|
||||
"""Print evaluation summary."""
|
||||
print("=" * 80)
|
||||
print("EVALUATION SUMMARY - REME V3")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Question Answering:")
|
||||
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ Time Metrics:")
|
||||
print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min")
|
||||
print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min")
|
||||
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
# ==================== Entry Point ====================
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
top_k: int = 20,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2
|
||||
):
|
||||
"""Main entry point for ReMe V3 evaluation."""
|
||||
config = EvalConfig(
|
||||
data_path=data_path,
|
||||
top_k=top_k,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency
|
||||
)
|
||||
|
||||
evaluator = HaluMemEvaluatorV3(config)
|
||||
asyncio.run(evaluator.run_evaluation())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate ReMe V3 on HaluMem benchmark (Question Answering)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=20,
|
||||
help="Number of memories to retrieve (default: 20)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--user_num",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of users to evaluate (default: 1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Maximum concurrent user processing (default: 2)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
top_k=args.top_k,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency
|
||||
)
|
||||
|
|
@ -1,171 +0,0 @@
|
|||
"""Evaluation tools for ReMe HaluMem benchmark."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
|
||||
from llms import llm_request_for_json
|
||||
|
||||
# Load prompts from YAML file
|
||||
_YAML_PATH = Path(__file__).parent / "halumem.yaml"
|
||||
with open(_YAML_PATH, "r", encoding="utf-8") as f:
|
||||
_PROMPTS = yaml.safe_load(f)
|
||||
|
||||
|
||||
async def evaluation_for_memory_integrity(
|
||||
extract_memories: str,
|
||||
target_memory: str,
|
||||
):
|
||||
"""
|
||||
Memory Integrity Evaluation
|
||||
extract_memories: A formatted string concatenating all memory points extracted by the memory system under evaluation.
|
||||
target_memory: The target key memory point.
|
||||
"""
|
||||
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY"].format(
|
||||
memories=extract_memories,
|
||||
expected_memory_point=target_memory,
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_memory_accuracy(
|
||||
dialogue: str,
|
||||
golden_memories: str,
|
||||
candidate_memory: str,
|
||||
):
|
||||
"""
|
||||
Memory Accuracy Evaluation
|
||||
dialogue: The complete human-machine dialogue record.
|
||||
golden_memories: The core memory points for this dialogue segment in the evaluation set (the correct reference memories).
|
||||
candidate_memory: A specific memory point extracted by the memory system being evaluated.
|
||||
"""
|
||||
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_ACCURACY"].format(
|
||||
dialogue=dialogue,
|
||||
golden_memories=golden_memories,
|
||||
candidate_memory=candidate_memory,
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_update_memory(
|
||||
extract_memories: str,
|
||||
target_update_memory: str,
|
||||
original_memory: str,
|
||||
):
|
||||
"""
|
||||
Memory Update Evaluation
|
||||
extract_memories: A formatted string concatenating all memory points extracted by the memory system under evaluation.
|
||||
target_update_memory: The target updated memory point.
|
||||
original_memory: str: A formatted string concatenating all original memory points corresponding to the target updated memory point (i.e., all memories before the update).
|
||||
"""
|
||||
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_UPDATE_MEMORY"].format(
|
||||
memories=extract_memories,
|
||||
updated_memory=target_update_memory,
|
||||
original_memory=original_memory,
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_question(
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
):
|
||||
"""
|
||||
Question-Answering Evaluation
|
||||
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.
|
||||
"""
|
||||
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def evaluation_for_question2(
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
dialogue: str,
|
||||
):
|
||||
"""
|
||||
Question-Answering Evaluation with Dialogue Context (Version 2)
|
||||
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).
|
||||
"""
|
||||
|
||||
# prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format(
|
||||
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
dialogue=dialogue,
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt, model_name="qwen3-max")
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def answer_question_with_memories(
|
||||
question: str,
|
||||
memories: str,
|
||||
user_id: str = None,
|
||||
):
|
||||
"""
|
||||
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
|
||||
|
||||
Args:
|
||||
question: The question to answer
|
||||
memories: The retrieved memories (formatted as context)
|
||||
user_id: Optional user ID for context formatting
|
||||
|
||||
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 llm_request_for_json(prompt, model_name="qwen3-max")
|
||||
result = await llm_request_for_json(prompt, model_name="qwen3-30b-a3b-instruct-2507")
|
||||
|
||||
return result
|
||||
|
|
@ -1,548 +0,0 @@
|
|||
TEMPLATE_MEMOS: |
|
||||
Memories for user {user_id}:
|
||||
{memories}
|
||||
|
||||
PROMPT_MEMZERO_JSON: |
|
||||
# CONTEXT:
|
||||
{context}
|
||||
|
||||
# CONTEXT PRIORITY:
|
||||
When the context contains information from multiple sources, follow this strict priority order:
|
||||
1. **Historical Dialogue** (highest priority) - Direct conversation content
|
||||
2. **Extracted Memories** (medium priority) - Summarized memory points
|
||||
3. **User Profile** (lowest priority) - General user information
|
||||
|
||||
# Question:
|
||||
{question}
|
||||
|
||||
# OUTPUT FORMAT:
|
||||
Do not hallucinate; strictly answer the user's question based on the content of the CONTEXT.
|
||||
Please provide your response in the following JSON format:
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "reasoning content",
|
||||
"answer": "Provide a detailed answer"
|
||||
}}
|
||||
```
|
||||
|
||||
PROMPT_MEMZERO_JSON2: |
|
||||
# CONTEXT:
|
||||
{context}
|
||||
|
||||
# CONTEXT PRIORITY:
|
||||
When the context contains information from multiple sources, follow this strict priority order:
|
||||
1. **Historical Dialogue** (highest priority) - Direct conversation content
|
||||
2. **Extracted Memories** (medium priority) - Summarized memory points
|
||||
3. **User Profile** (lowest priority) - General user information
|
||||
|
||||
# Question:
|
||||
{question}
|
||||
|
||||
# OUTPUT FORMAT:
|
||||
Do not hallucinate; strictly answer the user's question based on the content of the CONTEXT.
|
||||
Please provide your response in the following JSON format:
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "reasoning content",
|
||||
"answer": "Provide a detailed answer"
|
||||
}}
|
||||
```
|
||||
|
||||
PROMPT_MEMZERO: |
|
||||
You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories.
|
||||
|
||||
# CONTEXT:
|
||||
You have access to memories from two speakers in a conversation. These memories contain
|
||||
timestamped information that may be relevant to answering the question.
|
||||
|
||||
# INSTRUCTIONS:
|
||||
1. Carefully analyze all provided memories from both speakers
|
||||
2. Pay special attention to the timestamps to determine the answer
|
||||
3. If the question asks about a specific event or fact, look for direct evidence in the memories
|
||||
4. If the memories contain contradictory information, prioritize the most recent memory
|
||||
5. If there is a question about time references (like "last year", "two months ago", etc.),
|
||||
calculate the actual date based on the memory timestamp. For example, if a memory from
|
||||
4 May 2022 mentions "went to India last year," then the trip occurred in 2021.
|
||||
6. Always convert relative time references to specific dates, months, or years. For example,
|
||||
convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory
|
||||
timestamp. Ignore the reference while answering the question.
|
||||
7. Focus only on the content of the memories from both speakers. Do not confuse character
|
||||
names mentioned in memories with the actual users who created those memories.
|
||||
8. The answer should be less than 5-6 words.
|
||||
|
||||
# APPROACH (Think step by step):
|
||||
1. First, examine all memories that contain information related to the question
|
||||
2. Examine the timestamps and content of these memories carefully
|
||||
3. Look for explicit mentions of dates, times, locations, or events that answer the question
|
||||
4. If the answer requires calculation (e.g., converting relative time references), show your work
|
||||
5. Formulate a precise, concise answer based solely on the evidence in the memories
|
||||
6. Double-check that your answer directly addresses the question asked
|
||||
7. Ensure your final answer is specific and avoids vague time references
|
||||
|
||||
{context}
|
||||
|
||||
Question: {question}
|
||||
|
||||
Answer:
|
||||
|
||||
PROMPT_ZEP: |
|
||||
You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories.
|
||||
|
||||
# CONTEXT:
|
||||
You have access to memories from a conversation. These memories contain
|
||||
timestamped information that may be relevant to answering the question.
|
||||
|
||||
# INSTRUCTIONS:
|
||||
1. Carefully analyze all provided memories
|
||||
2. Pay special attention to the timestamps to determine the answer
|
||||
3. If the question asks about a specific event or fact, look for direct evidence in the memories
|
||||
4. If the memories contain contradictory information, prioritize the most recent memory
|
||||
5. If there is a question about time references (like "last year", "two months ago", etc.),
|
||||
calculate the actual date based on the memory timestamp. For example, if a memory from
|
||||
4 May 2022 mentions "went to India last year," then the trip occurred in 2021.
|
||||
6. Always convert relative time references to specific dates, months, or years. For example,
|
||||
convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory
|
||||
timestamp. Ignore the reference while answering the question.
|
||||
7. Focus only on the content of the memories. Do not confuse character
|
||||
names mentioned in memories with the actual users who created those memories.
|
||||
8. The answer should be less than 5-6 words.
|
||||
|
||||
# APPROACH (Think step by step):
|
||||
1. First, examine all memories that contain information related to the question
|
||||
2. Examine the timestamps and content of these memories carefully
|
||||
3. Look for explicit mentions of dates, times, locations, or events that answer the question
|
||||
4. If the answer requires calculation (e.g., converting relative time references), show your work
|
||||
5. Formulate a precise, concise answer based solely on the evidence in the memories
|
||||
6. Double-check that your answer directly addresses the question asked
|
||||
7. Ensure your final answer is specific and avoids vague time references
|
||||
|
||||
Context:
|
||||
|
||||
{context}
|
||||
|
||||
Question: {question}
|
||||
Answer:
|
||||
|
||||
PROMPT_MEMOS: |
|
||||
You are a knowledgeable and helpful AI assistant.
|
||||
|
||||
# CONTEXT:
|
||||
You have access to memories from two speakers in a conversation. These memories contain
|
||||
timestamped information that may be relevant to answering the question.
|
||||
|
||||
# INSTRUCTIONS:
|
||||
1. Carefully analyze all provided memories. Synthesize information across different entries if needed to form a complete answer.
|
||||
2. Pay close attention to the timestamps to determine the answer. If memories contain contradictory information, the **most recent memory** is the source of truth.
|
||||
3. If the question asks about a specific event or fact, look for direct evidence in the memories.
|
||||
4. Your answer must be grounded in the memories. However, you may use general world knowledge to interpret or complete information found within a memory (e.g., identifying a landmark mentioned by description).
|
||||
5. If the question involves time references (like "last year", "two months ago", etc.), you **must** calculate the actual date based on the memory's timestamp. For example, if a memory from 4 May 2022 mentions "went to India last year," then the trip occurred in 2021.
|
||||
6. Always convert relative time references to specific dates, months, or years in your final answer.
|
||||
7. Do not confuse character names mentioned in memories with the actual users who created them.
|
||||
8. The answer must be brief (under 5-6 words) and direct, with no extra description.
|
||||
|
||||
# APPROACH (Think step by step):
|
||||
1. First, examine all memories that contain information related to the question.
|
||||
2. Synthesize findings from multiple memories if a single entry is insufficient.
|
||||
3. Examine timestamps and content carefully, looking for explicit dates, times, locations, or events.
|
||||
4. If the answer requires calculation (e.g., converting relative time references), perform the calculation.
|
||||
5. Formulate a precise, concise answer based on the evidence from the memories (and allowed world knowledge).
|
||||
6. Double-check that your answer directly addresses the question asked and adheres to all instructions.
|
||||
7. Ensure your final answer is specific and avoids vague time references.
|
||||
|
||||
{context}
|
||||
|
||||
Question: {question}
|
||||
|
||||
Answer:
|
||||
|
||||
PROMPT_MEMOBASE: |
|
||||
You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories.
|
||||
|
||||
# CONTEXT:
|
||||
You have access to memories from two speakers in a conversation. These memories contain
|
||||
timestamped information that may be relevant to answering the question.
|
||||
|
||||
# INSTRUCTIONS:
|
||||
1. Carefully analyze all provided memories from both speakers
|
||||
2. Pay special attention to the timestamps to determine the answer
|
||||
3. If the question asks about a specific event or fact, look for direct evidence in the memories
|
||||
4. If the memories contain contradictory information, prioritize the most recent memory
|
||||
5. If there is a question about time references (like "last year", "two months ago", etc.), calculate the actual date based on the memory timestamp. For example, if a memory from 4 May 2022 mentions "went to India last year," then the trip occurred in 2021.
|
||||
6. Always convert relative time references to specific dates, months, or years. For example, convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory timestamp. Ignore the reference while answering the question.
|
||||
7. Focus only on the content of the memories from both speakers. Do not confuse character names mentioned in memories with the actual users who created those memories.
|
||||
8. The answer should be less than 5-6 words.
|
||||
|
||||
# APPROACH (Think step by step):
|
||||
1. First, examine all memories that contain information related to the question
|
||||
2. Examine the timestamps and content of these memories carefully
|
||||
3. Look for explicit mentions of dates, times, locations, or events that answer the question
|
||||
4. If the answer requires calculation (e.g., converting relative time references), show your work
|
||||
5. Formulate a precise, concise answer based solely on the evidence in the memories
|
||||
6. Double-check that your answer directly addresses the question asked
|
||||
7. Ensure your final answer is specific and avoids vague time references
|
||||
|
||||
{context}
|
||||
|
||||
Question: {question}
|
||||
|
||||
Answer:
|
||||
|
||||
|
||||
EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY: |
|
||||
You are a strict **"Memory Integrity" evaluator**.
|
||||
Your core task is to assess whether an AI memory system has **missed any key memory points** after processing a conversation. This evaluation measures the system’s **memory integrity**, i.e., its ability to resist **amnesia** or **omission**.
|
||||
|
||||
# Evaluation Context & Data:
|
||||
|
||||
1. **Extracted Memories:**
|
||||
These are all the memory items actually extracted by the memory system.
|
||||
{memories}
|
||||
|
||||
2. **Expected Memory Point:**
|
||||
The key memory point that *should* have been extracted.
|
||||
{expected_memory_point}
|
||||
|
||||
# Evaluation Instructions:
|
||||
|
||||
1. For each **Expected Memory Point**, search within the **Extracted Memories** list for corresponding or related information. Ignore unrelated items.
|
||||
2. Based on the following scoring rubric, rate how well the memory system captured the **Expected Memory Point** and provide a detailed explanation.
|
||||
|
||||
# Scoring Rubric:
|
||||
|
||||
* **2:** Fully covered or implied.
|
||||
One or more items in “Extracted Memories” fully cover or logically imply all information in the “Expected Memory Point.”
|
||||
|
||||
* **1:** Partially covered or mentioned.
|
||||
Some information in “Extracted Memories” mentions part of the “Expected Memory Point,” but key information is missing, inaccurate, or slightly incorrect.
|
||||
|
||||
* **0:** Not mentioned or incorrect.
|
||||
“Extracted Memories” contains no mention of the “Expected Memory Point,” or the corresponding information is entirely wrong.
|
||||
|
||||
# Scoring Notes:
|
||||
|
||||
* For **compound Expected Memory Points** (with multiple elements such as person/event/time/location/preference, etc.):
|
||||
|
||||
* All elements correct → **2 points**
|
||||
* Some elements correct / uncertain → **1 point**
|
||||
* Key elements missing or wrong → **0 points**
|
||||
|
||||
* Semantic matching is acceptable; exact wording is **not** required.
|
||||
|
||||
* If “Extracted Memories” contains **conflicting information**, assign the **best possible coverage score** and mention the conflict in your reasoning.
|
||||
|
||||
* Extra or stylistically different memories do **not** reduce the score; only the coverage of the **Expected Memory Point** matters.
|
||||
|
||||
* For uncertain wording (“might,” “probably,” “tends to,” etc.):
|
||||
|
||||
* If the Expected Memory Point is a definite statement, usually assign **1 point**.
|
||||
|
||||
* If critical fields (e.g., time, entity name, relationship) are partly wrong but others match → **1 point**.
|
||||
|
||||
* If all key fields are wrong or missing → **0 points**.
|
||||
|
||||
# Output Format:
|
||||
|
||||
Please output your result in the following JSON format:
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise justification for the score",
|
||||
"score": "2|1|0"
|
||||
}}
|
||||
```
|
||||
|
||||
EVALUATION_PROMPT_FOR_MEMORY_ACCURACY: |
|
||||
You are a **Dialogue Memory Accuracy Evaluator.** Your task is to evaluate the **accuracy** of a memory extracted by an AI memory system, based on three given inputs: the dialogue content, the *target (gold)* memory points (the correct annotated memories), and the *candidate* memory to be evaluated. The goal is to output a **structured evaluation result**.
|
||||
|
||||
# Input Content
|
||||
|
||||
* **Dialogue:**
|
||||
{dialogue}
|
||||
|
||||
* **Golden Memories (Target Memory Points):**
|
||||
The correct memory points pre-annotated for this dialogue in the evaluation dataset.
|
||||
{golden_memories}
|
||||
|
||||
* **Candidate Memory:**
|
||||
The memory extracted by the system to be evaluated.
|
||||
{candidate_memory}
|
||||
|
||||
# Evaluation Principles and Definitions
|
||||
|
||||
### 1) Support / Entailment
|
||||
|
||||
* An **information point** (atomic fact) in the candidate memory is considered *supported* if it can be directly stated or semantically entailed (via synonym, paraphrase, or equivalent expression) by the *Dialogue* or *Golden Memories*.
|
||||
* Only the given dialogue and golden memories can be used for judgment — **no external knowledge** or assumptions are allowed.
|
||||
Any information not appearing in or inferable from these two sources is considered *unsupported*.
|
||||
* Pay careful attention to **negation**, **quantities**, **time**, and **subjects**.
|
||||
If the candidate statement contradicts the dialogue or golden memories, it is considered a **conflict**.
|
||||
|
||||
### 2) Memory Accuracy Score (integer: 0 / 1 / 2)
|
||||
|
||||
* **2 points:** Every information point in the candidate memory is supported by the dialogue or golden memories, with **no contradictions or hallucinations**.
|
||||
* **1 point:** The candidate memory is *partially correct* (at least one supported information point) but also includes *unsupported* or *contradictory* content.
|
||||
* **0 points:** The candidate memory is **entirely unsupported or contradictory** to the sources (i.e., a “hallucinated memory”).
|
||||
|
||||
> Note:
|
||||
>
|
||||
> * If a candidate memory contains multiple information points, **any unsupported or contradictory element** prevents a full score (2).
|
||||
> * If both supported and unsupported/conflicting content appear, assign a score of **1**.
|
||||
|
||||
### 3) Inclusion in Golden Memories (Boolean field-level judgment)
|
||||
|
||||
**Definition:**
|
||||
|
||||
* **Atomic information point:** the smallest factual unit in the candidate memory (e.g., *name = Li Si*, *age = 25*, *location = Beijing*, *preference = coffee*, *budget ≤ 2000*, *meeting_time = Wednesday 10:00*, *tool = Zoom*, etc.).
|
||||
* **Field / Slot:** the semantic dimension of an information point (e.g., *name*, *age*, *residence*, *food preference*, *budget*, *meeting time*, *meeting tool*, etc.).
|
||||
|
||||
**Judgment Rules (independent of correctness):**
|
||||
|
||||
* **true:**
|
||||
Every atomic information point in the candidate memory has a corresponding **field** in the golden memories (allowing for synonyms, paraphrases, or equivalent expressions; ignore value, polarity, or quantity differences).
|
||||
|
||||
* Note: A single field in the gold list may match multiple candidate points (e.g., multiple “drink preference” facts can be covered by one “drink preference” field in gold).
|
||||
* **false:**
|
||||
If **any** atomic information point’s field in the candidate memory cannot be found in the golden memories, mark as *false*.
|
||||
|
||||
**Important Notes:**
|
||||
|
||||
* Field matching is restricted to fields that are **explicitly present or semantically recognizable** in the golden memories — no external knowledge may be used to expand the field set.
|
||||
* Differences in **values** (e.g., “Zhang San” vs. “Li Si”), **polarity** (like/dislike), or **exact number/time** do **not** affect this Boolean judgment.
|
||||
|
||||
# Evaluation Procedure
|
||||
|
||||
For each candidate memory:
|
||||
|
||||
1. **Decompose** it into atomic information points (e.g., name, number, location, preference).
|
||||
2. For each information point, **search** the dialogue and golden memories for supporting or contradictory evidence.
|
||||
3. Assign the **accuracy_score** (0 / 1 / 2) according to the rules above.
|
||||
4. Determine **is_included_in_golden_memories (true/false)**:
|
||||
|
||||
* Identify each information point’s field;
|
||||
* If *all* fields exist in the golden memories, mark as *true*; otherwise, *false*.
|
||||
5. Provide a **concise Chinese explanation** in `"reason"`, citing key evidence (short excerpts allowed), and clearly state any unsupported or contradictory parts if applicable.
|
||||
|
||||
# Output Format (strictly required)
|
||||
|
||||
Output **only one JSON object**, with the following three fields:
|
||||
|
||||
* `"accuracy_score"`: `"0"` or `"1"` or `"2"`
|
||||
* `"is_included_in_golden_memories"`: `"true"` or `"false"`
|
||||
* `"reason"`: `"brief explanation in Chinese"`
|
||||
|
||||
Do **not** include any other text, explanation, or fields.
|
||||
Do **not** include the candidate memory text inside the JSON.
|
||||
|
||||
Please output **only** the following JSON (in a code block):
|
||||
|
||||
```json
|
||||
{{
|
||||
"accuracy_score": "2 | 1 | 0",
|
||||
"is_included_in_golden_memories": "true | false",
|
||||
"reason": "Brief explanation in Chinese"
|
||||
}}
|
||||
```
|
||||
|
||||
EVALUATION_PROMPT_FOR_UPDATE_MEMORY: |
|
||||
Your task is to **evaluate the update accuracy** of an AI memory system.
|
||||
Based on the information provided below, determine whether the system-generated **“Generated Memories”** correctly **includes** the **Target Memory for Update**.
|
||||
|
||||
# Background Information
|
||||
|
||||
The following information is provided for evaluation:
|
||||
|
||||
1. **Generated Memories:**
|
||||
This is the list of memory points generated by the system after the current dialogue.
|
||||
{memories}
|
||||
|
||||
2. **Target Memory for Update:**
|
||||
This is the correct, updated version of the memory point that should have been produced — the one we focus on in this evaluation.
|
||||
{updated_memory}
|
||||
|
||||
3. **Original Memory Content:**
|
||||
This is the original version of the target memory before the update.
|
||||
{original_memory}
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
Please make your judgment **strictly based on the content update of the “Target Memory for Update.”**
|
||||
Use the following categories:
|
||||
|
||||
### Correct Update
|
||||
|
||||
* **Generated Memories** **contains all information points** from the “Target Memory for Update,” accurately and completely reflecting the intended update.
|
||||
* **Key fields** (e.g., date, time, values, proper nouns, etc.) must match exactly.
|
||||
* The **original memory** is effectively replaced or marked as outdated.
|
||||
* Synonymous or slightly rephrased expressions are acceptable.
|
||||
|
||||
### Hallucinated Update
|
||||
|
||||
* **Factual error:** The **Generated Memories** includes a new memory related to the “Target Memory for Update,” but its content contains factual mistakes or contradictions compared to the correct update.
|
||||
|
||||
### Omitted Update
|
||||
|
||||
* **Completely omitted:** The **Generated Memories** contains no new memory related to the “Target Memory for Update.”
|
||||
* **Partially omitted:** A related new memory was generated in **Generated Memories**, but it **misses key information** that should have been included.
|
||||
|
||||
### Other
|
||||
|
||||
Used for update failures that do **not clearly fall** into the above categories of “Hallucination” or “Omission.”
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please return your evaluation strictly in the following JSON format and provide a concise explanation.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reason": "Briefly explain your reasoning here and why it fits this category.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission | Other"
|
||||
}}
|
||||
```
|
||||
|
||||
EVALUATION_PROMPT_FOR_QUESTION: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.”
|
||||
* It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.”
|
||||
* It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.”
|
||||
* When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**.
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the “Reference Answer.”
|
||||
* It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.”
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**.
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *“unknown”* (without guessing), it may be **Correct**.
|
||||
* The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
EVALUATION_PROMPT_FOR_QUESTION2: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
|
||||
Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer."
|
||||
* It contains **no contradictions** with the "Key Memory Points" or "Reference Answer."
|
||||
* **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they:
|
||||
- Do not contradict the Key Memory Points or Reference Answer
|
||||
- Do not change or mislead the core conclusion
|
||||
- Are reasonable additional context that the memory system may have retained from the conversation
|
||||
* The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points."
|
||||
* The response provides information that **directly contradicts** known facts from the Key Memory Points.
|
||||
* When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information:
|
||||
- Directly contradicts the Key Memory Points or Reference Answer
|
||||
- Changes or misleads the core conclusion in a way that makes the answer incorrect
|
||||
- Provides a definitive answer when the Reference Answer indicates uncertainty
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the "Reference Answer."
|
||||
* It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points."
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead).
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *"unknown"* (without guessing), it may be **Correct**.
|
||||
* **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points.
|
||||
* Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
|
|
@ -1,88 +0,0 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
|
||||
from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log
|
||||
|
||||
from reme_ai.core.schema import Message
|
||||
from reme_ai.core.utils import load_env
|
||||
from reme_ai.reme import ReMe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
load_env()
|
||||
|
||||
WAIT_TIME_LOWER = 1
|
||||
WAIT_TIME_UPPER = 60
|
||||
RETRY_TIMES = 5
|
||||
|
||||
# Use ReMe singleton's LLM instead of creating a separate instance
|
||||
reme = ReMe()
|
||||
|
||||
@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(prompt, model_name: str = "qwen3-max", **kwargs) -> str:
|
||||
"""Make an LLM request using ReMe's LLM with optional model override.
|
||||
|
||||
Args:
|
||||
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(prompt, model_name: str = "qwen-flash", **kwargs):
|
||||
# async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs):
|
||||
"""Make an LLM request expecting JSON response using ReMe's LLM.
|
||||
|
||||
Args:
|
||||
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:
|
||||
Parsed JSON object from the LLM response
|
||||
|
||||
Raises:
|
||||
ValueError: If no JSON block is found in the model output
|
||||
"""
|
||||
content = await llm_request(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__":
|
||||
r = asyncio.run(llm_request_for_json('hello? answer in ```json\n{"answer": "..."}```'))
|
||||
print(r)
|
||||
|
|
@ -1,237 +0,0 @@
|
|||
"""
|
||||
Compute Question Answering statistics from evaluation results in tmp directory.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = hallucination = omission = valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
valid += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
valid += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
valid += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"correct_qa_ratio(valid)": correct / valid if valid > 0 else 0,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0,
|
||||
"omission_qa_ratio(valid)": omission / valid if valid > 0 else 0,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
def compute_time_metrics(users_data: list[dict]) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
add_duration = search_duration = 0
|
||||
|
||||
for user_data in users_data:
|
||||
for session in user_data.get("sessions", []):
|
||||
add_duration += session.get("add_dialogue_duration_ms", 0)
|
||||
eval_results = session.get("session", {}).get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
search_duration += qa.get("search_duration_ms", 0)
|
||||
|
||||
return {
|
||||
"add_dialogue_duration_time": add_duration / 1000 / 60,
|
||||
"search_memory_duration_time": search_duration / 1000 / 60,
|
||||
"total_duration_time": (add_duration + search_duration) / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
def load_from_tmp_dir(tmp_dir: str) -> list[dict]:
|
||||
"""Load data from tmp directory."""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
# Try flat file structure first (conversation_{user}_session_{idx}.json)
|
||||
json_files = [f for f in tmp_path.iterdir() if f.is_file() and f.suffix == ".json"]
|
||||
|
||||
if json_files:
|
||||
# Group files by user
|
||||
users_dict = defaultdict(list)
|
||||
|
||||
for json_file in json_files:
|
||||
with open(json_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
user_name = session_data.get("user_name")
|
||||
if user_name:
|
||||
users_dict[user_name].append(session_data)
|
||||
|
||||
# Sort sessions by session_idx for each user
|
||||
users_data = []
|
||||
for user_name, sessions in users_dict.items():
|
||||
sessions_sorted = sorted(sessions, key=lambda s: s.get("session_idx", 0))
|
||||
if sessions_sorted:
|
||||
user_data = {
|
||||
"uuid": sessions_sorted[0].get("uuid"),
|
||||
"user_name": user_name,
|
||||
"sessions": []
|
||||
}
|
||||
for session_data in sessions_sorted:
|
||||
session_copy = session_data.copy()
|
||||
session_copy.pop("uuid", None)
|
||||
session_copy.pop("user_name", None)
|
||||
user_data["sessions"].append(session_copy)
|
||||
users_data.append(user_data)
|
||||
|
||||
return users_data
|
||||
|
||||
# Fallback to directory structure (user_name/session_{idx}.json)
|
||||
user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()]
|
||||
|
||||
users_data = []
|
||||
for user_dir in user_dirs:
|
||||
session_files = sorted(
|
||||
[f for f in user_dir.iterdir() if "session_" in f.name and f.suffix == ".json"],
|
||||
key=lambda f: int(f.stem.split("_")[-1])
|
||||
)
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
with open(session_files[0], "r", encoding="utf-8") as f:
|
||||
first_session = json.load(f)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
users_data.append(user_data)
|
||||
|
||||
return users_data
|
||||
|
||||
|
||||
def main(tmp_dir: str):
|
||||
"""Main function to compute statistics from tmp directory."""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
if not tmp_path.exists() or not tmp_path.is_dir():
|
||||
print(f"❌ Error: Directory not found: {tmp_dir}")
|
||||
return
|
||||
|
||||
# Load data from tmp directory
|
||||
users_data = load_from_tmp_dir(tmp_dir)
|
||||
|
||||
# Collect QA records with metadata
|
||||
qa_records = []
|
||||
qa_with_metadata = []
|
||||
user_count = session_count = 0
|
||||
|
||||
for user_data in users_data:
|
||||
user_count += 1
|
||||
user_name = user_data.get("user_name", "Unknown")
|
||||
|
||||
valid_session_idx = 0
|
||||
for session in user_data.get("sessions", []):
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
session_count += 1
|
||||
eval_results = session.get("session", {}).get("evaluation_results", {})
|
||||
|
||||
for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])):
|
||||
qa_records.append(qa)
|
||||
qa_with_metadata.append({
|
||||
"user_name": user_name,
|
||||
"session_idx": valid_session_idx,
|
||||
"question_idx": qa_idx,
|
||||
"qa_record": qa
|
||||
})
|
||||
|
||||
valid_session_idx += 1
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = compute_qa_metrics(qa_records)
|
||||
time_metrics = compute_time_metrics(users_data)
|
||||
|
||||
# Save results
|
||||
output_dir = tmp_path.parent
|
||||
report_file = output_dir / "reme_eval_stat_result.json"
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
# Print summary
|
||||
print(f"\n📊 Data: {user_count} users, {session_count} sessions, {len(qa_records)} QA records")
|
||||
print(f"\n✅ Metrics:")
|
||||
print(f" Correct: {qa_metrics['correct_qa_ratio(all)']:.4f} | Hallucination: {qa_metrics['hallucination_qa_ratio(all)']:.4f} | Omission: {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Valid: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
print(f"\n⏱️ Time: {time_metrics['total_duration_time']:.2f} min (Add: {time_metrics['add_dialogue_duration_time']:.2f} | Search: {time_metrics['search_memory_duration_time']:.2f})")
|
||||
print(f"\n💾 Results saved: {report_file}")
|
||||
|
||||
# Print error records
|
||||
print(f"\n{'='*80}\n❌ ERROR RECORDS ({len([r for r in qa_with_metadata if r['qa_record'].get('result_type') not in ['Correct', '']])} errors)\n{'='*80}")
|
||||
|
||||
error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]]
|
||||
|
||||
if error_records:
|
||||
for idx, record in enumerate(error_records, 1):
|
||||
qa = record["qa_record"]
|
||||
print(f"\n[{idx}] {qa.get('result_type', 'Unknown')} - {record['user_name']} (S{record['session_idx']}/Q{record['question_idx']})")
|
||||
print(f" Q: {qa.get('question', 'N/A')}")
|
||||
print(f" Expected: {qa.get('answer', 'N/A')}")
|
||||
print(f" Got: {qa.get('system_response', 'N/A')}")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Compute QA statistics from tmp directory")
|
||||
parser.add_argument(
|
||||
"tmp_dir",
|
||||
nargs='?',
|
||||
default="./data",
|
||||
type=str,
|
||||
help="Path to tmp directory containing user session data (default: ./data)")
|
||||
|
||||
args = parser.parse_args()
|
||||
main(tmp_dir=args.tmp_dir)
|
||||
|
|
@ -1,145 +0,0 @@
|
|||
EVALUATION_PROMPT_FOR_QUESTION: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.”
|
||||
* It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.”
|
||||
* It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.”
|
||||
* When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**.
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the “Reference Answer.”
|
||||
* It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.”
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**.
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *“unknown”* (without guessing), it may be **Correct**.
|
||||
* The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
EVALUATION_PROMPT_FOR_QUESTION2: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
|
||||
Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer."
|
||||
* It contains **no contradictions** with the "Key Memory Points" or "Reference Answer."
|
||||
* **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they:
|
||||
- Do not contradict the Key Memory Points or Reference Answer
|
||||
- Do not change or mislead the core conclusion
|
||||
- Are reasonable additional context that the memory system may have retained from the conversation
|
||||
* The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points."
|
||||
* The response provides information that **directly contradicts** known facts from the Key Memory Points.
|
||||
* When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information:
|
||||
- Directly contradicts the Key Memory Points or Reference Answer
|
||||
- Changes or misleads the core conclusion in a way that makes the answer incorrect
|
||||
- Provides a definitive answer when the Reference Answer indicates uncertainty
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the "Reference Answer."
|
||||
* It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points."
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead).
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *"unknown"* (without guessing), it may be **Correct**.
|
||||
* **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points.
|
||||
* Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
|
|
@ -1,550 +0,0 @@
|
|||
"""
|
||||
Re-evaluate Question Answering results from data directory using LLM.
|
||||
|
||||
This script:
|
||||
1. Loads existing QA records from data directory
|
||||
2. Re-evaluates each system_response using multiple models in parallel
|
||||
3. Uses both EVALUATION_PROMPT_FOR_QUESTION and EVALUATION_PROMPT_FOR_QUESTION2
|
||||
4. Saves updated results with new evaluation metrics
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import yaml
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from reme_ai.core.schema import Message
|
||||
from reme_ai.core.utils import load_env
|
||||
from reme_ai.reme import ReMe
|
||||
from tenacity import retry, stop_after_attempt, wait_random_exponential
|
||||
|
||||
# Load environment
|
||||
load_env()
|
||||
|
||||
# Initialize ReMe singleton
|
||||
reme = ReMe()
|
||||
|
||||
# Load prompts from YAML file
|
||||
_YAML_PATH = Path(__file__).parent / "eval.yaml"
|
||||
with open(_YAML_PATH, "r", encoding="utf-8") as f:
|
||||
_PROMPTS = yaml.safe_load(f)
|
||||
|
||||
|
||||
@retry(
|
||||
wait=wait_random_exponential(min=1, max=60),
|
||||
stop=stop_after_attempt(3),
|
||||
reraise=True,
|
||||
)
|
||||
async def llm_request(prompt: str, model_name: str = "qwen3-max", **kwargs) -> str:
|
||||
"""Make an LLM request using ReMe's LLM."""
|
||||
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=1, max=60),
|
||||
stop=stop_after_attempt(5),
|
||||
reraise=True,
|
||||
)
|
||||
async def llm_request_for_json(prompt: str, model_name: str = "qwen3-max", **kwargs) -> dict:
|
||||
"""Make an LLM request expecting JSON response."""
|
||||
content = await llm_request(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)
|
||||
|
||||
|
||||
async def evaluate_qa_record(
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
dialogue: str = "",
|
||||
model_name: str = "qwen3-max",
|
||||
prompt_version: str = "v1"
|
||||
) -> dict:
|
||||
"""Evaluate a single QA record using LLM with specified prompt version.
|
||||
|
||||
Args:
|
||||
question: The question to evaluate
|
||||
reference_answer: The reference answer
|
||||
key_memory_points: Key memory points
|
||||
response: System response to evaluate
|
||||
dialogue: Dialogue context (optional)
|
||||
model_name: LLM model name
|
||||
prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION,
|
||||
"v2" for EVALUATION_PROMPT_FOR_QUESTION2
|
||||
|
||||
Returns:
|
||||
dict with evaluation_result and reasoning
|
||||
"""
|
||||
# Select prompt template
|
||||
if prompt_version == "v2":
|
||||
prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"]
|
||||
else:
|
||||
prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"]
|
||||
|
||||
# Format prompt
|
||||
prompt = prompt_template.format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
dialogue=dialogue or "N/A"
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt, model_name=model_name)
|
||||
return result
|
||||
|
||||
|
||||
def load_from_data_dir(data_dir: str) -> list[dict]:
|
||||
"""Load data from data directory (same as compute_qa_stats.py)."""
|
||||
data_path = Path(data_dir)
|
||||
|
||||
# Try flat file structure first (conversation_{user}_session_{idx}.json)
|
||||
json_files = [f for f in data_path.iterdir() if f.is_file() and f.suffix == ".json"]
|
||||
|
||||
if json_files:
|
||||
# Group files by user
|
||||
users_dict = defaultdict(list)
|
||||
|
||||
for json_file in json_files:
|
||||
with open(json_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
user_name = session_data.get("user_name")
|
||||
if user_name:
|
||||
users_dict[user_name].append({
|
||||
"file": json_file,
|
||||
"data": session_data
|
||||
})
|
||||
|
||||
# Sort sessions by session_idx for each user
|
||||
users_data = []
|
||||
for user_name, sessions in users_dict.items():
|
||||
sessions_sorted = sorted(
|
||||
sessions,
|
||||
key=lambda s: s["data"].get("session_idx", 0)
|
||||
)
|
||||
users_data.extend(sessions_sorted)
|
||||
|
||||
return users_data
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def format_dialogue_context(session_data: dict) -> str:
|
||||
"""Format dialogue context from session data."""
|
||||
dialogue = session_data.get("session", {}).get("dialogue", [])
|
||||
if not dialogue:
|
||||
return "N/A"
|
||||
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
role = turn.get("role", "unknown")
|
||||
content = turn.get("content", "")
|
||||
timestamp = turn.get("timestamp", "")
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\nContent: {content}\nTime: {timestamp}"
|
||||
)
|
||||
return "\n\n".join(formatted_turns)
|
||||
|
||||
|
||||
async def reevaluate_session(
|
||||
session_file: Path,
|
||||
session_data: dict,
|
||||
models: list[str],
|
||||
prompt_versions: list[str],
|
||||
parallel: bool = True
|
||||
) -> dict:
|
||||
"""Re-evaluate all QA records in a session using multiple models and prompts.
|
||||
|
||||
Args:
|
||||
session_file: Path to session file
|
||||
session_data: Session data dict
|
||||
models: List of model names to use for evaluation
|
||||
prompt_versions: List of prompt versions ("v1", "v2")
|
||||
parallel: If True, use asyncio.gather for parallel execution;
|
||||
if False, execute sequentially
|
||||
|
||||
Returns:
|
||||
Updated session data with evaluation results for each model+prompt combination
|
||||
|
||||
Note:
|
||||
Request rate limiting is handled by base_llm.py's request_interval mechanism.
|
||||
"""
|
||||
eval_results = session_data.get("session", {}).get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
|
||||
if not qa_records:
|
||||
print(f" ⏭️ No QA records found")
|
||||
return session_data
|
||||
|
||||
total_evals = len(models) * len(prompt_versions) * len(qa_records)
|
||||
print(f" 🔍 Re-evaluating {len(qa_records)} QA records with {len(models)} models × {len(prompt_versions)} prompts = {total_evals} evaluations...")
|
||||
|
||||
# Format dialogue context once
|
||||
dialogue_context = format_dialogue_context(session_data)
|
||||
|
||||
async def evaluate_single_combination(
|
||||
idx: int,
|
||||
qa: dict,
|
||||
model_name: str,
|
||||
prompt_version: str
|
||||
) -> tuple[int, str, str, dict]:
|
||||
"""Evaluate a single QA record with specific model and prompt.
|
||||
|
||||
Note: Rate limiting is handled by BaseLLM's request_interval mechanism.
|
||||
"""
|
||||
question = qa.get("question", "")
|
||||
reference_answer = qa.get("answer", "")
|
||||
|
||||
# Get key memory points from evidence
|
||||
evidence = qa.get("evidence", [])
|
||||
key_memory_points = "\n".join([e.get("memory_content", "") for e in evidence])
|
||||
|
||||
# Get system response
|
||||
system_response = qa.get("system_response", "")
|
||||
|
||||
try:
|
||||
# Call LLM for evaluation
|
||||
eval_result = await evaluate_qa_record(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=system_response,
|
||||
dialogue=dialogue_context,
|
||||
model_name=model_name,
|
||||
prompt_version=prompt_version
|
||||
)
|
||||
|
||||
result = {
|
||||
"result_type": eval_result.get("evaluation_result", "Invalid"),
|
||||
"reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
|
||||
return idx, model_name, prompt_version, result
|
||||
|
||||
except Exception as e:
|
||||
print(f" ❌ QA[{idx+1}] {model_name}/{prompt_version}: Error: {e}")
|
||||
return idx, model_name, prompt_version, {
|
||||
"result_type": "Error",
|
||||
"reasoning": f"Evaluation error: {str(e)}"
|
||||
}
|
||||
|
||||
# Create all evaluation tasks (all combinations of models, prompts, and QA records)
|
||||
tasks = []
|
||||
for idx, qa in enumerate(qa_records):
|
||||
for model_name in models:
|
||||
for prompt_version in prompt_versions:
|
||||
tasks.append(evaluate_single_combination(idx, qa, model_name, prompt_version))
|
||||
|
||||
# Execute evaluations based on parallel mode
|
||||
if parallel:
|
||||
print(f" ⚡ Starting {len(tasks)} parallel evaluations (rate limited by LLM layer)...")
|
||||
results = await asyncio.gather(*tasks)
|
||||
else:
|
||||
print(f" 🔄 Starting {len(tasks)} sequential evaluations...")
|
||||
results = []
|
||||
for i, task in enumerate(tasks, 1):
|
||||
result = await task
|
||||
results.append(result)
|
||||
if i % 10 == 0 or i == len(tasks):
|
||||
print(f" ⏳ Progress: {i}/{len(tasks)} evaluations completed")
|
||||
|
||||
# Organize results by QA index, then by model and prompt
|
||||
# Structure: qa_records[idx]["evaluations"][model][prompt_version] = {result_type, reasoning}
|
||||
for idx, qa in enumerate(qa_records):
|
||||
if "evaluations" not in qa:
|
||||
qa["evaluations"] = {}
|
||||
|
||||
# Initialize evaluations structure
|
||||
for model_name in models:
|
||||
if model_name not in qa["evaluations"]:
|
||||
qa["evaluations"][model_name] = {}
|
||||
|
||||
# Fill in results
|
||||
completed_count = 0
|
||||
for qa_idx, model_name, prompt_version, result in results:
|
||||
qa_records[qa_idx]["evaluations"][model_name][prompt_version] = result
|
||||
completed_count += 1
|
||||
if completed_count % 10 == 0 or completed_count == len(results):
|
||||
print(f" ✅ Completed {completed_count}/{len(results)} evaluations")
|
||||
|
||||
# Set default result_type to first model's v1 result for compatibility
|
||||
if models and prompt_versions:
|
||||
default_model = models[0]
|
||||
default_prompt = prompt_versions[0]
|
||||
for qa in qa_records:
|
||||
default_eval = qa["evaluations"].get(default_model, {}).get(default_prompt, {})
|
||||
qa["result_type"] = default_eval.get("result_type", "Invalid")
|
||||
qa["question_answering_reasoning"] = default_eval.get("reasoning", "")
|
||||
|
||||
# Update session data
|
||||
if "session" not in session_data:
|
||||
session_data["session"] = {}
|
||||
if "evaluation_results" not in session_data["session"]:
|
||||
session_data["session"]["evaluation_results"] = {}
|
||||
|
||||
session_data["session"]["evaluation_results"]["question_answering_records"] = qa_records
|
||||
|
||||
# Save updated session data
|
||||
with open(session_file, "w", encoding="utf-8") as f:
|
||||
json.dump(session_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print(f" 💾 Updated session saved with all evaluations")
|
||||
|
||||
return session_data
|
||||
|
||||
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = hallucination = omission = valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
valid += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
valid += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
valid += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"correct_qa_ratio(valid)": correct / valid if valid > 0 else 0,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0,
|
||||
"omission_qa_ratio(valid)": omission / valid if valid > 0 else 0,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
async def main(
|
||||
data_dir: str = "./data",
|
||||
models: list[str] = None,
|
||||
prompt_versions: list[str] = None,
|
||||
parallel: bool = True
|
||||
):
|
||||
"""Main function to re-evaluate QA records from data directory with multiple models and prompts.
|
||||
|
||||
Args:
|
||||
data_dir: Path to data directory
|
||||
models: List of model names (e.g., ["qwen3-max", "qwen-flash"])
|
||||
prompt_versions: List of prompt versions (e.g., ["v1", "v2"])
|
||||
parallel: If True, use parallel execution; if False, use sequential execution
|
||||
|
||||
Note:
|
||||
Request rate limiting is automatically handled by base_llm.py's request_interval mechanism.
|
||||
"""
|
||||
data_path = Path(data_dir)
|
||||
|
||||
if not data_path.exists() or not data_path.is_dir():
|
||||
print(f"❌ Error: Directory not found: {data_dir}")
|
||||
return
|
||||
|
||||
# Default values
|
||||
if models is None:
|
||||
models = ["qwen3-max"]
|
||||
if prompt_versions is None:
|
||||
prompt_versions = ["v1"]
|
||||
|
||||
print("=" * 80)
|
||||
print("RE-EVALUATING QA RECORDS WITH MULTIPLE MODELS & PROMPTS")
|
||||
print(f"Models: {', '.join(models)}")
|
||||
print(f"Prompts: {', '.join(['EVALUATION_PROMPT_FOR_QUESTION' if v == 'v1' else 'EVALUATION_PROMPT_FOR_QUESTION2' for v in prompt_versions])}")
|
||||
print(f"Execution mode: {'Parallel' if parallel else 'Sequential'}")
|
||||
print("Note: Request rate limiting handled by LLM layer (base_llm.py)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Load data from directory
|
||||
sessions = load_from_data_dir(data_dir)
|
||||
|
||||
if not sessions:
|
||||
print(f"❌ No session files found in {data_dir}")
|
||||
return
|
||||
|
||||
print(f"📂 Found {len(sessions)} session files\n")
|
||||
|
||||
# Process each session
|
||||
all_qa_records = []
|
||||
|
||||
for idx, session_info in enumerate(sessions, 1):
|
||||
session_file = session_info["file"]
|
||||
session_data = session_info["data"]
|
||||
user_name = session_data.get("user_name", "Unknown")
|
||||
session_idx = session_data.get("session_idx", 0)
|
||||
|
||||
print(f"[{idx}/{len(sessions)}] {user_name} - Session {session_idx}")
|
||||
|
||||
updated_session = await reevaluate_session(
|
||||
session_file=session_file,
|
||||
session_data=session_data,
|
||||
models=models,
|
||||
prompt_versions=prompt_versions,
|
||||
parallel=parallel
|
||||
)
|
||||
|
||||
# Collect QA records for metrics
|
||||
eval_results = updated_session.get("session", {}).get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
all_qa_records.extend(qa_records)
|
||||
|
||||
print()
|
||||
|
||||
# Compute and display metrics for each model+prompt combination
|
||||
print("=" * 80)
|
||||
print("UPDATED METRICS (BY MODEL & PROMPT)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for model_name in models:
|
||||
for prompt_version in prompt_versions:
|
||||
prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2"
|
||||
print(f"\n📊 {model_name} / {prompt_name}:")
|
||||
print("─" * 80)
|
||||
|
||||
# Extract QA records for this model+prompt combination
|
||||
model_qa_records = []
|
||||
for qa in all_qa_records:
|
||||
eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {})
|
||||
if eval_data:
|
||||
# Create a copy with the specific evaluation result
|
||||
qa_copy = {
|
||||
**qa,
|
||||
"result_type": eval_data.get("result_type", "Invalid"),
|
||||
"question_answering_reasoning": eval_data.get("reasoning", "")
|
||||
}
|
||||
model_qa_records.append(qa_copy)
|
||||
|
||||
if model_qa_records:
|
||||
metrics = compute_qa_metrics(model_qa_records)
|
||||
|
||||
print(f" Correct (all): {metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {metrics['qa_valid_num']}/{metrics['qa_num']}")
|
||||
|
||||
# Save detailed results with all evaluations
|
||||
report_file = data_path.parent / "reme_eval_stat_result_detailed.json"
|
||||
|
||||
# Create summary for each model+prompt combination
|
||||
evaluation_summary = {}
|
||||
for model_name in models:
|
||||
evaluation_summary[model_name] = {}
|
||||
for prompt_version in prompt_versions:
|
||||
prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2"
|
||||
|
||||
# Extract QA records for this combination
|
||||
model_qa_records = []
|
||||
for qa in all_qa_records:
|
||||
eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {})
|
||||
if eval_data:
|
||||
qa_copy = {
|
||||
**qa,
|
||||
"result_type": eval_data.get("result_type", "Invalid"),
|
||||
"question_answering_reasoning": eval_data.get("reasoning", "")
|
||||
}
|
||||
model_qa_records.append(qa_copy)
|
||||
|
||||
metrics = compute_qa_metrics(model_qa_records)
|
||||
evaluation_summary[model_name][prompt_name] = {
|
||||
"metrics": metrics,
|
||||
"qa_records": model_qa_records
|
||||
}
|
||||
|
||||
final_results = {
|
||||
"evaluation_summary": evaluation_summary,
|
||||
"all_qa_records_with_evaluations": all_qa_records
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print(f"\n💾 Detailed results saved: {report_file}")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Re-evaluate QA records from data directory using multiple models and prompts in parallel. "
|
||||
"Request rate limiting is automatically handled by base_llm.py's request_interval mechanism."
|
||||
)
|
||||
parser.add_argument(
|
||||
"data_dir",
|
||||
nargs='?',
|
||||
default="./data",
|
||||
type=str,
|
||||
help="Path to data directory containing user session files (default: ./data)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models",
|
||||
type=str,
|
||||
nargs='+',
|
||||
default=["gpt-5.1-2025-11-13", "gemini-3-pro-preview"],
|
||||
help="LLM model names for evaluation (space-separated, default: qwen3-max)"
|
||||
)
|
||||
# ["gpt-4o-2024-11-20"], ["qwen3-max", "qwen-flash", "qwen3-30b-a3b-instruct-2507", "qwen3-30b-a3b-instruct-2507"]
|
||||
parser.add_argument(
|
||||
"--prompts",
|
||||
type=str,
|
||||
nargs='+',
|
||||
choices=["v1", "v2"],
|
||||
default=["v1", "v2"],
|
||||
help="Prompt versions to use: v1=EVALUATION_PROMPT_FOR_QUESTION, v2=EVALUATION_PROMPT_FOR_QUESTION2 (default: v1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--serial",
|
||||
action="store_true",
|
||||
help="Use sequential execution instead of parallel (default: parallel)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
asyncio.run(main(
|
||||
data_dir=args.data_dir,
|
||||
models=args.models,
|
||||
prompt_versions=args.prompts,
|
||||
parallel=not args.serial
|
||||
))
|
||||
|
|
@ -1,237 +0,0 @@
|
|||
"""
|
||||
Compute Question Answering statistics from evaluation results in tmp directory.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = hallucination = omission = valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
valid += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
valid += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
valid += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"correct_qa_ratio(valid)": correct / valid if valid > 0 else 0,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0,
|
||||
"omission_qa_ratio(valid)": omission / valid if valid > 0 else 0,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
def compute_time_metrics(users_data: list[dict]) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
add_duration = search_duration = 0
|
||||
|
||||
for user_data in users_data:
|
||||
for session in user_data.get("sessions", []):
|
||||
add_duration += session.get("add_dialogue_duration_ms", 0)
|
||||
eval_results = session.get("session", {}).get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
search_duration += qa.get("search_duration_ms", 0)
|
||||
|
||||
return {
|
||||
"add_dialogue_duration_time": add_duration / 1000 / 60,
|
||||
"search_memory_duration_time": search_duration / 1000 / 60,
|
||||
"total_duration_time": (add_duration + search_duration) / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
def load_from_tmp_dir(tmp_dir: str) -> list[dict]:
|
||||
"""Load data from tmp directory."""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
# Try flat file structure first (conversation_{user}_session_{idx}.json)
|
||||
json_files = [f for f in tmp_path.iterdir() if f.is_file() and f.suffix == ".json"]
|
||||
|
||||
if json_files:
|
||||
# Group files by user
|
||||
users_dict = defaultdict(list)
|
||||
|
||||
for json_file in json_files:
|
||||
with open(json_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
user_name = session_data.get("user_name")
|
||||
if user_name:
|
||||
users_dict[user_name].append(session_data)
|
||||
|
||||
# Sort sessions by session_idx for each user
|
||||
users_data = []
|
||||
for user_name, sessions in users_dict.items():
|
||||
sessions_sorted = sorted(sessions, key=lambda s: s.get("session_idx", 0))
|
||||
if sessions_sorted:
|
||||
user_data = {
|
||||
"uuid": sessions_sorted[0].get("uuid"),
|
||||
"user_name": user_name,
|
||||
"sessions": []
|
||||
}
|
||||
for session_data in sessions_sorted:
|
||||
session_copy = session_data.copy()
|
||||
session_copy.pop("uuid", None)
|
||||
session_copy.pop("user_name", None)
|
||||
user_data["sessions"].append(session_copy)
|
||||
users_data.append(user_data)
|
||||
|
||||
return users_data
|
||||
|
||||
# Fallback to directory structure (user_name/session_{idx}.json)
|
||||
user_dirs = [d for d in tmp_path.iterdir() if d.is_dir()]
|
||||
|
||||
users_data = []
|
||||
for user_dir in user_dirs:
|
||||
session_files = sorted(
|
||||
[f for f in user_dir.iterdir() if "session_" in f.name and f.suffix == ".json"],
|
||||
key=lambda f: int(f.stem.split("_")[-1])
|
||||
)
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
with open(session_files[0], "r", encoding="utf-8") as f:
|
||||
first_session = json.load(f)
|
||||
|
||||
user_data = {
|
||||
"uuid": first_session["uuid"],
|
||||
"user_name": first_session["user_name"],
|
||||
"sessions": []
|
||||
}
|
||||
|
||||
for session_file in session_files:
|
||||
with open(session_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
session_data.pop("uuid", None)
|
||||
session_data.pop("user_name", None)
|
||||
user_data["sessions"].append(session_data)
|
||||
|
||||
users_data.append(user_data)
|
||||
|
||||
return users_data
|
||||
|
||||
|
||||
def main(tmp_dir: str):
|
||||
"""Main function to compute statistics from tmp directory."""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
if not tmp_path.exists() or not tmp_path.is_dir():
|
||||
print(f"❌ Error: Directory not found: {tmp_dir}")
|
||||
return
|
||||
|
||||
# Load data from tmp directory
|
||||
users_data = load_from_tmp_dir(tmp_dir)
|
||||
|
||||
# Collect QA records with metadata
|
||||
qa_records = []
|
||||
qa_with_metadata = []
|
||||
user_count = session_count = 0
|
||||
|
||||
for user_data in users_data:
|
||||
user_count += 1
|
||||
user_name = user_data.get("user_name", "Unknown")
|
||||
|
||||
valid_session_idx = 0
|
||||
for session in user_data.get("sessions", []):
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
session_count += 1
|
||||
eval_results = session.get("session", {}).get("evaluation_results", {})
|
||||
|
||||
for qa_idx, qa in enumerate(eval_results.get("question_answering_records", [])):
|
||||
qa_records.append(qa)
|
||||
qa_with_metadata.append({
|
||||
"user_name": user_name,
|
||||
"session_idx": valid_session_idx,
|
||||
"question_idx": qa_idx,
|
||||
"qa_record": qa
|
||||
})
|
||||
|
||||
valid_session_idx += 1
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = compute_qa_metrics(qa_records)
|
||||
time_metrics = compute_time_metrics(users_data)
|
||||
|
||||
# Save results
|
||||
output_dir = tmp_path.parent
|
||||
report_file = output_dir / "reme_eval_stat_result.json"
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
# Print summary
|
||||
print(f"\n📊 Data: {user_count} users, {session_count} sessions, {len(qa_records)} QA records")
|
||||
print(f"\n✅ Metrics:")
|
||||
print(f" Correct: {qa_metrics['correct_qa_ratio(all)']:.4f} | Hallucination: {qa_metrics['hallucination_qa_ratio(all)']:.4f} | Omission: {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Valid: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
print(f"\n⏱️ Time: {time_metrics['total_duration_time']:.2f} min (Add: {time_metrics['add_dialogue_duration_time']:.2f} | Search: {time_metrics['search_memory_duration_time']:.2f})")
|
||||
print(f"\n💾 Results saved: {report_file}")
|
||||
|
||||
# Print error records
|
||||
print(f"\n{'='*80}\n❌ ERROR RECORDS ({len([r for r in qa_with_metadata if r['qa_record'].get('result_type') not in ['Correct', '']])} errors)\n{'='*80}")
|
||||
|
||||
error_records = [r for r in qa_with_metadata if r["qa_record"].get("result_type") not in ["Correct", ""]]
|
||||
|
||||
if error_records:
|
||||
for idx, record in enumerate(error_records, 1):
|
||||
qa = record["qa_record"]
|
||||
print(f"\n[{idx}] {qa.get('result_type', 'Unknown')} - {record['user_name']} (S{record['session_idx']}/Q{record['question_idx']})")
|
||||
print(f" Q: {qa.get('question', 'N/A')}")
|
||||
print(f" Expected: {qa.get('answer', 'N/A')}")
|
||||
print(f" Got: {qa.get('system_response', 'N/A')}")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="Compute QA statistics from tmp directory")
|
||||
parser.add_argument(
|
||||
"tmp_dir",
|
||||
nargs='?',
|
||||
default="./data",
|
||||
type=str,
|
||||
help="Path to tmp directory containing user session data (default: ./data)")
|
||||
|
||||
args = parser.parse_args()
|
||||
main(tmp_dir=args.tmp_dir)
|
||||
|
|
@ -1,145 +0,0 @@
|
|||
EVALUATION_PROMPT_FOR_QUESTION: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.”
|
||||
* It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.”
|
||||
* It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.”
|
||||
* When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**.
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the “Reference Answer.”
|
||||
* It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.”
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**.
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *“unknown”* (without guessing), it may be **Correct**.
|
||||
* The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
|
||||
|
||||
EVALUATION_PROMPT_FOR_QUESTION2: |
|
||||
You are an **evaluation expert for AI memory system question answering**.
|
||||
|
||||
Based **only** on the provided **"Question"**, **"Reference Answer"**, and **"Key Memory Points"** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **"Memory System Response."** Classify it as one of **"Correct"**, **"Hallucination"**, or **"Omission."** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
|
||||
|
||||
# Evaluation Criteria
|
||||
|
||||
## Answer Type Classification
|
||||
|
||||
### 1. Correct
|
||||
|
||||
* The "Memory System Response" accurately answers the "Question," and its content is **semantically equivalent** to the "Reference Answer."
|
||||
* It contains **no contradictions** with the "Key Memory Points" or "Reference Answer."
|
||||
* **Extra details not present in the Key Memory Points are allowed and should not be penalized**, as long as they:
|
||||
- Do not contradict the Key Memory Points or Reference Answer
|
||||
- Do not change or mislead the core conclusion
|
||||
- Are reasonable additional context that the memory system may have retained from the conversation
|
||||
* The memory system may have stored additional information beyond the Key Memory Points. Such extra information should be treated as **supplementary context** rather than hallucination, provided it does not conflict with the core answer.
|
||||
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
|
||||
|
||||
### 2. Hallucination
|
||||
|
||||
* The "Memory System Response" includes information or facts that **contradict or are inconsistent** with the "Reference Answer" or the "Key Memory Points."
|
||||
* The response provides information that **directly contradicts** known facts from the Key Memory Points.
|
||||
* When the "Reference Answer" is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
|
||||
* **Important:** Extra information that is NOT in Key Memory Points is **NOT automatically a hallucination**. Only classify as hallucination if the extra information:
|
||||
- Directly contradicts the Key Memory Points or Reference Answer
|
||||
- Changes or misleads the core conclusion in a way that makes the answer incorrect
|
||||
- Provides a definitive answer when the Reference Answer indicates uncertainty
|
||||
|
||||
### 3. Omission
|
||||
|
||||
* The response is **incomplete** compared to the "Reference Answer."
|
||||
* It explicitly states "don't know," "can't remember," or "no related memory," even though relevant information exists in the "Key Memory Points."
|
||||
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
|
||||
|
||||
## Priority Rules (Conflict Handling)
|
||||
|
||||
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
|
||||
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
|
||||
* If the core answer is correct and complete, classify as **Correct** even if there are extra details not in Key Memory Points (as long as they don't contradict or mislead).
|
||||
|
||||
## Detailed Guidelines and Tolerance
|
||||
|
||||
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
|
||||
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
|
||||
* If the reference answer is *"unknown / cannot be determined"* and the system provides a definite fact, that is a **Hallucination**.
|
||||
If the system also answers *"unknown"* (without guessing), it may be **Correct**.
|
||||
* **Focus on evaluating whether the core answer to the question is correct**, not whether the response is limited to only the Key Memory Points.
|
||||
* Extra contextual information (e.g., additional preferences, related details) should be viewed as enrichment, not as errors, unless they contradict or mislead.
|
||||
|
||||
# Information for Evaluation
|
||||
|
||||
* **Question:**
|
||||
{question}
|
||||
|
||||
* **Reference Answer:**
|
||||
{reference_answer}
|
||||
|
||||
* **Key Memory Points:**
|
||||
{key_memory_points}
|
||||
|
||||
* **Memory System Response:**
|
||||
{response}
|
||||
|
||||
# Output Requirements
|
||||
|
||||
Please provide your evaluation result **strictly** in the JSON format below.
|
||||
Do **not** add any extra explanation or comments outside the JSON block.
|
||||
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Provide a concise and traceable evaluation rationale: first verify that the system's response correctly includes all required elements from the Reference Answer, then check if any information contradicts the Key Memory Points or Reference Answer. Extra details not in Key Memory Points should be noted but not penalized unless they contradict or mislead. Finally state the classification basis.",
|
||||
"evaluation_result": "Correct | Hallucination | Omission"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
|
|
@ -1,550 +0,0 @@
|
|||
"""
|
||||
Re-evaluate Question Answering results from data directory using LLM.
|
||||
|
||||
This script:
|
||||
1. Loads existing QA records from data directory
|
||||
2. Re-evaluates each system_response using multiple models in parallel
|
||||
3. Uses both EVALUATION_PROMPT_FOR_QUESTION and EVALUATION_PROMPT_FOR_QUESTION2
|
||||
4. Saves updated results with new evaluation metrics
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import yaml
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from reme_ai.core.schema import Message
|
||||
from reme_ai.core.utils import load_env
|
||||
from reme_ai.reme import ReMe
|
||||
from tenacity import retry, stop_after_attempt, wait_random_exponential
|
||||
|
||||
# Load environment
|
||||
load_env()
|
||||
|
||||
# Initialize ReMe singleton
|
||||
reme = ReMe()
|
||||
|
||||
# Load prompts from YAML file
|
||||
_YAML_PATH = Path(__file__).parent / "eval.yaml"
|
||||
with open(_YAML_PATH, "r", encoding="utf-8") as f:
|
||||
_PROMPTS = yaml.safe_load(f)
|
||||
|
||||
|
||||
@retry(
|
||||
wait=wait_random_exponential(min=1, max=60),
|
||||
stop=stop_after_attempt(3),
|
||||
reraise=True,
|
||||
)
|
||||
async def llm_request(prompt: str, model_name: str = "qwen3-max", **kwargs) -> str:
|
||||
"""Make an LLM request using ReMe's LLM."""
|
||||
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=1, max=60),
|
||||
stop=stop_after_attempt(5),
|
||||
reraise=True,
|
||||
)
|
||||
async def llm_request_for_json(prompt: str, model_name: str = "qwen3-max", **kwargs) -> dict:
|
||||
"""Make an LLM request expecting JSON response."""
|
||||
content = await llm_request(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)
|
||||
|
||||
|
||||
async def evaluate_qa_record(
|
||||
question: str,
|
||||
reference_answer: str,
|
||||
key_memory_points: str,
|
||||
response: str,
|
||||
dialogue: str = "",
|
||||
model_name: str = "qwen3-max",
|
||||
prompt_version: str = "v1"
|
||||
) -> dict:
|
||||
"""Evaluate a single QA record using LLM with specified prompt version.
|
||||
|
||||
Args:
|
||||
question: The question to evaluate
|
||||
reference_answer: The reference answer
|
||||
key_memory_points: Key memory points
|
||||
response: System response to evaluate
|
||||
dialogue: Dialogue context (optional)
|
||||
model_name: LLM model name
|
||||
prompt_version: "v1" for EVALUATION_PROMPT_FOR_QUESTION,
|
||||
"v2" for EVALUATION_PROMPT_FOR_QUESTION2
|
||||
|
||||
Returns:
|
||||
dict with evaluation_result and reasoning
|
||||
"""
|
||||
# Select prompt template
|
||||
if prompt_version == "v2":
|
||||
prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"]
|
||||
else:
|
||||
prompt_template = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"]
|
||||
|
||||
# Format prompt
|
||||
prompt = prompt_template.format(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=response,
|
||||
dialogue=dialogue or "N/A"
|
||||
)
|
||||
|
||||
result = await llm_request_for_json(prompt, model_name=model_name)
|
||||
return result
|
||||
|
||||
|
||||
def load_from_data_dir(data_dir: str) -> list[dict]:
|
||||
"""Load data from data directory (same as compute_qa_stats.py)."""
|
||||
data_path = Path(data_dir)
|
||||
|
||||
# Try flat file structure first (conversation_{user}_session_{idx}.json)
|
||||
json_files = [f for f in data_path.iterdir() if f.is_file() and f.suffix == ".json"]
|
||||
|
||||
if json_files:
|
||||
# Group files by user
|
||||
users_dict = defaultdict(list)
|
||||
|
||||
for json_file in json_files:
|
||||
with open(json_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
user_name = session_data.get("user_name")
|
||||
if user_name:
|
||||
users_dict[user_name].append({
|
||||
"file": json_file,
|
||||
"data": session_data
|
||||
})
|
||||
|
||||
# Sort sessions by session_idx for each user
|
||||
users_data = []
|
||||
for user_name, sessions in users_dict.items():
|
||||
sessions_sorted = sorted(
|
||||
sessions,
|
||||
key=lambda s: s["data"].get("session_idx", 0)
|
||||
)
|
||||
users_data.extend(sessions_sorted)
|
||||
|
||||
return users_data
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def format_dialogue_context(session_data: dict) -> str:
|
||||
"""Format dialogue context from session data."""
|
||||
dialogue = session_data.get("session", {}).get("dialogue", [])
|
||||
if not dialogue:
|
||||
return "N/A"
|
||||
|
||||
formatted_turns = []
|
||||
for turn in dialogue:
|
||||
role = turn.get("role", "unknown")
|
||||
content = turn.get("content", "")
|
||||
timestamp = turn.get("timestamp", "")
|
||||
formatted_turns.append(
|
||||
f"Role: {role}\nContent: {content}\nTime: {timestamp}"
|
||||
)
|
||||
return "\n\n".join(formatted_turns)
|
||||
|
||||
|
||||
async def reevaluate_session(
|
||||
session_file: Path,
|
||||
session_data: dict,
|
||||
models: list[str],
|
||||
prompt_versions: list[str],
|
||||
parallel: bool = True
|
||||
) -> dict:
|
||||
"""Re-evaluate all QA records in a session using multiple models and prompts.
|
||||
|
||||
Args:
|
||||
session_file: Path to session file
|
||||
session_data: Session data dict
|
||||
models: List of model names to use for evaluation
|
||||
prompt_versions: List of prompt versions ("v1", "v2")
|
||||
parallel: If True, use asyncio.gather for parallel execution;
|
||||
if False, execute sequentially
|
||||
|
||||
Returns:
|
||||
Updated session data with evaluation results for each model+prompt combination
|
||||
|
||||
Note:
|
||||
Request rate limiting is handled by base_llm.py's request_interval mechanism.
|
||||
"""
|
||||
eval_results = session_data.get("session", {}).get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
|
||||
if not qa_records:
|
||||
print(f" ⏭️ No QA records found")
|
||||
return session_data
|
||||
|
||||
total_evals = len(models) * len(prompt_versions) * len(qa_records)
|
||||
print(f" 🔍 Re-evaluating {len(qa_records)} QA records with {len(models)} models × {len(prompt_versions)} prompts = {total_evals} evaluations...")
|
||||
|
||||
# Format dialogue context once
|
||||
dialogue_context = format_dialogue_context(session_data)
|
||||
|
||||
async def evaluate_single_combination(
|
||||
idx: int,
|
||||
qa: dict,
|
||||
model_name: str,
|
||||
prompt_version: str
|
||||
) -> tuple[int, str, str, dict]:
|
||||
"""Evaluate a single QA record with specific model and prompt.
|
||||
|
||||
Note: Rate limiting is handled by BaseLLM's request_interval mechanism.
|
||||
"""
|
||||
question = qa.get("question", "")
|
||||
reference_answer = qa.get("answer", "")
|
||||
|
||||
# Get key memory points from evidence
|
||||
evidence = qa.get("evidence", [])
|
||||
key_memory_points = "\n".join([e.get("memory_content", "") for e in evidence])
|
||||
|
||||
# Get system response
|
||||
system_response = qa.get("system_response", "")
|
||||
|
||||
try:
|
||||
# Call LLM for evaluation
|
||||
eval_result = await evaluate_qa_record(
|
||||
question=question,
|
||||
reference_answer=reference_answer,
|
||||
key_memory_points=key_memory_points,
|
||||
response=system_response,
|
||||
dialogue=dialogue_context,
|
||||
model_name=model_name,
|
||||
prompt_version=prompt_version
|
||||
)
|
||||
|
||||
result = {
|
||||
"result_type": eval_result.get("evaluation_result", "Invalid"),
|
||||
"reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
|
||||
return idx, model_name, prompt_version, result
|
||||
|
||||
except Exception as e:
|
||||
print(f" ❌ QA[{idx+1}] {model_name}/{prompt_version}: Error: {e}")
|
||||
return idx, model_name, prompt_version, {
|
||||
"result_type": "Error",
|
||||
"reasoning": f"Evaluation error: {str(e)}"
|
||||
}
|
||||
|
||||
# Create all evaluation tasks (all combinations of models, prompts, and QA records)
|
||||
tasks = []
|
||||
for idx, qa in enumerate(qa_records):
|
||||
for model_name in models:
|
||||
for prompt_version in prompt_versions:
|
||||
tasks.append(evaluate_single_combination(idx, qa, model_name, prompt_version))
|
||||
|
||||
# Execute evaluations based on parallel mode
|
||||
if parallel:
|
||||
print(f" ⚡ Starting {len(tasks)} parallel evaluations (rate limited by LLM layer)...")
|
||||
results = await asyncio.gather(*tasks)
|
||||
else:
|
||||
print(f" 🔄 Starting {len(tasks)} sequential evaluations...")
|
||||
results = []
|
||||
for i, task in enumerate(tasks, 1):
|
||||
result = await task
|
||||
results.append(result)
|
||||
if i % 10 == 0 or i == len(tasks):
|
||||
print(f" ⏳ Progress: {i}/{len(tasks)} evaluations completed")
|
||||
|
||||
# Organize results by QA index, then by model and prompt
|
||||
# Structure: qa_records[idx]["evaluations"][model][prompt_version] = {result_type, reasoning}
|
||||
for idx, qa in enumerate(qa_records):
|
||||
if "evaluations" not in qa:
|
||||
qa["evaluations"] = {}
|
||||
|
||||
# Initialize evaluations structure
|
||||
for model_name in models:
|
||||
if model_name not in qa["evaluations"]:
|
||||
qa["evaluations"][model_name] = {}
|
||||
|
||||
# Fill in results
|
||||
completed_count = 0
|
||||
for qa_idx, model_name, prompt_version, result in results:
|
||||
qa_records[qa_idx]["evaluations"][model_name][prompt_version] = result
|
||||
completed_count += 1
|
||||
if completed_count % 10 == 0 or completed_count == len(results):
|
||||
print(f" ✅ Completed {completed_count}/{len(results)} evaluations")
|
||||
|
||||
# Set default result_type to first model's v1 result for compatibility
|
||||
if models and prompt_versions:
|
||||
default_model = models[0]
|
||||
default_prompt = prompt_versions[0]
|
||||
for qa in qa_records:
|
||||
default_eval = qa["evaluations"].get(default_model, {}).get(default_prompt, {})
|
||||
qa["result_type"] = default_eval.get("result_type", "Invalid")
|
||||
qa["question_answering_reasoning"] = default_eval.get("reasoning", "")
|
||||
|
||||
# Update session data
|
||||
if "session" not in session_data:
|
||||
session_data["session"] = {}
|
||||
if "evaluation_results" not in session_data["session"]:
|
||||
session_data["session"]["evaluation_results"] = {}
|
||||
|
||||
session_data["session"]["evaluation_results"]["question_answering_records"] = qa_records
|
||||
|
||||
# Save updated session data
|
||||
with open(session_file, "w", encoding="utf-8") as f:
|
||||
json.dump(session_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print(f" 💾 Updated session saved with all evaluations")
|
||||
|
||||
return session_data
|
||||
|
||||
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = hallucination = omission = valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
valid += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
valid += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
valid += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"correct_qa_ratio(valid)": correct / valid if valid > 0 else 0,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid if valid > 0 else 0,
|
||||
"omission_qa_ratio(valid)": omission / valid if valid > 0 else 0,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
async def main(
|
||||
data_dir: str = "./data",
|
||||
models: list[str] = None,
|
||||
prompt_versions: list[str] = None,
|
||||
parallel: bool = True
|
||||
):
|
||||
"""Main function to re-evaluate QA records from data directory with multiple models and prompts.
|
||||
|
||||
Args:
|
||||
data_dir: Path to data directory
|
||||
models: List of model names (e.g., ["qwen3-max", "qwen-flash"])
|
||||
prompt_versions: List of prompt versions (e.g., ["v1", "v2"])
|
||||
parallel: If True, use parallel execution; if False, use sequential execution
|
||||
|
||||
Note:
|
||||
Request rate limiting is automatically handled by base_llm.py's request_interval mechanism.
|
||||
"""
|
||||
data_path = Path(data_dir)
|
||||
|
||||
if not data_path.exists() or not data_path.is_dir():
|
||||
print(f"❌ Error: Directory not found: {data_dir}")
|
||||
return
|
||||
|
||||
# Default values
|
||||
if models is None:
|
||||
models = ["qwen3-max"]
|
||||
if prompt_versions is None:
|
||||
prompt_versions = ["v1"]
|
||||
|
||||
print("=" * 80)
|
||||
print("RE-EVALUATING QA RECORDS WITH MULTIPLE MODELS & PROMPTS")
|
||||
print(f"Models: {', '.join(models)}")
|
||||
print(f"Prompts: {', '.join(['EVALUATION_PROMPT_FOR_QUESTION' if v == 'v1' else 'EVALUATION_PROMPT_FOR_QUESTION2' for v in prompt_versions])}")
|
||||
print(f"Execution mode: {'Parallel' if parallel else 'Sequential'}")
|
||||
print("Note: Request rate limiting handled by LLM layer (base_llm.py)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Load data from directory
|
||||
sessions = load_from_data_dir(data_dir)
|
||||
|
||||
if not sessions:
|
||||
print(f"❌ No session files found in {data_dir}")
|
||||
return
|
||||
|
||||
print(f"📂 Found {len(sessions)} session files\n")
|
||||
|
||||
# Process each session
|
||||
all_qa_records = []
|
||||
|
||||
for idx, session_info in enumerate(sessions, 1):
|
||||
session_file = session_info["file"]
|
||||
session_data = session_info["data"]
|
||||
user_name = session_data.get("user_name", "Unknown")
|
||||
session_idx = session_data.get("session_idx", 0)
|
||||
|
||||
print(f"[{idx}/{len(sessions)}] {user_name} - Session {session_idx}")
|
||||
|
||||
updated_session = await reevaluate_session(
|
||||
session_file=session_file,
|
||||
session_data=session_data,
|
||||
models=models,
|
||||
prompt_versions=prompt_versions,
|
||||
parallel=parallel
|
||||
)
|
||||
|
||||
# Collect QA records for metrics
|
||||
eval_results = updated_session.get("session", {}).get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
all_qa_records.extend(qa_records)
|
||||
|
||||
print()
|
||||
|
||||
# Compute and display metrics for each model+prompt combination
|
||||
print("=" * 80)
|
||||
print("UPDATED METRICS (BY MODEL & PROMPT)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for model_name in models:
|
||||
for prompt_version in prompt_versions:
|
||||
prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2"
|
||||
print(f"\n📊 {model_name} / {prompt_name}:")
|
||||
print("─" * 80)
|
||||
|
||||
# Extract QA records for this model+prompt combination
|
||||
model_qa_records = []
|
||||
for qa in all_qa_records:
|
||||
eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {})
|
||||
if eval_data:
|
||||
# Create a copy with the specific evaluation result
|
||||
qa_copy = {
|
||||
**qa,
|
||||
"result_type": eval_data.get("result_type", "Invalid"),
|
||||
"question_answering_reasoning": eval_data.get("reasoning", "")
|
||||
}
|
||||
model_qa_records.append(qa_copy)
|
||||
|
||||
if model_qa_records:
|
||||
metrics = compute_qa_metrics(model_qa_records)
|
||||
|
||||
print(f" Correct (all): {metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {metrics['qa_valid_num']}/{metrics['qa_num']}")
|
||||
|
||||
# Save detailed results with all evaluations
|
||||
report_file = data_path.parent / "reme_eval_stat_result_detailed.json"
|
||||
|
||||
# Create summary for each model+prompt combination
|
||||
evaluation_summary = {}
|
||||
for model_name in models:
|
||||
evaluation_summary[model_name] = {}
|
||||
for prompt_version in prompt_versions:
|
||||
prompt_name = "EVALUATION_PROMPT_FOR_QUESTION" if prompt_version == "v1" else "EVALUATION_PROMPT_FOR_QUESTION2"
|
||||
|
||||
# Extract QA records for this combination
|
||||
model_qa_records = []
|
||||
for qa in all_qa_records:
|
||||
eval_data = qa.get("evaluations", {}).get(model_name, {}).get(prompt_version, {})
|
||||
if eval_data:
|
||||
qa_copy = {
|
||||
**qa,
|
||||
"result_type": eval_data.get("result_type", "Invalid"),
|
||||
"question_answering_reasoning": eval_data.get("reasoning", "")
|
||||
}
|
||||
model_qa_records.append(qa_copy)
|
||||
|
||||
metrics = compute_qa_metrics(model_qa_records)
|
||||
evaluation_summary[model_name][prompt_name] = {
|
||||
"metrics": metrics,
|
||||
"qa_records": model_qa_records
|
||||
}
|
||||
|
||||
final_results = {
|
||||
"evaluation_summary": evaluation_summary,
|
||||
"all_qa_records_with_evaluations": all_qa_records
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print(f"\n💾 Detailed results saved: {report_file}")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Re-evaluate QA records from data directory using multiple models and prompts in parallel. "
|
||||
"Request rate limiting is automatically handled by base_llm.py's request_interval mechanism."
|
||||
)
|
||||
parser.add_argument(
|
||||
"data_dir",
|
||||
nargs='?',
|
||||
default="./data",
|
||||
type=str,
|
||||
help="Path to data directory containing user session files (default: ./data)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models",
|
||||
type=str,
|
||||
nargs='+',
|
||||
default=["qwen3-max", "qwen-flash", "qwen-plus", "qwen3-30b-a3b-instruct-2507", "qwen3-235b-a22b-instruct-2507"],
|
||||
help="LLM model names for evaluation (space-separated, default: qwen3-max)"
|
||||
)
|
||||
# ["gpt-4o-2024-11-20"], ["qwen3-max", "qwen-flash", "qwen3-30b-a3b-instruct-2507", "qwen3-30b-a3b-instruct-2507"]
|
||||
parser.add_argument(
|
||||
"--prompts",
|
||||
type=str,
|
||||
nargs='+',
|
||||
choices=["v1", "v2"],
|
||||
default=["v1", "v2"],
|
||||
help="Prompt versions to use: v1=EVALUATION_PROMPT_FOR_QUESTION, v2=EVALUATION_PROMPT_FOR_QUESTION2 (default: v1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--serial",
|
||||
action="store_true",
|
||||
help="Use sequential execution instead of parallel (default: parallel)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
asyncio.run(main(
|
||||
data_dir=args.data_dir,
|
||||
models=args.models,
|
||||
prompt_versions=args.prompts,
|
||||
parallel=not args.serial
|
||||
))
|
||||
|
|
@ -41,6 +41,7 @@ class EvalConfig:
|
|||
max_concurrency: int = 2
|
||||
batch_size: int = 20
|
||||
output_dir: str = "bench_results/reme"
|
||||
reme_model_name: str = "qwen-flash"
|
||||
eval_model_name: str = "qwen3-max"
|
||||
algo_version: str = "v1"
|
||||
|
||||
|
|
@ -218,7 +219,7 @@ async def answer_question_with_memories(
|
|||
|
||||
result = await reme.llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name="qwen3-30b-a3b-instruct-2507"
|
||||
model_name="qwen-flash"
|
||||
)
|
||||
|
||||
return result
|
||||
|
|
@ -301,6 +302,8 @@ class MemoryProcessor:
|
|||
user_name=user_id,
|
||||
version=self.algo_version,
|
||||
return_dict=True,
|
||||
enable_time_filter=True,
|
||||
enable_thinking_params=False
|
||||
)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
|
|
@ -333,6 +336,8 @@ class MemoryProcessor:
|
|||
user_name=user_id,
|
||||
version=self.algo_version,
|
||||
return_dict=True,
|
||||
enable_time_filter=True,
|
||||
enable_thinking_params=False
|
||||
)
|
||||
|
||||
# Extract memories from response
|
||||
|
|
@ -404,6 +409,16 @@ class QuestionAnsweringEvaluator:
|
|||
model_name=self.eval_model_name
|
||||
)
|
||||
|
||||
eval_result_original_answer = await evaluation_for_question(
|
||||
reme=self.reme,
|
||||
question=qa["question"],
|
||||
reference_answer=qa["answer"],
|
||||
key_memory_points=evidence_text,
|
||||
response=retrieved_memories,
|
||||
dialogue=formatted_dialogue,
|
||||
model_name=self.eval_model_name
|
||||
)
|
||||
|
||||
# Build result record
|
||||
qa_result = {
|
||||
**qa,
|
||||
|
|
@ -416,7 +431,9 @@ class QuestionAnsweringEvaluator:
|
|||
"retrieve_messages": agent_messages,
|
||||
"search_duration_ms": duration_ms,
|
||||
"result_type": eval_result.get("evaluation_result"),
|
||||
"question_answering_reasoning": eval_result.get("reasoning", "")
|
||||
"question_answering_reasoning": eval_result.get("reasoning", ""),
|
||||
"original_result_type": eval_result_original_answer.get("evaluation_result"),
|
||||
"original_question_answering_reasoning": eval_result_original_answer.get("reasoning", ""),
|
||||
}
|
||||
results.append(qa_result)
|
||||
|
||||
|
|
@ -427,8 +444,8 @@ class MetricsAggregator:
|
|||
"""Aggregates evaluation metrics."""
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
def _compute_single_metric(qa_records: list[dict], result_key: str) -> dict[str, Any]:
|
||||
"""Compute metrics for a single result type key."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
|
|
@ -448,7 +465,7 @@ class MetricsAggregator:
|
|||
valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
result_type = qa.get(result_key, "")
|
||||
|
||||
if result_type in ["Correct", "Hallucination", "Omission"]:
|
||||
valid += 1
|
||||
|
|
@ -482,6 +499,14 @@ class MetricsAggregator:
|
|||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics for both result_type and original_result_type."""
|
||||
return {
|
||||
"with_llm_answer": MetricsAggregator._compute_single_metric(qa_records, "result_type"),
|
||||
"with_original_memories": MetricsAggregator._compute_single_metric(qa_records, "original_result_type")
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
|
|
@ -516,7 +541,7 @@ class HaluMemEvaluator:
|
|||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.reme = ReMe()
|
||||
self.reme = ReMe(llm={"model_name": self.config.reme_model_name})
|
||||
|
||||
# Load evaluation prompts into ReMe's prompt handler
|
||||
prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml"
|
||||
|
|
@ -536,6 +561,10 @@ class HaluMemEvaluator:
|
|||
)
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
# For real-time updates
|
||||
self._update_lock: asyncio.Lock | None = None
|
||||
self._output_file: str | None = None
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry."""
|
||||
return self
|
||||
|
|
@ -626,8 +655,20 @@ class HaluMemEvaluator:
|
|||
|
||||
self.file_manager.save_session(user_name, idx, session_data)
|
||||
|
||||
# Update results file after each session completes
|
||||
await self._trigger_update()
|
||||
|
||||
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
|
||||
|
||||
async def _trigger_update(self):
|
||||
"""Trigger real-time update of results and statistics."""
|
||||
if self._update_lock is None or self._output_file is None:
|
||||
return
|
||||
|
||||
async with self._update_lock:
|
||||
self.file_manager.combine_results(self._output_file)
|
||||
self._update_statistics(self._output_file)
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the complete evaluation pipeline using ReMe."""
|
||||
start_time = time.time()
|
||||
|
|
@ -661,6 +702,12 @@ class HaluMemEvaluator:
|
|||
print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Output file path for real-time updates
|
||||
self._output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
|
||||
# Lock for thread-safe file updates
|
||||
self._update_lock = asyncio.Lock()
|
||||
|
||||
# Process users with concurrency control
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
|
||||
|
|
@ -671,11 +718,14 @@ class HaluMemEvaluator:
|
|||
# Check cache
|
||||
if self.file_manager.user_has_cache(user_name):
|
||||
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
|
||||
return {"user_name": user_name, "status": "cached"}
|
||||
result = {"user_name": user_name, "status": "cached"}
|
||||
# Also trigger update for cached users
|
||||
await self._trigger_update()
|
||||
else:
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
result = await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
|
||||
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
result = await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
|
||||
return result
|
||||
|
||||
tasks = [
|
||||
|
|
@ -684,16 +734,57 @@ class HaluMemEvaluator:
|
|||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Combine results
|
||||
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
self.file_manager.combine_results(output_file)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
print(f"📁 Results: {output_file}\n")
|
||||
print(f"📁 Results: {self._output_file}\n")
|
||||
|
||||
# Aggregate metrics
|
||||
await self.aggregate_and_report(output_file)
|
||||
# Final aggregation and report
|
||||
await self.aggregate_and_report(self._output_file)
|
||||
|
||||
def _update_statistics(self, results_file: str):
|
||||
"""Update statistics file based on current results (for real-time monitoring)."""
|
||||
if not os.path.exists(results_file):
|
||||
return
|
||||
|
||||
# Collect all QA records
|
||||
qa_records = []
|
||||
try:
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
qa_records.extend(
|
||||
eval_results.get("question_answering_records", [])
|
||||
)
|
||||
except (json.JSONDecodeError, KeyError):
|
||||
return
|
||||
|
||||
if not qa_records:
|
||||
return
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
|
||||
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
# Save statistics
|
||||
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
async def aggregate_and_report(self, results_file: str):
|
||||
"""Aggregate results and generate final report."""
|
||||
|
|
@ -746,14 +837,27 @@ class HaluMemEvaluator:
|
|||
print("EVALUATION SUMMARY - REME")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Question Answering:")
|
||||
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
# Print metrics for LLM-generated answer (result_type)
|
||||
llm_metrics = qa_metrics["with_llm_answer"]
|
||||
print("📊 Question Answering (with LLM answer):")
|
||||
print(f" Correct (all): {llm_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {llm_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {llm_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {llm_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {llm_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {llm_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {llm_metrics['qa_valid_num']}/{llm_metrics['qa_num']}")
|
||||
|
||||
# Print metrics for original retrieved memories (original_result_type)
|
||||
orig_metrics = qa_metrics["with_original_memories"]
|
||||
print("\n📊 Question Answering (with original memories):")
|
||||
print(f" Correct (all): {orig_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {orig_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {orig_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {orig_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {orig_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {orig_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {orig_metrics['qa_valid_num']}/{orig_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ Time Metrics:")
|
||||
print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min")
|
||||
|
|
@ -769,6 +873,7 @@ async def main_async(
|
|||
top_k: int,
|
||||
user_num: int,
|
||||
max_concurrency: int,
|
||||
reme_model_name: str= "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
algo_version: str = "v1"
|
||||
):
|
||||
|
|
@ -778,6 +883,7 @@ async def main_async(
|
|||
top_k=top_k,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency,
|
||||
reme_model_name=reme_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version
|
||||
)
|
||||
|
|
@ -792,6 +898,7 @@ def main(
|
|||
top_k: int,
|
||||
user_num: int,
|
||||
max_concurrency: int,
|
||||
reme_model_name: str= "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
algo_version: str = "v1"
|
||||
):
|
||||
|
|
@ -801,6 +908,7 @@ def main(
|
|||
top_k=top_k,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency,
|
||||
reme_model_name=reme_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version
|
||||
))
|
||||
|
|
@ -815,7 +923,8 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
# required=True,
|
||||
default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/halumem/HaluMem-Medium.jsonl",
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
|
|
@ -833,20 +942,25 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=100,
|
||||
default=1,
|
||||
help="Maximum concurrent user processing (default: 100)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reme_model_name",
|
||||
type=str,
|
||||
default="qwen-flash",
|
||||
help="Model name for ReMe (default: qwen-flash)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eval_model_name",
|
||||
type=str,
|
||||
default="qwen3-max",
|
||||
# default="qwen3-235b-a22b-instruct-2507",
|
||||
help="Model name for evaluation (default: qwen3-max)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--algo_version",
|
||||
type=str,
|
||||
default="v1",
|
||||
default="halumem",
|
||||
help="Algorithm version for summary and retrieval (default: v1)"
|
||||
)
|
||||
|
||||
|
|
@ -857,6 +971,7 @@ if __name__ == "__main__":
|
|||
top_k=args.top_k,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency,
|
||||
reme_model_name=args.reme_model_name,
|
||||
eval_model_name=args.eval_model_name,
|
||||
algo_version=args.algo_version
|
||||
)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from .personal.personal_retriever import PersonalRetriever
|
|||
from .personal.personal_summarizer import PersonalSummarizer
|
||||
from .personal.personal_v1_retriever import PersonalV1Retriever
|
||||
from .personal.personal_v1_summarizer import PersonalV1Summarizer
|
||||
from .personal.personal_halumem_retriever import PersonalHalumemRetriever
|
||||
from .personal.personal_halumem_summarizer import PersonalHalumemSummarizer
|
||||
from .procedural.procedural_retriever import ProceduralRetriever
|
||||
from .procedural.procedural_summarizer import ProceduralSummarizer
|
||||
from .reme_retriever import ReMeRetriever
|
||||
|
|
@ -19,6 +21,8 @@ __all__ = [
|
|||
"PersonalSummarizer",
|
||||
"PersonalV1Retriever",
|
||||
"PersonalV1Summarizer",
|
||||
"PersonalHalumemRetriever",
|
||||
"PersonalHalumemSummarizer",
|
||||
"ProceduralRetriever",
|
||||
"ProceduralSummarizer",
|
||||
"ReMeRetriever",
|
||||
|
|
|
|||
92
reme/agent/memory/personal/personal_halumem_retriever.py
Normal file
92
reme/agent/memory/personal/personal_halumem_retriever.py
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
"""Personal memory retriever agent for retrieving personal memories through vector search."""
|
||||
|
||||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import Role, MemoryType
|
||||
from ....core.op import BaseTool
|
||||
from ....core.schema import Message
|
||||
from ....core.utils import format_messages
|
||||
|
||||
|
||||
class PersonalHalumemRetriever(BaseMemoryAgent):
|
||||
"""Retrieve personal memories through vector search and history reading."""
|
||||
|
||||
memory_type: MemoryType = MemoryType.PERSONAL
|
||||
|
||||
async def build_messages(self) -> list[Message]:
|
||||
if self.context.get("query"):
|
||||
context = self.context.query
|
||||
elif self.context.get("messages"):
|
||||
context = self.description + "\n" + format_messages(self.context.messages)
|
||||
else:
|
||||
raise ValueError("input must have either `query` or `messages`")
|
||||
|
||||
read_all_profiles_tool: BaseTool | None = self.pop_tool("read_all_profiles")
|
||||
if read_all_profiles_tool is not None:
|
||||
all_profiles = await read_all_profiles_tool.call(
|
||||
memory_target=self.memory_target,
|
||||
service_context=self.service_context,
|
||||
)
|
||||
else:
|
||||
all_profiles = ""
|
||||
|
||||
return [
|
||||
Message(
|
||||
role=Role.SYSTEM,
|
||||
content=self.prompt_format(
|
||||
prompt_name="system_prompt",
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
user_profile=all_profiles,
|
||||
context=context.strip(),
|
||||
),
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=self.prompt_format(
|
||||
prompt_name="user_message",
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
user_profile=all_profiles,
|
||||
context=context.strip(),
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
async def _acting_step(
|
||||
self,
|
||||
assistant_message: Message,
|
||||
tools: list[BaseTool],
|
||||
step: int,
|
||||
stage: str = "",
|
||||
**kwargs,
|
||||
) -> tuple[list[BaseTool], list[Message]]:
|
||||
"""Execute tool calls with memory context."""
|
||||
return await super()._acting_step(
|
||||
assistant_message,
|
||||
tools,
|
||||
step,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
retrieved_nodes=self.retrieved_nodes,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
result = await super().execute()
|
||||
answer = result["answer"]
|
||||
if "MEMORY_NOT_FOUND" in answer:
|
||||
result["answer"] = "\n".join(
|
||||
[
|
||||
n.format(
|
||||
include_memory_id=False,
|
||||
include_when_to_use=False,
|
||||
include_content=True,
|
||||
include_message_time=False,
|
||||
ref_memory_id_key="",
|
||||
)
|
||||
for n in self.retrieved_nodes
|
||||
],
|
||||
)
|
||||
|
||||
result["retrieved_nodes"] = self.retrieved_nodes
|
||||
return result
|
||||
91
reme/agent/memory/personal/personal_halumem_retriever.yaml
Normal file
91
reme/agent/memory/personal/personal_halumem_retriever.yaml
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
system_prompt: |
|
||||
# Role Definition:
|
||||
You are a Memory Retrieval Agent specialized in retrieving {memory_type} memories about {memory_target}.
|
||||
|
||||
## Multi-Phase Retrieval Strategy
|
||||
Follow these phases sequentially to gather comprehensive information:
|
||||
|
||||
### Tool Rules
|
||||
**Tool**: `retrieve_memory` (without time constraints)
|
||||
**Objective**: Cast a wide net to find potentially relevant memories
|
||||
**Approach**:
|
||||
- Execute 3-5 diverse search queries using different formulations:
|
||||
* Original question verbatim
|
||||
* Rephrased variations (different wording, synonyms)
|
||||
* Entity-focused queries (extract and search specific names, places, events)
|
||||
* Keyword-based searches (core concepts, topics)
|
||||
* Related context queries (broader themes)
|
||||
- Review all results before proceeding to next phase
|
||||
|
||||
**Tool**: `retrieve_memory` (with time filter)
|
||||
**When to use**: Only if the user question contains temporal references
|
||||
**Time Filter Format**:
|
||||
- Single date: `20200101`
|
||||
- Date range: `20200101,20200102` (inclusive: 20200101 ≤ time ≤ 20200102)
|
||||
- Before date: `0,20200102` (up to and including 20200102)
|
||||
- After date: `20200101,99999999` (from 20200101 onwards)
|
||||
**Approach**:
|
||||
- Identify temporal constraints from the user question
|
||||
- Refine Phase 1 queries with appropriate time filters
|
||||
- Try multiple time ranges if initial searches yield no results
|
||||
|
||||
**Tool**: `read_history`
|
||||
**When to use**: After exhausting retrieval attempts OR when specific conversation context is needed
|
||||
**Approach**:
|
||||
- Extract `history_id` from retrieved memory references
|
||||
- Prioritize histories that are most relevant or recent
|
||||
- Read multiple histories if necessary for complete context
|
||||
- Use this to understand the full conversation surrounding a memory
|
||||
|
||||
|
||||
user_message: |
|
||||
## User Profile
|
||||
{user_profile}
|
||||
|
||||
## User Question
|
||||
{context}
|
||||
|
||||
# Core Objective:
|
||||
Before responding to the user, you must strictly follow the **[Memory Retrieve -> Original Source Tracing -> Broad Search Fallback]** retrieval strategy. It is strictly forbidden to directly opt for an indiscriminate search of massive historical original texts.
|
||||
|
||||
## Retrieval Strategy & Workflow (Strictly Enforced Chain of Thought)
|
||||
### Phase 1: Intent Decomposition and Primary Retrieval (Summary First)
|
||||
1. **Analyze Intent**: Analyze the user's current Query, decomposing it into 1-3 core search intents.2. **Summary Priority**: First, retrieve from **high-level memories**.
|
||||
- **Action**: Call `vector_retrieve_memory` using at least two different `query`.
|
||||
- **Filter**: (Optional) Set metadata filter {{"timestamp": "YYYY-MM-DD"}}
|
||||
- **Goal**: Obtain refined conclusions such as entity attributes, task status, user preferences, or environmental information.
|
||||
|
||||
### Phase 2: Memory Evaluation and Deep Tracing (Drill Down)
|
||||
Check the retrieval results of Phase 1:
|
||||
- **Case A (Sufficient Information)**: If the summarized memory contains all the details needed for the answer, proceed directly to Phase 4 for the response.
|
||||
- **Case B (Vague/Complex Information)**: If summarized memory exists (e.g., 'discussed project architecture') but lacks specific details (e.g., 'specific parameter configuration'), use clues from the summary to trace the original text.
|
||||
- **Action**: Call `read_history` using ref_memory_id from the retrieved memory.
|
||||
- **Goal**: Obtain the specific conversation context at that time.
|
||||
|
||||
### Phase 3: Fallback Retrieval and Strategy Adjustment (Fallback & Expand)
|
||||
If no valid information is found in both Phase 1 and Phase 2 (result is empty or similarity is too low): Rewrite the Query based on the context (remove non-keywords, synonym substitution), and search again.
|
||||
|
||||
### Phase 4: Result Compilation and Response
|
||||
- Combine the retrieved content (summary or original text) with the current conversation context.
|
||||
- If all retrieved results are irrelevant, **it is strictly forbidden to fabricate memories**; directly inform the user that no relevant information was found.
|
||||
|
||||
## Output Format
|
||||
Before the final reply, ensure at least 3 tool calls for retrieval, and then output your answer in ten words.
|
||||
- Base your answer EXCLUSIVELY on retrieved memories, user profile, and history data
|
||||
- Never infer, assume, or hallucinate information
|
||||
- Always cite sources with timestamps: `[timestamp] Memory content`
|
||||
- Present conflicting information transparently with respective timestamps
|
||||
- Exhaust all search strategies before concluding information doesn't exist
|
||||
|
||||
Before the final reply, ensure at least 3 tool calls for retrieval, and then output the most relevant JSON retrieval summary, followed by your answer:
|
||||
```json
|
||||
{{
|
||||
"retrieved_memories": [
|
||||
{{"type": "profile", "timestamp":"...", "content": "..."}},
|
||||
{{"type": "personal", "timestamp":"...", "content": "..."}},
|
||||
{{"type": "history", "timestamp":"...", "content": "..."}},
|
||||
....
|
||||
],
|
||||
"summary": "Fill in your summarized answer here."
|
||||
}}
|
||||
```
|
||||
140
reme/agent/memory/personal/personal_halumem_summarizer.py
Normal file
140
reme/agent/memory/personal/personal_halumem_summarizer.py
Normal file
|
|
@ -0,0 +1,140 @@
|
|||
"""Personal memory summarizer agent for two-phase personal memory processing."""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import Role, MemoryType
|
||||
from ....core.op import BaseTool
|
||||
from ....core.schema import Message
|
||||
|
||||
|
||||
class PersonalHalumemSummarizer(BaseMemoryAgent):
|
||||
"""Two-phase personal memory processor: retrieve/add memories then update profile."""
|
||||
|
||||
memory_type: MemoryType = MemoryType.PERSONAL
|
||||
|
||||
async def _build_s1_messages(self) -> list[Message]:
|
||||
return [
|
||||
Message(
|
||||
role=Role.SYSTEM,
|
||||
content=self.prompt_format(
|
||||
prompt_name="system_prompt_s1",
|
||||
context=self.context.history_node.content,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
),
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
# content=self.get_prompt("user_message_s1"),
|
||||
content=self.prompt_format(
|
||||
prompt_name="user_message_s1",
|
||||
context=self.context.history_node.content,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
async def _build_s2_messages(self, user_profile: str) -> list[Message]:
|
||||
return [
|
||||
Message(
|
||||
role=Role.SYSTEM,
|
||||
content=self.prompt_format(
|
||||
prompt_name="system_prompt_s2",
|
||||
context=self.context.history_node.content,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
user_profile=user_profile,
|
||||
),
|
||||
),
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=self.prompt_format(
|
||||
prompt_name="user_message_s2",
|
||||
context=self.context.history_node.content,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
user_profile=user_profile,
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
async def _acting_step(
|
||||
self,
|
||||
assistant_message: Message,
|
||||
tools: list[BaseTool],
|
||||
step: int,
|
||||
stage: str = "",
|
||||
**kwargs,
|
||||
) -> tuple[list[BaseTool], list[Message]]:
|
||||
"""Execute tool calls with memory context."""
|
||||
return await super()._acting_step(
|
||||
assistant_message,
|
||||
tools,
|
||||
step,
|
||||
stage=stage,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
history_node=self.history_node,
|
||||
author=self.author,
|
||||
retrieved_nodes=self.retrieved_nodes,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
memory_tools = []
|
||||
profile_tools = []
|
||||
for i, tool in enumerate(self.tools):
|
||||
tool_name = tool.tool_call.name
|
||||
if "_memory" in tool_name:
|
||||
memory_tools.append(tool)
|
||||
elif "_profile" in tool_name:
|
||||
profile_tools.append(tool)
|
||||
else:
|
||||
raise ValueError(f"[{self.__class__.__name__}] unknown tool_name={tool_name}")
|
||||
logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}")
|
||||
|
||||
stage = "s1-memory"
|
||||
messages_s1 = await self._build_s1_messages()
|
||||
for i, message in enumerate(messages_s1):
|
||||
role = message.name or message.role
|
||||
logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}")
|
||||
tools_s1, messages_s1, success_s1 = await self.react(messages_s1, memory_tools, stage=stage)
|
||||
|
||||
if profile_tools:
|
||||
|
||||
read_all_profiles_tool: BaseTool | None = self.pop_tool("read_all_profiles")
|
||||
if read_all_profiles_tool is not None:
|
||||
all_profiles = await read_all_profiles_tool.call(
|
||||
memory_target=self.memory_target,
|
||||
service_context=self.service_context,
|
||||
)
|
||||
else:
|
||||
all_profiles = ""
|
||||
|
||||
stage = "s2-profile"
|
||||
messages_s2 = await self._build_s2_messages(user_profile=all_profiles)
|
||||
for i, message in enumerate(messages_s2):
|
||||
role = message.name or message.role
|
||||
logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}")
|
||||
tools_s2, messages_s2, success_s2 = await self.react(messages_s2, profile_tools, stage=stage)
|
||||
else:
|
||||
tools_s2, messages_s2, success_s2 = [], [], True
|
||||
|
||||
answer = (messages_s1[-1].content if success_s1 else "") + (messages_s2[-1].content if success_s2 else "")
|
||||
success = success_s1 and success_s2
|
||||
messages = messages_s1 + messages_s2
|
||||
tools = tools_s1 + tools_s2
|
||||
memory_nodes = []
|
||||
for tool in tools:
|
||||
if tool.memory_nodes:
|
||||
memory_nodes.extend(tool.memory_nodes)
|
||||
|
||||
return {
|
||||
"answer": answer,
|
||||
"success": success,
|
||||
"messages": messages,
|
||||
"tools": tools,
|
||||
"memory_nodes": memory_nodes,
|
||||
}
|
||||
96
reme/agent/memory/personal/personal_halumem_summarizer.yaml
Normal file
96
reme/agent/memory/personal/personal_halumem_summarizer.yaml
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
system_prompt_s1: |
|
||||
You are a Memory Agent responsible for managing {memory_type} memories about {memory_target}.
|
||||
|
||||
## Tool Rules
|
||||
1. `add_and_retrieve_similar_memory`: Create a memory in the vector store.
|
||||
- Use this tool to add memories, and it will return the relevant content related to the added memories.
|
||||
- Use actual names from the conversation (e.g., "Bob likes apples") instead of generic references (e.g., "user likes apples")
|
||||
- The tool will retrieve similar historical memories via vector search to help you consolidate in Step 2
|
||||
|
||||
2. `update_memory`: Update memories in the vector store.
|
||||
**What to Delete** (via `memory_ids_to_delete`):
|
||||
- Duplicate memories with identical or highly similar content
|
||||
- Memories that should be merged into a single consolidated entry
|
||||
|
||||
**What to Add** (via `memories_to_add` with message_time and memory_content):
|
||||
- For each topic with changes: add ONE consolidated memory that merges related information
|
||||
- New distinct memories that don't overlap with existing ones
|
||||
- Updated memories that capture the latest state while preserving temporal evolution
|
||||
|
||||
user_message_s1: |
|
||||
## Latest Conversation
|
||||
Format: round<index> [<timestamp>] <role/name>: <content>
|
||||
{context}
|
||||
|
||||
## Task
|
||||
### Step 1: Create Memory
|
||||
- At this step, you can call the tool multiple times to store memories, or you can call it once to store multiple memories.
|
||||
|
||||
### Step 2: Update Memory Store
|
||||
- Update the vector store using `update_memory` to keep it well-organized and consolidated.
|
||||
|
||||
## Storage Scope (Biographical & Behavioral ONLY)
|
||||
- Personal Biography: Significant milestones, past experiences, and life events.
|
||||
- Behavioral Patterns: How the agent reacts, specific actions taken, and recurring habits.
|
||||
- **EXCLUSION**: DO NOT record objective world facts, general knowledge, or user-specific health/states.
|
||||
|
||||
Extraction & Formatting Rules
|
||||
- Fact Filtering: Only extract information that builds the biography of **{memory_target}**.
|
||||
- Subject Splitting: If a conversation mentions multiple subject (e.g., the User's childhood and their Father's career), create separate memory entries for each subject.
|
||||
- Atomic Content: Each entry should focus on one specific event or trait. Keep descriptions concise to ensure efficient retrieval.
|
||||
|
||||
|
||||
system_prompt_s2: |
|
||||
You are a Profile Agent responsible for managing profiles about {memory_target}.
|
||||
|
||||
## Tool Rules
|
||||
1. Update Profile with `update_profile`
|
||||
Synchronize profile with new information from the conversation:
|
||||
- `profile_ids_to_delete`: Remove conflicting, or redundant entries (array of profile IDs).
|
||||
- `profiles_to_add`:
|
||||
- `conversation_time`: Time of conversation (format: `YYYY-MM-DD HH:MM:SS`, e.g., `2024-01-15 14:30:00`)
|
||||
- `profile_content`: Complete, self-contained profile description with full context
|
||||
Update user profile using `update_profile` based on the conversation and current profile.
|
||||
|
||||
|
||||
user_message_s2: |
|
||||
You are a memory agent managing **{memory_type}** memories about **{memory_target}**.
|
||||
|
||||
## Latest Conversation:
|
||||
{context}
|
||||
|
||||
Message format: `round<index> [<timestamp>] <role/name>: <content>` (timestamp: YYYY-MM-DD HH:MM:SS).
|
||||
|
||||
**CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate.
|
||||
|
||||
## Current User Profile:
|
||||
{user_profile}
|
||||
|
||||
## Task
|
||||
### Step 1: ADD Profile
|
||||
- Add new, relevant, and up-to-date information to the user profile using `update_profile` (via `profiles_to_add`).
|
||||
|
||||
### Step 2: DELETE Profile
|
||||
- Delete outdated, redundant, or resolved states from the user profile using `update_profile` (via `profile_ids_to_delete`).
|
||||
|
||||
## Storage Scope (Current States ONLY)
|
||||
- **EXCLUSION PRINCIPLE**: DO NOT record any user *actions*, *requests*, *queries*, or *interactions with the system* (e.g., "asked for code", "solved a puzzle", "requested translation"). These are interaction logs, not user states.
|
||||
- Identity: Geography, job title, work content, income.
|
||||
- Background: Education, family, relationships, hobbies, interests, and other personal preferences.
|
||||
- Temporary States: Physical health (e.g., "Has a cold"), emotional mood, stress levels, and specific prohibitions (e.g., "Cannot drink alcohol due to medication").
|
||||
|
||||
## Profile Management Rules
|
||||
- Subject Splitting (CRITICAL): If the conversation mentions multiple subjects (e.g., the User's job and their Spouse's health), you MUST create separate profile entries for each unique subject.
|
||||
- Conflict Resolution: Use profile_ids_to_delete to remove outdated, redundant, or resolved states (e.g., if a user is "Recovered," delete the "Illness" entry).
|
||||
- Each profile entry MUST describe a **persistent or temporary state of the user themselves** (e.g., who they are, what they like, what they’re dealing with), NOT an event they participated in or a request they made.
|
||||
|
||||
## Profile Format
|
||||
- **ONLY record what the user EXPLICITLY STATES about themselves as a state or preference.**
|
||||
- The key in the record represents the category of memory, and the value should record the specific content. For example:
|
||||
{{"message_time": "YYYY-MM-DD HH:MM:SS", "profile_key": "the category of memory", "profile_value": "content" }}
|
||||
- When there is no information conflict or outdated information, you don't need to delete any of the memory. If there is no information that meets the requirements, it is also acceptable not to add it.
|
||||
|
||||
## Forbidden Case
|
||||
1.There is no need to record user behavior: {{ "profile_key": "workouts", "profile_content": "confident in new running shoes' suitability for chosen route; they have significantly improved morning jogs"}}
|
||||
2. There is no need to record the users' plans or requirements: {{ "profile_key": "plans.vacation", "profile_content": "planning to go to a nearby city for a week and ask for a job change"}}
|
||||
|
||||
|
|
@ -16,7 +16,6 @@ llm:
|
|||
default:
|
||||
backend: openai
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
# model_name: qwen-flash
|
||||
request_interval: 1
|
||||
temperature: 0.0001
|
||||
|
||||
|
|
@ -46,3 +45,4 @@ token_counter:
|
|||
backend: hf
|
||||
model_name: Qwen/Qwen3-Coder-30B-A3B-Instruct
|
||||
use_mirror: true
|
||||
|
||||
|
|
|
|||
|
|
@ -80,7 +80,7 @@ class ESVectorStore(BaseVectorStore):
|
|||
|
||||
async def _get_client(self) -> AsyncElasticsearch:
|
||||
"""Create or return the existing AsyncElasticsearch client.
|
||||
|
||||
|
||||
This lazy initialization ensures the client is created in the correct event loop.
|
||||
"""
|
||||
if self._client is None:
|
||||
|
|
@ -93,7 +93,7 @@ class ESVectorStore(BaseVectorStore):
|
|||
headers=self.headers,
|
||||
)
|
||||
logger.info("AsyncElasticsearch client initialized")
|
||||
|
||||
|
||||
return self._client
|
||||
|
||||
async def list_collections(self) -> list[str]:
|
||||
|
|
|
|||
|
|
@ -109,7 +109,7 @@ class QdrantVectorStore(BaseVectorStore):
|
|||
|
||||
async def _get_client(self) -> AsyncQdrantClient:
|
||||
"""Create or return the existing AsyncQdrantClient.
|
||||
|
||||
|
||||
This lazy initialization ensures the client is created in the correct event loop.
|
||||
"""
|
||||
if self._client is None:
|
||||
|
|
@ -125,7 +125,7 @@ class QdrantVectorStore(BaseVectorStore):
|
|||
**self.client_kwargs,
|
||||
)
|
||||
logger.info("AsyncQdrantClient initialized")
|
||||
|
||||
|
||||
return self._client
|
||||
|
||||
async def list_collections(self) -> list[str]:
|
||||
|
|
|
|||
76
reme/reme.py
76
reme/reme.py
|
|
@ -9,6 +9,8 @@ from .agent.memory import (
|
|||
ReMeRetriever,
|
||||
PersonalV1Summarizer,
|
||||
PersonalV1Retriever,
|
||||
PersonalHalumemSummarizer,
|
||||
PersonalHalumemRetriever,
|
||||
PersonalSummarizer,
|
||||
PersonalRetriever,
|
||||
ProceduralSummarizer,
|
||||
|
|
@ -26,10 +28,12 @@ from .tool.memory import (
|
|||
ReadHistory,
|
||||
ProfileHandler,
|
||||
MemoryHandler,
|
||||
AddDraftAndRetrieveSimilarMemory,
|
||||
AddAndRetrieveSimilarMemory,
|
||||
UpdateMemoryV2,
|
||||
AddDraftAndReadAllProfiles,
|
||||
UpdateProfile,
|
||||
# DeleteProfile,
|
||||
# AddProfile,
|
||||
AddHistory,
|
||||
ReadAllProfiles,
|
||||
AddMemory,
|
||||
|
|
@ -122,7 +126,7 @@ class ReMe(Application):
|
|||
if version == "default":
|
||||
personal_summarizer = PersonalSummarizer(
|
||||
tools=[
|
||||
AddDraftAndRetrieveSimilarMemory(
|
||||
AddAndRetrieveSimilarMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
top_k=retrieve_top_k,
|
||||
),
|
||||
|
|
@ -141,8 +145,7 @@ class ReMe(Application):
|
|||
elif version == "v1":
|
||||
personal_summarizer = PersonalV1Summarizer(
|
||||
tools=[
|
||||
AddDraftAndRetrieveSimilarMemory(
|
||||
top_k=retrieve_top_k,
|
||||
AddAndRetrieveSimilarMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_when_to_use=False,
|
||||
|
|
@ -168,18 +171,51 @@ class ReMe(Application):
|
|||
),
|
||||
],
|
||||
)
|
||||
|
||||
elif version == "halumem":
|
||||
personal_summarizer = PersonalHalumemSummarizer(
|
||||
tools=[
|
||||
AddAndRetrieveSimilarMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
top_k=retrieve_top_k,
|
||||
),
|
||||
UpdateMemoryV2(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
),
|
||||
# RetrieveMemory(
|
||||
# enable_thinking_params=enable_thinking_params,
|
||||
# top_k=retrieve_top_k,
|
||||
# enable_time_filter=enable_time_filter,
|
||||
# ),
|
||||
# 处理userprofile
|
||||
ReadAllProfiles(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
UpdateProfile(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
# AddProfile(
|
||||
# enable_thinking_params=enable_thinking_params,
|
||||
# profile_dir=self.profile_dir,
|
||||
# ),
|
||||
# DeleteProfile(
|
||||
# enable_thinking_params=enable_thinking_params,
|
||||
# profile_dir=self.profile_dir,
|
||||
# ),
|
||||
],
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
procedural_summarizer: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
procedural_summarizer = ProceduralSummarizer(tools=[])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
tool_summarizer: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
tool_summarizer = ToolSummarizer(tools=[])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
@ -221,7 +257,7 @@ class ReMe(Application):
|
|||
memory_agents = [personal_summarizer, procedural_summarizer, tool_summarizer]
|
||||
|
||||
reme_summarizer: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
reme_summarizer = ReMeSummarizer(tools=[AddHistory(), DelegateTask(memory_agents=memory_agents)])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
@ -284,23 +320,37 @@ class ReMe(Application):
|
|||
top_k=retrieve_top_k,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_time_filter=enable_time_filter,
|
||||
enable_multiple=True
|
||||
enable_multiple=True,
|
||||
),
|
||||
ReadHistory(enable_thinking_params=enable_thinking_params),
|
||||
],
|
||||
)
|
||||
elif version == "halumem":
|
||||
personal_retriever = PersonalHalumemRetriever(
|
||||
tools=[
|
||||
ReadAllProfiles(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
RetrieveMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
top_k=retrieve_top_k,
|
||||
enable_time_filter=enable_time_filter,
|
||||
),
|
||||
ReadHistory(enable_thinking_params=enable_thinking_params),
|
||||
],
|
||||
)
|
||||
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
procedural_retriever: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
procedural_retriever = ProceduralRetriever(tools=[])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
tool_retriever: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
tool_retriever = ToolRetriever(tools=[])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
@ -340,7 +390,7 @@ class ReMe(Application):
|
|||
memory_agents = [personal_retriever, procedural_retriever, tool_retriever]
|
||||
|
||||
reme_retriever: BaseMemoryAgent
|
||||
if version in ["default", "v1"]:
|
||||
if version in ["default", "v1", "halumem"]:
|
||||
reme_retriever = ReMeRetriever(tools=[DelegateTask(memory_agents=memory_agents)])
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -7,8 +7,10 @@ from .history.read_history import ReadHistory
|
|||
from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
|
||||
from .profiles.profile_handler import ProfileHandler
|
||||
from .profiles.read_all_profiles import ReadAllProfiles
|
||||
from .profiles.add_profile import AddProfile
|
||||
from .profiles.update_profile import UpdateProfile
|
||||
from .vector.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory
|
||||
from .profiles.delete_profile import DeleteProfile
|
||||
from .vector.add_draft_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory
|
||||
from .vector.add_memory import AddMemory
|
||||
from .vector.delete_memory import DeleteMemory
|
||||
from .vector.memory_handler import MemoryHandler
|
||||
|
|
@ -27,11 +29,13 @@ __all__ = [
|
|||
"ReadHistory",
|
||||
# Profiles
|
||||
"AddDraftAndReadAllProfiles",
|
||||
"AddProfile",
|
||||
"ProfileHandler",
|
||||
"ReadAllProfiles",
|
||||
"UpdateProfile",
|
||||
"DeleteProfile",
|
||||
# Vector
|
||||
"AddDraftAndRetrieveSimilarMemory",
|
||||
"AddAndRetrieveSimilarMemory",
|
||||
"AddMemory",
|
||||
"DeleteMemory",
|
||||
"MemoryHandler",
|
||||
|
|
|
|||
71
reme/tool/memory/profiles/add_profile.py
Normal file
71
reme/tool/memory/profiles/add_profile.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
"""Add user profile tool"""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .profile_handler import ProfileHandler
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.schema import ToolCall
|
||||
|
||||
|
||||
class AddProfile(BaseMemoryTool):
|
||||
"""Tool to add a single profile entry"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
"""Build and return the tool call schema"""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Add a new profile entry for the user.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"message_time": {
|
||||
"type": "string",
|
||||
"description": "Message time, e.g. '2020-01-01 00:00:00'",
|
||||
},
|
||||
"profile_key": {
|
||||
"type": "string",
|
||||
"description": "Profile key or category, e.g. 'name'",
|
||||
},
|
||||
"profile_value": {
|
||||
"type": "string",
|
||||
"description": "Profile value or content, e.g. 'John Smith'",
|
||||
},
|
||||
},
|
||||
"required": ["message_time", "profile_key", "profile_value"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
|
||||
|
||||
# Get parameters
|
||||
message_time = self.context.get("message_time", "")
|
||||
profile_key = self.context.get("profile_key", "")
|
||||
profile_value = self.context.get("profile_value", "")
|
||||
|
||||
if not profile_key or not profile_value:
|
||||
return "Missing required parameters (profile_key or profile_value), operation cancelled."
|
||||
|
||||
# Build profile dict
|
||||
profile = {
|
||||
"message_time": message_time,
|
||||
"profile_key": profile_key,
|
||||
"profile_value": profile_value,
|
||||
}
|
||||
|
||||
# Add profile using ProfileHandler
|
||||
new_nodes = profile_handler.add_batch(profiles=[profile], ref_memory_id=self.history_id)
|
||||
self.memory_nodes.extend(new_nodes)
|
||||
|
||||
if new_nodes:
|
||||
output = f"Successfully added profile: [{profile_key}] = {profile_value}"
|
||||
logger.info(output)
|
||||
return output
|
||||
else:
|
||||
output = "Failed to add profile."
|
||||
logger.warning(output)
|
||||
return output
|
||||
53
reme/tool/memory/profiles/delete_profile.py
Normal file
53
reme/tool/memory/profiles/delete_profile.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""Delete user profile tool"""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .profile_handler import ProfileHandler
|
||||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ....core.schema import ToolCall
|
||||
|
||||
|
||||
class DeleteProfile(BaseMemoryTool):
|
||||
"""Tool to delete a single profile entry by ID"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
"""Build and return the tool call schema"""
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Delete a profile entry by profile ID.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"profile_id": {
|
||||
"type": "string",
|
||||
"description": "The unique ID of the profile to delete.",
|
||||
},
|
||||
},
|
||||
"required": ["profile_id"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
|
||||
|
||||
# Get profile_id parameter
|
||||
profile_id = self.context.get("profile_id", "")
|
||||
|
||||
if not profile_id:
|
||||
return "No profile_id provided, operation cancelled."
|
||||
|
||||
# Delete profile using ProfileHandler
|
||||
success = profile_handler.delete(profile_id)
|
||||
|
||||
if success:
|
||||
output = f"Successfully deleted profile with ID: {profile_id}"
|
||||
logger.info(output)
|
||||
return output
|
||||
else:
|
||||
output = f"Profile with ID '{profile_id}' not found."
|
||||
logger.warning(output)
|
||||
return output
|
||||
|
|
@ -12,7 +12,7 @@ from ....core.utils import CacheHandler, deduplicate_memories
|
|||
class ProfileHandler:
|
||||
"""User profile CRUD handler"""
|
||||
|
||||
def __init__(self, profile_path: str | Path, memory_target: str, max_capacity: int = 100):
|
||||
def __init__(self, profile_path: str | Path, memory_target: str, max_capacity: int = 50):
|
||||
"""init"""
|
||||
self.memory_target: str = memory_target
|
||||
self.cache_key: str = self.memory_target.replace(" ", "_").lower()
|
||||
|
|
@ -96,6 +96,12 @@ class ProfileHandler:
|
|||
ref_memory_id=ref_memory_id,
|
||||
)
|
||||
|
||||
# Remove existing nodes with the same when_to_use (profile_key)
|
||||
original_count = len(nodes)
|
||||
nodes = [n for n in nodes if n.when_to_use != profile_key]
|
||||
if len(nodes) < original_count:
|
||||
logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with key: {profile_key}")
|
||||
|
||||
nodes.append(new_node)
|
||||
self._save_nodes(nodes)
|
||||
logger.info(f"Added profile: {profile_key}={profile_value}")
|
||||
|
|
@ -120,6 +126,13 @@ class ProfileHandler:
|
|||
for p in profiles
|
||||
]
|
||||
|
||||
# Remove existing nodes with the same when_to_use (profile_key)
|
||||
new_keys = {n.when_to_use for n in new_nodes}
|
||||
original_count = len(nodes)
|
||||
nodes = [n for n in nodes if n.when_to_use not in new_keys]
|
||||
if len(nodes) < original_count:
|
||||
logger.info(f"Removed {original_count - len(nodes)} duplicate profile(s) with matching keys")
|
||||
|
||||
nodes.extend(new_nodes)
|
||||
self._save_nodes(nodes)
|
||||
logger.info(f"Batch added {len(new_nodes)} profiles")
|
||||
|
|
|
|||
|
|
@ -90,6 +90,7 @@ class UpdateProfile(BaseMemoryTool):
|
|||
if self.enable_memory_target:
|
||||
# Group profiles by memory_target
|
||||
from collections import defaultdict
|
||||
|
||||
profiles_by_target = defaultdict(list)
|
||||
for profile in profiles_to_add:
|
||||
target = profile.get("memory_target", self.memory_target)
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from ....core.schema import ToolCall, MemoryNode
|
|||
from ....core.utils import deduplicate_memories
|
||||
|
||||
|
||||
class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool):
|
||||
class AddAndRetrieveSimilarMemory(BaseMemoryTool):
|
||||
"""Tool to add draft memory and retrieve similar memories"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -60,7 +60,7 @@ class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool):
|
|||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Add draft memory and retrieve similar memories from the vector store.",
|
||||
"description": "Add memory and retrieve similar memories from the vector store.",
|
||||
"parameters": self._build_query_parameters(),
|
||||
},
|
||||
)
|
||||
|
|
@ -68,24 +68,24 @@ class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool):
|
|||
def _build_multiple_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": "Add draft memory and retrieve similar memories from the vector store.",
|
||||
"description": "Add memory and retrieve similar memories from the vector store.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"draft_items": {
|
||||
"items": {
|
||||
"type": "array",
|
||||
"description": "draft_items",
|
||||
"description": "items",
|
||||
"items": self._build_query_parameters(),
|
||||
},
|
||||
},
|
||||
"required": ["draft_items"],
|
||||
"required": ["items"],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
if self.enable_multiple:
|
||||
draft_items = self.context.get("draft_items", [])
|
||||
draft_items = self.context.get("items", [])
|
||||
else:
|
||||
draft_items = [self.context]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue