remove(bench): 删除旧的基准测试分析脚本

- 移除 HaluMem 数据集统计分析脚本 (analyze_dataset_stats.py)
- 移除评估结果分析脚本 (analyze_results.py)
- 移除人工回路问答统计计算脚本 (compute_qa_stats.py)
- 清理重复的人工回路2问答统计脚本
- 删除相关的数据分析和结果统计功能模块
This commit is contained in:
方应 2026-01-30 16:30:28 +08:00
parent 053c537845
commit a0aec042c9
20 changed files with 0 additions and 7491 deletions

View file

@ -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,
)

View file

@ -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
)

View file

@ -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)

View file

@ -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")

View file

@ -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)

View file

@ -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
)

View file

@ -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,
)

View file

@ -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
)

View file

@ -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
)

View file

@ -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

View file

@ -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"
}}
```
"""

View file

@ -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)

View file

@ -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)

View file

@ -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"
}}
```
"""

View file

@ -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
))

View file

@ -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)

View file

@ -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"
}}
```
"""

View file

@ -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
))