diff --git a/bench/halumem/analyze_results.py b/bench/halumem/analyze_results.py new file mode 100644 index 00000000..4814001e --- /dev/null +++ b/bench/halumem/analyze_results.py @@ -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) diff --git a/bench/halumem/eval_baseline_simple.py b/bench/halumem/eval_baseline_simple.py new file mode 100644 index 00000000..441ffe28 --- /dev/null +++ b/bench/halumem/eval_baseline_simple.py @@ -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 + ) diff --git a/bench/halumem/eval_reme_simple.py b/bench/halumem/eval_reme_simple.py index a7ac6964..aff0934d 100644 --- a/bench/halumem/eval_reme_simple.py +++ b/bench/halumem/eval_reme_simple.py @@ -1,516 +1,662 @@ """ -Simplified evaluation script for ReMe on HaluMem benchmark - Question Answering only. +HaluMem Benchmark Evaluator for ReMe - Question Answering -This script performs a simplified evaluation pipeline: -1. Load HaluMem data -2. Process each user's sessions with ReMe (summary + retrieve) -3. Evaluate question answering only -4. Generate metrics and statistics +A modular evaluation pipeline that: +1. Loads HaluMem benchmark data +2. Processes user sessions through ReMe (summarization + retrieval) +3. Evaluates question answering performance +4. Generates comprehensive metrics Usage: - python bench/halumem/eval_reme_simple.py --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Medium.jsonl \ + python bench/halumem/eval_reme_simple.py \ + --data_path /path/to/HaluMem-Medium.jsonl \ --top_k 20 --user_num 100 --max_concurrency 20 """ import asyncio -import copy 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 ( - _PROMPTS, - evaluation_for_question, - evaluation_for_question2, -) -from llms import llm_request +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 -# Initialize ReMe -reme: ReMe = 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" -def extract_user_name(persona_info: str): - """Extract user name from persona info.""" - match = re.search(r"Name:\s*(.*?); Gender:", persona_info) - if match: - username = match.group(1).strip() - return username - else: - raise ValueError("No name found.") +# ==================== Utilities ==================== - -def iter_jsonl(file_path: str): - """Iterate over lines in a JSONL file.""" - with open(file_path, "r", encoding="utf-8") as f: - for line in f: - line = line.strip() - if line: - yield json.loads(line) - - -# ==================== Main Processing ==================== - - -async def add_memory_async(user_id: str, messages: list[dict]) -> tuple[list[MemoryNode], list, bool, float]: - """Add memory to ReMe system asynchronously.""" - start = time.time() - memory_nodes, agent_messages, success = await reme.summary_v2(messages=messages, user_id=user_id) - duration_ms = (time.time() - start) * 1000 - return memory_nodes, agent_messages, success, duration_ms - - -async def search_memory_async(query: str, user_id: str, top_k: int = 20): - """Search memory and get LLM response directly.""" - start = time.time() - memories, agent_messages, success = await reme.retrieve_v2(query=query, user_id=user_id, top_k=top_k) +class DataLoader: + """Handles loading and parsing of HaluMem data.""" - # Format the context - context = f"User: {user_id}\nMemories:\n{memories}" + @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()] - # Get LLM response directly - prompt = f"Based on the following context, answer the question.\n\nContext:\n{context}\n\nQuestion: {query}\n\nAnswer:" - response = await llm_request(prompt) + @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() - duration_ms = (time.time() - start) * 1000 - return response, agent_messages, success, duration_ms - - -async def process_user_stage1( - user_data: dict, - top_k_value: int, - save_path: str, -): - """Process user data through ReMe (summary + retrieve + QA evaluation only).""" - user_name = extract_user_name(user_data["persona_info"]) - sessions = user_data["sessions"] - - tmp_dir = os.path.join(save_path, "tmp") - os.makedirs(tmp_dir, exist_ok=True) - tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json") - - new_user_data = { - "uuid": user_data["uuid"], - "user_name": user_name, - "sessions": [], - } - - for idx, session in enumerate(sessions): - logger.info(f"Processing user {user_name}: session {idx}/{len(sessions)}") - new_session = { - "memory_points": session["memory_points"], - "dialogue": session["dialogue"], - } - - # Format dialogue - dialogue = session["dialogue"] - formatted_dialogue = [ + @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") + "time_created": datetime.strptime( + turn["timestamp"], "%b %d, %Y, %H:%M:%S" + ) .replace(tzinfo=timezone.utc) .strftime("%Y-%m-%d %H:%M:%S"), } for turn in dialogue ] - - # Process in batches - result = [] - all_agent_messages = [] - all_success_flags = [] - total_duration_ms = 0 - batch_size = 20 - - for i in range(0, len(formatted_dialogue), batch_size): - batch = formatted_dialogue[i : i + batch_size] - batch_result, agent_messages, success, duration_ms = await add_memory_async( - user_id=user_name, - messages=batch, + + @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}" ) - if batch_result: - result.extend(batch_result) - all_agent_messages.append({ - "batch_index": i // batch_size, - "messages": [msg.model_dump() if hasattr(msg, 'model_dump') else str(msg) for msg in agent_messages] if agent_messages else [] - }) - all_success_flags.append({ - "batch_index": i // batch_size, - "success": success - }) + return "\n\n".join(formatted_turns) + + +class FileManager: + """Manages file I/O operations.""" + + def __init__(self, base_dir: str): + self.base_dir = Path(base_dir) + self.tmp_dir = self.base_dir / "tmp" + self.tmp_dir.mkdir(parents=True, exist_ok=True) + + def get_user_dir(self, user_name: str) -> Path: + """Get the directory path for a user.""" + user_dir = self.tmp_dir / user_name + user_dir.mkdir(parents=True, exist_ok=True) + return user_dir + + def get_session_file(self, user_name: str, session_id: int) -> Path: + """Get the file path for a specific session.""" + return self.get_user_dir(user_name) / f"session_{session_id}.json" + + def save_session(self, user_name: str, session_id: int, data: dict): + """Save session data to file.""" + file_path = self.get_session_file(user_name, session_id) + with open(file_path, "w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + logger.info(f"✅ Saved session {session_id} to {file_path}") + + def load_session(self, user_name: str, session_id: int) -> dict | None: + """Load session data from file.""" + file_path = self.get_session_file(user_name, session_id) + if not file_path.exists(): + return None + with open(file_path, "r", encoding="utf-8") as f: + return json.load(f) + + def user_has_cache(self, user_name: str) -> bool: + """Check if user has cached results.""" + user_dir = self.get_user_dir(user_name) + return any(f.name.startswith("session_") and f.suffix == ".json" + for f in user_dir.iterdir()) + + def combine_results(self, output_file: str): + """Combine all user session files into a single JSONL file.""" + with open(output_file, "w", encoding="utf-8") as f_out: + for user_dir in self.tmp_dir.iterdir(): + if not user_dir.is_dir(): + continue + + session_files = sorted([ + f for f in user_dir.iterdir() + if f.name.startswith("session_") and f.suffix == ".json" + ]) + + if not session_files: + continue + + # Load first session to get user metadata + with open(session_files[0], "r", encoding="utf-8") as f_in: + first_session = json.load(f_in) + + user_data = { + "uuid": first_session["uuid"], + "user_name": first_session["user_name"], + "sessions": [] + } + + # Load all sessions + for session_file in session_files: + with open(session_file, "r", encoding="utf-8") as f_in: + session_data = json.load(f_in) + # Remove redundant user metadata + session_data.pop("uuid", None) + session_data.pop("user_name", None) + user_data["sessions"].append(session_data) + + f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n") + + +# ==================== Memory Operations ==================== + +class MemoryProcessor: + """Handles ReMe memory operations.""" + + def __init__(self, reme: ReMe): + self.reme = reme + + async def add_memories( + self, + user_id: str, + messages: list[dict], + batch_size: int = 20 + ) -> tuple[list[str], list[list[dict]], float]: + """ + Add memories in batches and return extracted memory contents. + + Returns: + tuple: (extracted_memories, agent_messages, total_duration_ms) + """ + added_memories: list[MemoryNode] = [] + deleted_memories: list[str] = [] + all_agent_messages: list = [] + total_duration_ms = 0 + for i in range(0, len(messages), batch_size): + batch = messages[i:i + batch_size] + start = time.time() + + memory_nodes, agent_messages, success = await self.reme.summary_v2( + messages=batch, + user_id=user_id + ) + + duration_ms = (time.time() - start) * 1000 total_duration_ms += duration_ms + + # Save agent messages for this batch + if agent_messages: + all_agent_messages.extend(agent_messages) + + if memory_nodes: + for node in memory_nodes: + if isinstance(node, MemoryNode) and node.memory_type == MemoryType.HISTORY: + continue - duration_ms = total_duration_ms + if isinstance(node, MemoryNode): + added_memories.append(node) - # Extract memory content - memories = [] - for memory_node in result: - if isinstance(memory_node, MemoryNode) and memory_node.memory_type is not MemoryType.HISTORY: - memories.append(memory_node.content) + if isinstance(node, str): + deleted_memories.append(node) - if session.get("is_generated_qa_session", False): - new_session["add_dialogue_duration_ms"] = duration_ms - new_session["summary_agent_messages"] = all_agent_messages - new_session["summary_success_flags"] = all_success_flags - new_session["is_generated_qa_session"] = True - del new_session["dialogue"] - del new_session["memory_points"] - new_user_data["sessions"].append(new_session) - continue + extracted_memories = deleted_memories + extracted_memories += ["[delete]" + n.format_memory() for n in added_memories if n.memory_id in deleted_memories] + extracted_memories += ["[add]" + n.format_memory() for n in added_memories if n.memory_id not in deleted_memories] + return extracted_memories, all_agent_messages, total_duration_ms + + async def search_memory( + self, + query: str, + user_id: str, + top_k: int = 20 + ) -> tuple[str, list, float]: + """ + Search memory and return response. + + Returns: + tuple: (response, agent_messages, duration_ms) + """ + start = time.time() + response, agent_messages, success = await self.reme.retrieve_v2( + query=query, + user_id=user_id, + top_k=top_k + ) + duration_ms = (time.time() - start) * 1000 + return response, agent_messages, duration_ms - # Store extracted memories - new_session["extracted_memories"] = memories - new_session["add_dialogue_duration_ms"] = duration_ms - new_session["summary_agent_messages"] = all_agent_messages - new_session["summary_success_flags"] = all_success_flags - # Process questions - if "questions" not in session: - new_user_data["sessions"].append(new_session) - continue +# ==================== Evaluation ==================== - new_session["questions"] = [] - - for qa in session["questions"]: - response, agent_messages, success, duration_ms = await search_memory_async( +class QuestionAnsweringEvaluator: + """Evaluates question answering performance.""" + + def __init__(self, memory_processor: MemoryProcessor, top_k: int): + self.memory_processor = memory_processor + self.top_k = top_k + + async def evaluate_questions( + self, + questions: list[dict], + user_name: str, + uuid: str, + session_id: int, + formatted_dialogue: str + ) -> list[dict]: + """Evaluate all questions for a session.""" + results = [] + + for qa in questions: + # Search memory for answer + response, agent_messages, duration_ms = await self.memory_processor.search_memory( query=qa["question"], user_id=user_name, - top_k=top_k_value, + top_k=self.top_k ) - - new_qa = copy.deepcopy(qa) - new_qa["system_response"] = response - new_qa["search_duration_ms"] = duration_ms - new_qa["retrieve_agent_messages"] = [msg.model_dump() if hasattr(msg, 'model_dump') else str(msg) for msg in agent_messages] if agent_messages else [] - new_qa["retrieve_success"] = success - - new_session["questions"].append(new_qa) - - # ==================== Evaluation for this session ==================== - session_eval_results = { - "question_answering_records": [], - } - - uuid = user_data["uuid"] - - # Evaluate Question Answering - if "questions" in new_session: - logger.info(f"Evaluating Question Answering for session {idx}...") - # Format dialogue for evaluation - # Format: Each turn contains role, content, and time_created - dialogue_for_eval = [] - for turn in dialogue: - dialogue_for_eval.append( - f"Role: {turn['role']}\n" - f"Content: {turn['content']}\n" - f"Time: {datetime.strptime(turn['timestamp'], '%b %d, %Y, %H:%M:%S').replace(tzinfo=timezone.utc).strftime('%Y-%m-%d %H:%M:%S')}" - ) - formatted_dialogue_str = "\n\n".join(dialogue_for_eval) + # 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 + ) - for qa in new_session["questions"]: - new_qa = copy.deepcopy(qa) - new_qa["uuid"] = uuid - new_qa["session_id"] = idx - - result = await evaluation_for_question2( - qa["question"], - qa["answer"], - "\n".join([i["memory_content"] for i in qa["evidence"]]), - qa["system_response"], - formatted_dialogue_str, - ) - result_type = result.get("evaluation_result") - reasoning = result.get("reasoning", "") - new_qa["result_type"] = result_type - new_qa["question_answering_reasoning"] = reasoning - session_eval_results["question_answering_records"].append(new_qa) - - # Store evaluation results in session - new_session["evaluation_results"] = session_eval_results - - new_user_data["sessions"].append(new_session) - - # Save results - with open(tmp_file, "w", encoding="utf-8") as f: - json.dump(new_user_data, f, ensure_ascii=False, indent=2) - session_size = len(new_user_data["sessions"]) - logger.info(f"✅ Saved user {user_name} to {tmp_file} session_size={session_size}") - - logger.info(f"✅ Saved user {user_name} to {tmp_file} all!") - return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file} + # 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 -# ==================== Evaluation Aggregation ==================== - - -def aggregate_eval_results(eval_results): - """Aggregate evaluation results and compute metrics (QA only).""" +class MetricsAggregator: + """Aggregates evaluation metrics.""" - # Question-Answering Evaluation - correct_qa_num = 0 - hallucination_qa_num = 0 - omission_qa_num = 0 - qa_num = 0 - qa_valid_num = 0 - - for item in eval_results["question_answering_records"]: - item["is_valid"] = True - qa_num += 1 - - if item["result_type"] not in ["Correct", "Hallucination", "Omission"]: - item["is_valid"] = False - continue - - if item["result_type"] == "Correct": - correct_qa_num += 1 - elif item["result_type"] == "Hallucination": - hallucination_qa_num += 1 - elif item["result_type"] == "Omission": - omission_qa_num += 1 - - qa_valid_num += 1 - - if qa_num > 0: - eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = correct_qa_num / qa_num - eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = ( - hallucination_qa_num / qa_num - ) - eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = omission_qa_num / qa_num - else: - eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = 0 - eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = 0 - eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = 0 - - if qa_valid_num > 0: - eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = correct_qa_num / qa_valid_num - eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = ( - hallucination_qa_num / qa_valid_num - ) - eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = omission_qa_num / qa_valid_num - else: - eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = 0 - eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = 0 - eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = 0 - - eval_results["overall_score"]["question_answering"]["qa_valid_num"] = qa_valid_num - eval_results["overall_score"]["question_answering"]["qa_num"] = qa_num - - return eval_results + @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 ==================== - -async def main_async( - data_path: str, - top_k: int = 20, - user_num: int = 1, - max_concurrency: int = 2, -): - """Main evaluation pipeline - simplified for QA only.""" - frame = "reme_simple" - save_path = f"bench_results/{frame}/" - os.makedirs(save_path, exist_ok=True) - - output_file_stage1 = os.path.join(save_path, f"{frame}_eval_results.jsonl") - output_file_final = os.path.join(save_path, f"{frame}_eval_stat_result.json") - - start_time = time.time() - await reme.vector_store.delete_all() - - # ==================== Stage 1: Data Processing ==================== - print("\n" + "=" * 80) - print("PROCESSING DATA WITH ReMe (Simplified - QA Only)") - print(f"Max Concurrency: {max_concurrency}") - print("=" * 80) - - tmp_dir = os.path.join(save_path, "tmp") - os.makedirs(tmp_dir, exist_ok=True) - - # Load all user data - user_data_list = list(iter_jsonl(data_path)) - total_users = min(len(user_data_list), user_num) - user_data_list = user_data_list[:total_users] - - print(f"Processing {total_users} users with max concurrency {max_concurrency}...") +class HaluMemEvaluator: + """Main evaluator orchestrating the entire pipeline.""" - # Create semaphore to limit concurrency - semaphore = asyncio.Semaphore(max_concurrency) - - async def process_single_user(idx: int, user_data: dict): - """Process a single user with semaphore control.""" - async with semaphore: - uuid = user_data['uuid'] - tmp_file = os.path.join(tmp_dir, f"{uuid}.json") + def __init__(self, config: EvalConfig): + self.config = config + self.reme = ReMe() + self.file_manager = FileManager(config.output_dir) + self.memory_processor = MemoryProcessor(self.reme) + self.qa_evaluator = QuestionAnsweringEvaluator( + self.memory_processor, + config.top_k + ) + self.data_loader = DataLoader() + + async def process_session( + self, + session: dict, + session_id: int, + user_name: str, + uuid: str + ) -> dict: + """Process a single session.""" + session_data = { + "uuid": uuid, + "user_name": user_name, + "session_id": session_id, + "memory_points": session["memory_points"] + } + + # Skip generated QA sessions + if session.get("is_generated_qa_session", False): + session_data["is_generated_qa_session"] = True + return session_data + + # Format and add dialogue to memory + dialogue = session["dialogue"] + formatted_messages = self.data_loader.format_dialogue_messages(dialogue) + + extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories( + user_id=user_name, + messages=formatted_messages, + batch_size=self.config.batch_size + ) + + session_data.update({ + "dialogue": dialogue, + "extracted_memories": extracted_memories, + "summary_messages": [m.model_dump() for m in agent_messages], + "add_dialogue_duration_ms": duration_ms + }) + + # Evaluate questions if present + if "questions" in session: + formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name) + qa_results = await self.qa_evaluator.evaluate_questions( + questions=session["questions"], + user_name=user_name, + uuid=uuid, + session_id=session_id, + formatted_dialogue=formatted_dialogue + ) - 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} + 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'])}") - print(f"[{idx}/{total_users}] Processing user {uuid}...") - result = await process_user_stage1(user_data, top_k, save_path) - print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})") - return result - - # Process users in parallel with controlled concurrency - tasks = [process_single_user(idx, user_data) for idx, user_data in enumerate(user_data_list, 1)] - await asyncio.gather(*tasks) - - # Combine all results into final output - with open(output_file_stage1, "w", encoding="utf-8") as f_out: - for file in os.listdir(tmp_dir): - if file.endswith(".json"): - file_path = os.path.join(tmp_dir, file) - with open(file_path, "r", encoding="utf-8") as f_in: - data = json.load(f_in) - f_out.write(json.dumps(data, ensure_ascii=False) + "\n") - - elapsed_stage1 = time.time() - start_time - print(f"\n✅ Processing completed in {elapsed_stage1:.2f}s") - print(f"✅ Results saved to: {output_file_stage1}") - - # ==================== Aggregate Results ==================== - print("\n" + "=" * 80) - print("AGGREGATING EVALUATION RESULTS") - print("=" * 80) - - # Calculate time consuming - add_dialogue_duration_time = 0 - search_memory_duration_time = 0 - - for user_data in iter_jsonl(output_file_stage1): - sessions = user_data["sessions"] - - for session in sessions: - if "add_dialogue_duration_ms" in session: - add_dialogue_duration_time += session["add_dialogue_duration_ms"] - - if "questions" in session: - for question in session["questions"]: - if "search_duration_ms" in question: - search_memory_duration_time += question["search_duration_ms"] - - add_dialogue_duration_time = add_dialogue_duration_time / 1000 / 60 - search_memory_duration_time = search_memory_duration_time / 1000 / 60 - - print("\n🔄 Aggregating all user results...") - - eval_results = { - "overall_score": { - "question_answering": {}, - "time_consuming": { - "add_dialogue_duration_time": add_dialogue_duration_time, - "search_memory_duration_time": search_memory_duration_time, - "total_duration_time": add_dialogue_duration_time + search_memory_duration_time, + session_data = await self.process_session( + session=session, + session_id=idx, + user_name=user_name, + uuid=uuid + ) + + self.file_manager.save_session(user_name, idx, session_data) + + return {"uuid": uuid, "user_name": user_name, "status": "ok"} + + async def run_evaluation(self): + """Run the complete evaluation pipeline.""" + start_time = time.time() + + # Clear existing data + await self.reme.vector_store.delete_all() + + # Load user data + all_users = self.data_loader.load_jsonl(self.config.data_path) + users_to_process = all_users[:self.config.user_num] + + print("\n" + "=" * 80) + print("HALUMEM EVALUATION - QUESTION ANSWERING") + print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}") + print("=" * 80 + "\n") + + # Process users with concurrency control + semaphore = asyncio.Semaphore(self.config.max_concurrency) + + async def process_with_cache_check(idx: int, user_data: dict): + async with semaphore: + user_name = self.data_loader.extract_user_name(user_data["persona_info"]) + + # Check cache + if self.file_manager.user_has_cache(user_name): + print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)") + return {"user_name": user_name, "status": "cached"} + + print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...") + result = await self.process_user(user_data) + print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}") + return result + + tasks = [ + process_with_cache_check(idx, user) + for idx, user in enumerate(users_to_process, 1) + ] + await asyncio.gather(*tasks) + + # Combine results + output_file = os.path.join(self.config.output_dir, "eval_results.jsonl") + self.file_manager.combine_results(output_file) + + elapsed = time.time() - start_time + print(f"\n✅ Processing completed in {elapsed:.2f}s") + print(f"📁 Results: {output_file}\n") + + # Aggregate metrics + await self.aggregate_and_report(output_file) + + async def aggregate_and_report(self, results_file: str): + """Aggregate results and generate final report.""" + print("=" * 80) + print("AGGREGATING METRICS") + print("=" * 80 + "\n") + + # Collect all QA records + qa_records = [] + with open(results_file, "r", encoding="utf-8") as f: + for line in f: + if not line.strip(): + continue + user_data = json.loads(line) + + for session in user_data["sessions"]: + if session.get("is_generated_qa_session"): + continue + + eval_results = session.get("evaluation_results", {}) + qa_records.extend( + eval_results.get("question_answering_records", []) + ) + + # Compute metrics + qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records) + time_metrics = MetricsAggregator.compute_time_metrics(results_file) + + final_results = { + "overall_score": { + "question_answering": qa_metrics, + "time_consuming": time_metrics }, - }, - "question_answering_records": [], - } + "question_answering_records": qa_records + } + + # Save final report + report_file = os.path.join(self.config.output_dir, "eval_statistics.json") + with open(report_file, "w", encoding="utf-8") as f: + json.dump(final_results, f, ensure_ascii=False, indent=4) + + print(f"📊 Statistics saved to: {report_file}\n") + + # Print summary + self._print_summary(qa_metrics, time_metrics) + + def _print_summary(self, qa_metrics: dict, time_metrics: dict): + """Print evaluation summary.""" + print("=" * 80) + print("EVALUATION SUMMARY") + print("=" * 80 + "\n") + + print("📊 Question Answering:") + print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}") + print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}") + print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}") + print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}") + print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}") + print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}") + print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}") + + print(f"\n⏱️ Time Metrics:") + print(f" Memory Addition: {time_metrics['add_dialogue_duration_time']:.2f} min") + print(f" Memory Search: {time_metrics['search_memory_duration_time']:.2f} min") + print(f" Total: {time_metrics['total_duration_time']:.2f} min") + print("\n" + "=" * 80) - # Extract QA records from all users - for user_data in iter_jsonl(output_file_stage1): - for session in user_data["sessions"]: - if session.get("is_generated_qa_session", False): - continue - - if "evaluation_results" in session: - eval_results["question_answering_records"].extend( - session["evaluation_results"].get("question_answering_records", []) - ) - - eval_results = aggregate_eval_results(eval_results) - - with open(output_file_final, "w", encoding="utf-8") as f: - json.dump(eval_results, f, ensure_ascii=False, indent=4) - - elapsed_total = time.time() - start_time - print(f"\n✅ All done in {elapsed_total:.2f}s. Results saved to {output_file_final}") - - # Print summary - print("\n" + "=" * 80) - print("EVALUATION SUMMARY (Question Answering Only)") - print("=" * 80) - - print(f"\n📊 Question Answering:") - print( - f" - Correct (all): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(all)', 0):.4f}", - ) - print( - f" - Hallucination (all): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(all)', 0):.4f}", - ) - print( - f" - Omission (all): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(all)', 0):.4f}", - ) - print( - f" - Correct (valid): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(valid)', 0):.4f}", - ) - print( - f" - Hallucination (valid): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(valid)', 0):.4f}", - ) - print( - f" - Omission (valid): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(valid)', 0):.4f}", - ) - print( - f" - Valid QA: {eval_results['overall_score']['question_answering'].get('qa_valid_num', 0)}/{eval_results['overall_score']['question_answering'].get('qa_num', 0)}", - ) - - print(f"\n⏱️ Time Consuming:") - print(f" - Add Dialogue: {add_dialogue_duration_time:.2f} min") - print(f" - Search Memory: {search_memory_duration_time:.2f} min") - print(f" - Total: {add_dialogue_duration_time + search_memory_duration_time:.2f} min") - print("=" * 80) +# ==================== Entry Point ==================== def main( data_path: str, top_k: int = 20, user_num: int = 1, - max_concurrency: int = 2, + max_concurrency: int = 2 ): - """Synchronous entry point.""" - asyncio.run(main_async(data_path, top_k, user_num, max_concurrency)) + """Main entry point.""" + config = EvalConfig( + data_path=data_path, + top_k=top_k, + user_num=user_num, + max_concurrency=max_concurrency + ) + + evaluator = HaluMemEvaluator(config) + asyncio.run(evaluator.run_evaluation()) if __name__ == "__main__": import argparse - - parser = argparse.ArgumentParser(description="Simplified evaluation for ReMe on HaluMem benchmark (QA only)") + + parser = argparse.ArgumentParser( + description="Evaluate ReMe on HaluMem benchmark (Question Answering)" + ) parser.add_argument( "--data_path", type=str, required=True, - help="Path to HaluMem data file (e.g., HaluMem-medium.jsonl)", + help="Path to HaluMem JSONL file" ) parser.add_argument( "--top_k", type=int, default=20, - help="Number of top memories to retrieve (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)", + help="Number of users to evaluate (default: 1)" ) parser.add_argument( "--max_concurrency", type=int, default=2, - help="Maximum concurrency for processing (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, + max_concurrency=args.max_concurrency ) diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index f2ff5587..fb52dc0d 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -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, diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py index 6eb8cf9d..13395212 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py @@ -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 diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.yaml b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.yaml index 9a70d51b..5f8a74e0 100644 --- a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.yaml +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.yaml @@ -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 [] : ` 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 [] : ` - - 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: diff --git a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py index a580806a..1aae4ad4 100644 --- a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py @@ -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": { diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index 2c2d4e50..6accfe2a 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -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 diff --git a/reme_ai/mem_tool/v2/add_memory_drafts.py b/reme_ai/mem_tool/v2/add_memory_drafts.py index ba66c0e1..93caad0c 100644 --- a/reme_ai/mem_tool/v2/add_memory_drafts.py +++ b/reme_ai/mem_tool/v2/add_memory_drafts.py @@ -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) diff --git a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py index 5aba25ab..38107ab0 100644 --- a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py @@ -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: diff --git a/reme_ai/mem_tool/v2/summary_and_hands_off.py b/reme_ai/mem_tool/v2/summary_and_hands_off.py index 4a595578..131b1101 100644 --- a/reme_ai/mem_tool/v2/summary_and_hands_off.py +++ b/reme_ai/mem_tool/v2/summary_and_hands_off.py @@ -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, diff --git a/reme_ai/mem_tool/v2/update_memories.py b/reme_ai/mem_tool/v2/update_memories.py index 48576e2e..03a9f494 100644 --- a/reme_ai/mem_tool/v2/update_memories.py +++ b/reme_ai/mem_tool/v2/update_memories.py @@ -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) diff --git a/reme_ai/mem_tool/v2/update_memories.yaml b/reme_ai/mem_tool/v2/update_memories.yaml index 717762af..9aeec685 100644 --- a/reme_ai/mem_tool/v2/update_memories.yaml +++ b/reme_ai/mem_tool/v2/update_memories.yaml @@ -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. diff --git a/reme_ai/reme.py b/reme_ai/reme.py index 3aed1afa..e3aaf54d 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -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