mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
feat(benchmark): add HaluMem baseline evaluation and analysis tools
This commit is contained in:
parent
6125dff01e
commit
b8124fe31a
14 changed files with 1430 additions and 540 deletions
180
bench/halumem/analyze_results.py
Normal file
180
bench/halumem/analyze_results.py
Normal file
|
|
@ -0,0 +1,180 @@
|
|||
"""
|
||||
分析 bench_results/reme_simple/tmp 目录下的评估结果
|
||||
|
||||
统计所有用户session中的result_type分布,并输出非Correct结果的详细位置信息。
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
|
||||
def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
|
||||
"""
|
||||
分析评估结果目录。
|
||||
|
||||
Args:
|
||||
tmp_dir: 临时结果目录路径
|
||||
"""
|
||||
tmp_path = Path(tmp_dir)
|
||||
|
||||
if not tmp_path.exists():
|
||||
print(f"❌ 目录不存在: {tmp_dir}")
|
||||
return
|
||||
|
||||
# 统计数据
|
||||
result_counter = Counter()
|
||||
non_correct_results = [] # 存储非Correct结果的详细信息
|
||||
|
||||
# 遍历所有用户目录
|
||||
user_dirs = sorted([d for d in tmp_path.iterdir() if d.is_dir()])
|
||||
|
||||
if not user_dirs:
|
||||
print(f"❌ {tmp_dir} 下没有用户目录")
|
||||
return
|
||||
|
||||
print(f"📁 找到 {len(user_dirs)} 个用户目录\n")
|
||||
print("=" * 80)
|
||||
print("开始分析...")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
total_sessions = 0
|
||||
total_questions = 0
|
||||
|
||||
# 遍历每个用户目录
|
||||
for user_dir in user_dirs:
|
||||
user_name = user_dir.name
|
||||
|
||||
# 获取该用户的所有session文件
|
||||
session_files = sorted([
|
||||
f for f in user_dir.iterdir()
|
||||
if f.name.startswith("session_") and f.suffix == ".json"
|
||||
])
|
||||
|
||||
if not session_files:
|
||||
continue
|
||||
|
||||
# 遍历每个session
|
||||
for session_file in session_files:
|
||||
try:
|
||||
with open(session_file, "r", encoding="utf-8") as f:
|
||||
session_data = json.load(f)
|
||||
|
||||
session_id = session_data.get("session_id", -1)
|
||||
total_sessions += 1
|
||||
|
||||
# 跳过生成的QA session
|
||||
if session_data.get("is_generated_qa_session", False):
|
||||
continue
|
||||
|
||||
# 获取评估结果
|
||||
eval_results = session_data.get("evaluation_results", {})
|
||||
qa_records = eval_results.get("question_answering_records", [])
|
||||
|
||||
# 分析每个问题的结果
|
||||
for qa_idx, qa_record in enumerate(qa_records):
|
||||
result_type = qa_record.get("result_type", "Unknown")
|
||||
|
||||
# 统计result_type
|
||||
result_counter[result_type] += 1
|
||||
total_questions += 1
|
||||
|
||||
# 如果不是Correct,记录详细信息
|
||||
if result_type != "Correct":
|
||||
non_correct_results.append({
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"question_id": qa_idx,
|
||||
"result_type": result_type,
|
||||
"question": qa_record.get("question", ""),
|
||||
"answer": qa_record.get("answer", ""),
|
||||
"system_response": qa_record.get("system_response", "")
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
print(f"⚠️ 读取文件失败: {session_file}, 错误: {e}")
|
||||
continue
|
||||
|
||||
# 输出统计结果
|
||||
print("\n" + "=" * 80)
|
||||
print("统计结果")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print(f"📊 总用户数: {len(user_dirs)}")
|
||||
print(f"📊 总Session数: {total_sessions}")
|
||||
print(f"📊 总问题数: {total_questions}\n")
|
||||
|
||||
if total_questions == 0:
|
||||
print("❌ 没有找到任何问题数据")
|
||||
return
|
||||
|
||||
# 输出result_type分布
|
||||
print("=" * 80)
|
||||
print("Result Type 分布")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# 按数量降序排列
|
||||
sorted_results = sorted(result_counter.items(), key=lambda x: x[1], reverse=True)
|
||||
|
||||
for result_type, count in sorted_results:
|
||||
ratio = count / total_questions * 100
|
||||
print(f" {result_type:20s}: {count:5d} ({ratio:6.2f}%)")
|
||||
|
||||
# 输出非Correct结果的详细信息
|
||||
if non_correct_results:
|
||||
print("\n" + "=" * 80)
|
||||
print(f"非 Correct 结果详情 (共 {len(non_correct_results)} 条)")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
for idx, result in enumerate(non_correct_results, 1):
|
||||
print(f"[{idx}] {result['result_type']}")
|
||||
print(f" 用户: {result['user_name']}")
|
||||
print(f" 位置: Session {result['session_id']}, Question {result['question_id']}")
|
||||
print(f" 问题: {result['question']}")
|
||||
print(f" 正确答案: {result['answer']}")
|
||||
print(f" 系统回答: {result['system_response'][:200]}{'...' if len(result['system_response']) > 200 else ''}")
|
||||
print()
|
||||
|
||||
else:
|
||||
print("\n🎉 所有问题都是 Correct!")
|
||||
|
||||
# 保存详细报告到文件
|
||||
report_file = Path(tmp_dir).parent / "analysis_report.json"
|
||||
report_data = {
|
||||
"summary": {
|
||||
"total_users": len(user_dirs),
|
||||
"total_sessions": total_sessions,
|
||||
"total_questions": total_questions,
|
||||
"result_type_distribution": dict(result_counter),
|
||||
"result_type_ratio": {
|
||||
result_type: count / total_questions
|
||||
for result_type, count in result_counter.items()
|
||||
}
|
||||
},
|
||||
"non_correct_results": non_correct_results
|
||||
}
|
||||
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(report_data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
print("=" * 80)
|
||||
print(f"📄 详细报告已保存到: {report_file}")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="分析 ReMe 评估结果中的 result_type 分布"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tmp_dir",
|
||||
type=str,
|
||||
default="bench_results/reme_simple/tmp",
|
||||
help="临时结果目录路径 (默认: bench_results/reme_simple/tmp)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
analyze_results(args.tmp_dir)
|
||||
604
bench/halumem/eval_baseline_simple.py
Normal file
604
bench/halumem/eval_baseline_simple.py
Normal file
|
|
@ -0,0 +1,604 @@
|
|||
"""
|
||||
HaluMem Benchmark Evaluator - Baseline (Direct QA without Memory System)
|
||||
|
||||
A simple baseline evaluation pipeline that:
|
||||
1. Loads HaluMem benchmark data
|
||||
2. Directly uses dialogue history to answer questions (no memory system)
|
||||
3. Evaluates question answering performance
|
||||
4. Generates comprehensive metrics
|
||||
|
||||
Usage:
|
||||
python bench/halumem/eval_baseline_simple.py \
|
||||
--data_path /path/to/HaluMem-Medium.jsonl \
|
||||
--user_num 100 --max_concurrency 20
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from eval_tools import evaluation_for_question2
|
||||
from llms import llm_request_for_json
|
||||
|
||||
|
||||
# ==================== Configuration ====================
|
||||
|
||||
@dataclass
|
||||
class EvalConfig:
|
||||
"""Evaluation configuration parameters."""
|
||||
data_path: str
|
||||
user_num: int = 1
|
||||
max_concurrency: int = 2
|
||||
output_dir: str = "bench_results/baseline_simple"
|
||||
|
||||
|
||||
# ==================== Utilities ====================
|
||||
|
||||
class DataLoader:
|
||||
"""Handles loading and parsing of HaluMem data."""
|
||||
|
||||
@staticmethod
|
||||
def load_jsonl(file_path: str) -> list[dict]:
|
||||
"""Load all entries from a JSONL file."""
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
return [json.loads(line.strip()) for line in f if line.strip()]
|
||||
|
||||
@staticmethod
|
||||
def extract_user_name(persona_info: str) -> str:
|
||||
"""Extract user name from persona info string."""
|
||||
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
|
||||
if not match:
|
||||
raise ValueError(f"No name found in persona_info: {persona_info}")
|
||||
return match.group(1).strip()
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
|
||||
"""Format dialogue into ReMe message format."""
|
||||
return [
|
||||
{
|
||||
"role": turn["role"],
|
||||
"content": turn["content"],
|
||||
"time_created": datetime.strptime(
|
||||
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
|
||||
)
|
||||
.replace(tzinfo=timezone.utc)
|
||||
.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
for turn in dialogue
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
|
||||
"""Format dialogue into string for evaluation."""
|
||||
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")
|
||||
|
||||
|
||||
# ==================== Question Answering Prompt ====================
|
||||
|
||||
BASELINE_QA_PROMPT = """You are a helpful AI assistant. Based on the dialogue history provided below, please answer the question.
|
||||
|
||||
**Dialogue History:**
|
||||
{dialogue}
|
||||
|
||||
**Question:**
|
||||
{question}
|
||||
|
||||
**Instructions:**
|
||||
- Carefully read through the dialogue history
|
||||
- Answer the question based ONLY on information present in the dialogue
|
||||
- If the information needed to answer the question is NOT in the dialogue, respond with "I don't know" or "The information is not available in the dialogue"
|
||||
- Do NOT make up or hallucinate information that is not explicitly mentioned in the dialogue
|
||||
- Provide your reasoning process before giving the final answer
|
||||
|
||||
**Response Format:**
|
||||
Please respond in JSON format with the following structure:
|
||||
```json
|
||||
{{
|
||||
"reasoning": "Your step-by-step reasoning process",
|
||||
"answer": "Your final answer (or 'I don't know' if information is not available)"
|
||||
}}
|
||||
```"""
|
||||
|
||||
|
||||
# ==================== Evaluation ====================
|
||||
|
||||
class BaselineQuestionAnsweringEvaluator:
|
||||
"""Evaluates question answering performance using direct LLM inference (no memory system)."""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
async def answer_question(
|
||||
self,
|
||||
question: str,
|
||||
formatted_dialogue: str
|
||||
) -> tuple[str, str, float]:
|
||||
"""
|
||||
Answer a question using the dialogue history directly.
|
||||
|
||||
Returns:
|
||||
tuple: (answer, reasoning, duration_ms)
|
||||
"""
|
||||
start = time.time()
|
||||
|
||||
# Format prompt
|
||||
prompt = BASELINE_QA_PROMPT.format(
|
||||
dialogue=formatted_dialogue,
|
||||
question=question
|
||||
)
|
||||
|
||||
# Get answer from LLM
|
||||
try:
|
||||
# model_name = "qwen3-max"
|
||||
model_name = "qwen3-30b-a3b-instruct-2507"
|
||||
result = await llm_request_for_json(prompt, model_name=model_name)
|
||||
answer = result.get("answer", "I don't know")
|
||||
reasoning = result.get("reasoning", "")
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting answer from LLM: {e}")
|
||||
answer = "Error: Failed to get answer"
|
||||
reasoning = str(e)
|
||||
|
||||
duration_ms = (time.time() - start) * 1000
|
||||
return answer, reasoning, duration_ms
|
||||
|
||||
async def evaluate_questions(
|
||||
self,
|
||||
questions: list[dict],
|
||||
user_name: str,
|
||||
uuid: str,
|
||||
session_id: int,
|
||||
formatted_dialogue: str
|
||||
) -> list[dict]:
|
||||
"""Evaluate all questions for a session."""
|
||||
results = []
|
||||
|
||||
for qa in questions:
|
||||
# Get answer directly from LLM
|
||||
answer, reasoning, duration_ms = await self.answer_question(
|
||||
question=qa["question"],
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
# Evaluate response
|
||||
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
|
||||
eval_result = await evaluation_for_question2(
|
||||
qa["question"],
|
||||
qa["answer"],
|
||||
evidence_text,
|
||||
answer,
|
||||
formatted_dialogue
|
||||
)
|
||||
|
||||
# Build result record
|
||||
qa_result = {
|
||||
**qa,
|
||||
"uuid": uuid,
|
||||
"session_id": session_id,
|
||||
"system_response": answer,
|
||||
"reasoning": reasoning,
|
||||
"answer_duration_ms": duration_ms,
|
||||
"result_type": eval_result.get("evaluation_result"),
|
||||
"question_answering_reasoning": eval_result.get("reasoning", "")
|
||||
}
|
||||
results.append(qa_result)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class MetricsAggregator:
|
||||
"""Aggregates evaluation metrics."""
|
||||
|
||||
@staticmethod
|
||||
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
|
||||
"""Compute question answering metrics."""
|
||||
total = len(qa_records)
|
||||
if total == 0:
|
||||
return {
|
||||
"correct_qa_ratio(all)": 0,
|
||||
"hallucination_qa_ratio(all)": 0,
|
||||
"omission_qa_ratio(all)": 0,
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0,
|
||||
"qa_valid_num": 0,
|
||||
"qa_num": 0
|
||||
}
|
||||
|
||||
correct = 0
|
||||
hallucination = 0
|
||||
omission = 0
|
||||
valid = 0
|
||||
|
||||
for qa in qa_records:
|
||||
result_type = qa.get("result_type", "")
|
||||
|
||||
if result_type in ["Correct", "Hallucination", "Omission"]:
|
||||
valid += 1
|
||||
if result_type == "Correct":
|
||||
correct += 1
|
||||
elif result_type == "Hallucination":
|
||||
hallucination += 1
|
||||
elif result_type == "Omission":
|
||||
omission += 1
|
||||
|
||||
metrics = {
|
||||
"correct_qa_ratio(all)": correct / total,
|
||||
"hallucination_qa_ratio(all)": hallucination / total,
|
||||
"omission_qa_ratio(all)": omission / total,
|
||||
"qa_valid_num": valid,
|
||||
"qa_num": total
|
||||
}
|
||||
|
||||
if valid > 0:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": correct / valid,
|
||||
"hallucination_qa_ratio(valid)": hallucination / valid,
|
||||
"omission_qa_ratio(valid)": omission / valid
|
||||
})
|
||||
else:
|
||||
metrics.update({
|
||||
"correct_qa_ratio(valid)": 0,
|
||||
"hallucination_qa_ratio(valid)": 0,
|
||||
"omission_qa_ratio(valid)": 0
|
||||
})
|
||||
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
|
||||
"""Compute timing metrics from evaluation results."""
|
||||
answer_duration = 0
|
||||
|
||||
with open(eval_results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
for qa in eval_results.get("question_answering_records", []):
|
||||
answer_duration += qa.get("answer_duration_ms", 0)
|
||||
|
||||
# Convert to minutes
|
||||
return {
|
||||
"answer_duration_time": answer_duration / 1000 / 60,
|
||||
"total_duration_time": answer_duration / 1000 / 60
|
||||
}
|
||||
|
||||
|
||||
# ==================== Main Pipeline ====================
|
||||
|
||||
class HaluMemBaselineEvaluator:
|
||||
"""Main evaluator orchestrating the baseline evaluation pipeline."""
|
||||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.qa_evaluator = BaselineQuestionAnsweringEvaluator()
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
async def process_session(
|
||||
self,
|
||||
session: dict,
|
||||
session_id: int,
|
||||
user_name: str,
|
||||
uuid: str
|
||||
) -> dict:
|
||||
"""Process a single session."""
|
||||
session_data = {
|
||||
"uuid": uuid,
|
||||
"user_name": user_name,
|
||||
"session_id": session_id,
|
||||
"memory_points": session["memory_points"]
|
||||
}
|
||||
|
||||
# Skip generated QA sessions
|
||||
if session.get("is_generated_qa_session", False):
|
||||
session_data["is_generated_qa_session"] = True
|
||||
return session_data
|
||||
|
||||
# Store dialogue
|
||||
dialogue = session["dialogue"]
|
||||
session_data["dialogue"] = dialogue
|
||||
|
||||
# Evaluate questions if present
|
||||
if "questions" in session:
|
||||
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
|
||||
qa_results = await self.qa_evaluator.evaluate_questions(
|
||||
questions=session["questions"],
|
||||
user_name=user_name,
|
||||
uuid=uuid,
|
||||
session_id=session_id,
|
||||
formatted_dialogue=formatted_dialogue
|
||||
)
|
||||
|
||||
session_data["evaluation_results"] = {
|
||||
"question_answering_records": qa_results
|
||||
}
|
||||
|
||||
return session_data
|
||||
|
||||
async def process_user(self, user_data: dict) -> dict:
|
||||
"""Process all sessions for a user."""
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
uuid = user_data["uuid"]
|
||||
|
||||
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()
|
||||
|
||||
# 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)} | Concurrency: {self.config.max_concurrency}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Process users with concurrency control
|
||||
semaphore = asyncio.Semaphore(self.config.max_concurrency)
|
||||
|
||||
async def process_with_cache_check(idx: int, user_data: dict):
|
||||
async with semaphore:
|
||||
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
|
||||
|
||||
# Check cache
|
||||
if self.file_manager.user_has_cache(user_name):
|
||||
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
|
||||
return {"user_name": user_name, "status": "cached"}
|
||||
|
||||
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
|
||||
result = await self.process_user(user_data)
|
||||
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
|
||||
return result
|
||||
|
||||
tasks = [
|
||||
process_with_cache_check(idx, user)
|
||||
for idx, user in enumerate(users_to_process, 1)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Combine results
|
||||
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
|
||||
self.file_manager.combine_results(output_file)
|
||||
|
||||
elapsed = time.time() - start_time
|
||||
print(f"\n✅ Processing completed in {elapsed:.2f}s")
|
||||
print(f"📁 Results: {output_file}\n")
|
||||
|
||||
# Aggregate metrics
|
||||
await self.aggregate_and_report(output_file)
|
||||
|
||||
async def aggregate_and_report(self, results_file: str):
|
||||
"""Aggregate results and generate final report."""
|
||||
print("=" * 80)
|
||||
print("AGGREGATING METRICS")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
# Collect all QA records
|
||||
qa_records = []
|
||||
with open(results_file, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
user_data = json.loads(line)
|
||||
|
||||
for session in user_data["sessions"]:
|
||||
if session.get("is_generated_qa_session"):
|
||||
continue
|
||||
|
||||
eval_results = session.get("evaluation_results", {})
|
||||
qa_records.extend(
|
||||
eval_results.get("question_answering_records", [])
|
||||
)
|
||||
|
||||
# Compute metrics
|
||||
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
|
||||
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
|
||||
|
||||
final_results = {
|
||||
"overall_score": {
|
||||
"question_answering": qa_metrics,
|
||||
"time_consuming": time_metrics
|
||||
},
|
||||
"question_answering_records": qa_records
|
||||
}
|
||||
|
||||
# Save final report
|
||||
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
|
||||
with open(report_file, "w", encoding="utf-8") as f:
|
||||
json.dump(final_results, f, ensure_ascii=False, indent=4)
|
||||
|
||||
print(f"📊 Statistics saved to: {report_file}\n")
|
||||
|
||||
# Print summary
|
||||
self._print_summary(qa_metrics, time_metrics)
|
||||
|
||||
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
|
||||
"""Print evaluation summary."""
|
||||
print("=" * 80)
|
||||
print("EVALUATION SUMMARY")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
print("📊 Question Answering:")
|
||||
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
|
||||
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
|
||||
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
|
||||
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
|
||||
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
|
||||
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
|
||||
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
|
||||
|
||||
print(f"\n⏱️ Time Metrics:")
|
||||
print(f" Answer Duration: {time_metrics['answer_duration_time']:.2f} min")
|
||||
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
|
||||
print("\n" + "=" * 80)
|
||||
|
||||
|
||||
# ==================== Entry Point ====================
|
||||
|
||||
def main(
|
||||
data_path: str,
|
||||
user_num: int = 1,
|
||||
max_concurrency: int = 2
|
||||
):
|
||||
"""Main entry point."""
|
||||
config = EvalConfig(
|
||||
data_path=data_path,
|
||||
user_num=user_num,
|
||||
max_concurrency=max_concurrency
|
||||
)
|
||||
|
||||
evaluator = HaluMemBaselineEvaluator(config)
|
||||
asyncio.run(evaluator.run_evaluation())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate Baseline (Direct QA) on HaluMem benchmark"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--user_num",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of users to evaluate (default: 1)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_concurrency",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Maximum concurrent user processing (default: 2)"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
main(
|
||||
data_path=args.data_path,
|
||||
user_num=args.user_num,
|
||||
max_concurrency=args.max_concurrency
|
||||
)
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -142,6 +142,9 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
|
|||
if op.memory_nodes:
|
||||
self.memory_nodes.extend(op.memory_nodes)
|
||||
|
||||
if hasattr(op, "messages") and op.messages:
|
||||
self.messages.extend(op.messages)
|
||||
|
||||
tool_result = str(op.output)
|
||||
tool_message = Message(
|
||||
role=Role.TOOL,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ from ...core.utils import format_messages
|
|||
|
||||
@C.register_op()
|
||||
class PersonalSummarizerV2(BaseMemoryAgent):
|
||||
memory_type: MemoryType = MemoryType.PERSONAL
|
||||
|
||||
"""Simplified personal memory summarizer that uses v2 memory tools.
|
||||
|
||||
This summarizer follows a three-step workflow:
|
||||
|
|
@ -17,11 +19,6 @@ class PersonalSummarizerV2(BaseMemoryAgent):
|
|||
3. UpdateMemories: Delete outdated memories and add new ones
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
memory_type: MemoryType = MemoryType.PERSONAL
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
"""Build tool call schema for the agent."""
|
||||
return ToolCall(
|
||||
|
|
@ -83,28 +80,28 @@ class PersonalSummarizerV2(BaseMemoryAgent):
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
# Check if AddMemoryDrafts tool was executed
|
||||
exist_memory_drafts = False
|
||||
if assistant_message.tool_calls:
|
||||
for tool_call in assistant_message.tool_calls:
|
||||
if tool_call.name == "add_memory_drafts":
|
||||
exist_memory_drafts = True
|
||||
break
|
||||
|
||||
# If memory drafts were added, regenerate system prompt with simplified context
|
||||
if exist_memory_drafts:
|
||||
simplified_context = "The conversation context has been summarized in memory drafts."
|
||||
new_system_prompt = self.prompt_format(
|
||||
prompt_name="system_prompt",
|
||||
context=simplified_context,
|
||||
memory_type=self.memory_type.value,
|
||||
memory_target=self.memory_target,
|
||||
)
|
||||
|
||||
# Update the system message in the message history
|
||||
for i, msg in enumerate(self.messages):
|
||||
if msg.role == Role.SYSTEM:
|
||||
self.messages[i] = Message(role=Role.SYSTEM, content=new_system_prompt)
|
||||
break
|
||||
# # Check if AddMemoryDrafts tool was executed
|
||||
# exist_memory_drafts = False
|
||||
# if assistant_message.tool_calls:
|
||||
# for tool_call in assistant_message.tool_calls:
|
||||
# if tool_call.name == "add_memory_drafts":
|
||||
# exist_memory_drafts = True
|
||||
# break
|
||||
#
|
||||
# # If memory drafts were added, regenerate system prompt with simplified context
|
||||
# if exist_memory_drafts:
|
||||
# simplified_context = "The conversation context has been summarized in memory drafts."
|
||||
# new_system_prompt = self.prompt_format(
|
||||
# prompt_name="system_prompt",
|
||||
# context=simplified_context,
|
||||
# memory_type=self.memory_type.value,
|
||||
# memory_target=self.memory_target,
|
||||
# )
|
||||
#
|
||||
# # Update the system message in the message history
|
||||
# for i, msg in enumerate(self.messages):
|
||||
# if msg.role == Role.SYSTEM:
|
||||
# self.messages[i] = Message(role=Role.SYSTEM, content=new_system_prompt)
|
||||
# break
|
||||
|
||||
return messages
|
||||
|
|
|
|||
|
|
@ -3,58 +3,38 @@ tool: |
|
|||
Use this tool to analyze dialogues and extract important personal information about users,
|
||||
such as preferences, habits, personal background, relationships, and significant facts.
|
||||
|
||||
# - **Memory granularity**: Each memory should record ONE complete piece of information - don't pack multiple facts into one memory, and don't split a single fact into multiple memories.
|
||||
# - **Self-contained**: Each memory entry must be self-contained and understandable without additional context.
|
||||
|
||||
system_prompt: |
|
||||
You are a professional memory agent. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
|
||||
You are a professional memory agent managing **{memory_type}** memories about **{memory_target}** for the main agent.
|
||||
|
||||
**CRITICAL**: You must extract and store information STRICTLY based on what is explicitly stated in the context. DO NOT infer, assume, fabricate, or add any information that is not directly present in the dialogue. Only extract facts that are clearly and explicitly mentioned.
|
||||
|
||||
## Context:
|
||||
## Latest Conversation:
|
||||
The context below contains the most recent conversation. Each message is formatted as: `round<index> [<timestamp>] <role/name>: <content>` where timestamp is `YYYY-MM-DD HH:MM:SS`.
|
||||
{context}
|
||||
|
||||
**Context Format Explanation**:
|
||||
The context contains formatted conversation messages in the following structure:
|
||||
- Each message is formatted as: `round<index> [<timestamp>] <role/name>: <content>`
|
||||
- The timestamp is in format: `YYYY-MM-DD HH:MM:SS`
|
||||
- Content may include reasoning, tool calls
|
||||
- **Time metadata handling**: When extracting memories with time information, store year/month/day in the metadata. For relative time references (e.g., "last year", "two months ago"), calculate the actual date based on the message's timestamp and store the calculated year/month/day in metadata
|
||||
|
||||
## Memory Objective:
|
||||
You are managing **{memory_type}** memories about **{memory_target}** for the main agent. Focus on extracting and storing information directly related to this person's preferences, habits, personal background, and significant facts.
|
||||
**CRITICAL**: Extract information ONLY from what is explicitly stated. DO NOT infer, assume, or fabricate any information.
|
||||
|
||||
## Your Tasks - Three-Step Workflow:
|
||||
## Your Tasks
|
||||
|
||||
### Step 1: Generate Memory Drafts
|
||||
Use the `AddMemoryDrafts` tool to create initial memory drafts from the conversation context.
|
||||
- **Analyze the context**: Determine whether the conversation contains important, memorable information, including but not limited to: user preferences, habits, or personal details; key facts, decisions, or conclusions; relationships or contextual background related to people or topics.
|
||||
- **Extract key information**: Create memory drafts using clear and concise phrasing **strictly based on what is explicitly stated in the context**.
|
||||
- **Important**: DO NOT infer, assume, or add any information beyond what is directly mentioned in the conversation.
|
||||
- **Time references**: If the context involves time references (like "last year", "two months ago", etc.), you **must** calculate the actual date based on the context's timestamp metadata. For example, if a memory from May 4, 2022 mentions "went to India last year," then the trip occurred in 2021. Include this calculated time information in the memory metadata (year, month, day).
|
||||
- **Memory granularity**: Each memory should record ONE complete piece of information - don't pack multiple facts into one memory, and don't split a single fact into multiple memories.
|
||||
- **Self-contained**: Each memory entry must be self-contained and understandable without additional context.
|
||||
Use `AddMemoryDrafts` to extract key facts from the latest conversation.
|
||||
- Extract important information: preferences, habits, currentstatus, personal details, key facts, decisions, or conclusions.
|
||||
- Use clear, concise phrasing based strictly on explicit statements.
|
||||
- Record the timestamp of the source message for each memory including the year, month, and day.
|
||||
|
||||
### Step 2: Retrieve Similar and Recent Memories
|
||||
Use the `RetrieveRecentAndSimilarMemories` tool to find existing related memories.
|
||||
- **For EACH memory draft**, perform a semantic similarity search to find existing, potentially relevant memories.
|
||||
- **Example**: For "Person A was born on date X", search for "Person A birth date age".
|
||||
- **Retrieve comprehensively**: Retrieve all related memories for thorough comparison to prevent any duplication or conflicts.
|
||||
Use `RetrieveRecentAndSimilarMemories` to query historical memories for each draft.
|
||||
- Search for semantically similar memories and recent memories.
|
||||
- This ensures Step 3 avoids duplicates and properly updates existing memories.
|
||||
|
||||
### Step 3: Update Memories
|
||||
Use the `UpdateMemories` tool to finalize the memory updates.
|
||||
- **Compare and decide**: Compare the newly extracted memory drafts with the retrieved memories from Step 2.
|
||||
- **CRITICAL DEDUPLICATION CHECK**: Before adding ANY new memory:
|
||||
- Check if the SAME INFORMATION already exists in retrieved memories
|
||||
- Consider memories as duplicates even if wording differs, as long as they convey the SAME core fact
|
||||
- Examples of duplicate information:
|
||||
* "Person A was born on date X. He/She is N years old." vs "Person A is a gender born on date X. He/She is currently N years old." → DUPLICATES
|
||||
* "Lives in city" vs "Person A lives in city" → DUPLICATES
|
||||
* "Holds a Bachelor's degree in field" vs "Person A holds a Bachelor's degree in field" → DUPLICATES
|
||||
|
||||
- **Choose the appropriate operation**:
|
||||
- **If the information already exists and is consistent**: SKIP—fill empty array in `memory_ids_to_delete` and `memories_to_add`. Do NOT add duplicate memories.
|
||||
- **If existing memory needs supplementation with NEW details**: Delete the old memory (add its ID to `memory_ids_to_delete`), then add the enhanced consolidated version to `memories_to_add`.
|
||||
- **If existing memory is outdated or contradicted**: Delete it (add ID to `memory_ids_to_delete`), then add the corrected version to `memories_to_add`.
|
||||
- **If multiple memories contain similar/overlapping information**: Delete all duplicates (add IDs to `memory_ids_to_delete`), then add one merged memory to `memories_to_add`.
|
||||
- **If the information is entirely new**: Fill empty array in `memory_ids_to_delete`, and add the new memory to `memories_to_add`.
|
||||
Use `UpdateMemories` to update the memory store by combining drafts with historical memories.
|
||||
- **Delete conflicts**: Remove old memories that contradict the new drafts (keep most recent/accurate).
|
||||
- **Add new**: Add drafts that represent completely new information.
|
||||
- **Skip duplicates**: Do not add drafts that duplicate existing memories.
|
||||
- **Preserve others**: Keep unrelated historical memories unchanged.
|
||||
- Write concise memories using minimum words needed. Ensure no information loss.
|
||||
|
||||
user_message: |
|
||||
Please analyze the context and update the memory store following the three-step workflow:
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class ReMeSummarizerV2(BaseMemoryAgent):
|
|||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"description": self.prompt_format("tool"),
|
||||
"description": self.get_prompt("tool"),
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
|
|
|
|||
|
|
@ -118,7 +118,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
|
|||
metadata=metadata or {},
|
||||
)
|
||||
|
||||
logger.opt(depth=1).info(
|
||||
f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}",
|
||||
)
|
||||
# logger.opt(depth=1).info(
|
||||
# f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}",
|
||||
# )
|
||||
return node
|
||||
|
|
|
|||
|
|
@ -102,29 +102,5 @@ class AddMemoryDrafts(BaseMemoryTool):
|
|||
|
||||
async def execute(self):
|
||||
"""Execute add drafts operation: create memory drafts without persisting to vector store."""
|
||||
# Get memory drafts to add
|
||||
memory_drafts = self.context.get("memory_drafts", [])
|
||||
|
||||
# Validate input
|
||||
if not memory_drafts:
|
||||
self.output = "No memory drafts provided. Please provide at least one draft memory."
|
||||
return
|
||||
|
||||
# Build memory nodes (without persisting)
|
||||
memory_nodes = []
|
||||
for mem in memory_drafts:
|
||||
memory_content, when_to_use, metadata = self._extract_memory_data(mem)
|
||||
if not memory_content:
|
||||
logger.warning("Skipping memory draft with empty content")
|
||||
continue
|
||||
|
||||
memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata))
|
||||
|
||||
if memory_nodes:
|
||||
self.memory_nodes.extend(memory_nodes)
|
||||
draft_count = len(memory_nodes)
|
||||
self.output = f"Successfully created {draft_count} memory draft(s). These drafts are not yet persisted to the vector store."
|
||||
logger.info(self.output)
|
||||
else:
|
||||
self.output = "No valid memory drafts created. Please check your input."
|
||||
logger.warning(self.output)
|
||||
self.output = f"Successfully created memory draft(s). These drafts are not yet persisted to the vector store."
|
||||
logger.info(self.output)
|
||||
|
|
|
|||
|
|
@ -161,9 +161,6 @@ class RetrieveRecentAndSimilarMemories(BaseMemoryTool):
|
|||
# Update retrieved_nodes in context with new memories
|
||||
self.retrieved_nodes.extend(new_memory_nodes)
|
||||
|
||||
# Set output to new memories only (after deduplication)
|
||||
self.memory_nodes = new_memory_nodes
|
||||
|
||||
if not new_memory_nodes:
|
||||
self.output = "No new memory_nodes found (duplicates removed)."
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from loguru import logger
|
|||
from ..base_memory_tool import BaseMemoryTool
|
||||
from ...core.context import C
|
||||
from ...core.enumeration import MemoryType
|
||||
from ...core.schema import MemoryNode
|
||||
from ...core.schema import MemoryNode, Message
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ...mem_agent import BaseMemoryAgent
|
||||
|
|
@ -26,6 +26,7 @@ class SummaryAndHandsOff(BaseMemoryTool):
|
|||
from ...mem_agent import BaseMemoryAgent
|
||||
|
||||
self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)]
|
||||
self.messages: list[Message] = []
|
||||
|
||||
@property
|
||||
def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]:
|
||||
|
|
@ -145,6 +146,9 @@ class SummaryAndHandsOff(BaseMemoryTool):
|
|||
if agent.memory_nodes:
|
||||
self.memory_nodes.extend(agent.memory_nodes)
|
||||
|
||||
if agent.messages:
|
||||
self.messages.extend(agent.messages)
|
||||
|
||||
results.append(
|
||||
{
|
||||
"memory_type": memory_type.value,
|
||||
|
|
|
|||
|
|
@ -111,13 +111,15 @@ class UpdateMemories(BaseMemoryTool):
|
|||
# Get removal IDs
|
||||
memory_ids_to_delete = self.context.get("memory_ids_to_delete", [])
|
||||
memory_ids_to_delete = [m for m in memory_ids_to_delete if m]
|
||||
# Deduplicate memory IDs to avoid redundant deletions
|
||||
memory_ids_to_delete = list(dict.fromkeys(memory_ids_to_delete))
|
||||
|
||||
# Get memories to add
|
||||
memories_to_add = self.context.get("memories_to_add", [])
|
||||
|
||||
# Validate input
|
||||
if not memory_ids_to_delete and not memories_to_add:
|
||||
self.output = "No memories to remove or add. Please provide at least one operation."
|
||||
self.output = "No memories to remove or add. Operation has been done."
|
||||
return
|
||||
|
||||
removed_count = 0
|
||||
|
|
@ -164,6 +166,6 @@ class UpdateMemories(BaseMemoryTool):
|
|||
if operations:
|
||||
self.output = f"Successfully {' and '.join(operations)} in vector_store."
|
||||
else:
|
||||
self.output = "No valid operations performed. Please check your input."
|
||||
self.output = "Operation has been done."
|
||||
|
||||
logger.info(self.output)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ memory_ids_to_delete: |
|
|||
A list of unique identifiers (memory_ids) of the memories to remove.
|
||||
Each ID should be a valid memory_id obtained from previous memory retrieval or addition operations.
|
||||
These memories will be removed before adding the new updated memories.
|
||||
**IMPORTANT**: Do NOT add duplicate memory_ids. Each memory_id should appear only once in the list.
|
||||
|
||||
memories_to_add: |
|
||||
A list of new memory objects to add after removal.
|
||||
|
|
|
|||
|
|
@ -219,9 +219,9 @@ class ReMe(Application):
|
|||
|
||||
if user_id:
|
||||
metadata_desc = {
|
||||
"year": "The year when the memory content occurred.",
|
||||
"month": "The month when the memory content occurred.",
|
||||
"day": "The day when the memory content occurred.",
|
||||
"year": "The year when the message content occurred.",
|
||||
"month": "The month when the message content occurred.",
|
||||
"day": "The day when the message content occurred.",
|
||||
}
|
||||
meta_memories = [
|
||||
{
|
||||
|
|
@ -256,12 +256,12 @@ class ReMe(Application):
|
|||
],
|
||||
)
|
||||
|
||||
try:
|
||||
await reme_summarizer_v2.call(messages=messages, description=description, **kwargs)
|
||||
return personal_summarizer_v2.memory_nodes, personal_summarizer_v2.messages, personal_summarizer_v2.success
|
||||
except Exception as e:
|
||||
print(f"Warning: reme_summarizer_v2.call failed: {e}")
|
||||
return [], [], False
|
||||
# try:
|
||||
await reme_summarizer_v2.call(messages=messages, description=description, **kwargs)
|
||||
return reme_summarizer_v2.memory_nodes, reme_summarizer_v2.messages, reme_summarizer_v2.success
|
||||
# except Exception as e:
|
||||
# print(f"Warning: reme_summarizer_v2.call failed: {e}")
|
||||
# return [], [], False
|
||||
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
@ -305,12 +305,12 @@ class ReMe(Application):
|
|||
],
|
||||
)
|
||||
|
||||
try:
|
||||
await reme_retriever_v2.call(query=query, messages=messages, description=description, **kwargs)
|
||||
return reme_retriever_v2.output, reme_retriever_v2.messages, reme_retriever_v2.success
|
||||
except Exception as e:
|
||||
print(f"Warning: reme_retriever_v2.call failed: {e}")
|
||||
return "error, not retrieved", [], False
|
||||
# try:
|
||||
await reme_retriever_v2.call(query=query, messages=messages, description=description, **kwargs)
|
||||
return reme_retriever_v2.output, reme_retriever_v2.messages, reme_retriever_v2.success
|
||||
# except Exception as e:
|
||||
# print(f"Warning: reme_retriever_v2.call failed: {e}")
|
||||
# return "error, not retrieved", [], False
|
||||
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue