mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
commit
888ccea5f5
243 changed files with 8758 additions and 4280 deletions
4
.gitignore
vendored
4
.gitignore
vendored
|
|
@ -36,4 +36,6 @@ test_working_memory/*
|
|||
local_vector_store/*
|
||||
chroma_vector_store/*
|
||||
bench_results/*
|
||||
meta_memory/*
|
||||
meta_memory/*
|
||||
*.sqlite3
|
||||
**/data/*.json
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
285
bench/halumem/compute_qa_stats_v4.py
Normal file
285
bench/halumem/compute_qa_stats_v4.py
Normal 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")
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
676
bench/halumem/eval_reme_simple_v3.py
Normal file
676
bench/halumem/eval_reme_simple_v3.py
Normal 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
|
||||
)
|
||||
690
bench/halumem/eval_reme_simple_v4.py
Normal file
690
bench/halumem/eval_reme_simple_v4.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}}
|
||||
```
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
237
bench/human_in_the_loop/compute_qa_stats.py
Normal file
237
bench/human_in_the_loop/compute_qa_stats.py
Normal 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)
|
||||
145
bench/human_in_the_loop/eval.yaml
Normal file
145
bench/human_in_the_loop/eval.yaml
Normal 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"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
550
bench/human_in_the_loop/reevaluate_qa.py
Normal file
550
bench/human_in_the_loop/reevaluate_qa.py
Normal 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
|
||||
))
|
||||
237
bench/human_in_the_loop2/compute_qa_stats.py
Normal file
237
bench/human_in_the_loop2/compute_qa_stats.py
Normal 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)
|
||||
145
bench/human_in_the_loop2/eval.yaml
Normal file
145
bench/human_in_the_loop2/eval.yaml
Normal 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"
|
||||
}}
|
||||
```
|
||||
"""
|
||||
550
bench/human_in_the_loop2/reevaluate_qa.py
Normal file
550
bench/human_in_the_loop2/reevaluate_qa.py
Normal 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
3
docs/todo.md
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
1. 如何更好的注册class
|
||||
2. op的返回,使用return 还是 self.output
|
||||
3. 如何把agent的东西放出来
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
19
reme/__init__.py
Normal 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
7
reme/agent/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""A simple chatbot."""
|
||||
|
||||
from . import chat
|
||||
|
||||
__all__ = [
|
||||
"chat",
|
||||
]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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:
|
||||
|
|
@ -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:
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""Configuration parser for ReMe framework."""
|
||||
|
||||
from ..utils import PydanticConfigParser
|
||||
from ..core.utils import PydanticConfigParser
|
||||
|
||||
|
||||
class ReMeConfigParser(PydanticConfigParser):
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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",
|
||||
]
|
||||
363
reme/core/context/prompt_handler.py
Normal file
363
reme/core/context/prompt_handler.py
Normal 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)})"
|
||||
45
reme/core/context/registry_factory.py
Normal file
45
reme/core/context/registry_factory.py
Normal 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()
|
||||
|
|
@ -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."""
|
||||
230
reme/core/context/service_context.py
Normal file
230
reme/core/context/service_context.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
@ -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."""
|
||||
|
||||
38
reme/core/enumeration/json_schema_enum.py
Normal file
38
reme/core/enumeration/json_schema_enum.py
Normal 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()
|
||||
33
reme/core/enumeration/memory_type.py
Normal file
33
reme/core/enumeration/memory_type.py
Normal 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"
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
},
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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
58
reme/core/op/base_tool.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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.
|
||||
|
|
@ -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))
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)}")
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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,
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
@ -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
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue