Merge pull request #79 from agentscope-ai/dev_0121

Dev 0121
This commit is contained in:
jinliyl 2026-01-23 10:45:12 +08:00 • committed by GitHub
commit 888ccea5f5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
243 changed files with 8758 additions and 4280 deletions

4
.gitignore vendored
View file

@ -36,4 +36,6 @@ test_working_memory/*
local_vector_store/*
chroma_vector_store/*
bench_results/*
meta_memory/*
meta_memory/*
*.sqlite3
**/data/*.json

View file

@ -3,7 +3,7 @@ repos:
rev: v6.0.0
hooks:
- id: check-ast
exclude: ^(test/|cookbook/)
exclude: ^(test/|cookbook/|reme_ai/|bench)
- id: check-yaml
- id: check-xml
- id: check-toml
@ -14,18 +14,18 @@ repos:
rev: v4.0.0
hooks:
- id: add-trailing-comma
exclude: ^(test/|cookbook/)
exclude: ^(test/|cookbook/|reme_ai/|bench)
- repo: https://github.com/psf/black
rev: 25.9.0
hooks:
- id: black
exclude: ^(test/|cookbook/)
exclude: ^(test/|cookbook/|reme_ai/|bench)
args: [--line-length=120]
- repo: https://github.com/PyCQA/flake8
rev: 7.3.0
hooks:
- id: flake8
exclude: ^(test/|cookbook/)
exclude: ^(test/|cookbook/|reme_ai/|bench)
args: [
"--extend-ignore=E203",
"--max-line-length=120"
@ -44,6 +44,8 @@ repos:
| \.demo$
| \.md$
| \.html$
| reme_ai/
| bench
)
args: [
--disable=W0511,
@ -76,6 +78,7 @@ repos:
--disable=C3001,
--disable=R1702,
--disable=R0912,
--max-statements=75,
--max-line-length=120,
]
- repo: https://github.com/regebro/pyroma

View file

@ -29,6 +29,7 @@ class UserStats:
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
@ -37,29 +38,29 @@ 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 数量
@ -68,14 +69,14 @@ class DatasetStats:
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 中提取用户名"""
@ -83,7 +84,7 @@ class DatasetAnalyzer:
if not match:
return "Unknown"
return match.group(1).strip()
@staticmethod
def calculate_dialogue_length(dialogue: list[dict]) -> int:
"""计算对话的总长度(字符数)"""
@ -92,7 +93,7 @@ class DatasetAnalyzer:
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:
"""
@ -101,23 +102,23 @@ class DatasetAnalyzer:
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)
@ -136,93 +137,100 @@ class DatasetAnalyzer:
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
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 = (
@ -232,43 +240,43 @@ class DatasetAnalyzer:
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,
@ -292,16 +300,16 @@ class DatasetAnalyzer:
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 = []
@ -310,36 +318,36 @@ class DatasetAnalyzer:
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}")
@ -347,34 +355,34 @@ class DatasetAnalyzer:
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)
@ -384,24 +392,50 @@ class DatasetAnalyzer:
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 = (
@ -410,17 +444,17 @@ class DatasetAnalyzer:
)
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 = {
@ -469,15 +503,19 @@ class DatasetAnalyzer:
if u.dialogue_lengths_per_session else 0
),
"dialogues_per_session": u.dialogues_per_session,
"dialogue_lengths_per_session": u.dialogue_lengths_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}")
@ -487,24 +525,27 @@ 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)
@ -516,7 +557,7 @@ def main(data_path: str, output_path: str = None, show_per_user: bool = False):
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Analyze HaluMem dataset statistics"
)
@ -537,9 +578,9 @@ if __name__ == "__main__":
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,

View file

@ -13,73 +13,73 @@ 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()
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({
@ -91,42 +91,42 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
"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']}")
@ -135,10 +135,10 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
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 = {
@ -148,16 +148,16 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
"total_questions": total_questions,
"result_type_distribution": dict(result_counter),
"result_type_ratio": {
result_type: count / total_questions
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)
@ -165,7 +165,7 @@ def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="分析 ReMe 评估结果中的 result_type 分布"
)
@ -175,6 +175,6 @@ if __name__ == "__main__":
default="bench_results/reme_simple/tmp",
help="临时结果目录路径 (默认: bench_results/reme_simple/tmp)"
)
args = parser.parse_args()
analyze_results(args.tmp_dir)

View file

@ -0,0 +1,285 @@
"""
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

@ -9,7 +9,7 @@ are available in the tmp directory. It will:
4. Aggregate results and compute metrics
Usage:
python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme/tmp
python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme_simple_v4/tmp
"""
import asyncio
@ -358,7 +358,7 @@ async def main_async(tmp_dir: str):
# 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")
@ -392,7 +392,7 @@ async def main_async(tmp_dir: str):
# 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")

View file

@ -44,13 +44,13 @@ class EvalConfig:
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."""
@ -58,7 +58,7 @@ class DataLoader:
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."""
@ -74,7 +74,7 @@ class DataLoader:
}
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)."""
@ -83,14 +83,14 @@ class DataLoader:
# 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"
@ -101,29 +101,29 @@ class DataLoader:
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)
@ -131,38 +131,38 @@ class FileManager:
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"
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()
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:
@ -171,7 +171,7 @@ class FileManager:
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")
@ -206,10 +206,10 @@ Please respond in JSON format with the following structure:
class BaselineQuestionAnsweringEvaluator:
"""Evaluates question answering performance using direct LLM inference (no memory system)."""
def __init__(self):
pass
async def answer_question(
self,
question: str,
@ -217,18 +217,18 @@ class BaselineQuestionAnsweringEvaluator:
) -> 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"
@ -240,10 +240,10 @@ class BaselineQuestionAnsweringEvaluator:
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],
@ -254,14 +254,14 @@ class BaselineQuestionAnsweringEvaluator:
) -> 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(
@ -271,7 +271,7 @@ class BaselineQuestionAnsweringEvaluator:
answer,
formatted_dialogue
)
# Build result record
qa_result = {
**qa,
@ -284,13 +284,13 @@ class BaselineQuestionAnsweringEvaluator:
"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."""
@ -306,15 +306,15 @@ class MetricsAggregator:
"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":
@ -323,7 +323,7 @@ class MetricsAggregator:
hallucination += 1
elif result_type == "Omission":
omission += 1
metrics = {
"correct_qa_ratio(all)": correct / total,
"hallucination_qa_ratio(all)": hallucination / total,
@ -331,7 +331,7 @@ class MetricsAggregator:
"qa_valid_num": valid,
"qa_num": total
}
if valid > 0:
metrics.update({
"correct_qa_ratio(valid)": correct / valid,
@ -344,25 +344,25 @@ class MetricsAggregator:
"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,
@ -374,13 +374,13 @@ class MetricsAggregator:
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,
@ -395,16 +395,16 @@ class HaluMemBaselineEvaluator:
"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)
@ -415,25 +415,25 @@ class HaluMemBaselineEvaluator:
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(
@ -442,65 +442,65 @@ class HaluMemBaselineEvaluator:
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:
@ -508,20 +508,20 @@ class HaluMemBaselineEvaluator:
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,
@ -529,23 +529,23 @@ class HaluMemBaselineEvaluator:
},
"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}")
@ -554,7 +554,7 @@ class HaluMemBaselineEvaluator:
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")
@ -574,14 +574,14 @@ def main(
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"
)
@ -603,9 +603,9 @@ if __name__ == "__main__":
default=2,
help="Maximum concurrent user processing (default: 2)"
)
args = parser.parse_args()
main(
data_path=args.data_path,
user_num=args.user_num,

View file

@ -685,7 +685,7 @@ async def main_async(
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)
@ -694,11 +694,11 @@ async def main_async(
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']})")
@ -733,7 +733,7 @@ async def main_async(
# 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")

View file

@ -48,13 +48,13 @@ class EvalConfig:
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."""
@ -62,7 +62,7 @@ class DataLoader:
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."""
@ -78,7 +78,7 @@ class DataLoader:
}
for turn in dialogue
]
@staticmethod
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
"""Format dialogue into string for evaluation."""
@ -87,10 +87,10 @@ class DataLoader:
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"
@ -101,29 +101,29 @@ class DataLoader:
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)
@ -131,38 +131,38 @@ class FileManager:
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"
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()
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:
@ -171,7 +171,7 @@ class FileManager:
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")
@ -179,19 +179,19 @@ class FileManager:
class MemoryProcessor:
"""Handles ReMe memory operations."""
def __init__(self, reme: ReMe):
self.reme = reme
async def add_memories(
self,
user_id: str,
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)
"""
@ -202,19 +202,19 @@ class MemoryProcessor:
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,
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:
@ -230,23 +230,23 @@ class MemoryProcessor:
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,
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,
query=query,
user_id=user_id,
top_k=top_k
)
duration_ms = (time.time() - start) * 1000
@ -257,11 +257,11 @@ class MemoryProcessor:
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],
@ -272,7 +272,7 @@ class QuestionAnsweringEvaluator:
) -> 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(
@ -280,7 +280,7 @@ class QuestionAnsweringEvaluator:
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(
@ -290,7 +290,7 @@ class QuestionAnsweringEvaluator:
response,
formatted_dialogue
)
# Build result record
qa_result = {
**qa,
@ -303,13 +303,13 @@ class QuestionAnsweringEvaluator:
"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."""
@ -325,15 +325,15 @@ class MetricsAggregator:
"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":
@ -342,7 +342,7 @@ class MetricsAggregator:
hallucination += 1
elif result_type == "Omission":
omission += 1
metrics = {
"correct_qa_ratio(all)": correct / total,
"hallucination_qa_ratio(all)": hallucination / total,
@ -350,7 +350,7 @@ class MetricsAggregator:
"qa_valid_num": valid,
"qa_num": total
}
if valid > 0:
metrics.update({
"correct_qa_ratio(valid)": correct / valid,
@ -363,28 +363,28 @@ class MetricsAggregator:
"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,
@ -397,18 +397,18 @@ class MetricsAggregator:
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,
self.memory_processor,
config.top_k
)
self.data_loader = DataLoader()
async def process_session(
self,
session: dict,
@ -423,29 +423,29 @@ class HaluMemEvaluator:
"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)
@ -456,90 +456,90 @@ class HaluMemEvaluator:
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)
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:
@ -547,20 +547,20 @@ class HaluMemEvaluator:
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,
@ -568,23 +568,23 @@ class HaluMemEvaluator:
},
"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}")
@ -593,7 +593,7 @@ class HaluMemEvaluator:
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")
@ -616,14 +616,14 @@ def main(
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)"
)
@ -651,9 +651,9 @@ if __name__ == "__main__":
default=2,
help="Maximum concurrent user processing (default: 2)"
)
args = parser.parse_args()
main(
data_path=args.data_path,
top_k=args.top_k,

View file

@ -0,0 +1,676 @@
"""
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

@ -0,0 +1,690 @@
"""
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_v4.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, answer_question_with_memories
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_v4"
# ==================== 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 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 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_v4(
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[dict, list, float]:
"""
Search memory using ReMe and return structured answer with reasoning.
Returns:
tuple: (answer_dict, agent_messages, duration_ms)
answer_dict contains: {"reasoning": str, "answer": str, "memories": str}
"""
start = time.time()
# Retrieve memories from ReMe
memories_response, agent_messages, success = await self.reme.retrieve_v4(
query=query,
user_id=user_id,
top_k=top_k
)
# Use LLM to generate structured answer from memories
answer_result = await answer_question_with_memories(
question=query,
memories=memories_response,
user_id=user_id
)
# Add original memories to the result
answer_result["memories"] = memories_response
duration_ms = (time.time() - start) * 1000
return answer_result, 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:
answer_dict, agent_messages, duration_ms = await self.memory_processor.search_memory(
query=qa["question"],
user_id=user_name,
top_k=self.top_k
)
# Extract answer and reasoning from the structured response
system_answer = answer_dict.get("answer", "")
system_reasoning = answer_dict.get("reasoning", "")
retrieved_memories = answer_dict.get("memories", "")
# 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,
system_answer,
formatted_dialogue
)
# Build result record
qa_result = {
**qa,
"uuid": uuid,
"session_id": session_id,
"system_response": system_answer,
"system_reasoning": system_reasoning,
"retrieved_memories": retrieved_memories,
"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 HaluMemEvaluatorV4:
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."""
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
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."""
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 - 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")
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 evaluation."""
config = EvalConfig(
data_path=data_path,
top_k=top_k,
user_num=user_num,
max_concurrency=max_concurrency
)
evaluator = HaluMemEvaluatorV4(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

@ -120,7 +120,8 @@ async def evaluation_for_question2(
dialogue: The formatted dialogue history (role, content, time_created).
"""
prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format(
# 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,
@ -128,6 +129,43 @@ async def evaluation_for_question2(
dialogue=dialogue,
)
result = await llm_request_for_json(prompt)
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

@ -2,6 +2,54 @@ 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.
@ -419,15 +467,12 @@ EVALUATION_PROMPT_FOR_QUESTION: |
"evaluation_result": "Correct | Hallucination | Omission"
}}
```
EVALUATION_PROMPT_FOR_QUESTION2: |
You are an **evaluation expert for AI memory system question answering**.
**Dialogue:**
{dialogue}
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.
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
@ -435,36 +480,45 @@ EVALUATION_PROMPT_FOR_QUESTION2: |
### 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.
* 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.”
* 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**.
* 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.”
* 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**.
* 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**.
* 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.
* 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
@ -487,7 +541,7 @@ EVALUATION_PROMPT_FOR_QUESTION2: |
```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.",
"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

@ -28,12 +28,12 @@ reme = ReMe()
)
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
"""
@ -58,17 +58,18 @@ async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str:
reraise=True,
before_sleep=before_sleep_log(logger, logging.WARNING),
)
async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs):
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
"""

View file

@ -0,0 +1,237 @@
"""
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

@ -0,0 +1,145 @@
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

@ -0,0 +1,550 @@
"""
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

@ -0,0 +1,237 @@
"""
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

@ -0,0 +1,145 @@
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

@ -0,0 +1,550 @@
"""
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
))

3
docs/todo.md Normal file
View file

@ -0,0 +1,3 @@
1. 如何更好的注册class
2. op的返回,使用return 还是 self.output
3. 如何把agent的东西放出来

View file

@ -96,7 +96,7 @@ When context grows too large, model performance degrades significantly—a pheno
### Usage Pattern
For complete working examples of how to use MessageOffloadOp in practice, please refer to:
[test_message_offload_op.py](../../test_op/test_message_offload_op.py)
[test_message_offload_op.py](../../test/test_message_offload_op.py)
This test file demonstrates:
- **Compact mode**: How to configure and use compaction-only strategy

View file

@ -114,7 +114,7 @@ Example: Reading `/workspace/context_store/tool_call_123.txt` with `offset=0` an
## Usage Pattern: Combining Grep and ReadFile
For a complete working example of how to use these operations in practice, please refer to:
[test_agentic_retrieve_op.py](../../test_op/test_agentic_retrieve_op.py)
[test_agentic_retrieve_op.py](../../test/test_agentic_retrieve_op.py)
This test file demonstrates:
- How to configure the system prompt to guide AI in using Grep and ReadFile operations

View file

@ -57,7 +57,7 @@ full = [
[tool.setuptools.packages.find]
where = ["."]
include = ["reme_ai*"]
include = ["reme_ai*", "reme*"]
exclude = ["test*", "cookbook*", "doc*", "library*", "dist*"]
[tool.setuptools.package-data]

19
reme/__init__.py Normal file
View file

@ -0,0 +1,19 @@
"""ReMe"""
from . import agent
from . import config
from . import core
from . import tool
from . import workflow
from .reme_app import ReMeApp
__all__ = [
"agent",
"config",
"core",
"tool",
"workflow",
"ReMeApp",
]
__version__ = "0.3.0.0a1"

7
reme/agent/__init__.py Normal file
View file

@ -0,0 +1,7 @@
"""A simple chatbot."""
from . import chat
__all__ = [
"chat",
]

View file

@ -1,11 +1,13 @@
"""chat agent"""
from .remy_agent import ReMyAgent
from .simple_chat import SimpleChat
from .stream_chat import StreamChat
from ...core import R
__all__ = [
"ReMyAgent",
"StreamChat",
"SimpleChat",
]
R.op.register("simple_chat")(SimpleChat)
R.op.register("stream_chat")(StreamChat)

View file

@ -2,14 +2,12 @@
from loguru import logger
from ...core.context import C
from ...core.enumeration import Role
from ...core.op import BaseOp
from ...core.op import BaseTool
from ...core.schema import Message, ToolCall
@C.register_op()
class SimpleChat(BaseOp):
class SimpleChat(BaseTool):
"""Simple chat agent that handles non-streaming conversations."""
def _build_tool_call(self) -> ToolCall:
@ -59,4 +57,4 @@ class SimpleChat(BaseOp):
logger.info(f"messages={messages}")
assistant_message = await self.llm.chat(messages=messages)
logger.info(f"assistant_message={assistant_message.simple_dump()}")
self.output = assistant_message.content
return assistant_message.content

View file

@ -2,14 +2,12 @@
from loguru import logger
from ...core.context import C
from ...core.enumeration import Role, ChunkEnum
from ...core.op import BaseOp
from ...core.op import BaseTool
from ...core.schema import Message, ToolCall
@C.register_op()
class StreamChat(BaseOp):
class StreamChat(BaseTool):
"""Streaming chat agent that handles real-time conversation streaming."""
def _build_tool_call(self) -> ToolCall:

View file

@ -16,13 +16,14 @@ llm:
default:
backend: openai
model_name: qwen3-30b-a3b-instruct-2507
max_concurrency: 20
# model_name: qwen-flash
request_interval: 1
temperature: 0.0001
qwen3_max_instruct:
backend: openai
model_name: qwen3-max
# temperature: 0.6
max_concurrency: 20
request_interval: 2
embedding_model:
default:

View file

@ -1,6 +1,6 @@
"""Configuration parser for ReMe framework."""
from ..utils import PydanticConfigParser
from ..core.utils import PydanticConfigParser
class ReMeConfigParser(PydanticConfigParser):

View file

@ -1,9 +1,5 @@
"""Core module for ReMe AI framework."""
"""Core"""
# pylint: disable=wrong-import-position
# flake8: noqa: F401
from . import config
from . import context
from . import embedding
from . import enumeration
@ -15,3 +11,19 @@ from . import service
from . import token_counter
from . import utils
from . import vector_store
from .context import R
__all__ = [
"context",
"embedding",
"enumeration",
"flow",
"llm",
"op",
"schema",
"service",
"token_counter",
"utils",
"vector_store",
"R",
]

View file

@ -2,15 +2,14 @@
from .base_context import BaseContext
from .prompt_handler import PromptHandler
from .registry import Registry
from .registry_factory import R
from .runtime_context import RuntimeContext
from .service_context import ServiceContext, C
from .service_context import ServiceContext
__all__ = [
"BaseContext",
"PromptHandler",
"Registry",
"R",
"RuntimeContext",
"ServiceContext",
"C",
]

View file

@ -0,0 +1,363 @@
"""Module for managing and formatting prompt templates from files or dictionaries.
This module provides a PromptHandler class that:
- Loads prompts from YAML/JSON files or dictionaries
- Supports multi-language prompts with automatic suffix handling
- Provides conditional line filtering using boolean flags
- Formats prompts with template variable substitution
- Validates format strings and provides helpful error messages
"""
import json
from pathlib import Path
from string import Formatter
from typing import Any, Dict, Optional, Union
import yaml
from loguru import logger
from .base_context import BaseContext
class PromptNotFoundError(KeyError):
"""Exception raised when a requested prompt template is not found."""
def __init__(self, prompt_name: str, available_prompts: list[str]):
self.prompt_name = prompt_name
self.available_prompts = available_prompts
super().__init__(
f"Prompt '{prompt_name}' not found. "
f"Available prompts: {', '.join(available_prompts[:10])}"
f"{'...' if len(available_prompts) > 10 else ''}",
)
class PromptFormattingError(ValueError):
"""Exception raised when prompt formatting fails."""
class PromptHandler(BaseContext):
"""A context-aware handler for loading, retrieving, and formatting prompt templates.
This handler supports:
- Loading prompts from YAML/JSON files or dictionaries
- Multi-language prompt support with automatic language suffix
- Conditional line filtering using boolean flags (e.g., [debug], [verbose])
- Template variable substitution with validation
- Method chaining for fluent API
Examples:
>>> handler = PromptHandler(language="en")
>>> handler.load_prompt_dict({
... "greeting_en": "Hello, {name}!",
... "farewell_en": "[debug]Debug mode\\nGoodbye, {name}!"
... })
>>> handler.prompt_format("greeting", name="Alice")
'Hello, Alice!'
>>> handler.prompt_format("farewell", name="Bob", debug=False)
'Goodbye, Bob!'
"""
def __init__(self, language: str = "", **kwargs):
"""Initialize the PromptHandler with optional language configuration.
Args:
language: Language code to append as suffix (e.g., "en", "zh", "ja").
If provided, get_prompt will automatically try to find
prompts with this suffix (e.g., "greeting" -> "greeting_en").
**kwargs: Additional key-value pairs to initialize the context.
"""
super().__init__(**kwargs)
self.language: str = language.strip()
def load_prompt_by_file(
self,
prompt_file_path: Optional[Union[Path, str]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Load prompt configurations from a YAML or JSON file into the context.
Supports both YAML (.yaml, .yml) and JSON (.json) file formats.
Non-existent files are silently skipped.
Args:
prompt_file_path: Path to the prompt configuration file.
If None, returns self without changes.
overwrite: If True, allows overwriting existing prompts with warnings.
If False, skips existing prompts without overwriting.
Returns:
Self for method chaining.
Raises:
ValueError: If file format is not supported.
yaml.YAMLError: If YAML parsing fails.
json.JSONDecodeError: If JSON parsing fails.
"""
if prompt_file_path is None:
return self
if isinstance(prompt_file_path, str):
prompt_file_path = Path(prompt_file_path)
if not prompt_file_path.exists():
logger.warning(f"Prompt file not found: {prompt_file_path}")
return self
suffix = prompt_file_path.suffix.lower()
try:
with prompt_file_path.open(encoding="utf-8") as f:
if suffix in [".yaml", ".yml"]:
prompt_dict = yaml.safe_load(f)
elif suffix == ".json":
prompt_dict = json.load(f)
else:
raise ValueError(
f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json",
)
logger.info(f"Loaded {len(prompt_dict or {})} prompts from {prompt_file_path}")
self.load_prompt_dict(prompt_dict, overwrite=overwrite)
except (yaml.YAMLError, json.JSONDecodeError) as e:
logger.error(f"Failed to parse prompt file {prompt_file_path}: {e}")
raise
return self
def load_prompt_dict(
self,
prompt_dict: Optional[Dict[str, Any]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Merge a dictionary of prompt strings into the current context.
Only string values are stored as prompts. Non-string values are skipped.
Args:
prompt_dict: Dictionary mapping prompt names to prompt template strings.
overwrite: If True, allows overwriting existing prompts with warnings.
If False, skips existing prompts without overwriting.
Returns:
Self for method chaining.
"""
if not prompt_dict:
return self
for key, value in prompt_dict.items():
if not isinstance(value, str):
logger.debug(f"Skipping non-string prompt: key={key}, type={type(value)}")
continue
if key in self:
if overwrite:
logger.warning(
f"Overwriting prompt '{key}': " f"old length={len(self[key])}, new length={len(value)}",
)
self[key] = value
else:
logger.debug(f"Skipping existing prompt: key={key}")
else:
logger.debug(f"Adding new prompt: key={key}, length={len(value)}")
self[key] = value
return self
def get_prompt(self, prompt_name: str, fallback_to_base: bool = True) -> str:
"""Retrieve a prompt by name with automatic language suffix handling.
If a language is configured, this method will:
1. First try to find the prompt with language suffix (e.g., "greeting_en")
2. If not found and fallback_to_base is True, try the base name (e.g., "greeting")
3. Otherwise, raise PromptNotFoundError
Args:
prompt_name: Name of the prompt to retrieve.
fallback_to_base: If True and language-specific prompt not found,
fallback to prompt without language suffix.
Returns:
The prompt template string, stripped of leading/trailing whitespace.
Raises:
PromptNotFoundError: If the prompt is not found.
"""
# Try with language suffix first
if self.language and not prompt_name.endswith(f"_{self.language}"):
key_with_lang = f"{prompt_name}_{self.language}"
if key_with_lang in self:
return self[key_with_lang].strip()
# Try base name
if prompt_name in self:
return self[prompt_name].strip()
# Try fallback if enabled
if fallback_to_base and self.language:
# Check if prompt_name already has language suffix, try without it
if prompt_name.endswith(f"_{self.language}"):
base_name = prompt_name[: -(len(self.language) + 1)]
if base_name in self:
return self[base_name].strip()
# Not found, raise error with helpful message
available = list(self.keys())
raise PromptNotFoundError(prompt_name, available)
def has_prompt(self, prompt_name: str) -> bool:
"""Check if a prompt exists (with or without language suffix).
Args:
prompt_name: Name of the prompt to check.
Returns:
True if the prompt exists, False otherwise.
"""
try:
self.get_prompt(prompt_name)
return True
except PromptNotFoundError:
return False
def list_prompts(self, language_filter: Optional[str] = None) -> list[str]:
"""List all available prompt names.
Args:
language_filter: If provided, only return prompts for this language.
If None, return all prompts.
Returns:
List of prompt names.
"""
if language_filter is None:
return list(self.keys())
suffix = f"_{language_filter.strip()}"
return [key for key in self.keys() if key.endswith(suffix)]
@staticmethod
def _extract_format_fields(template: str) -> set[str]:
"""Extract all format field names from a template string.
Args:
template: Template string with {variable} placeholders.
Returns:
Set of field names used in the template.
"""
return {field_name for _, field_name, _, _ in Formatter().parse(template) if field_name is not None}
@staticmethod
def _filter_conditional_lines(prompt: str, flags: Dict[str, bool]) -> str:
"""Filter lines based on boolean flags.
Lines starting with [flag_name] are conditionally included based on
the value of flags[flag_name]. If True, the line is included (without
the flag marker). If False, the line is excluded.
Args:
prompt: The prompt text with conditional markers.
flags: Dictionary of flag names to boolean values.
Returns:
Filtered prompt text.
"""
filtered_lines = []
for line in prompt.split("\n"):
# Check each flag
matched_flag = None
for flag_name in flags:
marker = f"[{flag_name}]"
if line.startswith(marker):
matched_flag = flag_name
break
if matched_flag is None:
# No flag marker, always include
filtered_lines.append(line)
elif flags[matched_flag]:
# Flag is True, include without marker
marker = f"[{matched_flag}]"
filtered_lines.append(line[len(marker) :])
# else: Flag is False, skip this line
return "\n".join(filtered_lines)
def prompt_format(
self,
prompt_name: str,
validate: bool = True,
**kwargs,
) -> str:
"""Format a prompt with conditional line filtering and variable substitution.
This method performs two-stage formatting:
1. Conditional line filtering: Lines marked with [flag] are included only
if the corresponding boolean kwarg is True.
2. Variable substitution: Template variables {var} are replaced with
provided values.
Args:
prompt_name: Name of the prompt to format.
validate: If True, check that all required template variables are provided.
**kwargs: Keyword arguments for formatting. Boolean values are treated as
conditional flags, other values are used for template substitution.
Returns:
Formatted prompt string.
Raises:
PromptNotFoundError: If the prompt is not found.
PromptFormattingError: If validation fails or formatting errors occur.
Examples:
>>> handler = PromptHandler()
>>> handler["test"] = "[debug]Debug: {info}\\nResult: {value}"
>>> handler.prompt_format("test", debug=False, info="test", value=42)
'Result: 42'
>>> handler.prompt_format("test", debug=True, info="test", value=42)
'Debug: test\\nResult: 42'
"""
# Get the prompt template
prompt = self.get_prompt(prompt_name)
# Separate boolean flags from format variables
flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
format_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
# Step 1: Filter conditional lines
if flag_kwargs:
prompt = self._filter_conditional_lines(prompt, flag_kwargs)
# Step 2: Validate required fields if requested
if validate:
required_fields = self._extract_format_fields(prompt)
missing_fields = required_fields - set(format_kwargs.keys())
if missing_fields:
raise PromptFormattingError(
f"Missing required format variables for prompt '{prompt_name}': "
f"{', '.join(sorted(missing_fields))}",
)
# Step 3: Format with variables
try:
if format_kwargs:
prompt = prompt.format(**format_kwargs)
except KeyError as e:
raise PromptFormattingError(
f"Format error in prompt '{prompt_name}': missing variable {e}",
) from e
except (ValueError, IndexError) as e:
raise PromptFormattingError(
f"Format error in prompt '{prompt_name}': {e}",
) from e
return prompt.strip()
def __repr__(self) -> str:
"""Return a string representation of the PromptHandler."""
return f"PromptHandler(language='{self.language}', " f"num_prompts={len(self)})"

View file

@ -0,0 +1,45 @@
"""Module providing a registry class for managing class-to-name mappings via decorators."""
import inspect
from typing import Callable, TypeVar
from .base_context import BaseContext
from ..utils import singleton
T = TypeVar("T")
class Registry(BaseContext):
"""A registry container that uses decorators to map and store class references."""
def register(self, name: str | type = "") -> Callable[[type[T]], type[T]] | type[T]:
"""Return a decorator that registers a class under a specific name in the registry."""
if inspect.isclass(name):
self[name.__name__] = name
return name
else:
def decorator(cls):
key: str = name if isinstance(name, str) and name else cls.__name__
self[key] = cls
return cls
return decorator
@singleton
class RegistryFactory:
"""A factory class for creating registries."""
def __init__(self):
self.llm = Registry()
self.embedding_model = Registry()
self.vector_store = Registry()
self.op = Registry()
self.flow = Registry()
self.service = Registry()
self.token_counter = Registry()
R = RegistryFactory()

View file

@ -3,6 +3,7 @@
import asyncio
from .base_context import BaseContext
from .service_context import ServiceContext
from ..enumeration import ChunkEnum
from ..schema import Response, StreamChunk
@ -14,21 +15,24 @@ class RuntimeContext(BaseContext):
self,
response: Response | None = None,
stream_queue: asyncio.Queue | None = None,
service_context: ServiceContext | None = None,
**kwargs,
):
"""Initialize the context with optional response and queue."""
super().__init__(**kwargs)
self.response = response or Response()
self.stream_queue = stream_queue
self.response: Response | None = response or Response()
self.stream_queue: asyncio.Queue | None = stream_queue
self.service_context: ServiceContext | None = service_context
@classmethod
def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext":
"""Create a new context from an existing instance or keywords."""
if context is None:
return cls(**kwargs)
context.update(kwargs)
return context
else:
if kwargs:
context.update(kwargs)
return context
async def _enqueue(self, chunk: StreamChunk) -> None:
"""Internal helper to put a chunk into the queue if it exists."""

View file

@ -0,0 +1,230 @@
"""Service context."""
import os
from concurrent.futures import ThreadPoolExecutor
from loguru import logger
from .base_context import BaseContext
from .registry_factory import R
from ..schema import ServiceConfig
from ..utils import MCPClient, print_logo, PydanticConfigParser, init_logger, load_env, run_coro_safely
class ServiceContext(BaseContext):
"""Service context."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_api_base: str | None = None,
embedding_api_key: str | None = None,
embedding_api_base: str | None = None,
service_config: ServiceConfig | None = None,
parser: type[PydanticConfigParser] | None = None,
config_path: str | None = None,
enable_logo: bool = True,
llm: dict | None = None,
embedding_model: dict | None = None,
vector_store: dict | None = None,
token_counter: dict | None = None,
**kwargs,
):
super().__init__()
# Set environment variables
load_env()
self._update_env("REME_LLM_API_KEY", llm_api_key)
self._update_env("REME_LLM_BASE_URL", llm_api_base)
self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base)
# Use default parser if not provided
parser_class = parser if parser is not None else PydanticConfigParser
self.parser = parser_class(ServiceConfig)
# Service configuration
if service_config is None:
input_args = []
if config_path:
input_args.append(f"config={config_path}")
if args:
input_args.extend(args)
if kwargs:
input_args.extend([f"{k}={v}" for k, v in kwargs.items()])
service_config = self.parser.parse_args(*input_args)
self.service_config: ServiceConfig = service_config
# Initialize logger
if self.service_config.init_logger:
init_logger()
# Update service config with provided arguments
if llm:
self.update_section_config("llm", **llm)
if embedding_model:
self.update_section_config("embedding_model", **embedding_model)
if token_counter:
self.update_section_config("token_counter", **token_counter)
if vector_store:
self.update_section_config("vector_store", **vector_store)
# Print the ReMe logo if enabled in configuration.
self.service_config.enable_logo = enable_logo
if self.service_config.enable_logo:
print_logo(service_config=self.service_config)
# Service configuration and runtime settings
self.language: str = self.service_config.language
self.thread_pool: ThreadPoolExecutor = ThreadPoolExecutor(max_workers=service_config.thread_pool_max_workers)
# Initialize Ray for distributed computing if configured
if self.service_config.ray_max_workers > 1:
import ray
ray.init(num_cpus=self.service_config.ray_max_workers)
from ..llm import BaseLLM
from ..embedding import BaseEmbeddingModel
from ..vector_store import BaseVectorStore
from ..token_counter import BaseTokenCounter
from ..flow import BaseFlow, ExpressionFlow
from ..service import BaseService
# Initialize LLM instances
self.llms: dict[str, BaseLLM] = {}
for name, config in self.service_config.llm.items():
self.llms[name] = R.llm[config.backend](model_name=config.model_name, **config.model_extra)
# Initialize Embedding model instances
self.embedding_models: dict[str, BaseEmbeddingModel] = {}
for name, config in self.service_config.embedding_model.items():
self.embedding_models[name] = R.embedding_model[config.backend](
model_name=config.model_name,
**config.model_extra,
)
# Initialize Token counter instances
self.token_counters: dict[str, BaseTokenCounter] = {}
for name, config in self.service_config.token_counter.items():
self.token_counters[name] = R.token_counter[config.backend](
model_name=config.model_name,
**config.model_extra,
)
# Initialize Vector store instances
self.vector_stores: dict[str, BaseVectorStore] = {}
for name, config in self.service_config.vector_store.items():
self.vector_stores[name] = R.vector_store[config.backend](
collection_name=config.collection_name,
embedding_model=self.embedding_models[config.embedding_model],
thread_pool=self.thread_pool,
**config.model_extra,
)
# Initialize flow instances
self.flows: dict[str, BaseFlow] = {}
for name, flow_cls in R.flow.items():
if not self._filter_flows(name):
continue
flow: "BaseFlow" = flow_cls(name=name, service_context=self)
self.flows[flow.name] = flow
# Initialize flow instances from service config
for name, flow_config in self.service_config.flow.items():
if not self._filter_flows(name):
continue
flow_config.name = name
flow: BaseFlow = ExpressionFlow(flow_config=flow_config, service_context=self)
self.flows[flow.name] = flow
# Initialize service instance
self.service: BaseService = R.service[self.service_config.backend](service_context=self)
# MCP server mapping: maps server_name -> {tool_name: ToolCall}
if self.service_config.mcp_servers:
self.mcp_server_mapping: dict[str, dict] = run_coro_safely(self.prepare_mcp_servers())
else:
self.mcp_server_mapping: dict[str, dict] = {}
@staticmethod
def _update_env(key: str, value: str | None):
"""Update environment variable if value is provided."""
if value:
os.environ[key] = value
def update_section_config(self, section_name: str, **kwargs):
"""Update a specific section of the service config with new values."""
section_dict: dict = getattr(self.service_config, section_name)
if "default" not in section_dict:
raise KeyError(f"Default `{section_name}` config not found")
current_config = section_dict["default"]
section_dict["default"] = current_config.model_copy(update=kwargs, deep=True)
def _filter_flows(self, name: str) -> bool:
"""Filter flows based on enabled_flows and disabled_flows configuration."""
if self.service_config.enabled_flows:
return name in self.service_config.enabled_flows
elif self.service_config.disabled_flows:
return name not in self.service_config.disabled_flows
else:
return True
async def prepare_mcp_servers(self):
"""Prepare and initialize MCP server connections."""
mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers})
for server_name in self.service_config.mcp_servers.keys():
try:
# Retrieve all available tool calls from this MCP server
tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False)
# Build mapping: tool_name -> ToolCall for quick lookup
self.mcp_server_mapping[server_name] = {tool_call.name: tool_call for tool_call in tool_calls}
# Log discovered tools for debugging
for tool_call in tool_calls:
logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}")
except Exception as e:
logger.exception(f"list_tool_calls: {server_name} error: {e}")
async def close(self):
"""Close all service components asynchronously."""
for _, vector_store in self.vector_stores.items():
await vector_store.close()
for _, llm in self.llms.items():
await llm.close()
for _, embedding_model in self.embedding_models.items():
await embedding_model.close()
self.shutdown_thread_pool()
self.shutdown_ray()
def close_sync(self):
"""Close all service components synchronously."""
for _, vector_store in self.vector_stores.items():
run_coro_safely(vector_store.close())
for _, llm in self.llms.items():
llm.close_sync()
for _, embedding_model in self.embedding_models.items():
embedding_model.close_sync()
self.shutdown_thread_pool()
self.shutdown_ray()
def shutdown_thread_pool(self, wait: bool = True):
"""Shutdown the thread pool executor."""
if self.thread_pool:
self.thread_pool.shutdown(wait=wait)
def shutdown_ray(self, wait: bool = True):
"""Shutdown Ray cluster if it was initialized."""
if self.service_config and self.service_config.ray_max_workers > 1:
import ray
ray.shutdown(_exiting_interpreter=not wait)

View file

@ -3,9 +3,13 @@
from .base_embedding_model import BaseEmbeddingModel
from .openai_embedding_model import OpenAIEmbeddingModel
from .openai_embedding_model_sync import OpenAIEmbeddingModelSync
from ..context import R
__all__ = [
"BaseEmbeddingModel",
"OpenAIEmbeddingModel",
"OpenAIEmbeddingModelSync",
]
R.embedding_model.register("openai")(OpenAIEmbeddingModel)
R.embedding_model.register("openai_sync")(OpenAIEmbeddingModelSync)

View file

@ -6,10 +6,8 @@ from typing import Literal
from openai import AsyncOpenAI
from .base_embedding_model import BaseEmbeddingModel
from ..context import C
@C.register_embedding_model("openai")
class OpenAIEmbeddingModel(BaseEmbeddingModel):
"""Asynchronous embedding model implementation compatible with OpenAI-style APIs."""

View file

@ -3,10 +3,8 @@
from openai import OpenAI
from .openai_embedding_model import OpenAIEmbeddingModel
from ..context import C
@C.register_embedding_model("openai_sync")
class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel):
"""Synchronous embedding model implementation that extends the asynchronous OpenAI model."""

View file

@ -0,0 +1,38 @@
"""Defines the standard data types supported by JSON Schema.
This enum maps common JSON Schema primitive types to their corresponding
Python runtime types, and provides a convenient string representation
compatible with JSON Schema (`"string"`, `"number"`, etc.).
"""
from enum import Enum
class JsonSchemaEnum(Enum):
"""Enumeration of valid JSON Schema data types.
The enum value is the corresponding Python type, while the string
representation (`str(...)`) is the canonical JSON Schema type name.
"""
# Textual data
STRING = str
# Numeric values, including integers and floats
NUMBER = float
# Integer-only numeric values
INTEGER = int
# JSON objects (key-value mappings)
OBJECT = dict
# Ordered JSON lists/arrays
ARRAY = list
# Boolean values: true / false
BOOLEAN = bool
def __str__(self) -> str:
"""Return the lowercase JSON Schema type name for this enum member."""
return self.name.lower()

View file

@ -0,0 +1,33 @@
"""Defines the high-level categories of memory managed by ReMe.
This enumeration is used across the system to tag, route, and store different
kinds of memories (identity, personal context, procedures, tools, etc.).
"""
from enum import Enum
class MemoryType(str, Enum):
"""Enumeration of memory categories used by the memory subsystem.
These types describe *what* a piece of memory is about, which guides
storage, retrieval, and summarization strategies.
"""
# Long‑term, relatively stable attributes about the user (name, roles, etc.)
IDENTITY = "identity"
# User-specific preferences, habits, and evolving personal context
PERSONAL = "personal"
# How‑to knowledge, workflows, and step‑by‑step instructions
PROCEDURAL = "procedural"
# Information learned about tools, APIs, and their usage patterns
TOOL = "tool"
# Condensed representation of larger memory collections
SUMMARY = "summary"
# Raw chronological interaction history, typically before summarization
HISTORY = "history"

View file

@ -3,11 +3,9 @@
from .base_flow import BaseFlow
from .cmd_flow import CmdFlow
from .expression_flow import ExpressionFlow
from .simple_flow import SimpleFlow
__all__ = [
"BaseFlow",
"CmdFlow",
"ExpressionFlow",
"SimpleFlow",
]

View file

@ -7,30 +7,25 @@ from abc import ABC, abstractmethod
from loguru import logger
from ..context import C, RuntimeContext
from ..enumeration import ChunkEnum, RegistryEnum
from ..context import RuntimeContext, ServiceContext, R
from ..enumeration import ChunkEnum
from ..op import BaseOp, SequentialOp, ParallelOp
from ..schema import Response, ToolCall, ToolAttr
from ..schema import Response, ToolCall
from ..utils import camel_to_snake, CacheHandler
class BaseFlow(ABC):
"""Abstract base class for flow execution with caching, streaming, and operation tree management.
BaseFlow provides a framework for building complex workflows by composing operations
into executable trees. It supports both synchronous and asynchronous execution modes,
response caching, streaming outputs, and automatic tool call schema generation.
"""
"""Abstract base class for flow execution with caching, streaming, and operation tree management."""
def __init__(
self,
name: str = "",
flow_op: BaseOp | None = None,
stream: bool = False,
raise_exception: bool = True,
enable_cache: bool = False,
cache_path: str = "cache/flow",
cache_expire_hours: float = 0.1,
service_context: ServiceContext | None = None,
**kwargs,
):
"""Initialize flow configuration and execution state."""
@ -42,11 +37,12 @@ class BaseFlow(ABC):
self.enable_cache: bool = enable_cache
self.cache_path: str = cache_path
self.cache_expire_hours: float = cache_expire_hours
self.service_context: ServiceContext | None = service_context
self.flow_params: dict = kwargs
self._flow_op: BaseOp | None = flow_op
self._cache: CacheHandler | None = None
self._flow_printed: bool = False
self._flow_op: BaseOp | None = None
self._tool_call: ToolCall | None = None
def _build_tool_call(self) -> ToolCall | None:
@ -82,11 +78,7 @@ class BaseFlow(ABC):
return
if key := self._compute_cache_key(params):
self.cache.save(
key,
response.model_dump(exclude_none=True),
expire_hours=self.cache_expire_hours,
)
self.cache.save(key, response.model_dump(exclude_none=True), expire_hours=self.cache_expire_hours)
def _print_operation_tree(self, name: str, op: BaseOp, indent: int):
"""Recursively log the hierarchy of the flow's operation tree."""
@ -100,19 +92,13 @@ class BaseFlow(ABC):
@property
def tool_call(self) -> ToolCall | None:
"""Lazily construct the ToolCall schema describing this flow."""
if self.flow_op.tool_call:
if hasattr(self.flow_op, "tool_call"):
return self.flow_op.tool_call
if self._tool_call is None:
self._tool_call = self._build_tool_call()
if self._tool_call:
self._tool_call.name = self._tool_call.name or self.name
self._tool_call.output = self._tool_call.output or {
f"{self.name}_result": ToolAttr(
type="string",
description=f"The execution result of the {self.name}",
),
}
return self._tool_call
@property
@ -130,12 +116,6 @@ class BaseFlow(ABC):
self._flow_op = self._build_flow()
return self._flow_op
@flow_op.setter
def flow_op(self, op: BaseOp):
"""Set the root operation of the flow."""
self._flow_op = op
self._flow_printed = False
@property
def async_mode(self) -> bool:
"""Check if the current flow operation tree is asynchronous."""
@ -148,11 +128,10 @@ class BaseFlow(ABC):
if not lines:
raise ValueError("Expression is empty")
env: dict = C.registry_dict[RegistryEnum.OP]
if len(lines) > 1:
exec("\n".join(lines[:-1]), {"__builtins__": {}}, env)
exec("\n".join(lines[:-1]), {"__builtins__": {}}, R.op)
result = eval(lines[-1], {"__builtins__": {}}, env)
result = eval(lines[-1], {"__builtins__": {}}, R.op)
if not isinstance(result, BaseOp):
raise TypeError(f"Expression evaluated to {type(result)}, expected BaseOp")
return result
@ -172,30 +151,34 @@ class BaseFlow(ABC):
if cached := self._maybe_load_cached(kwargs):
return cached
context = RuntimeContext(**kwargs)
context = RuntimeContext(service_context=self.service_context, **kwargs)
try:
self.print_flow()
flow_op: BaseOp = self._build_flow()
assert self.flow_op.async_mode, "Async call requires an async flow operation."
await flow_op.call(context=context)
result = context.stream_queue if self.stream else context.response
if self.stream:
await context.add_stream_done()
return context.stream_queue
else:
self._maybe_save_cache(kwargs, context.response)
return context.response
self._maybe_save_cache(kwargs, result)
return result
except Exception as e:
logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}")
if self.raise_exception:
raise e
if self.stream:
await context.add_stream_chunk_and_type(str(e), ChunkEnum.ERROR)
await context.add_stream_done()
return context.stream_queue
context.add_response_error(e)
return context.response
else:
context.add_response_error(e)
return context.response
def call_sync(self, **kwargs) -> Response:
"""Execute the flow synchronously with parameter caching."""
@ -204,18 +187,20 @@ class BaseFlow(ABC):
if cached := self._maybe_load_cached(kwargs):
return cached
context = RuntimeContext(**kwargs)
context = RuntimeContext(service_context=self.service_context, **kwargs)
try:
self.print_flow()
flow_op: BaseOp = self._build_flow()
assert not self.flow_op.async_mode, "Sync call requires a sync flow operation."
flow_op.call_sync(context=context)
self._maybe_save_cache(kwargs, context.response)
return context.response
except Exception as e:
logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}")
if self.raise_exception:
raise e
context.add_response_error(e)
return context.response

View file

@ -1,6 +1,7 @@
"""Expression-based flow implementation driven by configuration objects."""
from .base_flow import BaseFlow
from ..context import ServiceContext
from ..op import BaseOp
from ..schema import FlowConfig, ToolCall
@ -8,7 +9,7 @@ from ..schema import FlowConfig, ToolCall
class ExpressionFlow(BaseFlow):
"""A flow implementation that constructs operations from a FlowConfig definition."""
def __init__(self, flow_config: FlowConfig):
def __init__(self, flow_config: FlowConfig, service_context: ServiceContext):
"""Initialize the flow using settings and metadata from a FlowConfig instance."""
self.flow_config: FlowConfig = flow_config
super().__init__(
@ -18,6 +19,7 @@ class ExpressionFlow(BaseFlow):
enable_cache=self.flow_config.enable_cache,
cache_path=self.flow_config.cache_path,
cache_expire_hours=self.flow_config.cache_expire_hours,
service_context=service_context,
**flow_config.model_extra,
)
@ -27,4 +29,9 @@ class ExpressionFlow(BaseFlow):
def _build_tool_call(self) -> ToolCall:
"""Construct a tool call representation based on configuration parameters."""
return ToolCall(**{"description": self.flow_config.description, "parameters": self.flow_config.parameters})
return ToolCall(
**{
"description": self.flow_config.description,
"parameters": self.flow_config.parameters,
},
)

View file

@ -5,6 +5,7 @@ from .lite_llm import LiteLLM
from .lite_llm_sync import LiteLLMSync
from .openai_llm import OpenAILLM
from .openai_llm_sync import OpenAILLMSync
from ..context import R
__all__ = [
"BaseLLM",
@ -13,3 +14,8 @@ __all__ = [
"OpenAILLM",
"OpenAILLMSync",
]
R.llm.register("litellm")(LiteLLM)
R.llm.register("litellm_sync")(LiteLLMSync)
R.llm.register("openai")(OpenAILLM)
R.llm.register("openai_sync")(OpenAILLMSync)

View file

@ -1,4 +1,4 @@
"""Abstract base interface for ReMe LLM implementations."""
"""Base interface for LLM implementations."""
import asyncio
import json
@ -15,37 +15,42 @@ from ..schema import ToolCall
class BaseLLM(ABC):
"""Abstract base class defining the standard interface for LLM interactions."""
"""Base class for LLM interactions."""
def __init__(
self,
model_name: str,
max_retries: int = 10,
raise_exception: bool = False,
request_interval: float = 0.0,
**kwargs,
):
"""Initialize LLM client.
def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_concurrency: int | None = None, **kwargs):
"""Initialize the LLM client with model configurations and retry policies.
Args:
model_name: The name of the model to use
max_retries: Maximum number of retry attempts on failure
raise_exception: Whether to raise exceptions or return default values
max_concurrency: Maximum concurrent requests for async operations. If None, no concurrency limit is applied.
model_name: Model name to use
max_retries: Maximum retry attempts on failure
raise_exception: Raise exceptions or return default values
request_interval: Minimum seconds between requests (default: 0.0)
**kwargs: Additional model-specific parameters
"""
self.model_name: str = model_name
self.max_retries: int = max_retries
self.raise_exception: bool = raise_exception
self.max_concurrency: int | None = max_concurrency
self.request_interval: float = request_interval
self.kwargs: dict = kwargs
# Concurrency control for async operations
self._semaphore: asyncio.Semaphore | None = asyncio.Semaphore(max_concurrency) if max_concurrency else None
self._last_request_time: float = 0.0
self._request_lock: asyncio.Lock = asyncio.Lock()
@staticmethod
def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):
"""Assemble incremental tool call fragments into complete ToolCall objects."""
"""Assemble incremental tool call chunks into complete ToolCall objects."""
index = tool_call.index
# Ensure we have a ToolCall object at this index
while len(ret_tools) <= index:
ret_tools.append(ToolCall(index=index))
# Accumulate tool call parts (id, name, arguments)
if tool_call.id:
ret_tools[index].id += tool_call.id
@ -57,7 +62,7 @@ class BaseLLM(ABC):
@staticmethod
def _validate_and_serialize_tools(ret_tool_calls: list[ToolCall], tools: list[ToolCall]) -> list[dict]:
"""Validate tool call integrity and return serialized tool dictionaries."""
"""Validate and serialize tool calls."""
if not ret_tool_calls:
return []
@ -68,8 +73,9 @@ class BaseLLM(ABC):
if tool.name not in tool_dict:
continue
if not tool.check_argument():
raise ValueError(f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}")
if not tool.sanitize_and_check_argument():
logger.error(f"Invalid JSON arguments in {tool.name}: {tool.arguments}")
raise ValueError(f"Invalid JSON arguments in {tool.name}: {tool.arguments}")
validated_tools.append(tool.simple_output_dump())
return validated_tools
@ -83,15 +89,7 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> dict:
"""Construct provider-specific parameters for streaming API requests.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
log_params: Whether to log parameters
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
"""Build provider-specific streaming parameters."""
async def _stream_chat(
self,
@ -99,7 +97,7 @@ class BaseLLM(ABC):
tools: list[ToolCall] | None,
stream_kwargs: dict,
) -> AsyncGenerator[StreamChunk, None]:
"""Internal async generator for streaming raw response chunks."""
"""Async generator for streaming response chunks."""
raise NotImplementedError
def _stream_chat_sync(
@ -108,7 +106,7 @@ class BaseLLM(ABC):
tools: list[ToolCall] | None = None,
stream_kwargs: dict | None = None,
) -> Generator[StreamChunk, None, None]:
"""Internal synchronous generator for streaming raw response chunks."""
"""Sync generator for streaming response chunks."""
raise NotImplementedError
async def stream_chat(
@ -118,23 +116,18 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> AsyncGenerator[StreamChunk, None]:
"""Public async interface for streaming chat completions with retries.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
# Apply concurrency control if configured
if self._semaphore:
async with self._semaphore:
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
else:
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
"""Stream chat completions with retries."""
if self.request_interval > 0:
async with self._request_lock:
current_time = time.time()
elapsed = current_time - self._last_request_time
if elapsed < self.request_interval:
await asyncio.sleep(self.request_interval - elapsed)
self._last_request_time = time.time()
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
async def _stream_chat_impl(
self,
messages: list[Message],
@ -142,7 +135,7 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> AsyncGenerator[StreamChunk, None]:
"""Internal implementation of stream_chat with retry logic."""
"""Stream chat with retry logic."""
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
for i in range(self.max_retries):
@ -152,7 +145,7 @@ class BaseLLM(ABC):
return
except Exception as e:
logger.exception(f"stream chat with model={self.model_name} encounter error with e={e.args}")
logger.exception(f"Stream chat error (model={self.model_name}): {e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@ -170,14 +163,7 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Generator[StreamChunk, None, None]:
"""Public synchronous interface for streaming chat completions with retries.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
"""Stream chat completions synchronously with retries."""
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
for i in range(self.max_retries):
@ -186,7 +172,7 @@ class BaseLLM(ABC):
return
except Exception as e:
logger.exception(f"stream chat sync with model={self.model_name} encounter error with e={e.args}")
logger.exception(f"Stream chat sync error (model={self.model_name}): {e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@ -205,15 +191,7 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Message:
"""Internal async method to aggregate a full response by consuming the stream.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
enable_stream_print: Whether to print stream chunks
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
"""Aggregate full response by consuming the stream."""
state = {
"enter_think": False,
"enter_answer": False,
@ -224,7 +202,6 @@ class BaseLLM(ABC):
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
# Process stream chunk
if stream_chunk.chunk_type is ChunkEnum.USAGE:
if enable_stream_print:
print(
@ -273,15 +250,7 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Message:
"""Internal synchronous method to aggregate a full response by consuming the stream.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
enable_stream_print: Whether to print stream chunks
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
"""Aggregate full response synchronously by consuming the stream."""
state = {
"enter_think": False,
"enter_answer": False,
@ -292,7 +261,6 @@ class BaseLLM(ABC):
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
# Process stream chunk
if stream_chunk.chunk_type is ChunkEnum.USAGE:
if enable_stream_print:
print(
@ -343,24 +311,25 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Message | Any:
"""Perform an async chat completion with integrated retries and error handling.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
enable_stream_print: Whether to print stream chunks
callback_fn: Optional callback function to process the result
default_value: Default value to return on error
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
# Apply concurrency control if configured
if self._semaphore:
async with self._semaphore:
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
else:
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
"""Chat completion with retries and error handling."""
if self.request_interval > 0:
async with self._request_lock:
current_time = time.time()
elapsed = current_time - self._last_request_time
if elapsed < self.request_interval:
await asyncio.sleep(self.request_interval - elapsed)
self._last_request_time = time.time()
return await self._chat_impl(
messages,
tools,
enable_stream_print,
callback_fn,
default_value,
model_name,
**kwargs,
)
async def _chat_impl(
self,
messages: list[Message],
@ -371,10 +340,9 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Message | Any:
"""Internal implementation of chat with retry and error handling logic."""
# Use the provided model_name or fall back to self.model_name
"""Chat with retry and error handling logic."""
effective_model = model_name if model_name is not None else self.model_name
for i in range(self.max_retries):
try:
result = await self._chat(
@ -387,15 +355,16 @@ class BaseLLM(ABC):
return callback_fn(result) if callback_fn else result
except Exception as e:
# Check if this is an inappropriate content error
error_message = str(e.args[0]) if e.args else str(e)
is_inappropriate_content = "inappropriate content" in error_message.lower()
is_rate_limit_error = "request rate increased too quickly" in error_message.lower()
is_rate_limit_error = (
"request rate increased too quickly" in error_message.lower()
or "exceeded your current quota" in error_message.lower()
or "insufficient_quota" in error_message.lower()
)
if is_inappropriate_content:
logger.error(f"chat with model={effective_model} detected inappropriate content error")
logger.error("=" * 80)
logger.error("Full message content that triggered the error:")
logger.error(f"Inappropriate content detected (model={effective_model})")
logger.error("=" * 80)
for idx, msg in enumerate(messages):
logger.error(f"Message {idx + 1} [role={msg.role}]:")
@ -406,15 +375,16 @@ class BaseLLM(ABC):
logger.error(f"Tool calls: {msg.tool_calls}")
logger.error("-" * 80)
logger.error("=" * 80)
# Return empty Message immediately without retrying
return Message(role=Role.ASSISTANT, content="")
if is_rate_limit_error:
logger.warning(f"chat with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})")
logger.warning(
f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})",
)
await asyncio.sleep(60)
continue
logger.exception(f"chat with model={effective_model} encounter error with e={e.args}")
logger.exception(f"Chat error (model={effective_model}): {e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@ -434,20 +404,9 @@ class BaseLLM(ABC):
model_name: str | None = None,
**kwargs,
) -> Message | Any:
"""Perform a synchronous chat completion with integrated retries and error handling.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
enable_stream_print: Whether to print stream chunks
callback_fn: Optional callback function to process the result
default_value: Default value to return on error
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
# Use the provided model_name or fall back to self.model_name
"""Chat completion synchronously with retries and error handling."""
effective_model = model_name if model_name is not None else self.model_name
for i in range(self.max_retries):
try:
result = self._chat_sync(
@ -460,15 +419,16 @@ class BaseLLM(ABC):
return callback_fn(result) if callback_fn else result
except Exception as e:
# Check if this is an inappropriate content error
error_message = str(e.args[0]) if e.args else str(e)
is_inappropriate_content = "inappropriate content" in error_message.lower()
is_rate_limit_error = "request rate increased too quickly" in error_message.lower()
is_rate_limit_error = (
"request rate increased too quickly" in error_message.lower()
or "exceeded your current quota" in error_message.lower()
or "insufficient_quota" in error_message.lower()
)
if is_inappropriate_content:
logger.error(f"chat sync with model={effective_model} detected inappropriate content error")
logger.error("=" * 80)
logger.error("Full message content that triggered the error:")
logger.error(f"Inappropriate content detected (model={effective_model})")
logger.error("=" * 80)
for idx, msg in enumerate(messages):
logger.error(f"Message {idx + 1} [role={msg.role}]:")
@ -479,15 +439,16 @@ class BaseLLM(ABC):
logger.error(f"Tool calls: {msg.tool_calls}")
logger.error("-" * 80)
logger.error("=" * 80)
# Return empty Message immediately without retrying
return Message(role=Role.ASSISTANT, content="")
if is_rate_limit_error:
logger.warning(f"chat sync with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})")
logger.warning(
f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})",
)
time.sleep(60)
continue
logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}")
logger.exception(f"Chat sync error (model={effective_model}): {e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@ -498,7 +459,7 @@ class BaseLLM(ABC):
return default_value
async def close(self):
"""Release any asynchronous resources or connections held by the client."""
"""Release async resources."""
def close_sync(self):
"""Release any synchronous resources or connections held by the client."""
"""Release sync resources."""

View file

@ -7,14 +7,12 @@ import litellm
from loguru import logger
from .base_llm import BaseLLM
from ..context import C
from ..enumeration import ChunkEnum
from ..schema import Message
from ..schema import StreamChunk
from ..schema import ToolCall
@C.register_llm("litellm")
class LiteLLM(BaseLLM):
"""Async LLM implementation using LiteLLM to support multiple providers."""
@ -40,7 +38,7 @@ class LiteLLM(BaseLLM):
**kwargs,
) -> dict:
"""Construct and log the parameters dictionary for LiteLLM API calls.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
@ -50,7 +48,7 @@ class LiteLLM(BaseLLM):
"""
# Use the provided model_name or fall back to self.model_name
effective_model = model_name if model_name is not None else self.model_name
# Construct the API parameters by merging multiple sources
llm_kwargs = {
"model": effective_model,
@ -98,7 +96,7 @@ class LiteLLM(BaseLLM):
if not chunk.choices:
if hasattr(chunk, "usage") and chunk.usage:
yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
continue
continue
delta = chunk.choices[0].delta

View file

@ -5,14 +5,12 @@ from typing import Generator
import litellm
from .lite_llm import LiteLLM
from ..context import C
from ..enumeration import ChunkEnum
from ..schema import Message
from ..schema import StreamChunk
from ..schema import ToolCall
@C.register_llm("litellm_sync")
class LiteLLMSync(LiteLLM):
"""Synchronous LiteLLM client for executing chat completions and streaming responses."""
@ -31,7 +29,7 @@ class LiteLLMSync(LiteLLM):
if not chunk.choices:
if hasattr(chunk, "usage") and chunk.usage:
yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
continue
continue
delta = chunk.choices[0].delta

View file

@ -7,14 +7,12 @@ from loguru import logger
from openai import AsyncOpenAI
from .base_llm import BaseLLM
from ..context import C
from ..enumeration import ChunkEnum
from ..schema import Message
from ..schema import StreamChunk
from ..schema import ToolCall
@C.register_llm("openai")
class OpenAILLM(BaseLLM):
"""Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution."""
@ -45,7 +43,7 @@ class OpenAILLM(BaseLLM):
**kwargs,
) -> dict:
"""Construct the parameter dictionary for the OpenAI Chat Completions API call.
Args:
messages: List of conversation messages
tools: Optional list of tool calls
@ -55,7 +53,7 @@ class OpenAILLM(BaseLLM):
"""
# Use the provided model_name or fall back to self.model_name
effective_model = model_name if model_name is not None else self.model_name
# Construct the API parameters by merging multiple sources
llm_kwargs = {
"model": effective_model,
@ -93,7 +91,7 @@ class OpenAILLM(BaseLLM):
if not chunk.choices:
if hasattr(chunk, "usage") and chunk.usage:
yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
continue
continue
delta = chunk.choices[0].delta

View file

@ -5,14 +5,12 @@ from typing import Generator
from openai import OpenAI
from .openai_llm import OpenAILLM
from ..context import C
from ..enumeration import ChunkEnum
from ..schema import Message
from ..schema import StreamChunk
from ..schema import ToolCall
@C.register_llm("openai_sync")
class OpenAILLMSync(OpenAILLM):
"""Synchronous LLM client for OpenAI-compatible APIs, inheriting from OpenAILLM."""
@ -35,7 +33,7 @@ class OpenAILLMSync(OpenAILLM):
if not chunk.choices:
if hasattr(chunk, "usage") and chunk.usage:
yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
continue
continue
delta = chunk.choices[0].delta

View file

@ -2,14 +2,19 @@
from .base_op import BaseOp
from .base_ray_op import BaseRayOp
from .base_tool import BaseTool
from .mcp_tool import MCPTool
from .parallel_op import ParallelOp
from .sequential_op import SequentialOp
from ..context import R
__all__ = [
"BaseOp",
"BaseRayOp",
"BaseTool",
"MCPTool",
"ParallelOp",
"SequentialOp",
]
R.op.register("mcp_tool")(MCPTool)

View file

@ -3,22 +3,23 @@
import asyncio
import copy
import inspect
from abc import ABCMeta
from pathlib import Path
from typing import Callable, Optional
from typing import Callable, Optional, Any
from loguru import logger
from tqdm import tqdm
from ..context import RuntimeContext, PromptHandler, C
from ..context import RuntimeContext, PromptHandler, ServiceContext
from ..embedding import BaseEmbeddingModel
from ..llm import BaseLLM
from ..schema import ToolCall, ToolAttr, Response
from ..schema import Response
from ..token_counter import BaseTokenCounter
from ..utils import camel_to_snake, CacheHandler, timer
from ..vector_store import BaseVectorStore
class BaseOp:
class BaseOp(metaclass=ABCMeta):
"""Base operator class for LLM workflow execution and composition."""
def __new__(cls, *args, **kwargs):
@ -34,6 +35,7 @@ class BaseOp:
async_mode: bool = True,
language: str = "",
prompt_name: str = "",
prompt_path: str = "",
llm: str | BaseLLM = "default",
embedding_model: str | BaseEmbeddingModel = "default",
vector_store: str | BaseVectorStore = "default",
@ -44,7 +46,6 @@ class BaseOp:
sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None,
input_mapping: dict[str, str] | None = None,
output_mapping: dict[str, str] | None = None,
save_response_result: bool = False,
enable_sync_thread_pool: bool = True,
max_retries: int = 1,
raise_exception: bool = False,
@ -53,8 +54,8 @@ class BaseOp:
"""Initialize operator configurations and internal state."""
self.name = name or camel_to_snake(self.__class__.__name__)
self.async_mode = async_mode
self.language = language or C.language
self.prompt = self._get_prompt_handler(prompt_name)
self.language = language
self.prompt = self._get_prompt_handler(prompt_name, prompt_path)
self._llm = llm
self._embedding_model = embedding_model
@ -64,12 +65,12 @@ class BaseOp:
self.enable_cache = enable_cache
self.cache_path = cache_path
self.cache_expire_hours = cache_expire_hours
self.sub_ops: list[BaseOp] = []
self.sub_ops: list["BaseOp"] = []
self.add_sub_ops(sub_ops)
self.input_mapping = input_mapping
self.output_mapping = output_mapping
self.save_response_result = save_response_result
self.enable_sync_thread_pool = enable_sync_thread_pool
self.max_retries = max(1, max_retries)
self.raise_exception = raise_exception
@ -78,86 +79,29 @@ class BaseOp:
self._pending_tasks: list = []
self.context: RuntimeContext | None = None
self._cache: CacheHandler | None = None
self._tool_call: ToolCall | None = None
def _get_prompt_handler(self, prompt_name: str) -> PromptHandler:
def _get_prompt_handler(self, prompt_name: str, prompt_path: str) -> PromptHandler:
"""Load prompt configuration from the associated YAML file."""
path = Path(inspect.getfile(self.__class__))
path = path.with_stem(prompt_name) if prompt_name else path
if prompt_path:
path = Path(prompt_path)
else:
path = Path(inspect.getfile(self.__class__))
if prompt_name:
path = path.with_stem(prompt_name)
return PromptHandler(language=self.language).load_prompt_by_file(path.with_suffix(".yaml"))
def _build_tool_call(self) -> ToolCall | None:
"""Build and return the tool call schema; override in subclasses."""
def _validate_inputs(self):
"""Ensure all required tool inputs are present in context."""
if self.tool_call is not None:
parameters = self.tool_call.parameters
if parameters.type == "object" and parameters.properties:
required_list = parameters.required or []
required_keys = {k: (k in required_list) for k in parameters.properties.keys()}
self.context.validate_required_keys(required_keys, self.name)
def _handle_failure(self, e: Exception, attempt: int):
def _handle_failure(self, e: Exception, attempt: int) -> str | None:
"""Log failures and handle final retry logic."""
message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}"
if attempt == self.max_retries - 1:
logger.exception(message)
if self.raise_exception:
raise e
if self.tool_call is not None:
self.output = f"{self.name} failed: {e}"
return f"{self.name} failed: {e}"
else:
logger.warning(message)
@property
def tool_call(self) -> ToolCall | None:
"""Lazily construct and return the tool call metadata."""
if self._tool_call is None:
self._tool_call = self._build_tool_call()
if self._tool_call is None:
return None
self._tool_call.name = self._tool_call.name or self.name
if not self._tool_call.output.properties:
self._tool_call.output.properties = {
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
}
return self._tool_call
@property
def input_dict(self) -> dict:
"""Extract required and optional inputs from context based on schema."""
parameters = self.tool_call.parameters
if parameters.type != "object" or not parameters.properties:
return {}
required_keys = set(parameters.required or [])
return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)}
@property
def output(self):
"""Get the single output value from context."""
output_properties = self.tool_call.output.properties
if not output_properties:
return None
keys = list(output_properties.keys())
if len(keys) >= 1 and keys[0] in self.context:
return self.context[keys[0]]
else:
return None
@output.setter
def output(self, value):
"""Set the single output value into context."""
output_properties = self.tool_call.output.properties
if not output_properties:
return
keys = list(output_properties.keys())
self.context[keys[0]] = value
@property
def cache(self) -> CacheHandler:
"""Access the operator-specific cache handler."""
@ -166,119 +110,120 @@ class BaseOp:
self._cache = CacheHandler(f"{self.cache_path}/{self.name}")
return self._cache
@property
def service_context(self) -> ServiceContext:
"""Access the service context."""
return self.context.service_context
@property
def llm(self) -> BaseLLM:
"""Get the LLM instance from ServiceContext."""
if isinstance(self._llm, str):
self._llm = C.get_llm(self._llm)
self._llm = self.service_context.llms[self._llm]
return self._llm
@property
def embedding_model(self) -> BaseEmbeddingModel:
"""Get the embedding model instance from ServiceContext."""
if isinstance(self._embedding_model, str):
self._embedding_model = C.get_embedding_model(self._embedding_model)
self._embedding_model = self.service_context.embedding_models[self._embedding_model]
return self._embedding_model
@property
def vector_store(self) -> BaseVectorStore:
"""Lazily initialize and return the vector store instance."""
if isinstance(self._vector_store, str):
self._vector_store = C.get_vector_store(self._vector_store)
self._vector_store = self.service_context.vector_stores[self._vector_store]
return self._vector_store
@property
def token_counter(self) -> BaseTokenCounter:
"""Get the token counter instance from ServiceContext."""
if isinstance(self._token_counter, str):
self._token_counter = C.get_token_counter(self._token_counter)
self._token_counter = self.service_context.token_counters[self._token_counter]
return self._token_counter
@property
def service_metadata(self) -> dict:
"""Get service configuration metadata."""
return C.service_config.model_extra
return self.service_context.service_config.model_extra
@property
def response(self) -> Response:
"""Get the response object."""
"""Access the response object."""
return self.context.response
def set_tool_call(self, tool_call: ToolCall | dict):
"""Set the tool call."""
if isinstance(tool_call, dict):
self._tool_call = ToolCall(**tool_call)
elif isinstance(tool_call, ToolCall):
self._tool_call = tool_call
else:
raise ValueError(f"Invalid tool call: {tool_call}")
self._tool_call.name = self._tool_call.name or self.name
if not self._tool_call.output.properties:
self._tool_call.output.properties = {
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
}
def set_language(self, language: str):
"""Set the language."""
self.language = language
return self
def before_execute_sync(self):
"""Prepare context and validate before sync execution."""
self.context.apply_mapping(self.input_mapping)
self._validate_inputs()
async def before_execute(self):
"""Prepare context and validate before async execution."""
self.context.apply_mapping(self.input_mapping)
def execute_sync(self):
"""Define core sync logic in subclasses."""
def after_execute_sync(self):
"""Finalize context and mappings after sync execution."""
self.context.apply_mapping(self.output_mapping)
if self.tool_call is not None and self.save_response_result:
self.context.response.answer = self.output
async def before_execute(self):
"""Prepare context and validate before async execution."""
self.before_execute_sync()
async def execute(self):
"""Define core async logic in subclasses."""
async def after_execute(self):
def after_execute_sync(self, response: Any):
"""Finalize context and mappings after sync execution."""
self.context.apply_mapping(self.output_mapping)
if response is not None:
if isinstance(response, dict):
for k, v in response.items():
if k == "answer":
self.response.answer = v
elif k == "success":
self.response.success = v.lower() == "true"
else:
self.response.metadata[k] = v
else:
self.response.answer = response
return response
async def after_execute(self, output: Any):
"""Finalize context and mappings after async execution."""
self.after_execute_sync()
return self.after_execute_sync(output)
@timer
def call_sync(self, context: RuntimeContext = None, **kwargs):
"""Execute the operator synchronously with retry logic."""
self.context = RuntimeContext.from_context(context, **kwargs)
response = None
for i in range(self.max_retries):
try:
self.before_execute_sync()
self.execute_sync()
self.after_execute_sync()
response = self.execute_sync()
response = self.after_execute_sync(response)
break
except Exception as e:
self._handle_failure(e, i)
return self.output if self.tool_call is not None else None
response = self._handle_failure(e, i)
return response
@timer
async def call(self, context: RuntimeContext = None, **kwargs):
"""Execute the operator asynchronously with retry logic."""
self.context = RuntimeContext.from_context(context, **kwargs)
response = None
for i in range(self.max_retries):
try:
await self.before_execute()
await self.execute()
await self.after_execute()
response = await self.execute()
response = await self.after_execute(response)
break
except Exception as e:
self._handle_failure(e, i)
return self.output if self.tool_call is not None else None
response = self._handle_failure(e, i)
return response
def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp":
"""Submit a task to the thread pool or local queue."""
task = C.thread_pool.submit(fn, *args, **kwargs) if self.enable_sync_thread_pool else (fn, args, kwargs)
if self.enable_sync_thread_pool:
task = self.service_context.thread_pool.submit(fn, *args, **kwargs)
else:
task = (fn, args, kwargs)
self._pending_tasks.append(task)
return self
@ -292,26 +237,32 @@ class BaseOp:
"""Wait for all pending sync tasks and return flattened results."""
results = []
for task in tqdm(self._pending_tasks, desc=task_desc or self.name):
res = task.result() if self.enable_sync_thread_pool else task[0](*task[1], **task[2])
if res:
results.extend(res if isinstance(res, list) else [res])
if self.enable_sync_thread_pool:
result = task.result()
else:
result = task[0](*task[1], **task[2])
if result:
if isinstance(result, list):
results.extend(result)
else:
results.append(result)
self._pending_tasks.clear()
return results
async def join_async_tasks(self, return_exceptions: bool = True) -> list:
"""Wait for all pending async tasks and aggregate results."""
try:
raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions)
results = []
for res in raw_results:
if isinstance(res, Exception):
logger.error(f"[{self.__class__.__name__}] Async task failed: {res}")
continue
if res:
results.extend(res if isinstance(res, list) else [res])
return results
finally:
self._pending_tasks.clear()
raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions)
results = []
for result in raw_results:
if isinstance(result, Exception):
logger.error(f"[{self.__class__.__name__}] Async task failed: {result}")
elif result:
if isinstance(result, list):
results.extend(result)
else:
result.append(result)
self._pending_tasks.clear()
return results
def add_sub_ops(self, sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"]):
"""Add child operators to this operator's sub_ops."""

View file

@ -8,14 +8,14 @@ from loguru import logger
from tqdm import tqdm
from .base_op import BaseOp
from ..context import BaseContext, C
from ..context import BaseContext
_RAY_IMPORT_ERROR = None
try:
import ray
except ImportError as e:
_RAY_IMPORT_ERROR = e
except ImportError as _e:
_RAY_IMPORT_ERROR = _e
ray = None
@ -35,7 +35,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta):
def submit_and_join_ray_task(self, fn: Callable, parallel_key: str = "", task_desc: str = "", **kwargs) -> list:
"""Divide data into chunks and execute them across Ray workers."""
max_workers = C.service_config.ray_max_workers
max_workers = self.service_context.ray_max_workers
self._ray_task_list.clear()
# Automatically detect the key containing the list to parallelize
@ -94,7 +94,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta):
def submit_ray_task(self, fn, *args, **kwargs):
"""Submit a single Ray task to the task list for later execution."""
if not ray.is_initialized():
ray.init(num_cpus=C.service_config.ray_max_workers, ignore_reinit_error=True)
ray.init(num_cpus=self.service_context.ray_max_workers, ignore_reinit_error=True)
remote_fn = ray.remote(fn)
task = remote_fn.remote(*args, **kwargs)

58
reme/core/op/base_tool.py Normal file
View file

@ -0,0 +1,58 @@
"""Base class for tools"""
from abc import ABCMeta
from . import BaseOp
from ..schema import ToolCall
class BaseTool(BaseOp, metaclass=ABCMeta):
"""Base class for tools"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._tool_call: ToolCall | None = None
def _build_tool_call(self) -> ToolCall:
"""Build and return the tool call schema; override in subclasses."""
def _validate_inputs(self):
"""Validate the inputs."""
parameters = self.tool_call.parameters
if parameters.type == "object" and parameters.properties:
required_list = parameters.required or []
required_keys = {k: (k in required_list) for k in parameters.properties.keys()}
self.context.validate_required_keys(required_keys, self.name)
@property
def tool_call(self) -> ToolCall:
"""Get the tool call schema."""
if self._tool_call is None:
self._tool_call = self._build_tool_call()
self._tool_call.name = self._tool_call.name or self.name
return self._tool_call
def set_tool_call(self, tool_call: ToolCall | dict):
"""Set the tool call schema."""
if isinstance(tool_call, dict):
self._tool_call = ToolCall(**tool_call)
elif isinstance(tool_call, ToolCall):
self._tool_call = tool_call
else:
raise ValueError(f"Invalid tool call: {tool_call}")
self._tool_call.name = self._tool_call.name or self.name
@property
def input_dict(self) -> dict:
"""Get the input dict."""
parameters = self.tool_call.parameters
if parameters.type != "object" or not parameters.properties:
return {}
required_keys = set(parameters.required or [])
return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)}
def before_execute_sync(self):
"""Hook before execute"""
super().before_execute_sync()
self._validate_inputs()

View file

@ -2,25 +2,20 @@
from typing import List
from .base_op import BaseOp
from ..context import C
from mcp.types import CallToolResult, TextContent
from .base_tool import BaseTool
from ..schema import ToolCall
from ..utils import MCPClient
@C.register_op()
class MCPTool(BaseOp):
"""Operator for calling remote MCP (Model Context Protocol) tools.
This class enables integration with external MCP servers to execute tools
and retrieve their results. It supports parameter customization and retry logic.
"""
class MCPTool(BaseTool):
"""Operator for calling remote MCP (Model Context Protocol) tools."""
def __init__(
self,
mcp_server: str = "",
tool_name: str = "",
save_response_result: bool = True,
parameter_required: List[str] | None = None,
parameter_optional: List[str] | None = None,
parameter_deleted: List[str] | None = None,
@ -29,13 +24,7 @@ class MCPTool(BaseOp):
raise_exception: bool = False,
**kwargs,
):
super().__init__(
save_response_result=save_response_result,
max_retries=max_retries,
raise_exception=raise_exception,
**kwargs,
)
super().__init__(max_retries=max_retries, raise_exception=raise_exception, **kwargs)
self.mcp_server: str = mcp_server
self.tool_name: str = tool_name
@ -43,12 +32,12 @@ class MCPTool(BaseOp):
self.parameter_optional: List[str] | None = parameter_optional
self.parameter_deleted: List[str] | None = parameter_deleted
self.timeout: float | None = timeout
# Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market
self._client = MCPClient(C.service_config.mcp_servers)
# Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market
self._client = MCPClient(self.service_context.service_config.mcp_servers)
def _build_tool_call(self) -> ToolCall:
tool_call_dict = C.mcp_server_mapping[self.mcp_server]
tool_call_dict = self.service_context.mcp_server_mapping[self.mcp_server]
tool_call: ToolCall = tool_call_dict[self.tool_name].model_copy(deep=True)
# Initialize required list if not exists
@ -74,9 +63,16 @@ class MCPTool(BaseOp):
return tool_call
async def execute(self):
self.output = await self._client.call_tool(
tool_result: CallToolResult = await self._client.call_tool(
server_name=self.mcp_server,
tool_name=self.tool_name,
arguments=self.input_dict,
parse_text_result=True,
)
self.context.tool_result = tool_result
text_result = []
for block in tool_result.content:
if isinstance(block, TextContent):
text_result.append(block.text)
output: str = "\n".join(text_result)
return output

View file

@ -11,14 +11,14 @@ class ParallelOp(BaseOp):
for op in self.sub_ops:
assert op.async_mode
self.submit_async_task(op.call, context=self.context)
await self.join_async_tasks()
return await self.join_async_tasks()
def execute_sync(self):
"""Executes all sub-operations concurrently using synchronous task management."""
for op in self.sub_ops:
assert not op.async_mode
self.submit_sync_task(op.call_sync, context=self.context)
self.join_sync_tasks()
return self.join_sync_tasks()
def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp):
"""Raises RuntimeError as the shift operator is not supported for parallel operations."""

View file

@ -8,15 +8,19 @@ class SequentialOp(BaseOp):
async def execute(self):
"""Executes sub-operations sequentially using asynchronous awaits."""
result = None
for op in self.sub_ops:
assert op.async_mode
await op.call(context=self.context)
result = await op.call(context=self.context)
return result
def execute_sync(self):
"""Executes sub-operations sequentially in a synchronous blocking manner."""
result = None
for op in self.sub_ops:
assert not op.async_mode
op.call_sync(context=self.context)
result = op.call_sync(context=self.context)
return result
def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp):
"""Raises RuntimeError as the left shift operator is not supported."""

View file

@ -6,7 +6,6 @@ memories in the ReMe system.
import datetime
import hashlib
import json
from typing import Any
from pydantic import BaseModel, Field, model_validator
@ -146,30 +145,6 @@ class MemoryNode(BaseModel):
metadata=metadata,
)
def format_memory(self) -> str:
"""Format memory as human-readable string.
Returns:
str: Formatted string with when_to_use, content, and ref_memory_id.
"""
parts: list[str] = [
f"memory_id={self.memory_id}",
]
if self.when_to_use:
parts.append(self.when_to_use)
if self.content:
parts.append(self.content)
if self.metadata:
parts.append(f"metadata={json.dumps(self.metadata, ensure_ascii=False)}")
if self.ref_memory_id:
parts.append(f"ref_memory_id={self.ref_memory_id}")
return " ".join(parts)
@classmethod
def from_vector_node(cls, node: VectorNode) -> "MemoryNode":
"""Reconstruct MemoryNode from VectorNode.

View file

@ -132,7 +132,7 @@ class Message(BaseModel):
def strip_md_func(line):
if strip_markdown_headers:
line = re.sub(r'\n##+ +', '\n', line)
line = re.sub(r"\n##+ +", "\n", line)
return line
if add_reasoning and self.reasoning_content:
@ -143,8 +143,9 @@ class Message(BaseModel):
elif isinstance(self.content, list):
for block in self.content:
text = block.content if isinstance(block.content, str) else \
json.dumps(block.content, ensure_ascii=False)
text = (
block.content if isinstance(block.content, str) else json.dumps(block.content, ensure_ascii=False)
)
text = str(text)
lines.append(strip_md_func(text))

View file

@ -1,11 +1,13 @@
"""Defines the standardized data structure for model output responses."""
from typing import Any
from pydantic import Field, BaseModel
class Response(BaseModel):
"""Represents a structured response containing the execution result, status, and metadata."""
answer: str | dict | list = Field(default="")
answer: str | Any = Field(default="")
success: bool = Field(default=True)
metadata: dict = Field(default_factory=dict)

View file

@ -1,7 +1,6 @@
"""Configuration schemas for service components using Pydantic models."""
import os
from typing import Dict, List
from pydantic import BaseModel, Field, ConfigDict
@ -99,15 +98,15 @@ class ServiceConfig(BaseModel):
thread_pool_max_workers: int = Field(default=16)
ray_max_workers: int = Field(default=-1)
init_logger: bool = Field(default=True)
disabled_flows: List[str] = Field(default_factory=list)
enabled_flows: List[str] = Field(default_factory=list)
mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration")
disabled_flows: list[str] = Field(default_factory=list)
enabled_flows: list[str] = Field(default_factory=list)
mcp_servers: dict[str, dict] = Field(default_factory=dict)
mcp: MCPConfig = Field(default_factory=MCPConfig)
http: HttpConfig = Field(default_factory=HttpConfig)
cmd: CmdConfig = Field(default_factory=CmdConfig)
flow: Dict[str, FlowConfig] = Field(default_factory=dict)
llm: Dict[str, LLMConfig] = Field(default_factory=dict)
embedding_model: Dict[str, EmbeddingModelConfig] = Field(default_factory=dict)
vector_store: Dict[str, VectorStoreConfig] = Field(default_factory=dict)
token_counter: Dict[str, TokenCounterConfig] = Field(default_factory=dict)
flow: dict[str, FlowConfig] = Field(default_factory=dict)
llm: dict[str, LLMConfig] = Field(default_factory=dict)
embedding_model: dict[str, EmbeddingModelConfig] = Field(default_factory=dict)
vector_store: dict[str, VectorStoreConfig] = Field(default_factory=dict)
token_counter: dict[str, TokenCounterConfig] = Field(default_factory=dict)

View file

@ -1,6 +1,4 @@
"""
MCP Tool Schema definitions for recursive JSON Schema representation.
"""
"""MCP Tool Schema definitions for recursive JSON Schema representation."""
import json
from typing import Any, Dict, List, Optional, Union
@ -41,14 +39,14 @@ class ToolAttr(BaseModel):
if self.enum:
res["enum"] = self.enum
if self.type == "object" and self.properties:
if self.type == "object" and self.properties is not None:
res["properties"] = {
k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items()
}
if self.required:
if self.required is not None:
res["required"] = self.required
if self.type == "array" and self.items:
if self.type == "array" and self.items is not None:
res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items
return res
@ -105,11 +103,6 @@ class ToolCall(BaseModel):
description="Specification for input parameters",
)
output: ToolAttr = Field(
default_factory=lambda: ToolAttr(type="object", properties={}),
description="Specification for the execution result (Schema)",
)
@model_validator(mode="before")
@classmethod
def init_tool_call(cls, data: dict) -> dict:
@ -147,6 +140,68 @@ class ToolCall(BaseModel):
},
}
def simple_output_dump(self) -> dict:
"""Convert ToolCall to output format dictionary for API responses."""
return {
"index": self.index,
"id": self.id,
self.type: {
"arguments": self.arguments,
"name": self.name,
},
"type": self.type,
}
@property
def argument_dict(self) -> dict:
"""Parse and return arguments as a dictionary."""
return json.loads(self.arguments)
def check_argument(self) -> bool:
"""Check if arguments can be parsed as valid JSON."""
try:
_ = self.argument_dict
return True
except Exception:
return False
def sanitize_and_check_argument(self) -> bool:
"""
Attempt to sanitize and validate arguments JSON.
Common issues from LLM streaming:
- Extra closing brackets: }]}] -> }]
- Missing closing brackets
- Trailing commas
"""
if not self.arguments or not self.arguments.strip():
return False
try:
# First try parsing as-is
_ = json.loads(self.arguments)
return True
except json.JSONDecodeError:
pass
# Try to fix common issues
sanitized = self.arguments.strip()
# Remove trailing extra brackets/braces
# Pattern: if it ends with multiple closing chars, try removing extras
while len(sanitized) > 1:
try:
json.loads(sanitized)
self.arguments = sanitized # Update with sanitized version
return True
except json.JSONDecodeError:
# Try removing last character
if sanitized[-1] in "]}":
sanitized = sanitized[:-1].rstrip()
else:
break
return False
@classmethod
def from_mcp_tool(cls, tool: Tool) -> "ToolCall":
"""Creates a ToolCall instance from an MCP Tool object."""
@ -164,28 +219,3 @@ class ToolCall(BaseModel):
description=self.description,
inputSchema=self.parameters.simple_input_dump(),
)
@property
def argument_dict(self) -> dict:
"""Parse and return arguments as a dictionary."""
return json.loads(self.arguments)
def check_argument(self) -> bool:
"""Check if arguments can be parsed as valid JSON."""
try:
_ = self.argument_dict
return True
except Exception:
return False
def simple_output_dump(self) -> dict:
"""Convert ToolCall to output format dictionary for API responses."""
return {
"index": self.index,
"id": self.id,
self.type: {
"arguments": self.arguments,
"name": self.name,
},
"type": self.type,
}

View file

@ -4,6 +4,7 @@ from .base_service import BaseService
from .cmd_service import CmdService
from .http_service import HttpService
from .mcp_service import MCPService
from ..context import R
__all__ = [
"BaseService",
@ -11,3 +12,7 @@ __all__ = [
"HttpService",
"MCPService",
]
R.service.register("cmd")(CmdService)
R.service.register("http")(HttpService)
R.service.register("mcp")(MCPService)

View file

@ -5,7 +5,7 @@ from abc import ABC, abstractmethod
from loguru import logger
from pydantic import BaseModel
from ..context import C
from ..context import ServiceContext
from ..flow import BaseFlow
from ..schema import ToolCall
from ..utils import create_pydantic_model
@ -14,8 +14,10 @@ from ..utils import create_pydantic_model
class BaseService(ABC):
"""Abstract base class for services that integrate and execute flows."""
def __init__(self, **kwargs):
def __init__(self, service_context: ServiceContext, **kwargs):
"""Initialize the base service."""
self.service_context: ServiceContext = service_context
self.service_config = self.service_context.service_config
self.kwargs = kwargs
@abstractmethod
@ -32,10 +34,10 @@ class BaseService(ABC):
def run(self):
"""Initialize and integrate all flows registered in the global context."""
flow_names: list[str] = []
for _, flow in C.flow_dict.items():
for flow in self.service_context.flows.values():
flow_name = self.integrate_flow(flow)
if flow_name:
flow_names.append(flow_name)
if flow_names:
logger.info(f"integrate {','.join(flow_names)}")
logger.info(f"Integrated {','.join(flow_names)}")

View file

@ -3,12 +3,10 @@
from loguru import logger
from .base_service import BaseService
from ..context import C
from ..flow import CmdFlow, BaseFlow
from ..utils.common_utils import run_coro_safely
@C.register_service("cmd")
class CmdService(BaseService):
"""Service implementation for handling command flow execution logic."""
@ -19,16 +17,16 @@ class CmdService(BaseService):
def integrate_flow(self, flow: BaseFlow) -> str | None:
"""Integrate the workflow configuration into the command service."""
self._cmd_flow = CmdFlow(flow=C.service_config.flow)
self._cmd_flow = CmdFlow(flow=self.service_config.cmd.flow)
def run(self):
"""Execute the command flow in either asynchronous or synchronous mode."""
super().run()
kwargs = self.service_config.cmd.model_extra
if self._cmd_flow.async_mode:
response = run_coro_safely(self._cmd_flow.call(**C.service_config.cmd.model_extra))
response = run_coro_safely(self._cmd_flow.call(**kwargs))
else:
response = self._cmd_flow.call_sync(**C.service_config.cmd.model_extra)
response = self._cmd_flow.call_sync(**kwargs)
if response.answer:
logger.info(f"response.answer={response.answer}")

View file

@ -9,20 +9,18 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
from .base_service import BaseService
from ..context import C
from ..flow import BaseFlow
from ..schema import Response
from ..utils.common_utils import execute_stream_task
@C.register_service("http")
class HttpService(BaseService):
"""Expose flows via HTTP REST and SSE endpoints."""
def __init__(self, **kwargs):
"""Initialize FastAPI app with CORS and health checks."""
super().__init__(**kwargs)
self.app = FastAPI(title=C.service_config.app_name)
self.app = FastAPI(title=self.service_config.app_name)
self.app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
@ -75,7 +73,7 @@ class HttpService(BaseService):
def run(self):
"""Start the Uvicorn server."""
super().run()
cfg = C.service_config.http
cfg = self.service_config.http
uvicorn.run(
self.app,
host=cfg.host,

View file

@ -1,23 +1,19 @@
"""Model Context Protocol (MCP) service implementation."""
from typing import Any
from fastmcp import FastMCP
from fastmcp.tools import FunctionTool
from .base_service import BaseService
from ..context import C
from ..flow import BaseFlow
@C.register_service("mcp")
class MCPService(BaseService):
"""Expose flows as Model Context Protocol (MCP) tools."""
def __init__(self, **kwargs: Any):
def __init__(self, **kwargs):
"""Initialize FastMCP instance with service settings."""
super().__init__(**kwargs)
self.mcp = FastMCP(name=C.service_config.app_name)
self.mcp = FastMCP(name=self.service_config.app_name)
def integrate_flow(self, flow: BaseFlow) -> str | None:
"""Register a non-streaming flow as an MCP tool."""
@ -45,12 +41,8 @@ class MCPService(BaseService):
def run(self):
"""Run the MCP server with specified transport protocol."""
super().run()
cfg = C.service_config.mcp
cfg = self.service_config.mcp
run_args: dict = {"transport": cfg.transport, "show_banner": False, **cfg.model_extra}
# Add network settings for non-stdio transports
if cfg.transport != "stdio":
run_args.update({"host": cfg.host, "port": cfg.port})
self.mcp.run(**run_args)

View file

@ -3,9 +3,14 @@
from .base_token_counter import BaseTokenCounter
from .hf_token_counter import HFTokenCounter
from .openai_token_counter import OpenAITokenCounter
from ..context import R
__all__ = [
"BaseTokenCounter",
"HFTokenCounter",
"OpenAITokenCounter",
]
R.token_counter.register("base")(BaseTokenCounter)
R.token_counter.register("hf")(HFTokenCounter)
R.token_counter.register("openai")(OpenAITokenCounter)

View file

@ -2,13 +2,12 @@
import math
import re
from loguru import logger
from ..context import C
from ..schema import Message, ToolCall
@C.register_token_counter("base")
class BaseTokenCounter:
"""A rule-based token counter for Chinese and non-Chinese text."""

View file

@ -5,11 +5,9 @@ import os
from loguru import logger
from .base_token_counter import BaseTokenCounter
from ..context import C
from ..schema import Message, ToolCall
@C.register_token_counter("hf")
class HFTokenCounter(BaseTokenCounter):
"""Token counter using transformers.AutoTokenizer.apply_chat_template."""

View file

@ -1,13 +1,13 @@
"""Token counting implementation for OpenAI-compatible models."""
import json
from loguru import logger
from .base_token_counter import BaseTokenCounter
from ..context import C
from ..schema import Message, ToolCall
@C.register_token_counter("openai")
class OpenAITokenCounter(BaseTokenCounter):
"""Token counter for OpenAI models using tiktoken."""

View file

@ -4,7 +4,7 @@ from .cache_handler import CacheHandler
from .case_converter import snake_to_camel, camel_to_snake
from .common_utils import run_coro_safely, execute_stream_task
from .env_utils import load_env
from .execute_tuils import exec_code, run_shell_command
from .execute_utils import exec_code, run_shell_command
from .http_client import HttpClient
from .llm_utils import extract_content, format_messages, deduplicate_memories
from .logger_utils import init_logger

View file

@ -15,7 +15,7 @@ class CacheHandler:
_EXTENSIONS = {
pd.DataFrame: ".csv",
dict: ".json",
list: ".json",
list: ".jsonl",
str: ".txt",
}
@ -76,11 +76,17 @@ class CacheHandler:
data.to_csv(path, index=kwargs.get("index", False), encoding="utf-8")
return {"row_count": len(data), "file_size": path.stat().st_size}
if dtype in (dict, list):
if dtype is dict:
with open(path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
return {"item_count": len(data), "file_size": path.stat().st_size}
if dtype is list:
with open(path, "w", encoding="utf-8") as f:
for item in data:
f.write(json.dumps(item, ensure_ascii=False) + "\n")
return {"item_count": len(data), "file_size": path.stat().st_size}
if dtype is str:
path.write_text(data, encoding=kwargs.get("encoding", "utf-8"))
return {"char_count": len(data), "file_size": path.stat().st_size}
@ -92,9 +98,17 @@ class CacheHandler:
"""Execute type-specific load operations."""
if type_name == "DataFrame":
return pd.read_csv(path, encoding=kwargs.get("encoding", "utf-8"))
if type_name in ("dict", "list"):
if type_name == "dict":
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
if type_name == "list":
result = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
result.append(json.loads(line))
return result
if type_name == "str":
return path.read_text(encoding=kwargs.get("encoding", "utf-8"))
raise ValueError(f"Unknown data type in metadata: {type_name}")

Some files were not shown because too many files have changed in this diff Show more